[vllm] vLLM, Blackwell 아키텍처를 위한 디코드 성능 최적화: GLM-5.2 및 DeepSeek-V3.2 지원 강화
PR 링크: vllm-project/vllm#48597 상태: Merged | 변경: +2139 / -138
들어가며
최근 대규모 언어 모델(LLM)의 발전은 GPU 하드웨어의 성능 향상과 밀접한 관련이 있습니다. 특히 NVIDIA의 Blackwell 아키텍처(SM100, GB200, GB300)는 이전 세대 대비 상당한 성능 개선을 제공하며, LLM 추론 성능을 극대화하기 위한 최적화의 필요성이 더욱 커지고 있습니다. vLLM은 이러한 최신 하드웨어의 잠재력을 최대한 활용하기 위해 지속적으로 노력해왔습니다. 이번 PR은 vLLM이 GLM-5.2 및 DeepSeek-V3.2 모델의 디코드(decode) 성능을 Blackwell 아키텍처에서 최적화하는 데 중점을 둡니다. 이 글에서는 해당 PR에서 이루어진 주요 코드 변경 사항을 분석하고, 이러한 변경이 왜 성능 향상으로 이어지는지, 그리고 그 일반적인 교훈은 무엇인지 살펴보겠습니다.
코드 분석
이번 PR은 여러 파일에 걸쳐 다양한 최적화를 적용했습니다. 주요 변경 사항을 파일별로 나누어 살펴보겠습니다.
1. CMakeLists.txt: BF16 Skinny GEMM 빌드 설정 추가
이 변경은 새로운 커널인 bf16_skinny_gemm을 빌드하기 위한 CMake 설정을 추가합니다. 이 커널은 특정 조건(SM90+ 및 CUDA 12.0 이상)에서만 활성화되며, BF16 데이터 타입을 사용하는 'skinny' GEMM 연산을 지원합니다. 이는 특히 디코드 시 가중치 대역폭에 민감한 연산에 대한 성능 향상을 목표로 합니다.
Before:
--- a/CMakeLists.txt
+++ b/CMakeLists.txt
@@ -727,6 +727,23 @@
"(requires SM90+ and CUDA >= 12.0).")
endif()
+ # BF16 skinny GEMM (M<=32; weight-bandwidth-bound decode shapes).
+ # Requires SM90+.
+ cuda_archs_sm90plus(BF16_SKINNY_GEMM_ARCHS "${CUDA_ARCHS}")
+ if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.0 AND BF16_SKINNY_GEMM_ARCHS)
+ set(BF16_SKINNY_GEMM_SRCS
+ "csrc/libtorch_stable/bf16_skinny_gemm_entry.cu"
+ "csrc/libtorch_stable/bf16_skinny_gemm.cu")
+ set_gencode_flags_for_srcs(
+ SRCS "${BF16_SKINNY_GEMM_SRCS}"
+ CUDA_ARCHS "${BF16_SKINNY_GEMM_ARCHS}")
+ list(APPEND VLLM_STABLE_EXT_SRC "${BF16_SKINNY_GEMM_SRCS}")
+ message(STATUS "Building bf16_skinny_gemm for archs: ${BF16_SKINNY_GEMM_ARCHS}")
+ else()
+ message(STATUS "Not building bf16_skinny_gemm as no compatible archs found "
+ "(requires SM90+ and CUDA >= 12.0).")
+ endif()
+
# Only build AllSpark kernels if we are building for at least some compatible archs.
cuda_archs_loose_intersection(ALLSPARK_ARCHS "8.0;8.6;8.7;8.9" "${CUDA_ARCHS}")
if (ALLSPARK_ARCHS)
After:
새로운 bf16_skinny_gemm 관련 소스 파일(bf16_skinny_gemm_entry.cu, bf16_skinny_gemm.cu)이 VLLM_STABLE_EXT_SRC에 추가되었습니다. 이는 해당 커널이 컴파일되고 사용될 수 있도록 합니다.
2. csrc/libtorch_stable/bf16_skinny_gemm.cu: BF16 Skinny GEMM 커널 구현
이 파일은 새로운 bf16_skinny_gemm 커널의 핵심 로직을 구현합니다. 이 커널은 다음과 같은 특징을 가집니다:
- 목표: M이 작고(M <= 32) K 차원이 큰(weight-bandwidth-bound) 디코드 시나리오에 최적화된 GEMM 연산입니다. 특히 GLM/DeepSeek 모델의
eh_proj와 같은 연산에서 기존 cuBLAS splitK 방식보다 빠른 성능을 목표로 합니다. - 데이터 타입: 입력 행렬 A와 가중치 행렬 B 모두 BF16 타입을 사용하며, 출력은 BF16으로 생성됩니다. 연산 중간에는 FP32로 누적됩니다.
- 로드 방식:
load_bf16x8및load_bf16x8_cs함수를 사용하여 BF16 데이터를 효율적으로 로드하고 FP32로 변환합니다.load_bf16x8_cs는 스트리밍 로드를 사용하여 가중치 재사용성을 높입니다. - Reduction: 워프(warp) 수준의 버터플라이(butterfly) 감소와 공유 메모리(shared memory)를 활용하여 K 차원을 효율적으로 줄입니다.
- Prefetching:
kPF(kernel prefetch depth) 옵션을 통해 일부 가중치 로드를 미리 수행하여 메모리 로딩 지연 시간을 숨기려고 시도합니다. 이는 특히 M=1일 때 성능 향상에 기여합니다. - Grid/Block 구성: 각 블록은
kNPB개의 출력 열을 계산하며, 그리드는 전체 N 차원을 커버합니다. 스레드 블록은kBlockSize개의 스레드를 가집니다.
Before (Conceptual): 기존에는 cuBLAS와 같은 라이브러리의 일반적인 GEMM 구현을 사용했을 것입니다. 이는 다양한 시나리오에 범용적으로 적용 가능하지만, 특정 'skinny' 형태의 연산에서는 최적이 아닐 수 있습니다.
After:
--- /dev/null
+++ b/csrc/libtorch_stable/bf16_skinny_gemm.cu
@@ -0,0 +1,262 @@
+// SPDX-License-Identifier: Apache-2.0
+// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
+//
+// Skinny GEMM: activation(bf16) x weight(bf16)^T -> bf16, for decode-time
+// M <= 32 with a large reduction dim. Replaces cuBLAS splitK (GEMM +
+// splitKreduce) with a single block-per-output-column kernel; these shapes
+// are weight-bandwidth-bound, so one coalesced pass over the weight at
+// fp32 accumulation is optimal. Adapted from fp32_router_gemm.cu.
+//
+// First user: the DeepSeek-V32/GLM-5.2 MTP eh_proj (K=2*hidden=12288,
+// N=hidden/TP), whose cuBLAS splitK pick costs ~34.6us vs the ~19us
+// bandwidth floor per replicated read (and ~4us once column-parallel).
+
+#include <cuda_bf16.h>
+#include <cuda_runtime.h>
+
+// ... (kernel implementation details as described above) ...
+
+// ---------------------------------------------------------------------------
+// Explicit instantiations. M = 1..32; (N, K) pairs:
+// (768, 12288) eh_proj shard, TP8
+// (1536, 12288) eh_proj shard, TP4
+// (6144, 12288) eh_proj unsharded
+// kNPB (B300 sweep, M=1): 6144 -> 2 (3072 blocks, 23.2us vs 25.3 at kNPB=8;
+// narrow blocks minimize wave quantization); shards 768/1536 keep 4.
+// ---------------------------------------------------------------------------
+
+#define INSTANTIATE(M, NPB, N, K) \
+ template void invokeBf16SkinnyGemm<128, NPB, M, N, K>(\
+ __nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, int64_t, \
+ cudaStream_t);
+
+#define INSTANTIATE_ALL_M(NPB, N, K) \
+ INSTANTIATE(1, NPB, N, K) \
+ INSTANTIATE(2, NPB, N, K) \
+ // ... (instantiations for M=3 to 32) ...
+
+INSTANTIATE_ALL_M(4, 768, 12288)
+INSTANTIATE_ALL_M(4, 1536, 12288)
+INSTANTIATE_ALL_M(2, 6144, 12288)
+// LL-mode (M<=8 wiring guard) backbone shapes, B300 sweep vs cuBLAS:
+// q_b_proj (2048, 2048): 1.67x/1.29x/1.15x at M=4/6/8 (NPB=4 within
+// 0.1us of per-M best)
+// shared-expert gate_up (512, 6144): 1.95x/1.58x/1.40x at M=4/6/8
+// cuBLAS keeps qkv_a (2624,6144) and o_proj (6144,2048) — already at
+// 3.4-3.8 TB/s there; the GEMV loses on activation re-reads.
+INSTANTIATE_ALL_M(4, 2048, 2048)
+INSTANTIATE_ALL_M(4, 512, 6144)
+// fused_qkv_a (2624, 6144), 32MB: skinny wins ONLY at M<=2 (B300: M=1
+// 6.99us vs cuBLAS 9.21 = 1.32x, M=2 1.24x; M>=4 cuBLAS holds at 3.5TB/s
+// and every alternative loses — cublas
+```
이 파일은 `bf16_skinny_gemm_kernel`이라는 CUDA 커널을 정의하고, 다양한 `M`, `N`, `K` 차원에 대해 이 커널을 인스턴스화하여 사용할 수 있도록 합니다. 이는 특정 모델 구조와 하드웨어 특성에 맞춰 최적화된 연산을 제공합니다.
### 3. 기타 변경 사항 (요약)
PR 설명에 따르면, 이 작업은 다음과 같은 여러 영역의 최적화를 포함합니다:
* **Router / MoE GEMM:** FP32 라우터 GEMM의 GLM/MiniMax 형태 지원 및 최적화, BF16 라우터 GEMM, 라우터 가중치 FP32 업캐스팅 방지, 디코드 M GEMM 디스패치 최적화 (`skinny GEMV`, `min-latency fused_a`), GLM 퓨즈드 Q 커널 개선.
* **MLA / Attention / DSA:** DSA(Dynamic Sparse Attention)에서 스킵-탑K 레이어 간 변환된 물리 인덱스 캐싱, 논리 탑K 버퍼와 DSA 물리 인덱스 정렬 유지.
* **MTP / Spec decode:** MTP(Multi-Token Prediction)의 `get_top_tokens`에서 로컬-아르그맥스(local-argmax) 감소 최적화.
* **Fusion / All-reduce / Norm:** 퓨즈드 All-reduce에서 FlashInfer의 자동 선택 최적화, 퓨즈드 AR+RMSNorm에서 오버사이즈 시 폴백(fallback) 메커니즘, `fused_norm_rope`의 `num_warps=1` 설정 (B300에서 약 6% 커널 시간 단축).
* **Quantization / PDL:** 작은 커널에 대한 PDL(Parallel Data Layout) 활성화, TRTLLM-Gen MoE 행(hang) 문제 해결을 위한 `VLLM_DISABLE_FLASHINFER_PDL` 옵션 추가.
이러한 변경들은 GPU 커널 수준에서의 연산 효율성 증대, 메모리 접근 패턴 최적화, 불필요한 연산 제거 등을 통해 전반적인 디코드 성능을 향상시키는 것을 목표로 합니다.
## 왜 이게 좋은가?
이번 PR의 핵심은 최신 Blackwell 아키텍처의 특성을 고려하여 LLM 디코드 연산, 특히 MoE(Mixture of Experts) 및 Sparse Attention과 같은 복잡한 연산의 성능을 극대화하는 데 있습니다.
1. **하드웨어 특화 최적화:** `bf16_skinny_gemm`과 같은 새로운 커널은 Blackwell GPU의 BF16 연산 능력과 메모리 대역폭 특성을 활용하도록 설계되었습니다. 특히 M이 작고 K가 큰 'skinny' GEMM 연산은 기존 범용 GEMM 라이브러리보다 훨씬 효율적일 수 있습니다. 리뷰어의 벤치마크 결과에 따르면, GB300에서 M=1일 때 TPOT(Time Per Output Token)이 약 1.61ms로 매우 낮은 지연 시간을 달성했습니다. 이는 이전보다 훨씬 빠른 응답 속도를 의미합니다.
```diff
# 리뷰어 벤치마크 결과 (2x GB300, GLM-5.2-NVFP4, M=1)
# TPOT (ms): 1.61 (이 PR 적용 후)
```
2. **메모리 대역폭 최적화:** DSA에서 물리 인덱스를 캐싱하거나(`DSA: cache converted physical indices across skip_topk layers`), 가중치 로딩을 최적화하는(`load_bf16x8_cs`) 등의 기법은 메모리 접근 패턴을 개선하여 대역폭 병목 현상을 줄입니다. 이는 특히 LLM 추론에서 중요한 요소입니다.
3. **연산 효율성 증대:** MTP `get_top_tokens`에서의 로컬-아르그맥스 감소, `fused_norm_rope`의 `num_warps=1` 설정 등은 연산량을 줄이거나 GPU 코어 활용을 최적화하여 커널 실행 시간을 단축시킵니다. `fused_norm_rope`의 경우 B300에서 약 6%의 커널 시간 단축이 보고되었습니다.
```diff
# fused_norm_rope 커널 시간 단축 (B300)
# num_warps=1 (bit-exact, ~6% kernel time on B300)
```
4. **정확도 유지:** 이러한 성능 최적화 과정에서 모델의 정확도 손실이 없음을 확인했습니다. GSM8K 벤치마크에서 이전과 동일한 정확도를 유지했습니다.
```diff
# lm_eval gsm8k 정확도 (5-shot)
# flexible-extract: 0.9477 (변경 없음)
# strict-match: 0.9469 (변경 없음)
```
**일반적인 교훈:**
* **하드웨어 특화 코딩의 중요성:** 최신 GPU 아키텍처의 새로운 기능(예: BF16 연산, 특정 메모리 접근 패턴)을 활용하는 커널을 직접 구현하는 것이 성능 향상의 핵심입니다.
* **병목 지점 식별 및 해결:** 디코드 성능의 병목이 어디에 있는지(계산, 메모리 대역폭, 레이턴시 등) 정확히 파악하고, 해당 병목을 해결하는 데 집중해야 합니다.
* **작은 단위의 최적화 누적:** 여러 작은 커널 최적화, 데이터 로딩 개선, 연산 감소 등이 모여 전체 시스템 성능에 큰 영향을 미칩니다.
* **정확도와 성능의 균형:** 성능 최적화는 모델의 정확도를 희생시키지 않는 범위 내에서 이루어져야 합니다. 철저한 테스트를 통해 이를 검증해야 합니다.
## 리뷰 피드백 반영
리뷰 과정에서 몇 가지 중요한 논의가 있었습니다. 특히 `chaunceyjiang`님은 DSA에서 물리 인덱스를 전역 레지스트리 대신 어텐션 메타데이터 내에 유지하는 방안을 제안했습니다. 이는 코드를 더 깔끔하게 만들고 최적화를 통합하는 데 도움이 될 수 있습니다. 비록 이 PR에서는 기존 방식을 유지하고 후속 PR에서 개선하기로 했지만, 이러한 논의는 코드 품질과 유지보수성을 높이는 데 기여했습니다.
또한, PR의 크기가 너무 커서 병합이 어렵다는 피드백이 있었고, 이에 따라 일부 작업은 이미 메인 브랜치에 병합되었음을 명확히 했습니다. 이는 대규모 PR을 관리하는 데 있어 중요한 고려 사항입니다.
## 결론
이번 vLLM PR은 NVIDIA Blackwell 아키텍처에서 GLM-5.2 및 DeepSeek-V3.2 모델의 디코드 성능을 획기적으로 개선하는 중요한 작업을 수행했습니다. BF16 Skinny GEMM 커널 도입, 메모리 접근 패턴 최적화, 연산 효율성 증대 등 다양한 기법을 통해 낮은 지연 시간과 높은 처리량을 달성했습니다. 이는 최신 하드웨어의 잠재력을 최대한 활용하려는 vLLM의 지속적인 노력을 보여주는 좋은 예시이며, LLM 추론 성능 향상에 크게 기여할 것입니다.
## 참고 자료
- https://github.com/vllm-project/vllm/blob/main/CMakeLists.txt
- https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/bf16_skinny_gemm.cu
- https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/bf16_skinny_gemm_entry.cu
> ⚠️ **알림:** 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM, DeepSeek-V3.2/GLM-5.2 MTP 경로 최적화: All-Reduce 융합 및 로컬 Argmax 도입
- [vllm] vLLM, 캐싱 비활성화 시 불필요한 LRU 해시 분할 제거로 디코드 처리량 3.5% 향상
- [vllm] vLLM, DeepSeek-V3.2 모델의 ROCm 성능 최적화: CPU 측 마이크로 최적화 3가지 분석
- [vllm] vLLM, Cohere 임베딩 바이너리 압축 성능 4배 개선: NumPy를 활용한 최적화 분석
- [vllm] vLLM에 Dots3 NOTE 모델 네이티브 지원 추가: 멀티모달 및 하이브리드 MLA 최적화
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime: Blackwell (SM120+)에서 NVFP4 QMoE를 위한 네이티브 FP4xFP4 Prefill 최적화
- 현재글 : [vllm] vLLM, Blackwell 아키텍처를 위한 디코드 성능 최적화: GLM-5.2 및 DeepSeek-V3.2 지원 강화
- 다음글 [flashinfer] FlashInfer의 FP8 양자화 AllReduce를 통한 통신 대역폭 최적화
댓글