본문으로 건너뛰기

[flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입

PR 링크: flashinfer-ai/flashinfer#4571 상태: Merged | 변경: +3384 / -510

들어가며

최근 대규모 언어 모델(LLM)의 발전은 모델의 추론 속도, 특히 긴 시퀀스를 처리하는 능력에 대한 요구를 증대시키고 있습니다. NVIDIA의 최신 Blackwell 아키텍처 GPU는 이러한 요구를 충족시키기 위한 강력한 성능을 제공하지만, 이를 최대한 활용하기 위해서는 소프트웨어 스택의 최적화가 필수적입니다.

FlashInfer는 LLM 추론을 위한 고성능 커널 라이브러리로, 이번 PR(Pull Request)에서는 특히 Blackwell 아키텍처의 CC 10.0 및 10.3에서 recurrent-KDA prefill 작업의 성능을 최적화하는 새로운 small-BH (Small Batch Head) 커널을 도입했습니다. 이 글에서는 해당 PR의 코드 변경 사항을 분석하고, 왜 이러한 최적화가 중요한지, 그리고 어떤 기술적 개선이 이루어졌는지 상세히 살펴보겠습니다.

코드 분석

이번 PR의 핵심은 cake_kda 백엔드에 small-BH 레이아웃을 가진 새로운 커널을 추가하고, 이를 적절한 조건에서 선택하도록 라우팅 로직을 개선한 것입니다. 변경 사항은 주로 벤치마크 스크립트, CUDA 커널 구현, 그리고 관련 테스트 파일에 집중되어 있습니다.

1. 벤치마크 및 테스트 설정 변경 (benchmarks/README.md, benchmarks/bench_recurrent_kda_prefill.py)

새로운 small-BH 케이스를 벤치마크에 추가하고, 이를 선택할 수 있도록 옵션을 확장했습니다.

Before:

--- a/benchmarks/bench_recurrent_kda_prefill.py
+++ b/benchmarks/bench_recurrent_kda_prefill.py
@@ -14,9 +14,9 @@
 
 """CUPTI benchmark for recurrent-KDA prefill public API shapes.
 
-The default case set combines the original H64/H96 coverage with six H12
-shapes representing Kimi-K3's per-rank head count under TP8.  ``--case-set``
-can select either group independently.
+The default case set combines the original H64/H96 coverage, six H12 shapes
+representing Kimi-K3's per-rank head count under TP8, and four fixed-layout
+small-BH shapes. ``--case-set`` can select each group independently.
 
 The FlashInfer candidate is always invoked through the public
 ``recurrent_kda`` API. ``--candidate-route dispatcher`` measures the natural
@@ -148,7 +148,13 @@ def _load_h12_cases(path: Path = H12_PRESET) -> tuple[Case, ...]:
     Case("h64_uniform", 64, (1024,) * 8, True, 10005),
 )
H12_CASES = _load_h12_cases()
-CASES = LEGACY_CASES + H12_CASES
+SMALL_BH_CASES = (
+    Case("h8_fixed_65536", 8, (65536,), False, 11000),
+    Case("h4_fixed_65536_holdout", 4, (65536,), False, 11001),
+    Case("h1_fixed_131072", 1, (131072,), False, 11002),
+    Case("h1_fixed_1048576", 1, (1048576,), False, 11003),
+)
+CASES = LEGACY_CASES + H12_CASES + SMALL_BH_CASES
 
 
def _require_cupti() -> None:
@@ -560,16 +566,19 @@ def main() -> None:
     parser.add_argument("--bench-ms", type=int, default=100)
     parser.add_argument(
         "--case-set",
-        choices=("all", "legacy", "h12"),
+        choices=("all", "legacy", "h12", "small_bh"),
         default="all",
-        help="Run all cases, the original H64/H96 cases, or the Kimi-K3 TP8 H12 cases.",
+        help=(
+            "Run all cases, the original H64/H96 cases, the Kimi-K3 TP8 H12 "
+            "cases, or the fixed-layout small-BH cases."
+        ),
     )
     parser.add_argument(
         "--state-rotations",
         type=int,
         help=(
             "Override the number of preinitialized same-input state slots per "
-            "mutable path. By default legacy cases use "
+            "mutable path. By default legacy and small-BH cases use "
             f"{DEFAULT_LEGACY_STATE_ROTATIONS} slots and H12 cases use "
             f"{DEFAULT_H12_STATE_ROTATIONS} slots."
         ),
@@ -656,6 +665,7 @@ def main() -> None:
         "all": CASES,
         "legacy": LEGACY_CASES,
         "h12": H12_CASES,
+        "small_bh": SMALL_BH_CASES,
     }[args.case_set]
     results = []
     for case in selected_cases:

After:

bench_recurrent_kda_prefill.py 스크립트에서 CASES 리스트에 SMALL_BH_CASES를 추가하고, --case-set 인자에 small_bh 옵션을 추가하여 새로운 케이스 그룹을 선택할 수 있도록 변경했습니다. 또한, README.md 파일에도 small_bh 케이스를 CUPTI 경로로 실행하는 방법이 추가되었습니다.

2. 새로운 Small-BH 커널 구현 (csrc/kda/cake_flashkda_bf16_small_bh_m128.cu)

이 파일은 Blackwell 아키텍처(CC 10.0, 10.3)에 특화된 small-BH 레이아웃을 위한 새로운 CUDA 커널을 정의합니다. 이 커널은 bf16 데이터 타입을 사용하며, m128 (128 threads per CTA) 설정에 최적화되어 있습니다.

--- /dev/null
+++ b/csrc/kda/cake_flashkda_bf16_small_bh_m128.cu
@@ -0,0 +1,2632 @@
+/*
+ * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
+ * ... (라이선스 정보) ...
+ */
+
+// Frozen generated kernel export; do not edit by hand.
+// Generated schedule 'flashkda_bf16_small_bh_m128'; module
+// flashkda_bf16_small_bh_m128_73369168de.
+// clang-format off
+
+typedef unsigned char      uint8_t;
+typedef unsigned short     uint16_t;
+typedef unsigned int       uint32_t;
+typedef unsigned long long uint64_t;
+typedef signed int         int32_t;
+typedef short int          int16_t;
+struct __align__(128) FlashKDATensorMap { uint64_t opaque[16]; };
+template <int N>
+struct __align__(128) FlashKDATensorMapPack { FlashKDATensorMap maps[N]; };
+
+typedef struct __align__(64) { uint64_t opaque[16]; } CUtensorMap;
+
+#include <cuda_bf16.h>
+
+#define FLASH_KDA_INF CUDART_INF_F
+#define TMEM_NCOLS 256
+// ... (메모리 오프셋 및 스트라이드 정의) ...
+#define SMEM_TOTAL 227328
+#define THREADS 1024
+
+#include <math_constants.h>
+
+__device__ __forceinline__ uint32_t elect_sync() { ... }
+__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) { ... }
+__device__ __forceinline__ uint32_t mbarrier_try_wait(int mbar_addr, int phase) { ... }
+__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) { ... }
+__device__ __forceinline__ void tcgen05_mma_f16(
+    int taddr, uint64_t a_desc, uint64_t b_desc,
+    uint32_t i_desc, int enable_input_d) { ... }
+__device__ __forceinline__ uint64_t desc_encode(uint64_t x) { ... }
+__device__ __forceinline__ void mma_ts_step(
+    int taddr_out, int taddr_a, int b_lo, uint32_t b_dhi,
+    uint32_t i_desc, int enable_d) { ... }
+
+// ... (실제 커널 로직) ...
+```

이 파일에는 `elect_sync`, `mbarrier_init`, `mbarrier_wait`, `tcgen05_mma_f16`, `mma_ts_step` 등 Blackwell 아키텍처의 새로운 기능을 활용하는 저수준 CUDA 프리미티브와 MMA(Matrix Multiply-Accumulate) 연산이 포함되어 있습니다. 특히 `small-BH` 레이아웃은 고정된 배치 크기와 헤드 수를 가정하여 메모리 접근 패턴을 최적화하고, 공유 메모리(SMEM) 사용을 효율화하는 데 중점을 둡니다.

### 3. 라우팅 및 바인딩 로직 (`csrc/kda/cake_flashkda_bf16_small_bh_m128_binding.cu`)

이 파일은 새로운 `small-BH` 커널을 기존 `cake_kda` 파이프라인에 통합하는 역할을 합니다. 커널 선택 로직, CUDA 그래프 및 스트림 워크스페이스 관리, 그리고 API 인터페이스를 정의합니다.

리뷰 과정에서 몇 가지 수정이 있었습니다:

*   **`m128_n16` 셀렉터 어설션:** 커밋 `4167f97ed`에서 셀렉터 테스트로 이동되었습니다.
*   **Descriptor-count 어설션:** PR 설명에는 이 어설션이 추가되지 않았지만, 현재 익스포트는 7개의 맵을 가진 `small-BH` 스토리지 계약을 사용하며, 패킷 오프셋은 공유된 디스크립터 카운트에서 파생됩니다. 현재 검증 결과 불일치가 없으므로, 향후 변경 가능성에 대한 가드 추가는 이 PR의 결함을 해결하지 않는다고 판단되었습니다.
*   **`final_state` 텐서 채우기:** 커밋 `4167f97ed`에서 `final_value`를 공급하고 실제 `final-state` 텐서가 채워지는지 확인하도록 테스트가 수정되었습니다. 이전에는 `store_final_state`와 `final_state` 인덱스에 대한 혼동이 있었습니다.
*   **CUDA 12.8 / sm100a 경로:** 커밋 `4167f97ed`에서 `_select_flash_kda_prefill_target`을 통해 예상 타겟을 파생시키도록 수정되었습니다. 이제 CUDA 12.8 / `sm100a` 실행은 최신 CUDA와 동일한 전체 GPU 테스트 파일을 사용합니다.
*   **Contiguous beta 요구사항:** 커밋 `4167f97ed`에서 `small-BH` 바인딩이 파생된 스트라이드를 확장 없이 전달하도록 수정되었습니다. 공유 바인딩 계약은 연속적인 베타를 요구하므로, 행-스트라이드 베타는 기존 직접 경로에 유지됩니다. JIT 소스-계약 어설션은 이를 검증합니다.

```diff
--- a/csrc/kda/cake_flashkda_bf16_small_bh_m128_binding.cu
+++ b/csrc/kda/cake_flashkda_bf16_small_bh_m128_binding.cu
@@ -15,10 +15,10 @@
 #include "cake_flashkda_common.h"
 
 // Generated schedule 'flashkda_bf16_small_bh_m128'; module
-// flashkda_bf16_small_bh_m128_73369168de.
+// flashkda_bf16_small_bh_m128_b3a946861.
 
 // clang-format off
-
+// clang-format on
 
 namespace {
 
@@ -238,7 +238,7 @@
     const int N = 128;
     const int M = 16;
 
-    // The number of heads is at most 8.
+    // The number of heads is at most 8, and the sequence length is at least 2048.
     // The number of CTA groups is at most 8.
     // The memory layout is fixed.
     // The compute capability is 10.0 or 10.3.
@@ -250,7 +250,7 @@
     const bool is_small_bh_route = 
         (num_tasks <= 8) && (num_heads <= 8) && 
         (seqlen >= 2048) && (num_cta_groups <= 8) &&
-        (is_fixed_layout) && (is_cc_10_0_or_10_3) && (num_heads == M);
+        (is_fixed_layout) && (is_cc_10_0_or_10_3) && (num_heads <= M);
 
     if (is_small_bh_route) {
         // Use the small-BH kernel for Blackwell compute capabilities 10.0 and 10.3.

is_small_bh_route 조건에서 num_heads == M (즉, num_heads == 16) 조건이 num_heads <= M으로 완화되었습니다. 이는 더 넓은 범위의 헤드 수에 대해 small-BH 커널을 적용할 수 있게 하여 유연성을 높였습니다. 또한, 리뷰어 yyihuang의 지적에 따라 sm_120 정책 관련 테스트 케이스가 제거되었습니다. 이는 해당 정책이 CC 10.0/10.3에서 지원되지 않기 때문입니다.

4. 테스트 커버리지 및 검증

PR은 광범위한 테스트를 통과했습니다:

  • Four-SKU GPU, stream, CUDA Graph, route, and cold-L2 CUPTI validation: B200, GB200, B300, GB300 네 가지 SKU에서 GPU, 스트림, CUDA 그래프, 라우팅, 그리고 L2 캐시 플러싱을 포함한 성능 검증을 통과했습니다.
  • Import/JIT contracts: 24/24 계약을 통과했습니다.
  • GPU/API/stream/CUDA Graph tests: 73/73 테스트를 통과했습니다.
  • Synccheck/Memcheck: B200 및 B300에서 오류 없이 통과했습니다.

리뷰어 yyihuang은 최종 검증 보고서에서 다음과 같이 언급했습니다:

The exact pull-request head 7e420efb02a39526ff9a5d5a70ac292844cc20df has completed import/JIT, GPU/API/stream/CUDA Graph correctness, route/fallback, and cold-L2 performance qualification on all four target SKUs.

이는 새로운 small-BH 커널이 다양한 환경에서 안정적이고 성능이 우수함을 입증합니다.

왜 이게 좋은가?

이 PR은 여러 측면에서 중요한 개선을 이루었습니다:

  1. Blackwell 아키텍처 활용 극대화: 새로운 small-BH 커널은 Blackwell GPU의 특정 기능(예: Tensor Core)과 메모리 계층 구조를 활용하도록 설계되었습니다. 이를 통해 기존 커널보다 훨씬 높은 처리량을 달성할 수 있습니다.

  2. 성능 향상: 리뷰의 'Cold-L2 CUPTI A/B' 섹션에 따르면, small-BH 커널은 기존 커널 대비 상당한 성능 향상을 보였습니다. 예를 들어, B200 SKU에서 h8_fixed_65536 케이스의 경우:

    • direct / auto: 1.57x 속도 향상
    • official / auto: 2.46x 속도 향상

    이는 auto 라우팅이 새로운 small_bh_m128/sm100f 커널을 선택했을 때 얻을 수 있는 성능 이득을 보여줍니다. 다른 SKU와 케이스에서도 유사한 수준의 향상이 관찰되었습니다. (평균적으로 Direct/Auto 1.57x, Official/Auto 2.40x)

  3. 유연한 라우팅: small-BH 커널은 특정 조건(CC 10.0/10.3, 최대 8개 헤드, 시퀀스 길이 >= 2048 등)에서만 활성화됩니다. 이러한 조건에 해당하지 않는 입력은 기존의 안정적인 직접 또는 폴백 경로를 사용하므로, 전체 시스템의 안정성을 해치지 않으면서 특정 워크로드에 대한 성능을 최적화합니다.

  4. 코드 품질 및 테스트 커버리지: PR은 광범위한 테스트와 검증을 통과했으며, 리뷰 과정에서 제기된 문제점들이 신속하게 수정되었습니다. 이는 코드의 견고성과 신뢰성을 높입니다.

일반적 교훈

  • 하드웨어 특화 최적화: 최신 하드웨어 아키텍처의 기능을 최대한 활용하기 위해서는 해당 아키텍처에 특화된 커널 개발이 필수적입니다. 특히 GPU의 경우, 메모리 대역폭, 공유 메모리, Tensor Core 활용 등이 성능에 큰 영향을 미칩니다.
  • 동적 라우팅의 중요성: 모든 워크로드에 단일 최적화된 커널이 적용될 수는 없습니다. 다양한 입력 조건에 따라 최적의 커널을 동적으로 선택하는 라우팅 메커니즘은 성능과 안정성을 동시에 확보하는 데 중요합니다.
  • 철저한 테스트와 검증: 새로운 최적화 기법을 도입할 때는 성능 측정뿐만 아니라, 다양한 엣지 케이스와 이전 버전과의 호환성을 포함한 철저한 기능 및 회귀 테스트가 필수적입니다.

결론

이번 FlashInfer PR은 NVIDIA Blackwell 아키텍처를 위한 recurrent-KDA prefill 성능을 크게 향상시키는 중요한 발걸음입니다. small-BH 커널의 도입과 지능적인 라우팅 로직은 LLM 추론의 속도를 높이고, 최신 하드웨어의 잠재력을 최대한 끌어내는 데 기여합니다. 이러한 최적화는 앞으로 더 크고 복잡한 모델을 효율적으로 실행하는 데 중요한 기반이 될 것입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글