[vllm] [ROCm] DeepSeek V4 성능 극대화: FP8 WO_A 출력 프로젝션 최적화 분석
PR 링크: vllm-project/vllm#54894 상태: Merged | 변경: +217 / -11
들어가며
DeepSeek V4와 같은 최신 대규모 언어 모델(LLM)은 MLA(Multi-head Latent Attention) 구조를 채택하여 메모리 효율성을 극대화합니다. 하지만 이 과정에서 발생하는 wo_a (Output Projection Part A) 연산은 여전히 상당한 계산 비용을 차지합니다. 특히 ROCm 환경에서 기존의 BF16 기반 grouped einsum 방식은 하드웨어의 잠재력을 완전히 끌어내지 못하는 병목 지점이었습니다.
이번 PR은 AMD의 차세대 가속기(gfx950/MI355X)에서 DeepSeek V4의 wo_a 경로를 FP8로 전환하는 최적화를 담고 있습니다. 핵심은 AITER 라이브러리를 활용하여 Inverse RoPE와 Quantization을 하나로 묶고(Fusion), 체크포인트에 저장된 네이티브 FP8 가중치를 그대로 사용하여 데이터 전송량과 연산 속도를 동시에 개선한 것입니다.
코드 분석: 핵심 변경 사항
1. OCP MX E8M0 스케일 정규화 (_wo_a_block_scale_to_e8m0)
FP8 연산을 위해서는 가중치와 활성화 함수의 스케일을 정확히 관리해야 합니다. 이 PR에서는 OCP(Open Compute Project) MX 사양의 E8M0 포맷을 사용합니다. 이는 가수(Mantissa) 없이 지수(Exponent)만으로 구성된 8비트 스케일 포맷입니다.
# vllm/models/deepseek_v4/amd/rocm.py 에 추가된 유틸리티
def _wo_a_block_scale_to_e8m0(scale: torch.Tensor) -> torch.Tensor | None:
# ... (생략)
# E8M0은 bias 127을 가진 unsigned exponent 포맷입니다.
if scale.dtype == torch.float8_e8m0fnu:
return scale.view(torch.uint8).contiguous()
# FP32 입력을 무손실로 E8M0으로 변환할 수 있는지 확인합니다.
exponent = torch.round(torch.log2(scale_f32))
if not torch.equal(torch.exp2(exponent), scale_f32):
return None
encoded = exponent.to(torch.int32) + 127
return encoded.to(torch.uint8).contiguous()
이 함수는 체크포인트의 스케일 데이터를 하드웨어가 이해할 수 있는 uint8 형태의 E8M0 바이트로 정규화합니다. 특히 단순 반올림이 아닌 무손실 변환(Lossless conversion)이 가능한 경우에만 FP8 경로를 활성화하도록 설계되어 정밀도 저하를 방지합니다.
2. FP8 실행 경로 준비 (_prepare_fp8_wo_a)
모델 초기화 단계에서 현재 하드웨어와 라이브러리 환경이 FP8 최적화를 지원하는지 검증합니다.
# vllm/models/deepseek_v4/amd/rocm.py
def _prepare_fp8_wo_a(self) -> None:
try:
from aiter.ops.batched_gemm_op_a8w8 import batched_gemm_a8w8_mxscale
from aiter.ops.inverse_rope_group_quant import inverse_rope_group_quant
except ImportError:
logger.warning_once("AITER >= 0.1.20 필요; BF16으로 폴백합니다.")
return
# 가중치와 스케일의 레이아웃(Group-128)이 최적화 커널과 호환되는지 확인
if out_per_group % 128 != 0 or in_features % 128 != 0:
return
# FP8 가중치와 E8M0 스케일을 뷰(View) 형태로 캐싱
self._wo_a_fp8_weight = weight.view(groups, out_per_group, in_features)
self._wo_a_e8m0_scale = e8m0_scale.view(groups, out_per_group // 128, in_features // 128)
이 섹션은 런타임 오버헤드를 줄이기 위해 필요한 텐서들을 미리 준비하고, 조건이 맞지 않으면 안전하게 기존 BF16 경로로 폴백(Fallback)하도록 보장합니다.
3. 연산 퓨전 및 실행 (_o_proj)
가장 드라마틱한 변화는 실제 연산이 일어나는 _o_proj 함수입니다.
Before (BF16 Fallback):
# 기존 방식: Inverse RoPE 연산 후 별도의 einsum 수행
z = rocm_inv_rope_einsum(
self.rotary_emb, o, positions, self.rope_head_dim,
self.n_local_groups, self.o_lora_rank, self.wo_a,
)
zf = z.flatten(1)
After (FP8 Fast Path):
# 최적화 방식: Inverse RoPE와 Quantization을 퓨전하고 FP8 GEMM 수행
o_fp8, o_scale = inverse_rope_group_quant(
o.view(o.shape[0], self.n_local_heads, self.head_dim),
positions.to(torch.int64),
self._wo_a_cos_cache,
self._wo_a_sin_cache,
num_groups=self.n_local_groups,
quant_group_size=128,
)
zf = batched_gemm_a8w8_mxscale(
o_fp8, self._wo_a_fp8_weight, o_scale, self._wo_a_e8m0_scale, dtype=o.dtype
).flatten(1)
inverse_rope_group_quant 커널은 Inverse RoPE 연산을 수행함과 동시에 결과를 FP8로 양자화하여 메모리 쓰기 횟수를 줄입니다. 이후 batched_gemm_a8w8_mxscale을 통해 가중치와 활성화 함수 모두 FP8인 상태에서 고속 GEMM을 수행합니다.
왜 이게 좋은가?
1. 성능 지표 (Performance)
PR 설명에 포함된 벤치마크 결과는 놀랍습니다 (8 x MI355X 기준):
- TTFT (Time To First Token): -7.35% 감소 (2459.0ms -> 2278.4ms)
- Input Throughput: +7.93% 향상 (40,667 tok/s -> 43,891 tok/s)
- TPOT (Time Per Output Token): 약 2~3% 개선
2. 메모리 대역폭 효율성
기존 BF16 방식은 중간 결과값을 메모리에 썼다가 다시 읽어오는 과정이 필요했지만, FP8 퓨전 커널은 레지스터 수준에서 연산을 처리하고 데이터 크기를 절반으로 줄여 메모리 대역폭 병목을 해소합니다.
3. 정확도 유지 (Accuracy)
리뷰어 Fangzhou-Ai의 요청으로 진행된 GSM8K 테스트 결과, FP8 경로를 사용하더라도 정확도가 BF16 대비 동등하거나 오히려 미세하게 높게 측정되었습니다 (94.7688%). 이는 E8M0 스케일의 정밀한 관리 덕분입니다.
결론 및 교훈
이 PR은 단순한 데이터 타입 변경을 넘어, 하드웨어 특화 라이브러리(AITER)와 모델 아키텍처(MLA)의 깊은 이해가 결합되었을 때 어떤 성능 향상을 가져올 수 있는지 보여주는 좋은 사례입니다.
시니어 엔지니어로서 배울 수 있는 점은 다음과 같습니다:
- 안전한 폴백 설계: 최신 하드웨어 기능을 도입할 때 기존 환경과의 호환성을 위해 BF16 폴백을 유지한 점.
- 무손실 검증: 양자화 도입 시
_wo_a_block_scale_to_e8m0와 같이 수학적으로 엄격한 검증 로직을 포함하여 모델 품질을 보호한 점. - 데이터 기반 의사결정: 리뷰어의 피드백에 따라 대규모 데이터셋(GSM8K 1319문항)으로 정확도를 재검증하여 신뢰성을 확보한 점.
ROCm 생태계에서 DeepSeek V4를 운영하려는 팀에게 이 최적화는 필수적인 업데이트가 될 것입니다.
참고 자료
- https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf
- https://rocm.docs.amd.com/projects/HIP/en/latest/reference/low_fp_types.html
- https://pytorch.org/docs/stable/generated/torch.Tensor.view.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM, ROCm 환경에서 AITER MoE 연산 성능 최적화를 위한 환경 변수 노출
- [vllm] vLLM ROCm 환경에서 AITER를 활용한 Multi-Head Convolutions(MHC) 성능 최적화 및 안정성 개선
- [vllm] vLLM, ROCm 환경에서 FP8 GEMM 최적화로 성능 4-9% 향상
- [vllm] vLLM, DeepSeek V4 모델 성능 최적화: AITER MXFP4 BF16 백엔드 개선
- [vllm] [vLLM 분석] DeepSeek V4의 Sparse FP8 Compressor 커널 최적화: CuteDSL을 통한 성능 극대화
PR Analysis 의 다른글
- 이전글 [sglang] SGLang에서 SM120 GPU를 위한 SubBlock Sage FP8 어텐션 최적화
- 현재글 : [vllm] [ROCm] DeepSeek V4 성능 극대화: FP8 WO_A 출력 프로젝션 최적화 분석
- 다음글 [ultralytics] Ultralytics YOLOv10 TensorRT 엔진 성능 최적화: FP16 및 INT8 속도 향상 비결
댓글