[flashinfer] FlashInfer Kimi-K3 Fused MoE Router 최적화: Warp-per-row 전략 도입
PR 링크: flashinfer-ai/flashinfer#5564 상태: Merged | 변경: +1683 / -1621
들어가며
FlashInfer의 Kimi-K3 Fused MoE Router는 대규모 언어 모델의 MoE(Mixture of Experts) 연산 효율을 극대화하기 위한 핵심 컴포넌트입니다. 특히 토큰 수가 많은 대규모 배치(4096, 8192) 환경에서 기존의 'thread-per-expert' 방식은 병목 현상을 유발했습니다. 이번 업데이트에서는 이를 해결하기 위해 'warp-per-row' 기반의 새로운 라우팅 암(arm)인 GW를 도입하여, SM100 및 SM103 아키텍처에서의 연산 효율을 크게 개선했습니다.
코드 분석
1. cake_backend.py: 라우팅 전략 교체
기존 G 암에서 GW 암으로 대규모 배치의 라우팅 로직을 변경했습니다. 핵심은 ctas_per_sm 설정을 최적화하여 더 효율적인 그리드 실행을 보장하는 것입니다.
# Before
ARM_G_CTAS_PER_SM = {(10, 0): 4, (10, 3): 6}
# After
ARM_GW_CTAS_PER_SM = {(10, 0): 4, (10, 3): 4}
또한, SHAPE_ROUTES 매핑을 통해 대규모 배치(4096, 8192)가 새로운 GW 로직을 사용하도록 명시했습니다.
# After
(4096, 8): "GW",
(4096, 16): "GW",
(8192, 8): "GW",
(8192, 16): "GW",
2. cake_jit.py: JIT 모듈 업데이트
새로운 GW 암에 대응하는 JIT 모듈 정의를 업데이트했습니다. 기존 G 모듈을 GW로 교체하고, 각 아키텍처(sm_100a, sm_103a)에 맞는 커널 소스와 closure 해시를 갱신하여 최적화된 바이너리가 생성되도록 했습니다.
왜 이게 좋은가
이번 최적화의 핵심은 두 가지입니다.
- Warp-per-row Selection: 각 워프가 토큰 행을 전담하여 16개의 전문가를 선택합니다. 이는 기존의 스레드 단위 선택보다 공유 메모리 접근 효율이 높고, 정렬 과정에서 발생하는 지연 시간을 획기적으로 줄여줍니다.
- Dual-expert Sort Phase: 4개의 CTA가 SM당 배치될 때, 896개의 전문가 세그먼트를 2개씩 묶어 처리함으로써 정렬 단계를 2라운드에서 1라운드로 단축했습니다.
성능 지표:
- NVIDIA B200: 기하평균 기준 약 1.49배의 속도 향상.
- NVIDIA GB300: 기하평균 기준 약 1.87배의 속도 향상.
이러한 최적화는 대규모 MoE 모델을 서빙할 때 라우팅 오버헤드를 최소화하여 전체 추론 파이프라인의 처리량(Throughput)을 높이는 데 기여합니다. 특히 고성능 GPU 아키텍처에서 커널 실행 그리드를 아키텍처 특성에 맞게 튜닝하는 것이 얼마나 중요한지 보여주는 좋은 사례입니다.
참고 자료
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#launch-bounds
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#thread-hierarchy
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer NVFP4 QKV GEMM 최적화: SM103a Epilogue 통합 및 CUDA 런처 개선
- [flashinfer] FlashInfer: Blackwell 아키텍처를 위한 결정론적 BGMV MoE 최적화
- [flashinfer] DeepSeek-V3 라우팅의 혁신: FlashInfer의 Cake 백엔드 가속 분석
- [flashinfer] FlashInfer의 Mixture-of-Experts(MoE) 라우팅 성능 최적화 분석
- [flashinfer] FlashInfer MoE All-to-All 최적화: TRT-LLM의 성능 비결을 파헤치다
PR Analysis 의 다른글
- 이전글 [sglang] SGLang 성능 최적화: RTX 4090에서 MXFP4 MoE 추론 속도 6배 향상시키기
- 현재글 : [flashinfer] FlashInfer Kimi-K3 Fused MoE Router 최적화: Warp-per-row 전략 도입
- 다음글 [flashinfer] NVIDIA Blackwell의 잠재력을 극한으로: MiniMax-H3 NVFP4 양자화 및 GEMM 최적화 분석
댓글