본문으로 건너뛰기

[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.cucake_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_amat_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를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글