[vllm] [ROCm 성능 최적화] vLLM의 Fused Shared-Expert Gate GEMM 경로 개선 분석
PR 링크: vllm-project/vllm#54185 상태: Merged | 변경: +5 / -2
들어가며: MoE와 Shared Expert의 융합, 그리고 숨겨진 성능 저하
최근 대규모 언어 모델(LLM) 아키텍처에서 MoE(Mixture of Experts)는 효율성을 높이는 핵심 기술로 자리 잡았습니다. 특히 DeepSeek-V3와 같은 모델에서 사용되는 Shared Expert 구조는 모든 토큰이 공통적으로 거치는 전문가 층을 두어 성능을 보완합니다. vLLM은 이러한 구조를 최적화하기 위해 Router Gate와 Shared-Expert Gate의 가중치를 하나로 합쳐(fused) 한 번의 연산으로 처리하는 VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS 기능을 제공합니다.
하지만 최근 ROCm 환경에서 이 융합(fused) 경로가 오히려 성능 저하를 일으키는 문제가 발견되었습니다. 원인은 단순했습니다. 융합된 경로가 표준 torch.nn.functional.linear를 사용하면서, ROCm 전용의 최적화된 Skinny-GEMM 커널을 타지 못하고 일반적인 Tensile 커널을 사용했기 때문입니다.
이번 포스트에서는 이 문제를 해결하기 위해 GEMM 연산을 플랫폼 디스패처(Platform Dispatcher)로 라우팅하여 성능을 복구하고, 향후 추가 최적화 가능성까지 열어둔 PR을 분석해 보겠습니다.
코드 분석: F.linear에서 dispatch_unquantized_gemm으로
문제의 핵심은 vllm/model_executor/layers/fused_moe/runner/moe_runner.py 파일에 있었습니다. 기존 코드는 융합된 가중치에 대해 단순히 F.linear를 호출하고 있었습니다.
1. 문제의 코드 (Before)
# vllm/model_executor/layers/fused_moe/runner/moe_runner.py
if self.gate is not None:
if self._fse_fuse_gate:
self._maybe_fuse_gate_weights()
# 단순 F.linear 호출: ROCm에서 최적화된 커널 선택 기회를 상실함
router_logits = F.linear(hidden_states, self._combined_gate_weight)
else:
router_logits, _ = self.gate(hidden_states)
위 코드에서 F.linear는 ROCm 환경에서 Tensile 라이브러리의 Cijk_* 커널을 호출하게 됩니다. 이 커널은 일반적인 행렬 곱셈에는 적합하지만, 디코딩(Decode) 단계처럼 입력 토큰 수가 적은(Skinny) 상황에서는 오버헤드가 큽니다.
2. 개선된 코드 (After)
# vllm/model_executor/layers/fused_moe/runner/moe_runner.py
from vllm.model_executor.layers.utils import dispatch_unquantized_gemm
# ... 중략 ...
if self.gate is not None:
if self._fse_fuse_gate:
self._maybe_fuse_gate_weights()
# 플랫폼별 최적화된 GEMM 디스패처 사용
router_logits = dispatch_unquantized_gemm()(
self, hidden_states, self._combined_gate_weight, None
)
else:
router_logits, _ = self.gate(hidden_states)
변경 사항은 매우 명확합니다. F.linear 대신 dispatch_unquantized_gemm()을 호출합니다. 이 함수는 vLLM의 추상화 계층으로, 현재 실행 중인 하드웨어 플랫폼에 따라 최적의 GEMM 구현체를 반환합니다.
- ROCm 플랫폼:
rocm_unquantized_gemm으로 연결되어,wvSplitK와 같은 전용 Skinny-GEMM 커널을 사용하거나 AMD의 최적화 라이브러리인 AITER의 튜닝된 커널을 찾습니다. - 기타 플랫폼 (CUDA 등): 기본적으로
F.linear를 반환하므로 기존 동작에 영향을 주지 않는 No-op(동작 변화 없음) 변경이 됩니다.
왜 이게 좋은 최적화인가?
1. 극적인 성능 회복 (약 2배 속도 향상)
PR 설명에 포함된 벤치마크 결과에 따르면, 100번의 디코딩 스텝당 소요 시간은 다음과 같습니다.
- 기존 (Fused, Before PR): 110.6 ms (Tensile
Cijk_*커널 사용) - 개선 (Fused, After PR): 55.6 ms (Skinny-GEMM 커널 사용)
단순히 호출하는 함수를 바꿨을 뿐인데, 특정 GEMM 연산에서 약 2배의 성능 향상을 이끌어냈습니다. 이는 전체 서빙 처리량(Throughput)에서도 유의미한 차이를 만듭니다. 특히 동시 접속자 수가 적은(Concurrency=4) 환경에서 출력 처리량이 +3.5% 향상되는 결과를 보여주었습니다.
2. AITER 튜닝 경로 확보
더 중요한 점은 이 변경이 미래의 최적화를 위한 관문이라는 것입니다. 기존의 F.linear는 PyTorch 내부의 rocBLAS로 직접 연결되어 vLLM이 제어하는 AITER 튜닝 시스템을 완전히 우회했습니다.
이번 수정을 통해 GEMM 연산이 rocm_unquantized_gemm_impl -> aiter.tuned_gemm.mm 경로를 타게 되었습니다. 현재는 513(512+1)이라는 특이한 가중치 폭에 대한 튜닝 값이 AITER에 없을 수 있지만, 향후 AITER에서 이 형상(Shape)에 대한 최적화 커널을 제공하기만 하면 vLLM의 코드를 수정하지 않고도 즉시 성능 이득을 볼 수 있게 된 것입니다.
3. 플랫폼 추상화의 모범 사례
리뷰어인 amd-mghanimi는 이 PR을 보고 "양자화되지 않은 모든 F.linear 호출을 이 래퍼(wrapper)로 교체해야 하는가?"라는 유의미한 질문을 던졌습니다. 이는 vLLM과 같은 멀티 플랫폼 지원 프로젝트에서 매우 중요한 포인트입니다. 하드웨어 제조사마다 강점을 가진 커널이 다르기 때문에, 상위 레이어에서는 dispatch_unquantized_gemm과 같은 추상화된 API를 사용하고 하위 레이어에서 플랫폼별 최적화를 수행하는 구조가 유지보수와 성능 측면에서 모두 유리합니다.
마치며: 데이터 기반의 최적화
이번 PR은 단순히 코드를 깔끔하게 만드는 것을 넘어, 프로파일링 도구(PyTorch Profiler)를 통해 어떤 커널이 병목인지 정확히 짚어내고(Tensile vs wvSplitK), 이를 해결하기 위한 올바른 아키텍처적 선택을 보여준 훌륭한 사례입니다.
엔지니어로서 배울 수 있는 교훈은 명확합니다. "표준 API(F.linear)가 항상 모든 하드웨어에서 최적은 아니다"라는 점과, "플랫폼 디스패처를 통한 추상화가 성능 최적화의 유연성을 제공한다"는 것입니다. ROCm 환경에서 vLLM을 운영하거나 MoE 모델을 최적화하려는 분들에게 이 분석이 도움이 되기를 바랍니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.nn.functional.linear.html
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/utils.py
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM ROCM 최적화: GLM-4 MoE를 위한 Fused Shared Expert(FSE) 도입
- [vllm] vLLM ROCm 환경에서 FlyDSL을 활용한 MXFP8 MoE 성능 최적화
- [vllm] vLLM ROCm 환경에서 Shared-Expert Fusion을 통한 MoE 추론 성능 최적화
- [sglang] AMD에서 MoE Gate router gemm을 tgemm.mm으로 교체
- [vllm] [vLLM 성능 최적화] Nemotron-3 디코딩 속도를 13% 향상시킨 Latent-MoE All-Reduce 최적화 분석
PR Analysis 의 다른글
- 이전글 [sglang] HiCache 최적화: TMA를 활용한 Host-Device KV 캐시 전송 성능 2배 향상
- 현재글 : [vllm] [ROCm 성능 최적화] vLLM의 Fused Shared-Expert Gate GEMM 경로 개선 분석
- 다음글 [flashinfer] FlashInfer의 Context-Parallel Decode 최적화: Fused A2A + LSE Reduce
댓글