본문으로 건너뛰기

[flashinfer] FlashInfer BF16 KDA 성능 최적화: M64 Value Split 도입

PR 링크: flashinfer-ai/flashinfer#5363 상태: Merged | 변경: +8497 / -7

들어가며

최근 대규모 언어 모델(LLM)의 발전과 함께, Transformer 아키텍처의 핵심 연산인 Key-Query-Value (KVA) 어텐션 메커니즘의 효율성을 높이는 것은 매우 중요한 과제가 되었습니다. 특히, 추론(inference) 단계에서의 지연 시간 감소는 사용자 경험에 직접적인 영향을 미치기 때문에, GPU 커널 수준에서의 최적화는 필수적입니다.

이번 글에서는 flashinfer-ai/flashinfer 레포지토리의 "perf(cake_kda): route one-wave BF16 grids to the M64 value split" PR을 분석합니다. 이 PR은 BF16 데이터 타입을 사용하는 KDA 연산에서 특정 조건에 맞는 연산들을 기존의 M128 타일 대신 더 효율적인 M64 Value Split 방식으로 라우팅하여 성능을 개선하는 것을 목표로 합니다. 이전 PR(#5278)에서 제기되었던 B200 GPU에서의 BF16 KDA 사전 채우기(prefill) 성능 저하 문제를 해결하기 위한 후속 작업입니다.

코드 변경사항 분석

이번 PR의 핵심은 BF16 연산에서 최적의 커널을 선택하는 디스패처 로직을 개선하여, 특정 워크로드에 대해 더 효율적인 M64 Value Split 커널을 사용하도록 유도하는 것입니다. 주요 변경 사항은 다음과 같습니다.

1. flashinfer/cake_kda_tf32_runtime.py의 호스트 디스패처 업데이트

이 PR은 이전 PR(#5278)에서 소스 코드 변경 사항을 미러링하여 BF16 호스트 디스패처를 업데이트했습니다. 특히, _should_use_bf16_one_wave_dvsplit 함수의 로직이 수정되어, 특정 조건을 만족하는 경우 M64 Value Split 경로를 선택하도록 변경되었습니다.

주요 변경 조건:

  • One-wave rule: BF16 연산, bounded gate, 32-토큰 정렬된 체크포인트 간격, 그리고 2 * tasks <= sm_count (즉, M64 Value CTA가 각 (sequence, head) 태스크를 한 웨이브 내에서 처리하는 경우) 조건을 만족하면, 기존의 M128 경로 대신 M64 Value Split 경로를 사용합니다.
  • TF32 및 강제 타일: TF32 연산과 강제 타일(forced tiles) 설정은 변경되지 않습니다.
  • Logit beta: 토큰 피치가 TMA(Tensor Memory Accelerator) 인코딩이 불가능한 경우, 매 런치마다 패딩된 캐리어로 새로고침됩니다. 짧은 H12 그리드(<= 256 토큰)는 H12 직접 타일이 스칼라 베타를 로드하기 때문에 기존 직접 타일을 유지합니다.
  • Active FP32 beta (beta_is_logit=False): BF16 어파인 스플릿이 가능한 경우를 제외하고는, 매 런치마다 logit + BF16 복사 대신 네이티브 M64 경로를 사용합니다.

이 변경은 B200 및 GB300과 같은 최신 GPU 아키텍처에서 BF16 연산의 성능을 극대화하기 위해 설계되었습니다.

2. 새로 생성된 M64 모듈

변경된 디스패처 로직에 따라 선택되는 새로운 M64 모듈들이 추가되었습니다. 이는 csrc/kda/bf16/ 디렉토리에 .cu 파일로 포함됩니다. 예를 들어, cake_kda_bf16_712f0e663d1a87d89dba5b808eb25c3ef503950105de552c970a8a32a5774713_binding.cu 파일은 이러한 M64 커널의 바인딩 및 관련 CUDA 코드를 포함합니다.

--- /dev/null
+++ b/csrc/kda/bf16/cake_kda_bf16_712f0e663d1a87d89dba5b808eb25c3ef503950105de552c970a8a32a5774713_binding.cu
@@ -0,0 +1,759 @@
+/*
+ * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
+ * ... (코드 생략) ...
+ */
+
+// Generated by the Cake source exporter — do not edit.
+// tvm-ffi direct-source launcher for Cake kernel 'kernel_cake_kda_bf16_712f0e663d1a87d89dba5b808eb25c3ef503950105de552c970a8a32a5774713'.
+#include <cuda.h>
+#include <cuda_bf16.h>
+#include <cuda_fp16.h>
+#include <cuda_runtime.h>
+
+#include "tvm_ffi_utils.h"
+
+#include <atomic>
+#include <cstdint>
+#include <cstring>
+#include <mutex>
+#include <string>
+#include <unordered_map>
+#include <vector>
+
+struct CakeTensorMap;
+extern "C" __global__ void kernel_cake_kda_bf16_712f0e663d1a87d89dba5b808eb25c3ef503950105de552c970a8a32a5774713(__nv_bfloat16* __restrict__ q, CakeTensorMap const* q_tma, __nv_bfloat16* __restrict__ k, CakeTensorMap const* k_tma, __nv_bfloat16* __restrict__ v, CakeTensorMap const* v_tma, __nv_bfloat16* __restrict__ g, CakeTensorMap const* g_tma, __nv_bfloat16* __restrict__ beta, CakeTensorMap const* beta_tma, float* __restrict__ A_log, float* __restrict__ dt_bias, long long* __restrict__ cu_seqlens, int* __restrict__ seq_order, __nv_bfloat16* __restrict__ initial_state, __nv_bfloat16* __restrict__ out, CakeTensorMap const* out_tma, __nv_bfloat16* __restrict__ final_state, unsigned long long state_indices_addr, long long state_slot_stride, int use_state_indices, float* __restrict__ initial_state_f32, float* __restrict__ final_state_f32, unsigned long long state_checkpoints_addr, unsigned long long checkpoint_cu_starts_addr, float* __restrict__ beta_active_out, long long beta_token_stride, long long g_token_stride, int checkpoint_every_n_tokens, int num_heads, int use_initial_state, int store_final_state, float scale, float lower_bound);
+ ... (코드 생략) ...
+```

이 코드는 `kernel_cake_kda_bf16_712f0e663d1a87d89dba5b808eb25c3ef503950105de552c970a8a32a5774713`라는 이름의 CUDA 커널을 정의하고, 이를 TVM FFI를 통해 호출할 수 있도록 하는 래퍼(wrapper) 역할을 합니다. 이는 M64 Value Split 전략을 사용하는 새로운 최적화된 커널임을 나타냅니다.

### 3. TMA Descriptor 인코딩 및 업로드 로직

PR에는 TMA(Tensor Memory Accelerator) 관련 로직도 포함되어 있습니다. 이는 GPU 메모리에서 레지스터로 데이터를 효율적으로 로드하기 위한 메커니즘입니다. 특히, `EncodeTma_q_tma` 함수는 입력 텐서 `q_tma`에 대한 TMA 디스크립터를 생성하는 로직을 보여줍니다.

```diff
@@ -136,12 +136,12 @@
   // descriptor's std.Expr global_dim/global_strides/checks record.
   inline CUtensorMap EncodeTma_q_tma(const TensorView& t) {
   TVM_FFI_CHECK(t.ndim() >= 2, ValueError)
-      << "TMA source 'q_tma' must have at least 2 dimensions, got ndim=" << t.ndim();
+      << "TMA source 'q_tma' must have at least 2 dimensions, got ndim=" << t.ndim();
   TVM_FFI_CHECK(t.stride(-1) == 1, ValueError)
-      << "TMA source 'q_tma' must have unit innermost stride, got " << t.stride(-1);
+      << "TMA source 'q_tma' must have unit innermost stride, got " << t.stride(-1);
   int64_t d1 = t.size(t.ndim() - 1);
   int64_t d2 = t.size(t.ndim() - 2);
   TVM_FFI_CHECK(d1 > 0 && d2 > 0, ValueError)
-      << "TMA source 'q_tma' trailing dims must be positive";
+      << "TMA source 'q_tma' trailing dims must be positive";
   int64_t outer2 = t.numel() / (d1 * d2);
   TVM_FFI_CHECK(d1 % 64 == 0, ValueError)
-      << "TMA source 'q_tma' extent " << d1
+      << "TMA source 'q_tma' extent " << d1
       << " must divide exactly by " << 64;
   uint64_t global_dim[4] = {(uint64_t)(64), (uint64_t)(outer2), (uint64_t)(d2), (uint64_t)((d1 / 64))};
   TVM_FFI_CHECK(global_dim[0] > 0 && global_dim[1] > 0 && global_dim[2] > 0 && global_dim[3] > 0, ValueError)
-      << "TMA descriptor for 'q_tma' resolved a non-positive global dim";
+      << "TMA descriptor for 'q_tma' resolved a non-positive global dim";
   TVM_FFI_CHECK(64u <= global_dim[0] && 1u <= global_dim[2] && 2u <= global_dim[3], ValueError)
-      << "TMA box (64, 32, 1, 2) exceeds resolved global dims for 'q_tma'";
+      << "TMA box (64, 32, 1, 2) exceeds resolved global dims for 'q_tma'";
   int64_t carrier_stride_0 = (d2 * d1);
   TVM_FFI_CHECK(carrier_stride_0 >= 0, ValueError)
-      << "TMA descriptor for 'q_tma' resolved global stride 1 negative";
+      << "TMA descriptor for 'q_tma' resolved global stride 1 negative";
   TVM_FFI_CHECK(carrier_stride_0 != 0 || global_dim[1] == 1, ValueError)
-      << "TMA descriptor for 'q_tma' resolved global stride 1 zero while global dimension 1 is not 1";
+      << "TMA descriptor for 'q_tma' resolved global stride 1 zero while global dimension 1 is not 1";
   TVM_FFI_CHECK((carrier_stride_0 * 16) % 8 == 0, ValueError)
-      << "TMA descriptor for 'q_tma' resolved global stride 1 to a non-whole-byte offset";
+      << "TMA descriptor for 'q_tma' resolved global stride 1 to a non-whole-byte offset";
   int64_t carrier_stride_1 = d1;
   TVM_FFI_CHECK(carrier_stride_1 >= 0, ValueError)
-      << "TMA descriptor for 'q_tma' resolved global stride 2 negative";
+      << "TMA descriptor for 'q_tma' resolved global stride 2 negative";
   TVM_FFI_CHECK(carrier_stride_1 != 0 || global_dim[2] == 1, ValueError)
-      << "TMA descriptor for 'q_tma' resolved global stride 2 zero
+      << "TMA descriptor for 'q_tma' resolved global stride 2 zero

이 코드는 TMA 디스크립터의 global_dimglobal_strides를 계산하는 로직을 보여줍니다. 특히 d1 % 64 == 0과 같은 조건들은 M64 Value Split 커널이 특정 데이터 레이아웃에서 최적의 성능을 내도록 설계되었음을 시사합니다.

왜 이게 좋은가?

이 PR은 다음과 같은 이유로 좋은 최적화/개선이라고 할 수 있습니다.

  1. 성능 향상:

    • PR 설명에 따르면, B200 GPU에서 short/H12/mixed 케이스에서 0.935배의 속도 향상을 보였으며, 다른 여러 케이스에서도 TF32 대비 2-4배 빠른 성능을 달성했습니다.
    • "Per-shape effect on the 346-row inventory" 섹션의 표는 다양한 워크로드에서 성능 개선이 이루어졌음을 보여줍니다. 예를 들어, B200에서 short/H12/mixed 케이스는 63.6 us에서 29.8 us로 크게 단축되었습니다 (2.13x 향상).
    • GB300 GPU에서도 유사한 성능 개선이 관찰되었습니다. 예를 들어, H6 1x64 CP64 serving 케이스는 41.9 us에서 11.3 us로 3.71배 빨라졌습니다.
    • 전반적으로 346개의 테스트 케이스 중 138개(B200) 및 143개(GB300)에서 경로 변경이 발생했으며, 이들 대부분에서 상당한 성능 향상이 있었습니다. 어떤 경로도 느려지지 않았습니다.
  2. 정밀한 조건 기반 라우팅:

    • 이전에는 모든 BF16 KDA 연산이 동일한 방식으로 처리되었을 수 있지만, 이 PR은 특정 조건(one-wave rule, active FP32 beta 등)을 만족하는 경우에만 M64 Value Split을 사용하도록 세밀하게 제어합니다. 이는 불필요한 오버헤드를 줄이고, 해당 조건에 최적화된 커널을 사용함으로써 성능을 극대화합니다.
    • 특히, M64 Value Split이 더 느린 경우(예: B200 H12 96x128 케이스에서 N16 24.7 us vs M64 34.3 us)에는 기존의 직접 타일(direct tiles)을 유지하여 성능 저하를 방지합니다.
  3. 정확성 보장:

    • PR 설명에 따르면, M64 출력/상태/체크포인트는 기존 직접 타일과 비교했을 때 최대 2.7e-4 / 4.8e-3 / 5.9e-3 이내의 오차만 발생합니다. 이는 BF16의 허용 오차(1e-2) 범위 내에 있으며, 모델의 정확성에 영향을 미치지 않으면서 성능을 개선할 수 있음을 의미합니다.
    • "Export parity" 섹션에서는 B200 및 GB300에서 346개의 모든 테스트 케이스에 대해 스케줄 패리티와 비트 단위 출력/상태/체크포인트 패리티를 달성했음을 보여줍니다. 이는 새로운 최적화가 기능적으로 동일함을 보장합니다.
  4. 일반적인 교훈:

    • 하드웨어 특성 활용: 최신 GPU 아키텍처(B200, GB300)의 특정 기능(예: M64 Value Split)을 활용하여 성능을 개선할 수 있습니다.
    • 조건부 최적화: 모든 경우에 단일 최적화 기법을 적용하기보다는, 워크로드의 특성(데이터 타입, 시퀀스 길이, 게이트 조건 등)에 따라 가장 효율적인 알고리즘이나 커널로 동적으로 라우팅하는 것이 중요합니다.
    • 정확성 검증: 성능 최적화 과정에서 수치적 정확성 저하가 발생하지 않도록 철저한 검증이 필요합니다. 비트 단위 비교 및 허용 오차 내에서의 검증은 필수적입니다.
    • 테스트 커버리지: 다양한 시나리오(shape grid, short grid, logit beta, active beta 등)와 하드웨어(B200, GB300)에 대한 광범위한 테스트는 최적화의 효과와 안정성을 보장하는 데 중요합니다.

맺음말

이번 PR은 FlashInfer 라이브러리의 BF16 KDA 연산 성능을 크게 향상시키는 중요한 개선 사항을 포함하고 있습니다. M64 Value Split을 특정 조건에서 효과적으로 활용하도록 디스패처 로직을 수정함으로써, 최신 GPU 하드웨어에서 LLM 추론 속도를 높이는 데 기여했습니다. 이러한 종류의 세밀한 커널 최적화는 AI 모델의 효율성을 극대화하는 데 핵심적인 역할을 합니다.

References

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글