본문으로 건너뛰기

[flashinfer] FlashInfer의 NVFP4 KV 캐시 성능 최적화: FP4 연산의 병목 현상 해소

PR 링크: flashinfer-ai/flashinfer#4746 상태: Merged | 변경: +170 / -24

들어가며

최근 대규모 언어 모델(LLM)의 발전과 함께 추론 성능 최적화는 매우 중요한 과제가 되었습니다. 특히, KV 캐시는 LLM 추론 시 메모리 대역폭과 연산량을 크게 차지하는 요소 중 하나입니다. FlashInfer는 이러한 KV 캐시 처리를 효율적으로 수행하기 위한 라이브러리로, 다양한 양자화 기법을 지원합니다. 본 글에서는 FlashInfer의 NVFP4 (4-bit Normal-Float) KV 캐시 연산에서 발생하는 심각한 성능 병목 현상을 분석하고, 이를 해결하기 위한 코드 변경 사항을 자세히 살펴보겠습니다. 특히, 특정 CUDA 아키텍처에서 NVFP4 연산이 FP8 연산보다 훨씬 느린 이유를 파악하고, 이를 개선하기 위한 최적화 기법들을 실제 코드 diff와 함께 설명합니다.

이 Pull Request(PR)는 NVFP4 KV 캐시가 FP8 KV 캐시보다 훨씬 느린 문제를 해결하는 데 중점을 둡니다. NVFP4는 FP8보다 적은 바이트를 읽음에도 불구하고, 특정 하드웨어(SM100 이전)에서는 FP8 대비 prefill 시 6.8–8.2배, decode 시 4.4–5.4배 느린 성능을 보였습니다. 이는 메모리 대역폭의 문제가 아니라, NVFP4 연산 자체의 비효율성 때문임을 ncu 프로파일링 결과가 시사합니다. 이 PR은 이러한 비효율성을 제거하고 NVFP4 연산의 성능을 크게 향상시키는 네 가지 주요 변경 사항을 포함합니다.

코드 분석

이 PR은 크게 네 가지 최적화 기법을 적용하며, 각 기법은 특정 하드웨어 타겟에만 적용되도록 if constexpr 또는 매크로를 통해 제어됩니다. 모든 변경은 비트 단위 동일성(bit-identical)을 유지하면서 성능을 개선하는 것을 목표로 합니다.

1. vec_dtypes.cuh: LUT 기반의 느린 NVFP4 디퀀타이제이션 경로 개선

이전 코드에서는 SM100 이전 아키텍처에서 vec_cast<half|nv_bfloat16, __nv_fp4x2_e2m1> fallback이 로컬 메모리에 저장된 작은 LUT(Look-Up Table)를 사용하여 NVFP4 값을 디퀀타이즈했습니다. SASS 분석 결과, 이 LUT 접근이 각 인라인 호출마다 스레드 스택에 저장되고 로드되는 비효율적인 연산으로 이어졌습니다. 이는 커널의 스택 프레임 크기를 증가시키고, 레지스터 사용량을 늘려 성능 저하의 원인이 되었습니다.

Before:

// 로컬 메모리에 저장되는 LUT를 사용한 디퀀타이제이션
// ... (생략) ...
constexpr uint16_t lut[16] = { ... };
// ... (생략) ...
// 각 인라인 호출마다 STL.128 x2, LDL.U16 연산 발생

After:

이 PR에서는 LUT를 레지스터 전용 prmt (permute) 명령어를 활용하는 경로로 대체했습니다. FP16의 경우, NVFP4의 8가지 magnitude 값({0, .5, 1, 1.5, 2, 3, 4, 6})은 모두 하위 바이트가 0이므로, prmt 명령어를 사용하여 8개의 다른 상위 바이트를 효율적으로 인덱싱할 수 있습니다. BF16의 경우, magnitude 값이 더 복잡하여 두 개의 8바이트 테이블을 인덱싱하고 인터리빙하는 방식으로 처리됩니다.

// vec_dtypes.cuh
// ... (생략) ...
__device__ __forceinline__ void e2m1x8_to_f16x8(uint32_t packed, uint32_t* d) {
  // ... (prmt 기반 연산) ...
}

__device__ __forceinline__ void e2m1x8_to_bf16x8(uint32_t packed, uint32_t* d) {
  // ... (prmt 기반 연산 및 테이블 인터리빙)
}
// ... (생략) ...

이 변경으로 인해 FA2 NVFP4-KV 인스턴스에서 LDL이 36202에서 128로, STL이 9602에서 190으로 크게 감소했습니다. 또한, 스택 프레임이 없는 커널은 0/72에서 54/72로, 255 레지스터 제한에 걸렸던 커널은 25에서 19로 개선되었습니다.

2 & 3. prefill.cuh: NVFP4 스케일 팩터 변환 최적화

이전에는 각 프래그먼트의 E4M3 스케일 바이트가 static_cast<DTypeQ>(sf_fp8)를 통해 변환되었습니다. SM90 이하 아키텍처에서는 하드웨어 E4M3->FP16 변환 명령이 없어 소프트웨어 변환으로 확장되었는데, 이 과정에서 denormal 경로가 데이터에 따라 분기하는 while 루프를 포함했습니다. 이는 가장 안쪽 루프에서 스케일 바이트마다 분기(divergent branch)를 발생시켜 성능 병목의 주된 원인이었습니다.

Before (개념적):

// ... (생략) ...
// sf_fp8는 __nv_fp8_e4m3 타입
auto scale = static_cast<DTypeQ>(sf_fp8);
// ... (sf_fp8의 denormal 경로에서 while 루프 발생 가능)

After:

새로 추가된 nvfp4_sf4_to_packed2x2() 함수는 네 개의 스케일 바이트를 하나의 워드로 팩하고, 이를 인라인된 fast_dequant_f8f16x4 함수에 전달합니다. 이 함수는 prmt/lop3 연산을 사용하여 분기 없이(branch-free) 두 개의 half2/nv_bfloat162 스케일 연산자를 생성합니다. 이 함수는 FLASHINFER_HARDWARE_FP8_CONVERSION_ENABLED 매크로로 제어되어, SM90 이상에서는 기존의 static_cast를 유지하여 하드웨어 명령을 사용하도록 합니다.

// prefill.cuh
// ... (생략) ...
void nvfp4_sf4_to_packed2x2(uint8_t b0, uint8_t b1, uint8_t b2, uint8_t b3, uint2* out) {
#if !defined(FLASHINFER_HARDWARE_FP8_CONVERSION_ENABLED)
  // ... (prmt/lop3 기반 fast_dequant_f8f16x4 호출)
#else
  // ... (하드웨어 변환 사용)
#endif
}
// ... (compute_qk, load_fp4_k_frag_scaled, compute_sfm_v, vosplit_compute_pv 등에서 사용)

이 변경으로 인해 ncu 프로파일링에서 스케일 팩터 변환 관련 분기 연산이 사라지고, 전체 명령어 수가 249.6M에서 171.1M으로 감소했습니다.

4. vec_dtypes.cuh: E2M1 사인 비트 전파 최적화

이전에는 SM100 이전 아키텍처에서 E2M1 값을 16비트로 변환할 때, 각 값의 사인 비트를 MSB로 옮기기 위해 여러 번의 시프트, 마스크, OR 연쇄(chain)를 사용했습니다. 이는 8개의 값에 대해 8번의 시프트와 8번의 lop3 연산을 필요로 하여 상당한 오버헤드를 발생시켰습니다.

Before (개념적):

// ... (생략) ...
// 여러 시프트, 마스크, OR 연산으로 사인 비트 생성
// ... (생략) ...

After:

prmt 명령어의 기본 모드는 선택된 바이트의 사인 비트(MSB)를 출력 바이트 전체에 복제하는 기능이 있습니다. 이를 활용하여, 4개의 사인 비트를 한 번의 prmt 연산으로 효율적으로 수집하고, lop3 연산으로 최종 결과를 생성합니다. 이 최적화는 8개 값당 연산 수를 25개에서 14개(half 기준)로 크게 줄였습니다.

// vec_dtypes.cuh
// ... (생략) ...
__device__ __forceinline__ uint32_t e2m1_sign_bytes(uint32_t packed, uint32_t sel) {
  return prmt_b32(packed << 4, packed, sel) & 0x80808080u;
}
// ... (생략) ...

이 변경은 NVFP4 prefill 모듈 전체에서 cuobjdump 기준 명령어 수를 501575에서 448209로 감소시켰습니다. 또한, e2m1x8_to_f16x8e2m1x8_to_bf16x8 함수 내에서 magnitude 테이블 조회도 prmt를 통해 레지스터 전용으로 처리하도록 개선되었습니다.

왜 이게 좋은가?

이 PR의 핵심은 NVFP4 KV 캐시 연산에서 발생하는 명령어 바운드(instruction-bound) 병목 현상을 제거하는 것입니다. 이전에는 NVFP4 디퀀타이제이션 및 스케일 팩터 변환 과정에서 발생하는 과도한 비트 조작, 분기 연산, 그리고 비효율적인 메모리 접근이 성능 저하의 주된 원인이었습니다.

성능 향상

PR 설명에 제시된 결과는 이러한 최적화의 효과를 명확히 보여줍니다:

  • Prefill: 현재 NVFP4 대비 2.65–2.78배 향상
  • Decode: 현재 NVFP4 대비 2.88–3.12배 향상

특히, FP8 KV 캐시와의 비교에서 NVFP4 KV 캐시의 성능이 크게 개선되었습니다. 예를 들어, prefill b1 4096² 케이스에서 NVFP4 before는 FP8 대비 6.76배 느렸지만, NVFP4 after는 2.43배로 크게 단축되었습니다. decode b64 kv8192 40h 케이스에서는 4.41배에서 1.52배로 개선되었습니다.

일반적 교훈

  1. 하드웨어 기능 적극 활용: prmt, lop3와 같은 저수준 CUDA 명령어는 복잡한 비트 조작 및 데이터 재배열을 매우 효율적으로 처리할 수 있습니다. 이러한 명령어들을 활용하면 소프트웨어 루프나 LUT 기반 접근보다 훨씬 높은 성능을 얻을 수 있습니다.
  2. 분기 최소화: 특히 가장 안쪽 루프에서의 데이터 종속적인 분기(divergent branch)는 GPU 성능에 치명적입니다. 이러한 분기를 제거하고 조건부 연산(conditional computation)을 활용하거나, 데이터를 재구성하여 분기 없는 경로로 처리하는 것이 중요합니다.
  3. 메모리 접근 패턴 최적화: LUT를 로컬 메모리에 저장하고 반복적으로 로드하는 대신, 레지스터에 상주하는 연산으로 대체함으로써 메모리 접근 병목을 해소하고 레지스터 활용도를 높일 수 있습니다.
  4. 정밀도와 성능의 균형: NVFP4는 FP8보다 적은 메모리를 사용하지만, 디퀀타이제이션 과정이 복잡하면 오히려 성능이 저하될 수 있습니다. 이 PR은 NVFP4의 장점을 살리면서도 디퀀타이제이션 연산의 효율성을 극대화하여, FP8 대비 경쟁력 있는 성능을 확보했습니다.
  5. 정확한 벤치마킹: PR에서는 FP8 KV 캐시를 기준으로 성능을 측정했습니다. 이는 NVFP4가 FP8보다 적은 데이터를 읽는다는 점을 고려할 때, FP8의 디퀀타이제이션 비용이 이미 최적화되어 있다는 점을 감안한 합리적인 비교 기준입니다. 또한, 측정 시 _nvfp4_kv_requires_disabled_split_kv()와 같은 기존의 제약 조건이 NVFP4 decode 성능에 미치는 영향을 명확히 설명하여 측정 결과의 신뢰도를 높였습니다.

리뷰 피드백 반영

리뷰어 노트에서 언급된 fast_dequant_f8f16x4static_cast 간의 NaN 인코딩 차이는 이 PR에서 직접 수정되지 않았지만, 해당 이슈가 인지되고 있음을 명확히 했습니다. 이는 FP8 KV 경로 자체의 잠재적 이슈이며, 향후 별도의 PR에서 다루어질 수 있음을 시사합니다. 또한, 스케일 곱셈을 MMA 누산기(accumulator)로 옮기는 것이 왜 더 나은 선택이 아닌지에 대한 설명은 최적화 결정의 근거를 명확히 했습니다.

결론

이 PR은 FlashInfer의 NVFP4 KV 캐시 연산에서 발생하던 심각한 성능 병목 현상을 CUDA 하드웨어 기능을 적극적으로 활용하고, 비효율적인 소프트웨어 경로를 제거함으로써 성공적으로 해결했습니다. prmt 명령어와 같은 저수준 최적화를 통해 디퀀타이제이션 및 스케일 팩터 변환 과정을 혁신적으로 개선하여, NVFP4 KV 캐시의 성능을 이전 대비 2~3배 이상 향상시켰습니다. 이는 LLM 추론 성능 최적화에 있어 양자화 기법의 효율적인 구현이 얼마나 중요한지를 다시 한번 보여주는 사례입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글