[flashinfer] Gemma-4를 위한 Blackwell 최적화: FlashInfer의 비대칭 VO-Split NVFP4 구현 분석
PR 링크: flashinfer-ai/flashinfer#3684 상태: Merged | 변경: +1280 / -124
들어가며
최근 공개된 Gemma-4 모델은 기존 모델들과는 다른 독특한 어텐션 구조를 가지고 있습니다. 바로 Query-Key(QK)의 head_dim은 512인 반면, Value-Output(VO)의 head_dim은 256인 비대칭(Asymmetric) 구조입니다.
기존의 많은 어텐션 커널들은 QK와 VO의 차원이 동일하다고 가정하거나, 대칭적인 구조(head_dim <= 256)에 최적화되어 있었습니다. 특히 NVIDIA의 최신 Blackwell(SM120, SM121) 아키텍처에서 제공하는 4비트 부동소수점 형식인 NVFP4를 사용하여 KV 캐시를 서빙할 때, 이러한 비대칭성은 커널 수준에서 해결해야 할 큰 과제가 되었습니다.
이번 글에서는 flashinfer-ai/flashinfer 레포지토리에 올라온 PR을 통해, 어떻게 비대칭 VO-split NVFP4 paged prefill을 구현하고 Blackwell GPU에서의 성능과 정확도를 확보했는지 분석해 보겠습니다.
코드 분석: 핵심 변경 사항
1. K/V Stride의 독립적 관리
기존 코드에서는 K 캐시와 V 캐시의 stride(메모리 보폭)가 동일해야 한다는 엄격한 체크가 있었습니다. 하지만 비대칭 구조에서는 K와 V의 데이터 크기가 달라지므로 이를 분리해야 합니다.
Before:
// csrc/batch_decode.cu
for (int i = 0; i < k_strides.size(); ++i) {
TVM_FFI_ICHECK_EQ(k_strides[i], v_strides[i]);
}
kv_cache_strides = k_strides.data();
After:
// csrc/batch_decode.cu
for (int i = 0; i < k_strides.size(); ++i) {
TVM_FFI_ICHECK_EQ(k_strides[i], v_strides[i])
<< "K/V strides differ at dim " << i
<< ": the FA2 decode kernel addresses both K and V through a single set of "
"(K) strides... NVFP4/asymmetric decode with independent K/V strides is not yet supported.";
}
// paged_kv_t 생성 시 k_strides와 v_strides를 별도로 전달
paged_kv_t<DTypeKV, IdType> paged_kv(
..., k_strides.data(), v_strides.data(), ...);
Prefill 단계에서는 이제 K와 V의 stride를 완전히 분리하여 수용하며, paged_kv_t 구조체 내부에서 protective_get_k_offset과 protective_get_v_offset을 통해 각각의 메모리 주소를 정확히 계산합니다.
2. CTA Tile Q의 동적 결정 (Register Pressure 관리)
head_dim_qk가 512로 커지면 GPU 레지스터 압박(Register Pressure)이 심해집니다. 이를 해결하기 위해 CTA_TILE_Q 크기를 헤드 차원에 따라 다르게 할당하도록 Jinja 템플릿을 수정했습니다.
Before:
{# head_dim_vo가 512 이상일 때만 16, 32 선택 #}
{% for cta_tile_q in ([16, 32] if head_dim_vo | int >= 512 else [16, 64, 128]) %}
After:
{# head_dim_qk가 512 이상인 비대칭 케이스 대응 #}
{% for cta_tile_q in ([16, 32] if head_dim_vo | int >= 512 else ([16] if head_dim_qk | int >= 512 else [16, 64, 128])) %}
head_dim_qk >= 512이면서 VO가 작은 경우, 레지스터 부족으로 인해 CTA_TILE_Q를 16으로 제한하여 커널 실행 가능성을 확보한 것이 핵심입니다.
3. NVFP4 Split-KV(Flash-Decoding) 비활성화
이 PR에서 발견된 가장 흥미로운 버그 중 하나는 NVFP4 사용 시 Split-KV 모드에서의 정확도 하락입니다. NVFP4는 16개 요소마다 하나의 Scale Factor(SF)를 가지는 블록 구조를 사용하는데, Split-KV가 토큰 축을 따라 임의의 지점에서 분할될 경우 이 SF 블록 경계가 깨지면서 수치적 오류가 발생했습니다.
Fix:
// csrc/batch_prefill.cu (plan 함수 내부)
bool disable_split_kv_logic = (kv_data_type == QKVDataType::kNVFP4);
if (disable_split_kv_logic) {
disable_split_kv = true;
}
리뷰 과정에서 밝혀졌듯이, 이는 성능 최적화보다 정확도(Correctness)를 우선시한 결정입니다. 향후 SF 블록 경계에 맞춰 분할하는 로직으로 개선될 예정입니다.
4. LSE(Log-Sum-Exp) 초기화 문제 해결
Blackwell SM120 바인딩 코드에서 lse 텐서가 초기화되지 않은 채 사용되어 NaN이 발생하는 문제가 보고되었습니다. 이를 위해 커널 실행 전 cudaMemsetAsync로 초기화하는 로직이 추가되었습니다.
After:
// csrc/nvfp4_attention_sm120/nvfp4_attention_sm120_binding.cu
status = cudaMemsetAsync(lse.data_ptr(), 0, numel(lse) * get_element_size(lse), stream);
왜 이게 좋은 최적화인가?
- 메모리 효율성 극대화: Gemma-4와 같은 대형 모델을 서빙할 때 4비트(NVFP4) KV 캐시를 사용하면 메모리 사용량을 BF16 대비 1/4로 줄일 수 있습니다. 이는 더 긴 컨텍스트를 수용하거나 더 큰 배치 사이즈를 가능하게 합니다.
- 비대칭 아키텍처 지원: 단순히 대칭적인 구조만 지원하는 것이 아니라, 실제 최신 모델의 아키텍처 특성(QK=512, VO=256)을 커널 수준에서 직접 수용하여 불필요한 패딩이나 메모리 낭비를 방지했습니다.
- Blackwell 하드웨어 가속: SM120(RTX 5090 등)과 SM121(GB10) 하드웨어의 특성을 활용하여 최신 GPU의 잠재력을 끌어올렸습니다.
- 수치적 안정성: Prefix Caching 상황에서 발생할 수 있는 Split-KV의 미세한 오차를 발견하고 이를 차단함으로써, 긴 문맥에서도 모델의 출력이 오염되지 않도록 보장했습니다.
결론 및 교훈
이번 PR은 단순히 기능을 추가하는 것을 넘어, 하드웨어 아키텍처(Blackwell)와 모델 아키텍처(Gemma-4) 사이의 간극을 메우는 정교한 엔지니어링의 결과물입니다. 특히 "성능보다 정확도가 우선"이라는 원칙하에 수치적 오류를 유발하는 Split-KV를 과감히 제어한 점은 시니어 엔지니어로서 배울 점이 많은 대목입니다.
앞으로 NVFP4를 직접 소비하는 mma.sync 기반의 QKᵀ 연산(A4Q)이 추가되면, Blackwell에서의 추론 성능은 한 단계 더 도약할 것으로 기대됩니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#shared-memory-opt-in
- https://github.com/flashinfer-ai/flashinfer/blob/main/include/flashinfer/attention/prefill.cuh
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer: SM120/SM121 아키텍처를 위한 네이티브 MXFP4 W4A4 Fused MoE 지원
- [flashinfer] FlashInfer, Blackwell GPU를 위한 Gated MoE 커널 최적화로 성능 대폭 향상
- [flashinfer] FlashInfer의 Blackwell 아키텍처 최적화: CAKE 기반 TinyGEMM2 커널 도입
- [flashinfer] Blackwell NVFP4 양자화 최적화: TMA OOB Zero-fill을 이용한 메모리 복사 오버헤드 제거
- [flashinfer] FlashInfer: NVIDIA Blackwell(SM120)을 위한 고성능 FP8 MoE GEMM 최적화
PR Analysis 의 다른글
- 이전글 [sglang] SGLang의 Wan2.2-TI2V 최적화: Triton 커널을 통한 메모리 트래픽 병목 해결
- 현재글 : [flashinfer] Gemma-4를 위한 Blackwell 최적화: FlashInfer의 비대칭 VO-Split NVFP4 구현 분석
- 다음글 [sglang] AMD MI355X 환경에서 Triton 3.7 레지스터 스필링 최적화
댓글