본문으로 건너뛰기

[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_offsetprotective_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);

왜 이게 좋은 최적화인가?

  1. 메모리 효율성 극대화: Gemma-4와 같은 대형 모델을 서빙할 때 4비트(NVFP4) KV 캐시를 사용하면 메모리 사용량을 BF16 대비 1/4로 줄일 수 있습니다. 이는 더 긴 컨텍스트를 수용하거나 더 큰 배치 사이즈를 가능하게 합니다.
  2. 비대칭 아키텍처 지원: 단순히 대칭적인 구조만 지원하는 것이 아니라, 실제 최신 모델의 아키텍처 특성(QK=512, VO=256)을 커널 수준에서 직접 수용하여 불필요한 패딩이나 메모리 낭비를 방지했습니다.
  3. Blackwell 하드웨어 가속: SM120(RTX 5090 등)과 SM121(GB10) 하드웨어의 특성을 활용하여 최신 GPU의 잠재력을 끌어올렸습니다.
  4. 수치적 안정성: Prefix Caching 상황에서 발생할 수 있는 Split-KV의 미세한 오차를 발견하고 이를 차단함으로써, 긴 문맥에서도 모델의 출력이 오염되지 않도록 보장했습니다.

결론 및 교훈

이번 PR은 단순히 기능을 추가하는 것을 넘어, 하드웨어 아키텍처(Blackwell)와 모델 아키텍처(Gemma-4) 사이의 간극을 메우는 정교한 엔지니어링의 결과물입니다. 특히 "성능보다 정확도가 우선"이라는 원칙하에 수치적 오류를 유발하는 Split-KV를 과감히 제어한 점은 시니어 엔지니어로서 배울 점이 많은 대목입니다.

앞으로 NVFP4를 직접 소비하는 mma.sync 기반의 QKᵀ 연산(A4Q)이 추가되면, Blackwell에서의 추론 성능은 한 단계 더 도약할 것으로 기대됩니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글