Skip to content

Latest commit

 

History

3 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

onehotFIX — FlagGems one_hot 算子 inference_mode 无限递归修复

修复 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 # 完整报告:根因链、实验矩阵、自查清单

如果你想深究 bug 在哪一行

原版位置(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

About

FlagGems one_hot inference_mode RecursionError fix (issue #5802): device-aware dispatch replacing is_cuda fallback + unmasked zero-grid/value-validation defects

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages