修复 FlagGems issue #5802:在昇腾 NPU 上
开启 flag_gems 后,torch.inference_mode() 里调用 torch.nn.functional.one_hot 触发
RecursionError: maximum recursion depth exceeded。
一句话版本:旧代码用 is_cuda 判断设备,NPU 张量恒为 False,于是每个 NPU 调用都
fallback 回 F.one_hot 自己——inference_mode 把 dispatch 逼进这条回环,无限递归。
我们改成按 device.type 分发,让 NPU 第一次真正跑上 Triton kernel,并把 kernel
路径上被 fallback 掩盖多年的三个缺陷一并补齐。
修复前: RecursionError: maximum recursion depth exceeded (145 层)
修复后: 输出与原生逐位一致, normal / inference_mode / no_grad / CPU 全场景通过
issue 的最小复现只有几行:
import torch
import torch_npu
import flag_gems
flag_gems.enable()
x = torch.tensor([0, 1, 2, 0], device='npu:0')
torch.nn.functional.one_hot(x, num_classes=3) # 没问题
with torch.inference_mode():
torch.nn.functional.one_hot(x, num_classes=3) # RecursionError!同一个调用,普通模式和 inference_mode 表现天差地别,玄机在 aten 分发层:
aten::one_hot 是 "composite implicit autograd" 算子。普通模式下 Autograd key
在 dispatch keyset 里,PyTorch 用自己的 composite 分解(zeros + scatter_)拦截,
根本不会进 flag_gems(实测 GEMS ONE_HOT 日志触发 0 次)。inference_mode 会把
Autograd key 从 keyset 排除,dispatch 于是落到 PrivateUse1——flag_gems 注册的
one_hot kernel。然后:
if not tensor.is_cuda: # NPU 张量 is_cuda == False
return torch.nn.functional.one_hot(tensor, num_classes) # 回到自己!F.one_hot 内部还是调 aten::one_hot,还是 inference_mode,还是落 PrivateUse1,
还是 is_cuda=False……回环闭合,145 层后栈溢出。
三个分发对照实验坐实了这条链(tests/dispatch_probe.py):
| 场景 | GEMS 触发 | 结果 |
|---|---|---|
| no_grad + NPU | 0 | 正常(Autograd key 未排除,composite 拦截) |
| inference_mode + CPU | 0 | 正常(CPU keyset 到不了 PrivateUse1) |
| inference_mode + NPU | 145 | RecursionError |
# 修复前
if not tensor.is_cuda:
return torch.nn.functional.one_hot(tensor, num_classes)
# 修复后
if tensor.device.type == "cpu":
return torch.nn.functional.one_hot(tensor, num_classes)CPU 依然走原生(CPU dispatch 到不了 PrivateUse1,fallback 安全且永不回环); NPU 等所有 flag_gems 后端走自己的 Triton kernel。
但这只是撕开掩盖的第一步。 is_cuda fallback 的副作用是:这个 Triton kernel
在 NPU 上从未真正执行过(普通模式被 composite 拦截,其余模式走 fallback),kernel
路径上的既有缺陷全部被掩盖。分发修好后它们立刻暴露,必须一并处理:
| 掩盖的缺陷 | 原行为(kernel 直跑) | 原生行为 | 修复 |
|---|---|---|---|
空输入 grid=(0,) |
CANN coreDim=0 abort(#5743 同款) | OK, shape (0, C) | numel==0 早退 |
负类别值 [2,0,-1] |
静默输出全零行 | RuntimeError | 显式校验 |
| 值 ≥ num_classes / C≤0 | 静默丢弃 | RuntimeError | 显式校验 |
kernel 本体(one_hot_kernel)一字未动——实测数值本来就与原生逐位一致
(1d / 2d / 标量 / auto num_classes)。
前置:昇腾环境(本修复在 Ascend 910 + CANN 8.5 / torch 2.10 / torch_npu 2.10
- flag_gems 5.3.5 上验证;issue 报告环境 torch 2.9,同病)。
部署(唯一要动的文件):
cp /path/to/site-packages/flag_gems/ops/one_hot.py /path/to/backup/
cp src/one_hot.py /path/to/site-packages/flag_gems/ops/one_hot.py
find /path/to/site-packages/flag_gems/ops -name __pycache__ -exec rm -rf {} +或者 git apply one_hot.patch(基于 v5.3.5)。
验证(零依赖,直接 python 跑):
ASCEND_LAUNCH_BLOCKING=1 python tests/repro_5802.py # issue 原复现,应双 OK
ASCEND_LAUNCH_BLOCKING=1 python tests/test_red.py # 6 项语义契约,TOTAL: 6, failed: 0
ASCEND_LAUNCH_BLOCKING=1 python tests/acceptance.py # 7 场景端到端,ALL SCENARIOS PASS
ASCEND_LAUNCH_BLOCKING=1 python tests/dispatch_probe.py # 分发机制三对照(分析用)回归(仓库 pytest):
cd /root/FlagGems
ASCEND_LAUNCH_BLOCKING=1 python -m pytest tests/test_one_hot.py -q # 1 passed├── README.md # 本文
├── one_hot.patch # 最小 diff(含署名标记),git apply 用
├── src/
│ ├── one_hot.py # 修复后完整文件(含 wt 署名标记),直接替换
│ └── one_hot.py.orig # v5.3.5 原版备份
├── tests/
│ ├── repro_5802.py # issue 复现 + 修复验收
│ ├── test_red.py # 6 项语义契约(TDD 红转绿脚本)
│ ├── acceptance.py # 7 场景端到端
│ ├── dispatch_probe.py # 分发机制对照实验(根因证据)
│ ├── dispatch_probe2.py # aten 直调对照(普通模式也不进 gems 的证据)
│ ├── native_baseline.py # 原生 F.one_hot 边界行为基准(CPU)
│ └── test_one_hot_repo.py # FlagGems 仓库版 pytest(放回仓库 tests/ 跑)
└── docs/
└── ONE_HOT_5802_REPORT.md # 完整报告:根因链、实验矩阵、自查清单
| 原版位置(src/one_hot.py.orig) | 函数 | 与本次修复的关系 |
|---|---|---|
| L52-53 | one_hot() 的 is_cuda fallback |
肇事行:NPU 恒走 fallback,inference_mode 下成回环 |
| L56-58 | auto num_classes 推断 | 未动的主体,补了空张量报错(对齐原生) |
| L24-47 | one_hot_kernel |
一字未动(数值本就正确) |
| L60-77 | 输出分配 + grid + launch | 补 numel==0 早退(#5743 同款零 grid 守卫) |
- 验证只覆盖 Ascend 后端。issue 附注"其他芯片可能也存在类似问题"——
is_cuda在所有非 CUDA 后端(Cambricon/Sunrise/KunlunXin/Enflame…)都为 False,理论上 同病,但无实机验证。patch 里device.type == "cpu"的判断对它们同样成立, 理论上顺带修好,未实测 - 仓库测试
test_one_hot.py的错误契约断言(负值/越界 RuntimeError)原本是靠 bug 性 fallback 意外满足的(直调→fallback→原生报错);修复后由显式校验满足, 断言语义从"碰巧对"变成"真的对",但这也意味着原版该测试从未真正测过 kernel 的错误路径 - 值域校验用
tensor.min()/max()各一次 D2H 同步——inference_mode 场景可接受 (本来就是低频算子),性能敏感路径可改为 kernel 内校验,未做 torch.compile/ dynamo 图捕获下的行为未测
| 验证项 | 结果 |
|---|---|
| issue #5802 原复现(修复前) | RecursionError,145 层(与 issue 一致) |
| 分发三对照(no_grad / CPU / inference+NPU) | 根因链坐实 |
| TDD 红转绿(6 项语义契约) | 红 3 FAIL+递归崩 → 绿 6/6 |
| 端到端 7 场景(normal/inference/no_grad/aten直调/CPU/2d+auto/空张量) | 全过 |
| 数值正确性(1d/2d/标量/auto-C vs 原生) | 逐位一致 |
| 边界契约(负值/越界/C=0/空+autoC) | 与原生逐项对齐 |
| 仓库 pytest test_one_hot.py | 1 passed |