[flashinfer] FlashInfer MiniMax-H3 Attention 최적화: K/V-split을 통한 성능 향상 분석
PR 링크: flashinfer-ai/flashinfer#5530 상태: Merged | 변경: +18072 / -616
들어가며
최근 GPU 기반 딥러닝 연산 라이브러리들은 LLM(거대 언어 모델)의 폭발적인 성장과 함께 더욱 빠르고 효율적인 연산을 제공하기 위한 최적화 경쟁을 벌이고 있습니다. 특히 Attention 메커니즘은 LLM의 핵심 구성 요소로서, 그 연산 속도는 모델의 전체 성능에 지대한 영향을 미칩니다.
이번 글에서는 "perf(cake_minimax_h3): K/V-split the partial-wave units of the packed-varlen attention (SM100/SM103)" 라는 제목의 GitHub Pull Request(PR)를 심층 분석하여, FlashInfer 라이브러리가 어떻게 Attention 연산의 성능을 개선했는지 살펴보겠습니다. 이 PR은 특히 짧은 시퀀스 길이에서 발생하는 성능 병목 현상을 해결하기 위해 K/V-split이라는 새로운 기법을 도입했습니다. 실제 코드 변경 사항을 중심으로 이 최적화가 왜 효과적이며, 어떤 기술적 원리로 작동하는지 자세히 알아보겠습니다.
문제점: 짧은 시퀀스에서의 Wave Quantization 병목
기존의 MiniMax-H3 packed-varlen attention 구현에서는 짧은 packed row (예: Ulysses P=8 / P=4, 4k-14k 토큰)의 경우, 긴 시퀀스 길이에 비해 57-85% 수준의 처리량만을 보여주는 성능 저하가 관찰되었습니다. PR 설명에 따르면, 이러한 문제는 tail padding이 아닌 persistent grid의 wave quantization 때문이라고 지적합니다.
예를 들어, cu_seqlens=[0, 6096] x 7개의 헤드를 가진 경우, U = 84개의 동일한 비용을 가진 유닛들이 G = 74개의 클러스터에 분산되어 처리됩니다. 이는 평균 1.14 웨이브(wave)에 해당하며, 커널이 두 번의 전체 유닛 시간(unit time)을 소요하게 만들어 비효율을 야기합니다.
변경 사항: K/V-split 도입 및 최적화
이 PR은 이러한 문제를 해결하기 위해 다음과 같은 핵심 변경 사항을 도입했습니다:
-
K/V-split 전략 도입: 짧은 시퀀스에서 발생하는 웨이브 양자화(wave quantization) 문제를 완화하기 위해, 각 K/V 블록 범위를 여러 개의 유닛으로 분할(split)하는 방식을 채택했습니다. 이는 마치 "flash-decoding" 스타일과 유사하게, 부분적인 웨이브(partial wave)를 K/V 범위에 걸쳐 분할하고, 이후 작은
combine단계를 통해 이들을 병합하는 방식입니다. -
Host Planner (
choose_kv_splits):- 웨이브 수가 4개 미만인 경우, 가장 오래 걸리는 처리 시간 우선(longest-processing-time-first) 슬롯 할당을 시뮬레이션합니다.
U mod G개의 가장 비용이 많이 드는 유닛들을k = 2..8개 방식으로 분할하여 거의 동일한 K/V 블록 범위를 만듭니다.combine비용을 추가하고, 3% 이상의 예측된 성능 향상이 있을 경우에만 이 분할을 채택합니다.
diff --git a/flashinfer/experimental/minimax_h3_varlen_attention/README.md b/flashinfer/experimental/minimax_h3_varlen_attention/README.md index a711df9fab..1c2108723a 100644 --- a/flashinfer/experimental/minimax_h3_varlen_attention/README.md +++ b/flashinfer/experimental/minimax_h3_varlen_attention/README.md @@ -84,11 +84,23 @@ * `seg_begin[s]`, `seg_len[s]` (int32, `num_segments` entries), * `unit_table` (int32, `2 * total_tiles` entries): per slot the segment index and `head << 16 | cluster_in_segment`. +* `unit_table` (int32, `4 * total_tiles` entries): per slot the segment index, + `head << 16 | cluster_in_segment`, `kv_block_begin << 16 | kv_blocks` (the + unit's K/V block range inside the segment) and its partial slot (`-1` for + an unsplit unit, which writes the BF16 output directly), +* `combine_table` (int32, `4 * num_combine_units` entries): per K/V-split + unit the segment index, `head << 16 | cluster_in_segment`, its first partial + slot and the number of splits, plus the plan's partial workspace + `partial_O` (FP16, `num_partial_slots * 512 * 128`) and `partial_ML` (FP32, + `num_partial_slots * 512 * 2`). + Units are enumerated segment-major (heads slow, clusters fast, so a cluster's Q tiles reuse the segment's K/V from L2) and placed into slots longest-processing-time first over the per-unit cost `ceil(seg_len / 128) + 2` K/V blocks (`assign_unit_slots`): the partial tail round receives the cheapest units and, within every full round, the clusters that also own a tail unit receive that round's cheapest units; equal costs keep the enumeration order. +**K/V splits for the partial wave.** When the unit count leaves a partial +wave on the persistent grid (`total_units mod num_clusters != 0`, fewer than +four waves), the planner (`choose_kv_splits`) simulates the slot assignment +for splitting the `total_units mod num_clusters` most expensive units (and, +as the fallback, every unit) `k = 2..8` ways into near-equal K/V block ranges +and keeps the candidate with the lowest makespan plus combine cost when it +beats the unsplit plan by more than 3 %. A split unit's ranges are separate +units of the table; each runs the full online softmax over its range and +writes its rows normalized by its own softmax sum as FP16 into its partial +slot together with the FP32 `(scaled log2 row max, row sum)`. The `combine` +stage (one warp per output row, `128` CTAs of 128 threads per split unit) +then merges the slots with the exact FlashAttention formula +`O = sum_i 2^(m_i - m) l_i O_i / sum_i 2^(m_i - m) l_i` into the BF16 output. +It is launched after the attention kernel on the same stream and skipped when +the plan has no split units. Plans with at least as many units as clusters +per wave are unchanged (unsplit units are bitwise identical to the previous +kernel). The plan is a host-side function of `(cu_seqlens, num_heads, num_SMs)` and reproduces the Cake production plan table for table (`num_heads < 2^15`, fewer than `2^16` clusters per segment). K/V TMA loads that run past a segment boundary are handled by the `kv_block_begin`/`kv_blocks` fields in the `unit_table` and `tile_table` respectively (the `kv_block_begin` is the first K/V block of the unit's range, and `kv_blocks` is the number of K/V blocks in the range). The `kv_block_begin` is the first K/V block of the unit's range, and `kv_blocks` is the number of K/V blocks in the range. -
BF16
unit_table및 NVFP4tile_table업데이트:unit_table은 이제 세그먼트, 헤드, 클러스터 정보 외에도 K/V 블록 범위(kv_block_begin,kv_blocks) 및 부분 슬롯 정보(partial slot)를 포함합니다.- NVFP4
tile_table도 유사하게cl_kv_begin,cl_kv_blocks,cl_ws_slot필드를 추가했습니다.
-
combine커널: K/V-split된 유닛들은 각자의 소프트맥스 합으로 정규화된 FP16 결과와 FP32(scaled log2 max, sum)을 부분 작업 공간(partial workspace)에 기록합니다. 이후combine커널이 이 결과들을 FP32에서 정확하게 병합합니다. 이 커널은 어텐션 커널 이후에 실행되며, 분할된 유닛이 있을 때만 호출됩니다. -
NVFP4 프로그램 라우팅: NVFP4는 아키텍처당 두 개의 어텐션 프로그램을 사용합니다. 분할되지 않은 플랜을 위한 dense 프로그램과 분할된 플랜을 위한 split 프로그램입니다. 이를 통해 분할되지 않은 유닛의 소프트맥스 블록 루프는 기존 커널의 스케줄을 유지하면서 성능 저하를 방지합니다.
왜 이게 좋은가? (성능 향상 및 교훈)
이 PR의 가장 큰 성과는 짧은 시퀀스 길이에서 발생하는 성능 병목을 효과적으로 해소했다는 점입니다. PR에 제시된 결과는 다음과 같습니다:
- BF16 B200:
[0, 6096] x 7시퀀스에서 0.1585ms -> 0.1105ms로 1.43배 향상. - NVFP4 fp4 B200:
[0, 6096] x 7시퀀스에서 0.1524ms -> 0.1062ms로 1.44배 향상.
이러한 성능 향상은 다음과 같은 이유로 가능했습니다:
- 웨이브 양자화 문제 해결: K/V-split을 통해 웨이브당 처리해야 하는 유닛 수를 줄이고, 각 유닛이 더 효율적으로 작업을 완료할 수 있게 되었습니다. 특히
U < G(유닛 수가 클러스터 수보다 적은 경우) 상황에서 성능 향상이 두드러집니다. - 동적 프로그램 선택: 분할되지 않은 유닛은 기존의 최적화된 dense 프로그램을 사용하고, 분할된 유닛은 split 프로그램을 사용하여 각 상황에 맞는 최적의 실행 경로를 선택합니다. 이는 불필요한 오버헤드를 줄이고 성능을 극대화합니다.
- 호스트-GPU 연동 최적화:
quantize및attention단계 간의 호스트 시간(host time)을 줄이고, GPU 커널 간의 의존성 관리를 개선하여 전체 파이프라인의 지연 시간을 단축했습니다. 특히cuLaunchAttributeProgrammaticStreamSerialization속성을 활용하여 CUDA 그래프 캡처 가능성을 높였습니다.
일반적인 교훈:
- 병목 식별 및 타겟팅: 성능 병목이 발생하는 특정 시나리오(짧은 시퀀스, 웨이브 양자화)를 정확히 식별하고, 해당 문제에 특화된 해결책(K/V-split)을 적용하는 것이 중요합니다.
- 동적 실행 경로: 하드웨어 및 입력 데이터의 특성에 따라 최적의 실행 경로를 동적으로 선택하는 메커니즘은 복잡한 시스템에서 전반적인 성능을 향상시키는 데 효과적입니다.
- 호스트-GPU 상호작용 최적화: GPU 커널 자체의 최적화뿐만 아니라, 호스트 코드와 GPU 커널 간의 통신 및 동기화 오버헤드를 줄이는 것도 전체 성능에 큰 영향을 미칩니다.
코드 분석 (파일별)
flashinfer/experimental/minimax_h3_varlen_attention/README.md
이 파일은 새로운 기능과 사용법을 설명하는 문서입니다. K/V-split 메커니즘, choose_kv_splits 플래너의 작동 방식, unit_table 및 tile_table의 구조 변경, combine 커널의 역할, NVFP4의 동적 프로그램 라우팅 등이 상세하게 기술되어 있습니다. 특히 unit_table의 필드가 4개로 늘어난 점과 combine_table의 추가는 K/V-split 구현의 핵심을 보여줍니다.
# 위 "변경 사항" 섹션의 diff 참조
flashinfer/experimental/minimax_h3_varlen_attention/cake_backend.py
이 파일은 백엔드 로직을 담당합니다. prepare_minimax_h3_varlen_attention 함수는 K/V-split을 고려하여 unit_table 및 tile_table을 구성하고, choose_kv_splits 함수를 호출하여 최적의 분할 전략을 결정합니다. 또한, runner() 함수는 분할된 플랜에 따라 attention_split 또는 combine 단계를 포함한 전체 파이프라인을 실행하도록 수정되었습니다.
# 위 "변경 사항" 섹션의 diff 참조
리뷰 댓글 분석
리뷰 댓글은 주로 CI 파이프라인의 테스트 결과와 관련된 내용입니다. 초기에는 GB300 GPU 및 CUDA 12.9 환경에서 tests.experimental.test_cake_minimax_h3_varlen_attention 테스트가 실패하는 문제가 있었습니다. 이는 AssertionError: Tensor-likes are not close! 메시지와 함께 특정 인덱스에서 텐서 값의 미세한 차이를 보여줍니다.
[flashinfer-bot] [FAILED] Pipeline [#69666963](https://nv/flashinfer-ci/-/pipelines/69666963) — 24/25 executed test jobs passed
Compared with nightly [#69428841](https://nv/flashinfer-ci/-/pipelines/69428841) (different CI configuration).
### Unit Tests
| GPU | CUDA 12.9 | CUDA 13.0 | CUDA 13.4 | Notes |
|---|---|---|---|---|
| B200 | ✅ Pass | ✅ Pass | ✅ Pass | — |
| GB200 | ✅ Pass | ✅ Pass | ✅ Pass | — |
| GB300 | ❌ New | ✅ Pass | ✅ Pass | **PR-related:** `tests.experimental.test_cake_minimax_h3_varlen_attention` (1 failure; CUDA 12.9) |
<details>
<summary>Failure details</summary>
### PR-related regressions
- `tests.experimental.test_cake_minimax_h3_varlen_attention` — 1 failure on GB300 / CUDA 12.9
- AssertionError: Tensor-likes are not close! Mismatched elements: 1 / 806400 (0.0%) Greatest absolute difference: 1.2757673263549805 at index (23, 3, 11) (up to 1.0 allowed) Grea…
</details>
이후 여러 번의 CI 실행과 코드 수정(정확한 수정 내용은 PR diff에 명시되지 않았지만, 테스트 실패를 해결하기 위한 조정이 있었을 것으로 추정됩니다)을 거쳐 최종적으로 모든 테스트가 통과되었습니다. 이는 K/V-split 및 관련 최적화가 다양한 환경에서 정확성을 유지함을 보여줍니다.
결론
이번 PR은 FlashInfer의 MiniMax-H3 Attention 커널에서 짧은 시퀀스 길이로 인한 성능 병목 현상을 해결하기 위해 K/V-split이라는 정교한 최적화 기법을 성공적으로 도입했습니다. 호스트 플래너의 지능적인 분할 결정, 업데이트된 테이블 구조, 그리고 동적 프로그램 선택 메커니즘을 통해 상당한 성능 향상을 달성했습니다. 특히 1.4배 이상의 속도 개선은 LLM 추론 성능에 직접적인 영향을 미칠 수 있는 중요한 결과입니다.
이러한 최적화는 복잡한 GPU 커널 개발에서 병목 지점을 정확히 파악하고, 하드웨어 특성을 고려한 동적인 실행 전략을 수립하는 것이 얼마나 중요한지를 다시 한번 보여줍니다. 앞으로도 FlashInfer와 같은 라이브러리들이 Attention 연산의 효율성을 지속적으로 개선해 나갈 것으로 기대됩니다.
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer SM110 XQA 최적화: register_mma_split 도입으로 FP16 Paged Attention 성능 향상
- [flashinfer] FlashInfer, Qwen3-30B 모델의 성능 향상을 위한 CUDA 커널 최적화: L2 캐시 힌트 도입
- [flashinfer] FlashInfer NVFP4 QKV GEMM 최적화: SM103a Epilogue 통합 및 CUDA 런처 개선
- [flashinfer] FlashInfer의 실험적 NVFP4 어텐션 도입: SM103 최적화
- [flashinfer] FlashInfer의 SM100/SM103 최적화: CAKE 기반 블록 희소 어텐션(VSA) 도입
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer NVFP4 QKV GEMM 최적화: SM103a Epilogue 통합 및 CUDA 런처 개선
- 현재글 : [flashinfer] FlashInfer MiniMax-H3 Attention 최적화: K/V-split을 통한 성능 향상 분석
- 다음글 [Liger-Kernel] Liger-Kernel: Cross-Entropy와 Total Variation Distance를 하나로 융합하여 성능을 극대화하다
댓글