执行摘要
- 一句话:修复 XPU FP8 block scale 布局与 MLA 兼容
- 推荐动作:该 PR 值得精读,尤其是对 XPU 或量化相关开发人员。关键设计决策是使用
.t() 视图来同时满足不同消费者对 scale 形状的期望,避免了数据拷贝。建议阅读 process_weights_after_loading 和 apply_block_scaled_mm 的完整实现。
功能与动机
原始代码将 scale 直接存储为连续的 [k_blocks, n_blocks] 布局,但 MLA 的 process_weights_after_loading 调用 scaled_dequantize 时期望 scale 形状与权重布局 [N, K] 匹配,即 [n_blocks, k_blocks],导致形状不匹配和断言失败。该 PR 旨在解决此兼容性问题。
实现拆解
实现拆解
-
调整 scale 存储布局(vllm/model_executor/kernels/linear/scaled_mm/xpu.py):在 process_weights_after_loading 中,将 checkpoint 的 [n_blocks, k_blocks] scale 转置为连续的 [k_blocks, n_blocks] 缓冲区,但通过 replace_parameter 存储其 .t() 视图,使外部消费者(如 MLA 的 scaled_dequantize)看到 [n_blocks, k_blocks] 形状。
-
计算时恢复布局:在 apply_block_scaled_mm 中,对 Bs 调用 .t() 恢复连续的 [k_blocks, n_blocks] 缓冲区,以满足 oneDNN 的 fp8_gemm 要求,无需额外数据拷贝。
-
抽取 BMM 参数准备:将原本内联在 process_weights_after_loading 中的 BMM 参数预处理逻辑抽取为独立的 _prepare_bmm_params 方法,接收转置后的 scale([k_blocks, n_blocks]),按 batch 拆分 scale 和 weight 为 3D 张量,供 grouped fp8_bmm 使用。
-
测试与验证:本 PR 未增加新的单元测试,但作者在 PR body 中报告了相关测试通过情况(test_can_initialize_large_subset[DeepseekV3ForCausalLM])和 GSM8K 评估结果(94.92% vs 94.01%)。
关键文件:
vllm/model_executor/kernels/linear/scaled_mm/xpu.py(模块 内核层;类别 source;类型 data-contract;符号 _prepare_bmm_params): 核心变更文件,修正 scale 布局以兼容 MLA 和 oneDNN,并抽取 BMM 参数准备逻辑。
关键符号:process_weights_after_loading, _prepare_bmm_params, apply_block_scaled_mm
关键源码片段
vllm/model_executor/kernels/linear/scaled_mm/xpu.py
核心变更文件,修正 scale 布局以兼容 MLA 和 oneDNN,并抽取 BMM 参数准备逻辑。
关键源码片段
# vllm/model_executor/kernels/linear/scaled_mm/xpu.py
class XPUFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel):
def process_weights_after_loading(self, layer: torch.nn.Module):
super().process_weights_after_loading(layer)
scale_attr = (
"weight_scale_inv" if hasattr(layer, "weight_scale_inv") else "weight_scale"
)
scale = getattr(layer, scale_attr)
# Checkpoint scale is [n_blocks, k_blocks] (one value per 128x128 tile).
# oneDNN fp8_gemm requires contiguous [k_blocks, n_blocks] layout.
# 我们存储转置后的连续缓冲区作为 .t() 视图,使得:
# - MLA 的 scaled_dequantize 仍能看到 [n_blocks, k_blocks] 形状
# - apply_block_scaled_mm 通过 .t() 恢复连续缓冲区
scale_kn = scale.data.t().contiguous() # [k_blocks, n_blocks]
replace_parameter(layer, scale_attr, scale_kn.t()) # view: [n_blocks, k_blocks]
if getattr(layer, "is_bmm", False):
self._prepare_bmm_params(layer, scale_kn)
def _prepare_bmm_params(self, layer: torch.nn.Module, scale_kn: torch.Tensor) -> None:
"""Precompute batched weight and scale for grouped fp8_bmm (e.g. wo_a)."""
# 拆分 scale [k_blocks, n_blocks] 为 [G, k_blocks, n_blocks_per_group]
# 以及 weight [N_total, K] 为 [G, K, N_per_group] 用于批量 GEMM。
batch = layer.bmm_batch_size
k_blocks, n_blocks = scale_kn.shape
layer.bmm_scale = (
scale_kn.reshape(k_blocks, batch, n_blocks // batch)
.permute(1, 0, 2)
.contiguous()
)
w = layer.weight
N_total, K = w.shape
layer.bmm_weight = w.reshape(batch, N_total // batch, K).permute(
0, 2, 1
) # [G, K, N_per_group]
def apply_block_scaled_mm(self, A, B, As, Bs) -> torch.Tensor:
# B 为 [N, K];.t() 得到 [K, N] 视图(无拷贝)。
# Bs 存储为 [n_blocks, k_blocks] 视图;.t() 恢复 oneDNN 期望的连续 [k_blocks, n_blocks] 缓冲区。
return torch.ops._xpu_C.fp8_gemm(
A,
B.t(),
self.config.out_dtype,
As,
Bs.t(),
torch.Tensor(),
)
评论区精华
本 PR 的 review 评论较少,未发现实质性的技术争议。多数 reviewer(如 xinyu-intel、xwu-intel 等)均批准了该变更。yma11 仅评论“LGTM. Thanks.”,表明变更获得认可。
风险与影响
关联脉络
- PR #50434 [XPU] [BugFix] Add deepseek_v4_fp8 to xpu supported_quantization list: 同为 XPU 平台相关,涉及 FP8 量化支持,可能与本次变更相互影响。
- PR #46516 Enable gfx1250 ROCm architecture: 涉及量化内核和平台支持,与本 PR 在量化路径上有潜在关联。
参与讨论