执行摘要
- 一句话:XPU 支持 MXFP8 线性权重,优化 DeepSeek V4 输出投影
- 推荐动作:值得精读。核心亮点是零额外内存的布局兼容方案:不复制 scale,只用
.t() 视图同时满足 checkpoint 规范布局与 oneDNN 连续布局要求;_prepare_bmm_params 的预计算也体现了批 GEMM 场景下减少逐调用重排开销的思路。建议结合 #49596 的讨论理解“公共算子布局假设”带来的维护困境,并跟进 vpirogov 所要求的 oneDNN reproducer,推动上游能力补齐。
功能与动机
PR body 明确说明目标是让 DeepSeek V4 XPU 输出投影路径能够高效消费 compressed-tensors MXFP8 权重。此前 XPU 的 MXFP8 kernel 在加载时直接转置权重与 scale,破坏了 checkpoint 的规范布局,导致其他公共算子(如 MLA 反量化路径)出现 shape mismatch;本 PR 通过保留原始布局并运行时视图转置来解决兼容性问题,同时为 BMM 层预计算批 GEMM 参数以消除逐调用转置开销。
实现拆解
该 PR 只修改一个文件 vllm/model_executor/kernels/linear/mxfp8/xpu.py,实现分三步:
- 权重与 scale 布局策略调整:在
process_weights_after_loading 中,不再把 layer.weight 转置替换,也不再把 layer.weight_scale 直接替换为连续 [K//32, N] 缓冲;而是将 checkpoint 的 [N, K//32] scale 先转置并 contiguous() 得到 scale_kn,再以 scale_kn.t() 视图存回 layer.weight_scale。这样外部消费者看到的 shape 仍是 checkpoint 的 [N, K//32],而 apply_weights 里 .t() 就能零成本拿回 oneDNN 需要的连续 [K//32, N] 布局。
- 新增
_prepare_bmm_params:当 layer 带有 is_bmm=True 时(如 DeepSeek V4 的 wo_a),根据 bmm_batch_size 将 scale [K//32, N_total] reshape 为 [G, K//32, N_per_group] 并重排为连续 bmm_scale,同时把 weight [N_total, K] reshape 为连续 bmm_weight [G, K, N_per_group],避免每次前向都做 permute().contiguous()。
apply_weights 运行时视图转置:调用 torch.ops._xpu_C.fp8_gemm 时传 layer.weight.t() 与 layer.weight_scale.t(),两个视图都是零拷贝;同时保持 layer.weight 原始 [N, K] 不动,避免影响其他假设原始形状的算子。
- 测试与配套:本 PR 没有新增单元测试,模型侧验证为 GSM8K 8x B70 精度 0.948;CI 仅由维护者触发 Intel 相关流水线。
关键文件:
vllm/model_executor/kernels/linear/mxfp8/xpu.py(模块 MXFP8 内核;类别 source;类型 data-contract;符号 _prepare_bmm_params, process_weights_after_loading, apply_weights): 唯一改动文件,承载了布局策略调整、BMM 参数预计算与运行时视图转置三处核心逻辑,直接决定 XPU MXFP8 权重能否被 oneDNN 高效消费且不破坏其他算子的布局假设。
关键符号:_prepare_bmm_params, process_weights_after_loading, apply_weights
关键源码片段
vllm/model_executor/kernels/linear/mxfp8/xpu.py
唯一改动文件,承载了布局策略调整、BMM 参数预计算与运行时视图转置三处核心逻辑,直接决定 XPU MXFP8 权重能否被 oneDNN 高效消费且不破坏其他算子的布局假设。
# vllm/model_executor/kernels/linear/mxfp8/xpu.py
# XPUMxFp8LinearKernel:XPU 上 MXFP8 W8A8 GEMM 内核。
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
# 检查点中 scale 形状为 [N, K//32](每 32 个元素一个 E8M0 scale)。
# oneDNN 的 fp8_gemm 要求连续 [K//32, N] 布局。
# 这里把 transposed contiguous 缓冲存成 .t() 视图,使得:
# - 其他消费者(如反量化路径)仍然看到 checkpoint 的形状 [N, K//32];
# - apply_weights 里只需 .t() 即可零成本拿回 oneDNN 需要的连续布局。
weight_scale = layer.weight_scale.view(torch.float8_e8m0fnu)
scale_kn = weight_scale.data.t().contiguous()
replace_parameter(layer, "weight_scale", scale_kn.t())
# 对 BMM 类层(如 DeepSeek V4 的 wo_a)预计算批 GEMM 权重与 scale
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:
"""预计算批处理权重与 scale,供分组 fp8_bmm 使用(例如 wo_a)。
将 scale [K//32, N_total] 切分为 [G, K//32, N_per_group],
将 weight [N_total, K] 整理为连续 [G, K, N_per_group]。
"""
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).contiguous()
)
def apply_weights(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
out_dtype = x.dtype
x_fp8, x_scale = quant_mxfp8(x)
# weight 保存为 [N, K];.t() 得到 [K, N] 视图,不拷贝。
# scale 保存为 [N, K//32] 视图;.t() 恢复 oneDNN 期望的连续 [K//32, N] 缓冲。
return torch.ops._xpu_C.fp8_gemm(
x_fp8,
layer.weight.t(),
out_dtype,
x_scale,
layer.weight_scale.t(),
bias,
)
评论区精华
审阅中围绕 xpu.py:77 的 apply_weights 改动展开了多轮讨论:
风险与影响
关联脉络
参与讨论