MetalConfig quantization: batched generate silently corrupts all rows after the first — affine_qmm_t computes only batch element 0 of 3D inp

此报错发生在 Transformers 使用 MetalConfig 量化(bits=4/8)的模型上,当以 batch size > 1 进行生成时,只有第一行输出正确,后续所有行解码结果坍缩为 token 0(输出 !!!!... )。优先排查 affine_qmm_t Metal kernel

快速结论:此报错发生在 Transformers 使用 MetalConfig 量化(bits=4/8)的模型上,当以 batch size > 1 进行生成时,只有第一行输出正确,后续所有行解码结果坍缩为 token 0(输出 !!!!...)。优先排查 affine_qmm_t Metal kernel 对 3D 输入 [batch, 1, hidden] 的批处理维度处理;Issue 已给出 kernel 层修复验证,但尚未发布到正式版本。

适用环境:macOS (Apple Silicon, MPS 后端), Transformers 5.13.1 / 5.16.1, PyTorch 2.10–2.14 (MPS), kernels-community Metal 量化 kernel,Python 3.11/3.12;硬件涉及 Apple M4 Max / M4 Pro(36–64 GB)。

最快修复方案:暂无确认的一步修复方案。Issue 中指出在 MetalLinear.forward 调用 affine_qmm_t 前将 3D 输入展平为 2D 可使 batched generate 完全正确,但这是社区验证的本地补丁,尚未合并到 Transformers 主分支。

注意事项:该修复仅针对调用侧展平,kernel 本身的 host dispatch 代码(set_strides 中 stride 传参错误)仍未在官方发布包中修复;此外 Issue 还提到 affine_qmv(向量乘)在 batched 场景有独立 bug,batched grid.z 路径性能退化(例如 [32,1,K] 形状为 60ms vs 展平后 2.2ms),需分别处理。

问题场景

报错发生在用户使用 transformers 加载经 MetalConfig(4-bit 或 8-bit)量化的模型(如 Qwen/Qwen3-0.6B),并调用 model.generate() 进行批量推理时。单个样本(batch size=1)解码正常;一旦将相同的 prompt 复制为多条形成一个 batch(batch size>1),只有第一条样本输出正确,第二条及之后的样本 argmax 全部坍缩到 token 0,解码结果变成连续的感叹号串,且输出确定性复现、逐位一致。非量化(bf16 原模型)在同环境下所有 batch 行均正常。

报错原文

MetalConfig quantization: batched generate silently corrupts all rows after the first — affine_qmm_t computes only batch element 0 of 3D inputs

典型输出对比(batch=2,复制相同 prompt):

['" Paris. The capital of Italy is Rome. ..."', '"!!!!!!!!!!!!!!!!!!!!"']

原因分析

根本原因定位于 Metal 量化 kernel 的 host 端 dispatch 代码(kernels-community 库中 quantized.mm)。已验证的结论:affine_qmm_t kernel 在输入为 3D tensor [B, 1, K](B>1)时,只计算第一个 batch element,后续元素全部产生垃圾输出;kernel 错误地将 3D 输入当作 2D [shape[-2], K] 处理,完全忽略 leading batch 维度。set_strides 辅助函数把 scales/biases 的“行 stride”传给 kernel,而 kernel 期望的是“batch stride”,导致每个 batch element 之后读到的 scale 偏移 z(row)行,高行时甚至越过 tensor 末尾,故只有第 0 行天然正确。

补充发现(同为 contributor 验证):affine_qmv(bached 向量乘)有独立缺陷;host 端传入 ndim-1 个 batch 维度,而 kernel 期望 ndim-2,并从 x_shape[batch_ndim] 读取 M 导致越界。此外 batched grid.z 路径性能极差:即使正确,每个 batch 元素为 M=1 付出整 32-row tile 代价,实测 [32,1,K] lm_head 形状耗时 60ms vs 展平 2D gemm 的 2.2ms。

环境排查

  • 确认操作系统:macOS (arm64),问题在 macOS 26.5/26.6.2 上复现。
  • 确认 PyTorch:使用 MPS 后端,问题在 torch 2.10 与 2.13/2.14 上均复现——已排除 torch 版本差异。
  • 确认 Transformers 版本:问题存在于 5.13.1 和 5.16.1。
  • 确认 kernel 包:kernels-community 的 Metal 量化 kernel;有复现者核验为 native 2.14 Metal build(非旧版 fallback)。
  • 确认使用 Python ≥3.11 和 accelerate 不是必现条件。

解决步骤

  1. 确认复现:构造 batch>1 的 prompt 列表,对比 batch=1 输出;如果只有第一条正确,后续为 !!!!...,即可确认命中此 kernel 层 bug。
  2. 使用 kernel 级探针验证(可选,无模型下载):直接调用 _get_metal_kernel().affine_qmm_t,比较 [2,K][2,1,K] 的输出与 x @ dequant(W).T 的相对误差:2D 全部正确(误差~0.005),3D 时报错与参考差异达 ~9.874。
  3. 优先尝试本地补丁:MetalLinear.forward 中,将 3D 输入 [batch, 1, hidden] 在调用 affine_qmm_t 前 reshape 为 2D [batch, hidden],完成 kernel 调用后再恢复形状。Issue 作者已实测此方法使 batch=4 的每一行输出与 batch=1 完全一致。
  4. 关注官方 kernel 修复 PR:huggingface/kernels-community #1038 包含 qmv dispatch 修复、bached 测试和 dense batch collapse(将连续 bached 输入合并为一次 2D gemm,保留 strided 路径作为 fallback),并增加 contiguity 守卫。如果尚未合并,请勿在生产中依赖。
  5. 规避方案:在官方 kernel 修复发布前,将 bached generate 改为循环逐条调用(batch=1),或对输入做 flatten 预处理;注意速度影响。

验证方法

重新运行 generate(batch=4),并检查输出:每一行的解码结果必须与相同 prompt 的 batch=1 输出逐 token 完全一致,不应出现 !!!!... 或其他重复崩溃模式。kernel 级验证:对 [B,1,K] 形状输入,计算 kernel.affine_qmm_t 输出与 x @ dequant(W).T 的逐行最大相对误差,确认所有 batch element 误差都在 ~0.5% 以内(而非某一元素出现 >9 的误差)。Issue 证实修复后 batch=1 输出与原版二进位逐位一致,无回归。

参考来源

huggingface/transformers #47331

GamsGo AI

AI 工具推荐

想把多个 AI 模型放在一个入口?

GamsGo AI 集成 ChatGPT、DeepSeek、Gemini、Claude、Midjourney、Veo 等常用模型,适合写作、绘图、视频和日常 AI 工作流。

了解 GamsGo AI

推广链接:通过此链接购买,我可能获得佣金,不影响你的价格。

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 22011

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注