본문으로 건너뛰기

[flashinfer] FlashInfer의 GEMM 성능 혁신: cuTile 백엔드 도입과 최적화 여정

PR 링크: flashinfer-ai/flashinfer#4020 상태: Merged | 변경: +7921 / -114

들어가며FlashInfer는 고성능 딥러닝 추론을 위한 GPU 가속 라이브러리로, 특히 대규모 언어 모델(LLM)의 효율적인 연산에 중점을 둡니다. 행렬 곱셈(GEMM)은 이러한 모델의 핵심 연산으로, 성능에 지대한 영향을 미칩니다. 기존 FlashInfer는 CUTLASS, cuDNN, cuBLAS 등 다양한 백엔드를 활용했지만, 특정 시나리오, 특히 MoE(Mixture-of-Experts) 모델의 불균형한(ragged) 또는 마스킹된(masked) GEMM 연산, 그리고 FP8 정밀도 연산에서 성능 격차가 있거나 기능적 공백이 존재했습니다.이 PR은 NVIDIA의 저수준 GPU 프로그래밍 인터페이스인 cuTile API를 FlashInfer의 GEMM 연산에 도입하여 이러한 문제를 해결하고 성능을 극대화하는 것을 목표로 합니다. cuTile은 GPU 하드웨어의 기능을 최대한 활용하여 커스텀 커널을 작성할 수 있게 해주며, 이를 통해 기존 백엔드 대비 상당한 성능 향상과 새로운 기능(예: alpha-beta GEMM, masked-bmm, ragged-bmm, block-scaled FP8 GEMM)을 제공합니다. 이 글에서는 PR의 주요 코드 변경사항을 분석하고, 이러한 최적화가 왜 중요한지, 그리고 개발 과정에서 발견되고 해결된 흥미로운 문제점들을 살펴보겠습니다.## 핵심 변경사항 분석### 1. 새로운 cuTile GEMM 커널 및 벤치마크 추가 (benchmarks/bench_masked_scaled_bmm_backend_comparison.py 외)이 PR의 핵심은 cuTile을 활용한 새로운 GEMM 커널들을 FlashInfer에 통합하는 것입니다. gemm_alpha_beta, masked_bmm, ragged_bmm, ragged_block_scaled_bmm, gemm_fp8_nt_groupwise와 같은 다양한 커널이 추가되었으며, 이들의 성능을 검증하기 위한 벤치마크 파일도 함께 추가되었습니다. 특히 bench_masked_scaled_bmm_backend_comparison.pymasked_scaled_bmmcuTile 백엔드 성능을 기존 SOTA 백엔드와 비교합니다.```python# benchmarks/bench_masked_scaled_bmm_backend_comparison.py (부분 발췌)def make_call(provider, block_scale_type, num_groups, max_m, exp_m, N, K, out_dtype): # ... (생략) if provider ==

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글