본문으로 건너뛰기

[flashinfer] FlashInfer: B200용 최적화된 Recurrent KDA Prefill 백엔드 도입

PR 링크: flashinfer-ai/flashinfer#4262 상태: Merged | 변경: +8742 / -2

들어가며

최신 LLM 추론 환경에서 Recurrent KDA(Kimi Delta Attention)는 효율적인 상태 업데이트를 위해 필수적인 연산입니다. 기존 FlashInfer의 CuTe-DSL 기반 구현은 범용적이지만, 최신 NVIDIA B200(SM100a) 아키텍처의 하드웨어 잠재력을 100% 활용하기에는 한계가 있었습니다. 본 PR은 B200 아키텍처에 특화된 최적화된 Recurrent KDA Prefill 백엔드를 도입하여, 상태 업데이트를 커널 내에서 직접 수행하고 메모리 복사 오버헤드를 제거함으로써 성능을 극대화합니다.

코드 분석

1. 커널 내 In-place 상태 업데이트

기존 방식은 상태 업데이트를 위해 별도의 scratch 메모리를 할당하고 커널 실행 후 memcpy를 수행해야 했습니다. 이번 변경에서는 커널이 직접 initial_state 포인터를 수정하여 메모리 복사 단계를 완전히 제거했습니다.

// Before: 별도 할당 및 복사 필요
// After: 커널 내 직접 업데이트
// M128 assigns one CTA to each (sequence, head).
// Each CTA loads all initial-state rows it owns before storing those same final-state rows.

2. Tensor Map 및 CUDA Graph 최적화

B200의 TMA(Tensor Memory Accelerator) 성능을 극대화하기 위해 tensormap을 사전에 준비하고, 커널 실행 전 fence.proxy.tensormap::generic.acquire.gpu를 통해 메모리 일관성을 보장합니다.

// Tensor maps are prepared outside capture and published to stable global storage
// Every consumer CTA executes fence.proxy.tensormap::generic.acquire.gpu
fence.proxy.tensormap::generic.acquire.gpu;
__syncthreads();

왜 이게 좋은가

이번 최적화는 단순히 알고리즘을 개선한 것이 아니라, 하드웨어의 특성(SM100a)을 깊이 이해하고 적용한 결과입니다. 벤치마크 결과, 기존 FlashKDA 대비 기하 평균 2.05배의 성능 향상을 달성했습니다.

  • 성능 수치: H=64 mixed 케이스에서 최대 2.50배 속도 향상.
  • 교훈:
    1. 메모리 오버헤드 제거: 커널 내 in-place 업데이트는 고대역폭 메모리(HBM) 사용량을 줄이고 지연 시간을 획기적으로 낮춥니다.
    2. 하드웨어 특화: B200의 TMA와 같은 최신 기능을 활용하기 위해 커널 수준의 정밀한 제어가 필요합니다.
    3. ABI 준수: __restrict__와 같은 컴파일러 힌트를 적절히 사용하여 최적화된 기계어를 생성하도록 유도했습니다.

결론

이번 PR은 FlashInfer가 최신 하드웨어에서 최고의 성능을 낼 수 있도록 하는 중요한 이정표입니다. 특히 CUDA Graph와의 호환성을 유지하면서도 성능을 극대화한 점은 실무적인 관점에서 매우 훌륭한 설계입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글