已关闭
[Bug]: triton_experimental 后端 DeBERTa Bernoulli NPUGraph 捕获失败及 QA 标量 loss 精度异常 #3812
htchu创建于  8月7日关闭于  25 天前
htchu成员
8月7日 创建

在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。

⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:

  • API 令牌或密钥
  • 密码或身份验证凭证
  • 私有网址或接口地址
  • 个人或机密数据
  • ...

在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用 <TOKEN> 等占位符替代原有内容。

环境信息

  • NPU:Ascend 910B2(A2)
  • CANN:9.1.0
  • PyTorch:2.13 分支自编包
  • torch_npu:2.13 分支自编包
  • Transformers:4.36.0
  • 后端:triton_experimental
  • 模式:FP32 training accuracy,50 iterations

🐛 问题描述

使用 triton_experimental 后端编译训练 DeBERTa/DeBERTaV2 及其他 HuggingFace QA 模型时存在两个独立问题。

1. DeBERTa/DeBERTaV2 训练态 Bernoulli 无法被 NPUGraph 捕获

DeBERTa 训练态 dropout 会进入 aten.bernoulli fallback。当前 NPU Bernoulli op-api 仍通过 NPUGeneratorImpl::philox_engine_inputs() 获取 host 侧 seed/offset,CANN 尚无接收 PhiloxNpuState tensor seed/offset 的 Bernoulli ABI。NPUGraph 捕获期间执行该 fallback 会报错:

RuntimeError: Refactor this op to use NPUGeneratorImpl::philox_npu_state.
Cannot call NPUGeneratorImpl::philox_engine_inputs during NPU graph capture.
Current npuStreamCaptureStatus: npuStreamCaptureStatusActive

受影响模型:

  • DebertaForMaskedLM
  • DebertaForQuestionAnswering
  • DebertaV2ForMaskedLM
  • DebertaV2ForQuestionAnswering

2. QA 最终标量 loss 的 Grid1D 过量 launch 导致精度异常

DeBERTa V2、ALBERT、MobileBERT 和 RoBERTa QA 的 start/end logits、全部梯度及参数更新均通过精度检查,但最终返回的 loss 标量偏差达到 3~6,关闭 ACLGraph 后仍可复现。

最终 loss kernel 具有以下 metadata:

grid_type = "Grid1D"
npu_num_x_nodes = 0
xnumel = 1
mutated_arg_names = ["in_out_ptr0"]

旧 launcher 仅在 npu_num_x_nodes == 1 时按 xnumel 限制 Grid1D 的 grid_0。零 free-x-node 的标量 kernel 因此进入通用分支,在 Ascend 910B2 上使用 grid_0 = NPU_CU_COUNT = 48。生成 kernel 的 program_id(0) 没有参与地址计算,48 个 program 会同时读写原地复用的 in_out_ptr0[0],产生非确定性数据竞争。

复现步骤

export ASCEND_RT_VISIBLE_DEVICES=0
source env.sh

python -u benchmarks/torchbench/huggingface.py \
  --accuracy --cold-start-latency --train --float32 \
  --backend inductor --npu-backend triton \
  --only DebertaV2ForQuestionAnswering \
  --iterations 50 --accu-summary

Bernoulli 问题在默认 ACLGraph 配置下表现为捕获异常;临时规避捕获后,修复 scalar-grid 之前会表现为 fail_accuracy,且只有最终 loss 失败。

期望行为

  • 含当前不支持 graph-safe RNG ABI 的 Bernoulli fallback 的 compiled graph 应安全跳过 NPUGraph 捕获,同时保留 Inductor/Triton 编译及 eager RNG 序列。
  • 零/单 free-x-node 的非 reduction Grid1D kernel 应根据 xnumel 计算 launch grid;xnumel=1 的标量 kernel 只启动一个 program。
  • QA logits、loss、梯度、参数更新和 buffers 均应通过精度检查。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
Hhtchu成员
8月7日 关联了看板:FrameworkPTAdapter 版本issue看板
ascend-robotascend-robot成员
8月7日 添加了label:bug
Hhtchu成员
8月7日 将 htchu 设为负责人
TorchNPU-BotTorchNPU-Bot成员
8月7日 添加了label:bot-triaged
TorchNPU-Bot
TorchNPU-Bot成员
8月7日 评论:

检测到当前 issue 已关联 PR,自动添加标签:bot-triaged

likedislike
Hhtchu成员
26 天前 关联了pull request:fix: resolve Triton Experimental model correctness regressions
Hhtchu成员
25 天前 issue状态由 TODO 改变为 DONE
Hhtchu成员
25 天前 关闭了 issue
ascend-robotascend-robot成员
25 天前 添加了label:resolved