[flashinfer] NVIDIA Blackwell 아키텍처를 위한 FlashInfer의 Router GEMM 최적화
PR 링크: flashinfer-ai/flashinfer#4594 상태: Merged | 변경: +5653 / -5
들어가며
최근 대규모 언어 모델(LLM)의 발전과 함께 GPU 하드웨어의 성능 향상 또한 가속화되고 있습니다. 특히 NVIDIA의 최신 Blackwell 아키텍처는 이전 세대 대비 상당한 성능 개선을 약속하며, 이를 활용하기 위한 소프트웨어 최적화의 중요성이 더욱 커지고 있습니다.
이번 글에서는 LLM 추론을 위한 고성능 커널 라이브러리인 FlashInfer의 Pull Request(PR) #1276을 분석합니다. 이 PR은 NVIDIA Blackwell GPU 아키텍처에 맞춰 Router GEMM(General Matrix Multiply) 연산을 최적화하여, 특정 워크로드에서 최대 1.1x 이상의 성능 향상을 달성했습니다. 이 PR이 어떻게 이러한 성능 개선을 이루었는지, 코드 변경 사항을 중심으로 자세히 살펴보겠습니다.
코드 분석: Blackwell Router GEMM 추가
이번 PR의 핵심은 csrc/cake_router_gemm/ 디렉토리에 새로운 CUDA 커널 파일들을 추가한 것입니다. 이 파일들은 NVIDIA Blackwell 아키텍처의 특성을 활용하여 Router GEMM 연산을 효율적으로 처리하도록 설계되었습니다. 주요 변경 사항은 다음과 같습니다.
1. 새로운 커널 파일 추가 (cake_router_gemm_m10_k6144_device.cu, cake_router_gemm_m10_k7168_device.cu 등)
PR은 cake_router_gemm 디렉토리에 여러 .cu 파일을 추가했습니다. 예를 들어 cake_router_gemm_m10_k6144_device.cu와 cake_router_gemm_m10_k7168_device.cu는 각각 다른 K 차원을 가진 행렬 곱셈 연산을 위한 Blackwell GPU용 커널을 정의합니다. 이 파일들은 기존의 FlashInfer 커널과는 별개로, Blackwell 아키텍처에 특화된 최적화를 적용합니다.
Before (변경 없음 - 기존 커널은 이 PR에 포함되지 않음)
이 PR은 새로운 커널을 추가하는 것이므로, 'Before' 상태는 해당 기능이 존재하지 않았음을 의미합니다.
After (새로운 커널 코드 예시 - cake_router_gemm_m10_k6144_device.cu)
__global__ __launch_bounds__(128, 1) void kernel_cake_blackwell_router_gemm_m10_k6144(
__nv_bfloat16* __restrict__ mat_a, __nv_bfloat16* __restrict__ mat_b,
float* __restrict__ out_f32, __nv_bfloat16* __restrict__ out_bf16, int num_experts,
int out_is_bf16) {
const int tid = threadIdx.x;
const int warp = make_warp_uniform(tid / 32);
const int lane = tid % 32;
extern __shared__ __align__(1024) char smem_raw[];
int smem;
smem = (int)(unsigned long long)__cvta_generic_to_shared(smem_raw);
const int bid = blockIdx.x;
const int num_bids = gridDim.x;
// Kernel setup ops
float* red = reinterpret_cast<float*>(smem_raw + 0);
const int red_addr = smem + 0;
// === Task calls (dependency order) ===
int expert = blockIdx.x;
float acc[10];
#pragma unroll
for (int m = 0; m < 10; m++) {
acc[m] = 0.0f;
}
asm volatile("griddepcontrol.wait;" ::: "memory");
#pragma unroll
for (int ki = 0; ki < 6; ki++) { // K dimension loop
int k_base = ki * 1024 + tid * 8;
float _vec_load_0[8];
{ // Load from mat_b
const uint4* _vptr_0 = reinterpret_cast<const uint4*>(mat_b + (expert * 6144 + k_base) + 0);
// ... (vector load operations) ...
}
#pragma unroll
for (int m_1 = 0; m_1 < 10; m_1++) { // M dimension loop
float _vec_load_1[8];
{ // Load from mat_a
const uint4* _vptr_1 = reinterpret_cast<const uint4*>(mat_a + (m_1 * 6144 + k_base) + 0);
// ... (vector load operations) ...
}
#pragma unroll
for (int j = 0; j < 8; j++) {
float _fma_0 = __fmaf_rn(_vec_load_1[j], _vec_load_0[j], acc[m_1]);
acc[m_1] = _fma_0;
}
}
}
#pragma unroll
for (int m_2 = 0; m_2 < 10; m_2++) { // Warp reduction
float _warp_reduce_0 = acc[m_2];
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
_warp_reduce_0 += __shfl_xor_sync(0xFFFFFFFF, _warp_reduce_0, offset);
float warp_sum = _warp_reduce_0;
if (lane == 0) {
red[m_2 * 4 + warp] = warp_sum;
}
}
asm volatile("barrier.sync 2, 128;" ::: "memory");
if (warp == 0) { // Global reduction and store
if (lane < 10) {
int m_3 = lane;
float total = red[m_3 * 4];
#pragma unroll
for (int source_warp = 1; source_warp < 4; source_warp++) {
total = total + red[m_3 * 4 + source_warp];
}
int offset = m_3 * num_experts + expert;
if (out_is_bf16 == 0) {
*(reinterpret_cast<float*>(out_f32 + offset) + (0)) = total;
} else {
*(reinterpret_cast<__nv_bfloat16*>(out_bf16 + offset) + (0)) = __float2bfloat16_rn(total);
}
}
}
asm volatile("griddepcontrol.launch_dependents;" ::: "memory");
}
이 커널은 다음과 같은 특징을 가집니다:
__nv_bfloat16사용: 입력 행렬mat_a와mat_b가__nv_bfloat16타입으로 선언되어, FP16보다 넓은 동적 범위를 가지면서도 FP32보다 메모리 대역폭 및 연산 효율성을 높입니다. 이는 최신 GPU 아키텍처에서 권장되는 데이터 타입입니다.__fmaf_rn사용: Fused Multiply-Add 연산을 사용하여 연산 정밀도를 유지하면서 성능을 향상시킵니다._rn접미사는 round-to-nearest 모드를 의미합니다.- 벡터 로딩 및 FMA:
uint4를 사용하여 데이터를 로드하고, 이를float으로 변환하여 FMA 연산을 수행합니다. 이는 SIMD(Single Instruction, Multiple Data) 명령어를 활용하여 데이터 처리량을 극대화하려는 시도입니다. - Warp-level 및 Grid-level Reduction:
__shfl_xor_sync를 이용한 warp 내 reduction과 공유 메모리(smem)를 이용한 grid 간 reduction을 통해 최종 결과를 효율적으로 집계합니다. - Blackwell 특화 어셈블리: `asm volatile(
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입
- [flashinfer] FlashInfer, FP8 지원으로 장문 컨텍스트 추론 성능을 극적으로 향상시키다
- [flashinfer] FlashInfer SM12x MoE 최적화: 정적 MoE 경로 통합 및 성능 향상
- [sglang] SGLang, NVIDIA Blackwell GPU를 위한 Wan2.2 모델의 ML P 연산 최적화
- [flashinfer] FlashInfer SM120 MoE GEMM 최적화: 웨이브+잔여물 비용 모델 도입
PR Analysis 의 다른글
- 이전글 [flashinfer] DeepSeek-V3 라우팅의 혁신: FlashInfer의 Cake 백엔드 가속 분석
- 현재글 : [flashinfer] NVIDIA Blackwell 아키텍처를 위한 FlashInfer의 Router GEMM 최적화
- 다음글 [flashinfer] Blackwell 아키텍처를 위한 MXFP8/MXFP4 기반 고성능 MoE 추론 최적화
댓글