본문으로 건너뛰기

[flashinfer] FlashInfer의 Blackwell 아키텍처 최적화: CAKE 기반 TinyGEMM2 커널 도입

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

들어가며

최신 GPU 아키텍처인 NVIDIA Blackwell(SM100/103) 환경에서 LLM 추론 성능을 극대화하기 위해, FlashInfer는 기존 tinygemm2 구현을 대체할 새로운 커널을 도입했습니다. 이번 PR은 CAKE(Compiler-Assisted Kernel Engineering)를 통해 생성된 최적화된 커널을 통합하여, 비트 단위의 정확도(bit-identical)를 유지하면서도 추론 지연 시간을 대폭 단축하는 것을 목표로 합니다.

코드 분석

1. csrc/tinygemm2_sm100.cu (신규 커널 구현)

이 파일은 CAKE로 생성된 4가지 변형 커널(stage4/8, 각각 PDL 사용 여부에 따른 조합)을 하나의 Translation Unit(TU)으로 통합한 핵심 파일입니다. 기존의 tinygemm2.cu 구조를 따르며, 호스트 바인딩 섹션이 추가되었습니다.

Before (기존 방식):

// 기존 tinygemm2.cu는 범용적인 구현을 사용
// 특정 아키텍처에 최적화된 스케줄링이 부족함

After (개선된 방식):

// csrc/tinygemm2_sm100.cu
// 커널 심볼을 변형별로 명시적으로 분리하여 최적화된 스케줄링 적용
extern "C" {
__global__ __launch_bounds__(384, 1) void
kernel_tinygemm2_sm100_stage4(__grid_constant__ LoomTensorMap const tmap_wt, ...)
{
    // ... TMA(Tensor Memory Accelerator)를 활용한 고성능 데이터 이동 로직
    mbarrier_init_pred(smem + 0, 1, leader);
    // ...
}
}

핵심은 TMAmbarrier를 활용한 비동기 데이터 복사입니다. Blackwell의 하드웨어 가속 기능을 직접 호출하여 메모리 대역폭 효율을 극대화했습니다.

2. flashinfer/gemm/routergemm.py (라우팅 로직)

새로운 커널을 사용하기 위해 런타임 디스패처를 수정했습니다. compute capability가 10.0 이상인 경우 자동으로 최적화된 커널로 라우팅됩니다.

왜 이게 좋은가

이번 최적화의 핵심은 하드웨어 특화 스케줄링입니다.

  1. 성능 향상: B200 환경에서 벤치마크 결과, 특정 워크로드(64, 4096, 3072)에서 기존 대비 1.79배의 성능 향상을 보였습니다. 전체적인 기하 평균 기준 18-23%의 지연 시간 감소를 달성했습니다.
  2. 안정성: 비트 단위의 정확도(bitwise equality)를 보장하여, 기존 모델의 가중치나 결과값에 영향을 주지 않고 즉시 교체 가능합니다.
  3. 유연성: FLASHINFER_DISABLE_TINYGEMM2_SM100=1 환경 변수를 통해 문제가 발생할 경우 즉시 기존 구현으로 롤백할 수 있는 안전장치를 마련했습니다.

교훈: 범용 커널보다 특정 아키텍처(Blackwell)의 TMA 및 mbarrier 특성을 직접 활용하는 커널을 생성(CAKE)하여 사용하는 것이 최신 GPU에서 성능을 뽑아내는 가장 효과적인 방법임을 보여줍니다.

리뷰어 피드백

리뷰어들은 생성된 커널 코드 자체보다는 바인딩 섹션의 정확성과 라우팅 로직의 안전성에 집중했습니다. 특히 cuTensorMapEncodeTiled 호출이 기존 레퍼런스와 일치하는지 검증하는 과정이 중요하게 다뤄졌으며, CI를 통해 다양한 shape에 대한 parity 테스트를 통과함으로써 안정성을 확보했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글