본문으로 건너뛰기

[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 해시를 갱신하여 최적화된 바이너리가 생성되도록 했습니다.

왜 이게 좋은가

이번 최적화의 핵심은 두 가지입니다.

  1. Warp-per-row Selection: 각 워프가 토큰 행을 전담하여 16개의 전문가를 선택합니다. 이는 기존의 스레드 단위 선택보다 공유 메모리 접근 효율이 높고, 정렬 과정에서 발생하는 지연 시간을 획기적으로 줄여줍니다.
  2. Dual-expert Sort Phase: 4개의 CTA가 SM당 배치될 때, 896개의 전문가 세그먼트를 2개씩 묶어 처리함으로써 정렬 단계를 2라운드에서 1라운드로 단축했습니다.

성능 지표:

  • NVIDIA B200: 기하평균 기준 약 1.49배의 속도 향상.
  • NVIDIA GB300: 기하평균 기준 약 1.87배의 속도 향상.

이러한 최적화는 대규모 MoE 모델을 서빙할 때 라우팅 오버헤드를 최소화하여 전체 추론 파이프라인의 처리량(Throughput)을 높이는 데 기여합니다. 특히 고성능 GPU 아키텍처에서 커널 실행 그리드를 아키텍처 특성에 맞게 튜닝하는 것이 얼마나 중요한지 보여주는 좋은 사례입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글