[onnxruntime] ONNX Runtime CUDA 커널 최적화: Speculative Decoding을 위한 GEMV 확장
PR 링크: microsoft/onnxruntime#32289 상태: Merged | 변경: +469 / -200
들어가며
최신 LLM 추론 환경에서 Speculative Decoding은 지연 시간을 줄이기 위한 핵심 기술입니다. 하지만 이 과정에서 발생하는 다중 토큰 검증(Multi-token verification)은 행렬 연산의 행(row) 개수(M)를 기존의 단일 토큰 디코딩 범위를 넘어서게 만듭니다. 기존 ONNX Runtime의 CUDA 커널은 M이 작을 때(주로 M=8 이하) 최적화된 GEMV 경로를 타지만, 그 이상의 행렬 크기에서는 일반적인 cuBLAS GEMM으로 폴백(fallback)되어 dequantization 및 workspace 오버헤드가 발생했습니다. 본 PR은 이 임계치를 M=64까지 확장하여 성능 병목을 제거합니다.
코드 분석
1. matmul_block_scaled_fp4.cc: 동적 임계치 도입
기존에는 kGemvMaxM이 8로 고정되어 있었으나, 이제는 MatMulBlockQuantizedFp4WeightGemvMaxM 함수를 통해 디바이스 속성과 K 차원에 따라 동적으로 결정됩니다.
// Before
constexpr int kGemvMaxM = 8;
if (m_i > 0 && m_i <= kGemvMaxM && ...)
// After
const int gemv_max_m = MatMulBlockQuantizedFp4WeightGemvMaxM(k_i, GetDeviceProp());
if (m_i > 0 && m_i <= gemv_max_m && ...)
2. matmul_block_scaled_fp4.cu: 커널 내 루프 언롤링 및 타일링
핵심은 MatMulBlockQuantizedFp4WeightMmaGemvKernel 내에서 여러 M 타일을 처리하도록 루프를 확장한 것입니다. 이제 커널은 MTiles 템플릿 인자를 통해 8행 단위의 타일을 여러 개 처리하며, 레지스터 압박을 최소화하면서도 가중치 재사용 효율을 극대화합니다.
// After: 여러 M 타일을 처리하기 위한 루프 구조
#pragma unroll
for (int mt = 0; mt < MTiles; ++mt) {
const int row = g + (mt << 3);
a_ok[mt] = row < m;
a_row[mt] = a + static_cast<size_t>(a_ok[mt] ? row : 0) * k;
}
3. LaunchMatMulBlockQuantizedFp4WeightGemv: 분할 실행(Split-launch)
M이 32를 초과하는 경우, 하나의 거대한 커널 대신 두 번의 서브-런칭을 통해 기존의 고성능 GEMV 경로를 유지합니다.
if (m > kFp4MmaGemvTileM) {
// 첫 번째 32행 처리
ORT_RETURN_IF_ERROR(LaunchMatMulBlockQuantizedFp4WeightGemv(..., kFp4MmaGemvTileM, ...));
// 나머지 행 처리
return LaunchMatMulBlockQuantizedFp4WeightGemv(..., m - kFp4MmaGemvTileM, ...);
}
왜 이게 좋은가
- 오버헤드 제거: Speculative Decoding 시 발생하는 9, 17, 33, 64 행 크기에서 cuBLAS로의 폴백을 방지하여 메모리 대역폭과 연산 효율을 극대화했습니다.
- 유연한 확장성:
ORT_FP4_GEMV_MAX_M환경 변수를 통해 런타임에 최적화 임계치를 조정할 수 있게 설계되었습니다. - 리뷰 피드백 반영: 리뷰어들의 지적에 따라 BF16/FP16에 대한 테스트 케이스를 보강하고,
counter초기화 책임을 커널 내부로 이동시켜 API 사용성을 개선했습니다.
이 최적화는 특정 연산 패턴(메모리 바운드 GEMV)이 자주 발생하는 추론 워크로드에서 커널의 범용성보다 특수화된 경로(Specialized path)를 유지하는 것이 얼마나 중요한지 보여주는 좋은 사례입니다.
참고 자료
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [onnxruntime] ONNX Runtime의 FP4/FP8 GEMV 커널 최적화: Tensor Core와 M-Tiling을 통한 성능 극대화
- [sglang] SGLang NGRAM 성능 최적화: 호스트 기반 트리 링크 유도로 GPU 병목 제거하기
- [onnxruntime] [ONNX Runtime] SGEMM의 함정에서 벗어나기: GQA 전용 GEMV 커널을 통한 디코딩 최적화
- [onnxruntime] ONNX Runtime: MoE Router GEMV 최적화 및 Bias Fusion 구현
- [onnxruntime] ONNX Runtime CUDA Graph: 진정한 비동기 추론을 위한 동기화 지점 제거
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer의 Blackwell 아키텍처를 위한 Cake All-Gather Matmul 최적화 분석
- 현재글 : [onnxruntime] ONNX Runtime CUDA 커널 최적화: Speculative Decoding을 위한 GEMV 확장
- 다음글 [sglang] SGLang, NVIDIA Blackwell GPU를 위한 Wan2.2 모델의 ML P 연산 최적화
댓글