[vllm] vLLM, ROCm 환경에서 FP8을 활용한 DeepSeek-V4.1 모델 성능 최적화
PR 링크: vllm-project/vllm#58456 상태: Merged | 변경: +664 / -29
들어가며
최근 AI 모델의 발전 속도는 눈부시지만, 이러한 거대 모델을 효율적으로 서빙하는 것은 여전히 큰 도전 과제입니다. 특히 LLM(Large Language Model)은 방대한 연산량과 메모리 요구량으로 인해 고성능 하드웨어에서도 최적화 없이는 만족스러운 성능을 내기 어렵습니다. vLLM은 LLM 추론을 위한 고성능 서빙 엔진으로, 지속적인 최적화를 통해 다양한 하드웨어 환경에서 최고의 성능을 추구하고 있습니다. 이번 글에서는 vLLM이 AMD의 ROCm 환경에서 DeepSeek-V4.1 모델의 추론 성능을 극대화하기 위해 수행한 코드 변경 사항을 심층적으로 분석합니다. 핵심은 기존의 BF16(bfloat16) 데이터 타입 대신 FP8(8-bit Floating Point) 데이터 타입을 적극적으로 활용하여 메모리 대역폭 병목 현상을 완화하고 연산 속도를 높이는 것입니다.
이 PR은 특히 wo_a (weight-only attention output projection) 레이어의 연산 방식을 개선하는 데 초점을 맞춥니다. 기존에는 어텐션 계산 후 중간 결과를 BF16으로 저장하고, 이를 다시 읽어와서 wo_a 연산을 수행했습니다. 이 과정에서 불필요한 데이터 변환과 메모리 접근이 발생하여 성능 저하의 원인이 되었습니다. 본 PR은 이 중간 저장 및 로딩 과정을 제거하고, FP8 데이터 타입을 사용하여 연산의 효율성을 높이는 방안을 제시합니다.
코드 분석
이번 변경 사항은 주로 ROCm 환경에서의 어텐션 연산과 wo_a 프로젝션 부분을 최적화하는 데 집중되어 있습니다. 주요 변경 사항을 파일별로 살펴보겠습니다.
1. tests/kernels/attention/test_rocm_triton_attn_dsv4.py
이 파일은 ROCm 환경에서의 어텐션 커널 동작을 검증하는 테스트 코드입니다. 이번 PR에서는 FP8 데이터 타입을 활용하는 새로운 기능들을 테스트하기 위한 여러 함수가 추가되었습니다.
-
_launch_sparse_decode_reduce함수 수정: 이 함수는 sparse decode의 reduce 커널을 실행하는 역할을 합니다. 기존에는out_dtype이torch.bfloat16으로 고정되어 있었으나, 이제는out_dtype과out_scale인자를 받아 MXFP8 (e4m3) 형식으로 출력을 저장할 수 있도록 변경되었습니다. 이는QUANT_OUT옵션을 통해 제어되며, FP8 출력을 직접 생성하여 BF16 변환 및 저장 과정을 생략합니다.# Before out = torch.empty( (num_queries, num_heads, HEAD_DIM), dtype=torch.bfloat16, device=part_m.device, ) # ... # After out = torch.empty( (num_queries, num_heads, HEAD_DIM), dtype=out_dtype, # Now accepts out_dtype device=part_m.device, ) # ... # QUANT_OUT=out_scale is not None, -
test_sparse_attn_decode_reduce_mxfp8_epilogue함수 추가: 이 테스트는_sparse_attn_decode_reduce커널이 MXFP8 출력을 올바르게 생성하는지 검증합니다. FP32 출력에 대해 MXFP8 양자화를 수행한 결과와 커널의 FP8 출력 결과를 비교하여 비트 단위의 정확성을 확인합니다.# Expected calculation using torch function expected_data, expected_scale = _mxfp8_e4m3_quantize_torch( rotated.view(num_queries, -1) ) # Kernel output data = _launch_sparse_decode_reduce( *args, out_dtype=torch.float8_e4m3fn, out_scale=scale ) assert torch.equal(scale, expected_scale) assert torch.equal( data.view(num_queries, -1).view(torch.uint8), expected_data.view(torch.uint8) ) -
test_inverse_rope_mxfp8_rows함수 추가: 이 테스트는 prefill 단계에서 inverse RoPE와 MXFP8 양자화가 올바르게 수행되는지 검증합니다. 디코드 에필로그와 동일한 로직을 사용하여 FP8 출력을 생성합니다. -
test_sparse_attn_decode_mxfp8_output함수 추가: 이 테스트는out_mxfp8옵션을 사용하여 실제 디코드 과정에서 MXFP8 출력이 사용되는지 확인합니다. BF16 출력과 비교하여 수치적 정확성을 검증합니다. -
test_rocm_mxfp8_wo_a_bmm함수 추가: 이 테스트는 MXFP8 가중치와 활성화 값을 사용하여 그룹화된 FP8 GEMM 연산을 수행하는rocm_mxfp8_wo_a_bmm커널의 정확성을 검증합니다. 다양한 토큰 길이와 그룹 수에 대해 FP32 참조 결과와 비교합니다.
2. vllm/models/deepseek_v41/amd/rocm.py
이 파일은 DeepSeek-V4.1 모델의 ROCm 특화 구현을 담당합니다. 이번 PR에서는 FP8 연산을 지원하기 위해 어텐션 레이어의 동작 방식을 수정했습니다.
-
_alloc_attn_out함수 수정: 이 함수는 어텐션 연산의 출력을 저장할 버퍼를 할당합니다. 기존에는 BF16 타입의 텐서를 반환했지만, 이제_ON_GFX950(AMD의 최신 GPU 아키텍처인 GFX950 지원 여부) 조건 하에서QuantizedActivation객체를 반환하도록 변경되었습니다. 이 객체는 MXFP8 데이터(data)와 스케일(scale)을 포함하며, 이는wo_a연산에서 직접 사용됩니다.# Before (simplified) # return torch.empty(...) # After if not _ON_GFX950: return super()._alloc_attn_out(num_tokens, hidden_states) # ... allocate data and scale for MXFP8 return QuantizedActivation( data=data, scale=scale, # ... quant_key=kMxfp8Dynamic, ) -
_o_proj함수 수정: 이 함수는 어텐션 출력(attn_out)을 받아wo_a및wo_b레이어를 통과시키는 역할을 합니다. 이제attn_out이QuantizedActivation객체일 경우,rocm_mxfp8_wo_a_bmm커널을 직접 호출하여 FP8 GEMM 연산을 수행합니다. 이는 기존의 BF16 중간 저장 및 로딩 과정을 완전히 제거합니다.# Before # ... bf16 path with rocm_inv_rope_einsum ... # After if isinstance(attn_out, QuantizedActivation): z = rocm_mxfp8_wo_a_bmm( attn_out.data, # Use FP8 data attn_out.scale, # Use FP8 scale self.wo_a, # Use FP8 weight self.n_local_groups, self.o_lora_rank, ) return self._wo_b_after_wo_a(z) # ... existing bf16 path ... -
_wo_b_after_wo_a함수 추가:_o_proj함수 내에서wo_a연산 후wo_b연산을 수행하는 로직을 별도의 함수로 분리하여 코드의 가독성을 높였습니다. 이는 FP8 경로와 BF16 경로 모두에서 재사용됩니다.
왜 이게 좋은가?
이번 PR의 핵심 목표는 ROCm 환경에서 DeepSeek-V4.1 모델의 추론 성능을 향상시키는 것입니다. FP8 데이터 타입의 도입은 다음과 같은 이점을 제공합니다:
-
메모리 대역폭 감소: FP8은 BF16(16비트)에 비해 절반의 저장 공간을 차지합니다. 이는 메모리에서 가중치와 활성화 데이터를 읽고 쓰는 데 필요한 대역폭을 크게 줄여, 특히 메모리 대역폭이 병목 현상을 일으키는 작업에서 성능 향상으로 이어집니다. PR 설명에 따르면,
wo_a레이어에서 BF16 복사본을 제거함으로써 약 33.5MB/layer의 메모리를 절약할 수 있습니다. -
연산 속도 향상: 최신 AMD GPU 아키텍처(GFX950)는 FP8 연산을 위한 MX Matrix 코어를 갖추고 있습니다. FP8 데이터를 직접 활용하는 커널은 이러한 하드웨어 가속 기능을 활용하여 연산 속도를 높일 수 있습니다. PR의 성능 측정 결과에 따르면, TP2 설정에서 최대 1.59배의 속도 향상을 달성했습니다. TP4 설정에서도 최대 1.70배의 향상을 보였습니다.
-
연산 과정 간소화 및 융합: 기존에는 어텐션 출력 후 BF16으로 저장하고, 이를 다시 읽어와서 양자화 및 GEMM 연산을 수행했습니다. 이 PR은 이러한 중간 저장 및 로딩 단계를 제거하고, reduce 에필로그 내에서 직접 MXFP8 양자화를 수행하여 FP8 GEMM으로 바로 연결합니다. 이러한 연산 융합(fusion)은 커널 실행 횟수를 줄이고 데이터 이동을 최소화하여 전체적인 지연 시간을 단축시킵니다. PR의 성능 측정 결과에서, standalone quant kernel을 사용할 경우 성능 향상이 미미했던 반면, epilogue에 통합했을 때 상당한 성능 향상이 있었음을 보여줍니다.
-
성능 수치: PR 설명에 제시된 성능 측정 결과는 매우 고무적입니다. TP2, TP4 설정에서 다양한 시나리오(conc, spec)에 걸쳐 이전(BF16) 대비 최대 1.7배의 속도 향상을 보여줍니다. 특히, 어텐션 출력부터
wo_a까지의 전체 체인에서 37%의 지연 시간 감소를 기록했습니다. -
일반적 교훈:
- 데이터 타입의 중요성: 모델의 성능을 극대화하기 위해서는 하드웨어 아키텍처가 지원하는 데이터 타입을 적극적으로 활용해야 합니다. FP8은 LLM 추론에서 메모리 대역폭과 연산 속도 모두에서 큰 이점을 제공할 수 있습니다.
- 연산 융합(Fusion)의 힘: 불필요한 메모리 접근과 중간 저장 과정을 제거하고 여러 연산을 하나의 커널로 융합하는 것은 성능 향상의 핵심 전략입니다. 특히 LLM과 같이 연산 집약적인 모델에서는 이러한 융합이 더욱 중요합니다.
- 하드웨어 특화 최적화: 특정 하드웨어 아키텍처(예: ROCm의 GFX950, MX Matrix 코어)의 기능을 최대한 활용하는 최적화는 성능 향상의 지름길입니다.
-
리뷰 댓글 분석
리뷰 댓글들은 주로 코드의 정확성, 기존 코드와의 차별점, 그리고 잠재적인 개선점에 대한 논의를 포함하고 있습니다.
-
_ON_GFX950vson_gfx950():shen-shanshan이 제기한 질문으로, 두 조건이 동일한 값을 참조하는지 확인했습니다.Fangzhou-Ai의 답변에 따르면, 두 조건 모두 동일한 모듈 레벨의_ON_GFX950값을 사용하며, 이는 GCN 아키텍처를 기반으로 한 번 계산됩니다. 현재로서는 CI 통과 및 병합 준비 상태이므로 유지하되, 향후 PR에서 통일성을 고려할 예정이라고 합니다. -
기존 PR과의 차별점:
shen-shanshan은 이 PR이 #54894 PR과 유사한 FP8wo_a최적화를 다루는지 질문했습니다.Fangzhou-Ai는 이 PR이 #54894와 달리 다음과 같은 차이점을 가진다고 명확히 설명했습니다:- 양자화 시점: 이 PR은 reduce 에필로그 내에서 MXFP8 양자화를 수행하여 BF16 쓰기/읽기 및 추가 패스를 제거합니다. 반면 #54894는 별도의
inverse_rope_group_quant패스를 사용했습니다. - 스케일 레이아웃: 이 PR은 체크포인트의 네이티브 32x32 블록 스케일 레이아웃을 그대로 사용합니다. #54894는 128x128로 재배치했습니다.
- 성능: 실제 측정 결과, 이 PR의 방식이 #54894의 방식보다 더 나은 성능을 보였습니다 (예: T=1일 때 7.61us vs 5.82us).
- 양자화 시점: 이 PR은 reduce 에필로그 내에서 MXFP8 양자화를 수행하여 BF16 쓰기/읽기 및 추가 패스를 제거합니다. 반면 #54894는 별도의
-
향후 계획:
Fangzhou-Ai는 네이티브 32x32 스케일 레이아웃과 더 많은 융합을 적용하는 것이 다음 단계가 될 것이라고 언급했습니다. 이는 현재 PR이 FP8 최적화의 중요한 첫걸음을 내딛었음을 시사합니다.
결론
이번 vLLM의 PR은 AMD ROCm 환경에서 DeepSeek-V4.1 모델의 추론 성능을 획기적으로 개선하는 중요한 업데이트입니다. FP8 데이터 타입의 도입, 특히 MXFP8 형식을 활용하여 메모리 대역폭 병목을 완화하고, 연산 융합을 통해 지연 시간을 단축했습니다. 테스트 코드의 추가와 기존 커널의 수정은 이러한 최적화가 정확하고 효율적으로 이루어졌음을 보여줍니다. 리뷰 과정에서의 논의는 기존 작업과의 차별점을 명확히 하고 향후 개선 방향을 제시하며, vLLM이 지속적으로 LLM 서빙 성능을 극한까지 끌어올리고 있음을 증명합니다. 이러한 최적화는 더 빠르고 효율적인 AI 서비스 제공에 크게 기여할 것입니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://github.com/vllm-project/vllm/pull/57435
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/quantization/utils/mxfp8_utils.py
- https://github.com/vllm-project/vllm/blob/main/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/fusion/quant_activation.py
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/quantization/utils/quant_utils.py
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer의 plan() 함수 최적화: Python max()에서 Tensor.max()로의 전환
- 현재글 : [vllm] vLLM, ROCm 환경에서 FP8을 활용한 DeepSeek-V4.1 모델 성능 최적화
- 다음글 [flashinfer] FlashInfer, MiniMax-H3 어텐션 최적화: BF16 및 NVFP4 지원으로 성능 혁신
댓글