[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배 속도 향상.
- 교훈:
- 메모리 오버헤드 제거: 커널 내 in-place 업데이트는 고대역폭 메모리(HBM) 사용량을 줄이고 지연 시간을 획기적으로 낮춥니다.
- 하드웨어 특화: B200의 TMA와 같은 최신 기능을 활용하기 위해 커널 수준의 정밀한 제어가 필요합니다.
- ABI 준수:
__restrict__와 같은 컴파일러 힌트를 적절히 사용하여 최적화된 기계어를 생성하도록 유도했습니다.
결론
이번 PR은 FlashInfer가 최신 하드웨어에서 최고의 성능을 낼 수 있도록 하는 중요한 이정표입니다. 특히 CUDA Graph와의 호환성을 유지하면서도 성능을 극대화한 점은 실무적인 관점에서 매우 훌륭한 설계입니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#tma-descriptor
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer의 GDN 커널 런칭 오버헤드 80% 절감하기: 호스트 측 최적화 전략
- [flashinfer] FlashInfer FP8 Causal Attention 최적화: O(1) 디코딩과 글로벌 스케줄링의 힘
- [flashinfer] FlashInfer의 Mixture-of-Experts(MoE) 라우팅 성능 최적화 분석
- [flashinfer] FlashInfer의 Fused SwiGLU 및 NVFP4 양자화 최적화 분석
- [flashinfer] FlashInfer MoE 최적화: PDL 스케줄링 개선 및 GEMM2 균형 잡힌 스토어 구현
PR Analysis 의 다른글
- 이전글 [onnxruntime] [CUDA] QMoE MXFP4/NVFP4 가중치 역양자화 성능 최적화: Coalesced Memory Access의 힘
- 현재글 : [flashinfer] FlashInfer: B200용 최적화된 Recurrent KDA Prefill 백엔드 도입
- 다음글 [hermes-agent] Python 백엔드 콜드 스타트 성능 개선: 14초 GIL 스톨 해결하기
댓글