[vllm] vLLM, ROCm 환경에서 FP8 GEMM 최적화로 성능 4-9% 향상
PR 링크: vllm-project/vllm#51692 상태: Merged | 변경: +159 / -5
들어가며
최근 vLLM 프로젝트에서는 대규모 언어 모델(LLM)의 추론 성능을 극대화하기 위한 지속적인 노력을 기울이고 있습니다. 특히 AMD의 ROCm 환경에서 NVIDIA GPU 환경 못지않은 성능을 제공하는 것이 중요한 과제 중 하나입니다. 이번 PR([#PR_NUMBER])은 ROCm 환경에서 FP8(8비트 부동소수점) 정밀도를 활용하는 GEMM(General Matrix Multiply) 연산의 성능을 개선하는 데 초점을 맞추고 있습니다. 구체적으로, 기존의 'ck-tile' 커널 대신 'bpreshuffled' 블록 스케일 FP8 GEMM 커널을 도입하여 특정 연산에서 상당한 성능 향상을 달성했습니다.
이 PR은 특히 DSv3 모델을 1k/1k 시퀀스 길이로 사용할 때, TP8+DPA(Tensor Parallelism 8-way + Data Parallelism) 구성에서 4-8%, TP8+EP(Tensor Parallelism 8-way + Ensemble Parallelism) 구성에서 0-4%의 QPS(Queries Per Second) 향상을 목표로 합니다. 본 글에서는 이 PR이 어떤 기술적 변화를 가져왔고, 왜 이러한 변화가 성능 향상으로 이어졌는지, 그리고 이 최적화가 가지는 일반적인 교훈은 무엇인지 코드 diff와 함께 자세히 살펴보겠습니다.
코드 분석
이번 PR의 핵심 변경사항은 ROCm 환경에서 FP8 GEMM 연산을 위한 새로운 커널인 AiterPreshuffledFp8BlockScaledMMKernel을 도입하고, 이를 지원하기 위한 유틸리티 함수 및 커널 등록 로직을 수정하는 것입니다.
1. vllm/_aiter_ops.py 변경사항
_load_gemm_tuned_configs 함수가 제네릭하게 개선되었습니다. 기존에는 q_dtype_w만을 필터링 조건으로 사용했지만, 이제는 filters 인자를 통해 다양한 컬럼과 값을 기준으로 CSV 파일에서 튜닝된 설정을 로드할 수 있게 되었습니다. 또한, 반환되는 키 컬럼(key_cols)도 지정할 수 있게 되어 유연성이 크게 향상되었습니다.
Before:
@@ -175,18 +175,26 @@ def is_aiter_found_and_supported_on_rdna4() -> bool:
@functools.cache
def _load_gemm_tuned_configs(
- q_dtype_w: torch.dtype, csv_path: str
-) -> set[tuple[int, int, int]]:
+ csv_path: str,
+ filters: tuple[tuple[str, object], ...],
+ key_cols: tuple[str, ...] = ("N", "K", "M"),
+) -> set[tuple[int, ...]]:
try:
df = pd.read_csv(csv_path).drop_duplicates()
- df = df[df["q_dtype_w"] == str(q_dtype_w)]
- return set(zip(df["N"].astype(int), df["K"].astype(int), df["M"].astype(int)))
+ for col, val in filters:
+ if col not in df.columns:
+ continue
+ if isinstance(val, int):
+ df = df[df[col].astype(int) == val]
+ else:
+ df = df[df[col].astype(str) == str(val)]
+ return set(zip(*(df[c].astype(int) for c in key_cols)))
except Exception:
return set()
def _check_kernel_tuned(N: int, K: int, q_dtype_w: torch.dtype, csv_path: str) -> bool:
- configs = _load_gemm_tuned_configs(q_dtype_w, csv_path)
+ configs = _load_gemm_tuned_configs(csv_path, (("q_dtype_w", q_dtype_w),))
l_m = (
[1, 2, 4]
+ list(range(8, 513, 8))
After:
@@ -3200,6 +3208,20 @@ def is_per_token_w8a8_gemm_tuned(N: int, K: int, q_dtype_w: torch.dtype) -> bool
csv_path = aiter_gemm_a8w8_ops.AITER_CONFIGS.AITER_CONFIG_GEMM_A8W8_FILE
return _check_kernel_tuned(N, K, q_dtype_w, csv_path)
+ @staticmethod
+ def is_blockscale_bpreshuffle_tuned(n: int, k: int) -> bool:
+ """Whether (N, K) has a tuned aiter blockscale bpreshuffle config."""
+ if not current_platform.is_rocm():
+ return False
+ import aiter.ops.gemm_op_a8w8 as aiter_gemm_a8w8_ops
+
+ csv_path = aiter_gemm_a8w8_ops.AITER_CONFIGS.AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_BPRESHUFFLE_FILE
+ gfx = aiter_gemm_a8w8_ops.get_gfx()
+ cu_num = aiter_gemm_a8w8_ops.get_cu_num()
+ return (n, k) in _load_gemm_tuned_configs(
+ csv_path, (("gfx", gfx), ("cu_num", cu_num)), key_cols=("N", "K")
+ )
+
@staticmethod
def shuffle_weight(
tensor: torch.Tensor, layout: tuple[int, int] = (16, 16)
리뷰어 Rohan138의 제안에 따라 _load_gemm_tuned_configs 함수가 제네릭하게 변경되었고, is_blockscale_bpreshuffle_tuned 함수가 새로 추가되어 특정 조건(gfx, cu_num)에 맞는 튜닝된 설정을 확인하게 되었습니다. 이는 코드 재사용성을 높이고, 향후 AITER 라이브러리 자체적으로 튜닝된 설정을 관리하게 될 때 통합을 용이하게 합니다.
2. vllm/model_executor/kernels/linear/__init__.py 변경사항
새로운 커널 AiterPreshuffledFp8BlockScaledMMKernel이 _get_linear_backend 및 _resolve_backend_kernels 함수에서 등록되고, register_linear_kernel 함수를 통해 사용 가능하게 됩니다. 이는 vLLM이 FP8 연산 시 이 새로운 커널을 선택할 수 있도록 하는 기반을 마련합니다.
Before: (일부 발췌)
@@ -316,6 +317,7 @@ def _get_linear_backend() -> str:
"aiter": {
AiterInt8ScaledMMLinearKernel,
AiterFp8BlockScaledMMKernel,
+ AiterPreshuffledFp8BlockScaledMMKernel,
AiterPerTokenFp8ScaledMMLinearKernel,
AiterMxfp4LinearKernel,
@@ -456,6 +458,7 @@ def _resolve_backend_kernels(
BlockWiseTorchFP8ScaledMMLinearKernel,
],
PlatformEnum.ROCM: [
+ AiterPreshuffledFp8BlockScaledMMKernel,
AiterFp8BlockScaledMMKernel,
TritonFp8BlockScaledMMKernel,
],
After: (일부 발췌)
@@ -1210,6 +1213,7 @@ def register_linear_kernel(
"ScaledMMLinearLayerConfig",
"AiterHipbMMPerTokenFp8ScaledMMLinearKernel",
"AiterPreshuffledPerTokenFp8ScaledMMLinearKernel",
+ "AiterPreshuffledFp8BlockScaledMMKernel",
"AiterPerTokenFp8ScaledMMLinearKernel",
"NvFp4LinearKernel",
"NvFp4LinearLayerConfig",
3. vllm/model_executor/kernels/linear/scaled_mm/aiter.py 변경사항
이 파일에서 핵심적인 AiterPreshuffledFp8BlockScaledMMKernel 클래스가 정의됩니다. 이 클래스는 기존 Fp8BlockScaledMMLinearKernel을 상속받으며, 'bpreshuffled' 가중치 레이아웃을 사용합니다.
주요 구현 내용:
is_supported: ROCm 환경 및 AITER FP8 선형 커널 지원 여부를 확인합니다.can_implement: FP8 GEMM 구현 가능 여부를 상세하게 검증합니다. 특히, 활성화 양자화 그룹 크기, N/K 차원의 128 배수 여부, 그리고is_blockscale_bpreshuffle_tuned함수를 통한 튜닝된 설정 존재 여부를 확인합니다. 이는 잘못된 설정으로 인한 잠재적 오류나 성능 저하를 방지하기 위함입니다. 리뷰어 BadrBasowid와 tjtanaa의 피드백을 반영하여,apply_weights내의 조기 반환 대신can_implement에서 이러한 제약 조건을 명확히 하는 것이 더 적절하다는 의견이 반영되었습니다. 또한, 별도의 클래스로 분리함으로써 어떤 커널이 선택되었는지 로깅 및 추적이 용이해졌습니다.process_weights_after_loading: 로드된 가중치를 'bpreshuffle' 레이아웃으로 변환합니다.shuffle_weight함수를 사용하여 가중치를 (16, 16) 레이아웃으로 섞습니다.apply_weights: 실제 연산을 수행합니다. 입력 텐서x를 FP8으로 양자화하고, 'bpreshuffled' 가중치와 함께gemm_a8w8_blockscale_bpreshuffle함수를 호출하여 결과를 계산합니다.
코드 예시 (AiterPreshuffledFp8BlockScaledMMKernel 클래스):
class AiterPreshuffledFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel):
"""Aiter FP8 block-scaled GEMM using a pre-shuffled (bpreshuffle) weight."""
@classmethod
def is_supported(
cls, compute_capability: int | None = None
) -> tuple[bool, str | None]:
return AiterPreshuffledPerTokenFp8ScaledMMLinearKernel.is_supported(
compute_capability
)
@classmethod
def can_implement(cls, config: FP8ScaledMMLinearLayerConfig):
can_implement_base, reason = super().can_implement(config)
if not can_implement_base:
return can_implement_base, reason
act_quant_desc = config.activation_quant_key.scale
if act_quant_desc.group_shape != GroupShape(1, 128):
return (
False,
(
"Supports only dynamic per token group activation "
"quantization with group_shape=(1,128)."
),
)
# bpreshuffle GEMM requires aiter fp8 linear to be enabled and an fp8
# 2D weight with N and K divisible by 128. Require a tuned config as the
# fallback can trigger faults or numerical errors.
if not rocm_aiter_ops.is_linear_fp8_enabled():
return (
False,
(
"requires setting `VLLM_ROCM_USE_AITER=1` "
"and `VLLM_ROCM_USE_AITER_LINEAR=1`."
),
)
n, k = config.weight_shape
if not (n % 128 == 0 and k % 128 == 0):
return (
False,
(
f"requires N and K dimensions divisible by 128, received "
f"N={n} and K={k}."
),
)
if not rocm_aiter_ops.is_blockscale_bpreshuffle_tuned(n, k):
return (
False,
(
f"requires a tuned aiter blockscale bpreshuffle config for "
f"N={n} and K={k}."
),
)
return True, None
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
super().process_weights_after_loading(layer)
params = FP8BlockParams.from_layer(layer)
if params.weight_scale_inv is not None:
ws, attr = params.weight_scale_inv, params.WEIGHT_SCALE_INV
else:
ws, attr = params.weight_scale, params.WEIGHT_SCALE
if ws is not None and ws.dtype == torch.float8_e8m0fnu:
replace_parameter(layer, attr, _upcast_e8m0_to_fp32(ws).contiguous())
weight = params.weight
# runtime safety net
assert (
weight.dim() == 2
and weight.dtype == current_platform.fp8_dtype()
and weight.shape[0] % 128 == 0
and weight.shape[1] % 128 == 0
), (
"AiterPreshuffledFp8BlockScaledMMKernel requires a 2D fp8 weight "
"with N and K divisible by 128."
)
shuffled_weight = rocm_aiter_ops.shuffle_weight(
weight.contiguous(), layout=(16, 16)
)
replace_parameter(
layer,
params.WEIGHT,
torch.nn.Parameter(shuffled_weight.data, requires_grad=False),
)
def apply_weights(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
params = self._get_layer_params(layer)
Bs = (
params.weight_scale
if params.weight_scale_inv is None
else params.weight_scale_inv
)
x_2d = x.view(-1, x.shape[-1])
A, As = rocm_aiter_ops.group_fp8_quant(x_2d, transpose_scale=True)
output = rocm_aiter_ops.gemm_a8w8_blockscale_bpreshuffle(
A, params.weight, As, Bs, output_dtype=self.config.out_dtype
)
if bias is not None:
output = output + bias
return output.view(*x.shape[:-1], params.weight.shape[0])
def apply_block_scaled_mm(
self,
A: torch.Tensor,
B: torch.Tensor,
As: torch.Tensor,
Bs: torch.Tensor,
) -> torch.Tensor:
raise NotImplementedError(
"AiterPreshuffledFp8BlockScaledMMKernel overrides apply_weights and "
"does not use apply_block_scaled_mm."
)
왜 이게 좋은가?
1. 성능 향상
이 PR의 가장 큰 장점은 실제 성능 향상입니다. 제공된 테스트 결과에 따르면, 특히 TP8+DPA 구성에서 QPS가 최대 8.14%까지 향상되었습니다. TP8+EP 구성에서도 최대 9.01%의 QPS 향상이 관찰되었습니다. 이는 기존의 'ck-tile' 커널 대비 'bpreshuffled' GEMM 커널이 특정 연산(특히 o-proj 레이어)에서 훨씬 효율적임을 보여줍니다. 프로파일링 결과에서도 'bpreshuffled' 버전이 21us로 'ck-tile' 버전(52us)보다 2배 이상 빠름을 확인할 수 있습니다.
프로파일링 결과 비교:
- This branch (bpreshuffled): 21us
- Main branch (ck-tile): 52us
TP8+DPA QPS 변화:
| Concurrency | Branch | QPS | QPS change (%) |
|---|---|---|---|
| 2 | This branch | 0.1049 | +8.14 |
TP8+EP QPS 변화:
| Concurrency | Branch | QPS | QPS change (%) |
|---|---|---|---|
| 64 | This branch | 2.7676 | +9.01 |
2. FP8 활용 극대화
FP8은 FP16이나 BF16에 비해 메모리 대역폭 요구량을 줄이고 연산 속도를 높일 수 있는 잠재력을 가지고 있습니다. 하지만 FP8을 효과적으로 사용하기 위해서는 하드웨어 아키텍처에 최적화된 커널 구현이 필수적입니다. 'bpreshuffled' 기법은 가중치 데이터를 미리 특정 레이아웃으로 섞어둠으로써, GEMM 연산 시 데이터 접근 패턴을 개선하고 하드웨어의 병렬 처리 능력을 최대한 활용할 수 있게 합니다. 이는 특히 AMD ROCm 환경에서 FP8 연산 성능을 끌어올리는 데 중요한 역할을 합니다.
3. 코드 품질 및 유지보수성 향상
- 제네릭 유틸리티 함수:
_load_gemm_tuned_configs함수의 일반화는 코드 중복을 줄이고 유지보수성을 높입니다. 이는 향후 유사한 튜닝 설정 로딩 로직이 필요할 때 쉽게 확장될 수 있음을 의미합니다. - 명확한 커널 분리:
AiterPreshuffledFp8BlockScaledMMKernel클래스를 별도로 정의함으로써, 어떤 커널이 사용되는지 명확하게 알 수 있습니다. 이는 디버깅 및 로깅에 큰 도움이 되며, 코드의 가독성을 향상시킵니다. 리뷰어들의 피드백이 이러한 코드 품질 개선에 기여했습니다. - 조건부 최적화:
can_implement메서드에서 다양한 조건을 검사하여 해당 커널이 실제로 사용 가능한 경우에만 선택되도록 함으로써, 런타임 오류를 방지하고 안정성을 높입니다. 이는 특히 다양한 하드웨어 및 설정 환경에서 vLLM을 사용하는 경우에 중요합니다.
4. 일반적인 교훈
- 하드웨어 특화 최적화의 중요성: LLM 추론 성능은 하드웨어 아키텍처에 대한 깊은 이해와 이를 반영한 커널 최적화에 크게 좌우됩니다. 특히 FP8과 같은 저정밀도 연산은 하드웨어 특성에 민감하므로, 각 플랫폼(NVIDIA CUDA, AMD ROCm 등)에 맞는 최적화가 필요합니다.
- 데이터 레이아웃과 연산 효율: 데이터의 메모리 레이아웃은 연산 성능에 지대한 영향을 미칩니다. 'bpreshuffled'와 같이 데이터를 미리 재배열하는 기법은 메모리 접근 패턴을 개선하여 캐시 효율성을 높이고 병렬 처리 성능을 극대화할 수 있습니다.
- 점진적 개선과 커뮤니티 피드백: 이 PR은 기존 커널의 성능 병목을 식별하고 새로운 커널을 도입하는 점진적인 개선 과정을 보여줍니다. 또한, 리뷰 과정에서 제네릭 유틸리티 함수 사용, 커널 분리 등 코드 품질 향상에 대한 건설적인 피드백이 반영된 점은 오픈소스 프로젝트의 협업 모델이 어떻게 더 나은 결과물을 만들어내는지 잘 보여줍니다.
결론
이번 vLLM PR은 ROCm 환경에서 FP8 GEMM 연산의 성능을 크게 향상시키는 중요한 발걸음입니다. 'bpreshuffled' 블록 스케일 FP8 GEMM 커널의 도입은 특정 연산에서 2배 이상의 속도 향상을 가져왔으며, 이는 전체적인 QPS를 4-9%까지 끌어올리는 결과를 낳았습니다. 또한, 코드의 제네릭화와 명확한 커널 분리를 통해 유지보수성과 안정성 또한 개선되었습니다. 이러한 최적화는 vLLM이 다양한 하드웨어 환경에서 최고 수준의 추론 성능을 제공하려는 노력의 일환이며, 앞으로도 저정밀도 연산과 하드웨어 특화 최적화가 LLM 성능 향상의 핵심 동력이 될 것임을 시사합니다.
참고 자료
- https://github.com/vllm-project/vllm/issues/51957
- https://github.com/ROCm-Developer-Tools/HIP/blob/master/rocblas/samples/gemm_fp8.cpp
- https://github.com/ROCm-Developer-Tools/HIP/blob/master/rocblas/samples/gemm_strided_batched_fp8.cpp
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer의 GEMM 성능 혁신: cuTile 백엔드 도입과 최적화 여정
- [vllm] vLLM의 ROCm 환경에서 듀얼 스트림 디코드를 통한 성능 최적화
- [vllm] vLLM, DeepSeek-V4 사전 생성 처리량 향상을 위한 Sparse Top-K 메타데이터 커널 최적화
- [vllm] vLLM, DeepSeek-V3.2/GLM-5.2 MTP 경로 최적화: All-Reduce 융합 및 로컬 Argmax 도입
- [flashinfer] FlashInfer SM120 MoE GEMM 최적화: 웨이브+잔여물 비용 모델 도입
PR Analysis 의 다른글
- 이전글 [sglang] Ascend NPU 환경에서 HiCache L2 I/O 성능 최적화: Memfabric과 AscendC 활용
- 현재글 : [vllm] vLLM, ROCm 환경에서 FP8 GEMM 최적화로 성능 4-9% 향상
- 다음글 [sglang] SGLang aiter 백엔드의 Sliding Window Attention(SWA) 최적화 및 안정성 개선
댓글