[Bug] Ascend NPU: RMSNorm crashes with elementwise_affine=False; _native_npu FA rejects [B, N, 1, Skv] masks (LTX-2)

在 Ascend NPU 上以 attn_backend=_native_npu 运行 LTX-2( dg845/LTX-2.3-Diffusers )时,如果 RMSNorm 层配置了 elementwise_affine=False ,或交叉注意力使用了 [B, N, 1, Skv] 形状的掩码

快速结论:在 Ascend NPU 上以 attn_backend=_native_npu 运行 LTX-2(dg845/LTX-2.3-Diffusers)时,如果 RMSNorm 层配置了 elementwise_affine=False,或交叉注意力使用了 [B, N, 1, Skv] 形状的掩码,就会触发 NPU 后端不兼容而崩溃。优先排查这两处:RMSNorm 是否给 npu_rms_norm 传入了 weight=None,以及注意力掩码是否只有 [B, 1, 1, Skv] 会被自动扩展。

适用环境:OS Linux aarch64;硬件 Ascend NPU;Python 3.11;PyTorch 2.10.0;torch_npu 2.10.0;CANN 9.0.0;diffusers main;模型 dg845/LTX-2.3-Diffusers;运行场景为 verl-omni 的 LTX-2.3 text-to-audio-video FlowGRPO LoRA 训练。

最快修复方案:暂无确认的一步修复方案。Issue 中作者表示已在 PR #14288 准备了修复:RMSNorm 在 elementwise_affine=False 时回退到已有的 PyTorch RMSNorm 实现;同时把 [B, N, 1, Skv] 掩码扩展为 [B, N, Sq, Skv]。可优先尝试跟进该 PR 或在本地应用等价改动。

注意事项:该修复属于运行时正确性与 NPU 兼容性修复,不是输出质量调整。修复前不会产生推理输出,因为执行在崩溃处中断;修复后相同负载可正常完成受影响算子。上述结论为 Issue 作者在 PR 中给出的验证,落地前建议在自己的 NPU 环境与 torch/torch_npu 版本组合上复测。

问题场景

用户在 Ascend NPU 上运行 Diffusers 的 LTX-2,具体是 dg845/LTX-2.3-Diffusers 检查点,用于 LTX-2.3 text-to-audio-video LoRA FlowGRPO 训练(verl-omni 的 examples/flowgrpo_trainer/ltx2/run_ltx2_3_t2av_lora_npu.sh)。当注意力后端指定为 _native_npu 时,执行会进入两条 NPU 专用代码路径并报错:一是带 elementwise_affine=False 的 RMSNorm,二是 LTX 交叉注意力使用的 [B, N, 1, Skv] 形状掩码。

报错原文

# Error 1: RMSNorm(elementwise_affine=False)
npu_rms_norm called with weight=None
gamma is None

# Error 2: _native_npu fused attention mask [B, N, 1, Skv]
Ascend FA expects Sq on dim=-2, not a singleton 1

原因分析

两条报错都源于 Ascend NPU 融合算子对输入的限制比 PyTorch/SDPA 更严格:

  • RMSNorm:elementwise_affine=False 时,层的 weightNone,但 torch_npu.npu_rms_norm 要求提供 gamma 张量,因此以 gamma is None 崩溃。
  • 融合注意力掩码:Ascend FA 不会像 SDPA 那样广播单例的 query 长度维。LTX 交叉注意力传入的是 [B, N, 1, Skv](例如 [1, 32, 1, 1024]),而当前 _maybe_modify_attn_mask_npu 只会扩展 [B, 1, 1, Skv],因此该掩码被拒绝。可能原因是掩码扩展逻辑未覆盖这种形状。

环境排查

  • 确认硬件为 Ascend NPU,操作系统为 Linux aarch64。
  • 确认 Python 3.11。
  • 确认 PyTorch 2.10.0 与 torch_npu 2.10.0 匹配。
  • 确认 CANN 版本为 9.0.0。
  • 确认 diffusers 安装自 main
  • 确认注意力后端确实设置为 _native_npu,且模型检查点为 dg845/LTX-2.3-Diffusers
  • 确认 RMSNorm 层是否使用了 elementwise_affine=False,以及交叉注意力掩码形状是否为 [B, N, 1, Skv]

解决步骤

  1. 先按 Issue 的最小脚本复现两类失败:一个构造 RMSNorm(dim=64, eps=1e-6, elementwise_affine=False) 后在 NPU 上前向;另一个构造 [B, N, 1, Skv]attn_mask,在 attention_backend(AttentionBackendName._NATIVE_NPU) 下调用 dispatch_attention_fn
  2. 确认 RMSNorm 报错来自 npu_rms_norm 收到 weight=None。按 PR #14288 的思路,在 elementwise_affine=False 时回退到已有的 PyTorch RMSNorm 实现,而不是调用 NPU 融合算子。
  3. 确认注意力报错来自掩码形状。按 PR #14288 的思路,在 _maybe_modify_attn_mask_npu 中把 [B, N, 1, Skv] 扩展为 [B, N, Sq, Skv],而不仅是 [B, 1, 1, Skv]
  4. 在本地应用或拉取 PR #14288 后,重跑上述两个最小脚本,确认两个操作都能完成。
  5. 再跑端到端 verl-omni LTX-2.3 配方 run_ltx2_3_t2av_lora_npu.sh,确认 rollout/training 路径不再停在这两个错误上。

验证方法

按 Issue 作者给出的两种验证方式确认:

  • 最小脚本层面:修复前 RMSNorm(elementwise_affine=False)weight=None 传给 npu_rms_norm 并崩溃,修复后回退到 PyTorch 实现并成功完成;修复前 [B, N, 1, Skv] 掩码被 Ascend 融合注意力拒绝,扩展为 [B, N, Sq, Skv] 后注意力调用成功完成。
  • 端到端层面:使用 dg845/LTX-2.3-Diffusers 检查点跑 examples/flowgrpo_trainer/ltx2/run_ltx2_3_t2av_lora_npu.sh,修复前执行停在 RMSNorm 或融合注意力掩码错误,修复后 rollout/training 能通过这两条受影响代码路径且不再报这些错误。

参考来源

huggingface/diffusers #14380

GamsGo AI

AI 工具推荐

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

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

了解 GamsGo AI

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

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 24807

发表回复

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