Zamba-7B-v1 runs slow during torch.compile due to use_associative_scan=False

这个问题通常出现在用 torch.compile (inductor 后端)跑 Zyphra/Zamba-7B-v1 时,表现为编译/运行极其缓慢,容易被误判为“卡死”。优先排查 ZambaMambaMixer.forward 中写死的 use_associative_scan=False ,它会把

快速结论:这个问题通常出现在用 torch.compile(inductor 后端)跑 Zyphra/Zamba-7B-v1 时,表现为编译/运行极其缓慢,容易被误判为“卡死”。优先排查 ZambaMambaMixer.forward 中写死的 use_associative_scan=False,它会把 selective scan 强制切换到逐 token 的 Python 循环回退路径。注意:Issue 中已明确“hangs”这个说法具有误导性,实际是极慢,不是真的挂起。

适用环境:Issue 中报告者为 conda 环境,Python 3.10,torch 2.14.0.dev20260719+xpu、torchaudio 2.11.0.dev20260720+xpu、torchvision 0.29.0.dev20260720+xpu、transformers 5.15.0.dev0,设备为 XPU(--device xpu),dtype 为 torch.bfloat16。其他后端和依赖版本未在 Issue 中确认。

最快修复方案:暂无确认的一步修复方案。Issue 中“Proposed Fix (verified work)”提出的做法是:在 ZambaConfig 中新增 use_associative_scan: bool = True,在 __init__ 中保存 self.use_associative_scan,并在 forward() 中把它传给 selective scan。维护者回复称不介意提交 PR 为 zamba 补上该配置,但明确说明优先级不高、review 可能有延迟,因此该修复尚未合并,也未在 Issue 中得到官方确认。

注意事项:Issue 评论区指出这“不是正确性问题”,只是慢,属于当前实现的预期行为;有人怀疑是 torch 侧对循环没有做合适的转换,但这只是猜测(first intuition),未经验证。此外维护者提到 use_associative_scan 这行是随 #47630 引入的,zamba1 当初缺少对应 config 条目。若自行打补丁,需注意后续与上游合并时的冲突,以及该改动对显存/数值行为的影响未经 Issue 验证。

问题场景

用户在运行 Hugging Face Transformers 中的 Zyphra/Zamba-7B-v1 模型时触发该问题。具体场景是 torch.compile(inductor 后端)下跑 causal-lm 基准测试:eager 模式正常,Zamba2 系列模型不受影响,但 Zamba-7B-v1 在 torch.compile 下运行极慢,本地多次复现都没能看到编译完成,最初被描述为“hangs indefinitely”。

报错原文

Zamba-7B-v1 runs slow during torch.compile due to use_associative_scan=False

原因分析

最可能的原因在 src/transformers/models/zamba/modeling_zamba.py 的 ZambaMambaMixer.forward(约第 583 行):调用 mamba_selective_scan 时写死了 use_associative_scan=False,导致该函数走 recurrent fallback 路径,其中包含一个按 seq_len 迭代的 Python for 循环。在 torch.compile 下,这个循环被逐层展开/跟踪,代价极高,于是表现为编译或运行极慢。

评论区进一步指出:zamba1 当初没有对应的 config 条目来启用 associative scan,当时出于时间限制未补上;这属于性能问题而非正确性问题。至于为什么 torch 无法把这个循环优化掉,有维护者猜测是 torch 侧缺少对循环的合适转换,但明确表示不确定,仅为直觉判断。

环境排查

  • 确认 Python 版本(Issue 中为 3.10)。
  • 确认 torch / torchaudio / torchvision 版本,Issue 中为 XPU 开发版:2.14.0.dev20260719+xpu、2.11.0.dev20260720+xpu、0.29.0.dev20260720+xpu。
  • 确认 transformers 版本(Issue 中为 5.15.0.dev0)。
  • 确认设备与 dtype:--device xpu、--dtype torch.bfloat16。
  • 确认问题仅出现在 torch.compile(inductor)下,eager 模式是否正常。
  • 确认模型是 Zamba-7B-v1;Zamba2 未受影响。
  • 确认 ZambaMambaMixer.forward 中 use_associative_scan 的取值,以及 ZambaConfig 是否存在对应配置项。

解决步骤

  1. 先复现并确认现象。用 Issue 中的基准脚本跑一次,观察在 torch.compile 下是否极慢而 eager 正常,以区分“慢”与“真挂起”。
  2. 定位代码:查看 src/transformers/models/zamba/modeling_zamba.py 中 ZambaMambaMixer.forward 调用 mamba_selective_scan 的位置,确认 use_associative_scan=False 是硬编码的。
  3. 按 Issue 提出的方向打本地补丁(该方案被提交者称为“verified work”,但未合并、未经官方确认,可优先尝试):在 ZambaConfig 中新增 use_associative_scan: bool = True;在 ZambaMambaMixer.__init__ 中保存 self.use_associative_scan;在 forward() 中将其传入 selective scan 调用,替换原来的写死值。
  4. 用同一基准命令重新跑,对比 torch.compile 下的耗时与是否能跑完。
  5. 若本地改动有效,可按维护者建议向上游提 PR(维护者回复不介意提交,但提示优先级不高、review 可能延迟)。

验证方法

用 Issue 中的重现命令在 torch.compile 下重新运行 Zamba-7B-v1,确认运行时间大幅下降、不再出现长时间无法完成编译/运行的情况,且输出结果与 eager 模式或其他正确性基准一致(因为该问题被判定为非正确性问题)。如果补丁后仍然极慢,说明瓶颈可能不止这一处,需要回到 torch 侧进一步排查。

参考来源

huggingface/transformers #48080

GamsGo AI

AI 工具推荐

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

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

了解 GamsGo AI

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

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 25796

发表回复

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