Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

arithFIX — FlagGems pow 链式布尔运算兼容性 + 半精度指数查表缺陷修复

修复 FlagGems issue #5732: backends.yaml 为 ascend-cann850 钉住的 flagtree 0.6.0+ascend3.2(triton 3.2) 的代码生成器不支持链式布尔运算(a or b or c),而 _ascend/ops/pow.py 三个 kernel 正是这种写法——该环境下 x.pow(2) 直接编译崩:

UnsupportedLanguageConstruct: chained boolean operators (A or B or C)
are not supported; use parentheses to split the chain.

一句话版本:链式 or 改成加括号的嵌套写法(triton 3.2/3.5 都接受的等价 形式);顺带修复被它掩盖的第二个缺陷——fp16/bf16 指数路径的组合 (fp32, fp16) 查表 KeyError(geluFIX #5867 同族)。

修复前 (triton 3.2 环境): pow 全变体编译崩 (UnsupportedLanguageConstruct)
修复前 (triton 3.5 环境): fp16/bf16 指数路径 KeyError(fp32, fp16) 编译崩
修复后: 两环境均可编译; 14/14 数值矩阵 + 仓库 pytest 82f -> 11f(剩余为
        fp16 大指数精度边界, 见边界声明)

问题是怎么回事

缺陷一(issue 报的):pow.py 三个 jit kernel 用了三项链式 or:

if (
    tl.constexpr(exponent.dtype.is_fp32())
    or tl.constexpr(exponent.dtype.is_fp16())     # ← A or B or C
    or tl.constexpr(exponent.dtype.is_bf16())
):

backends.yaml 把 ascend850 环境钉在 triton 3.2(flagtree 0.6.0+ascend3.2), 该版本的 code_generator 对 ast.BoolOp 只支持二元——三项链直接抛 UnsupportedLanguageConstruct。AST 全仓库扫描(tests/scan_chain.py): jit kernel 内共 4 处——pow.py ×3(issue 报的)+ spacemit/argmin.py ×1 (主动排查发现,同模式 is or is or is)。host 层 225 处链式 or 是合法 Python,不受影响。

缺陷二(修复过程中暴露的):链式 or 在本机(triton 3.5.1,支持链式) 不崩,于是跑到了下一层——fp16/bf16 指数分支只把 x 升 fp32、exponent 原样传 fp16,CANN libdevice pow 查表只有 (fp32,fp32)/(fp16,fp16)/(bf16,bf16) 组合,(fp32, fp16) 直接 KeyError。git stash 对照证实原版同样崩(82 failed 里大量此错)——即 issue 环境(3.2)永远死在缺陷一,从没跑到过缺陷二; 3.5 环境则直接踩缺陷二。这与 geluFIX #5867 是同族问题(gelu 家族的 int 指数 KeyError,这里是半精度指数组合 KeyError)。

修复思路

  1. 链式 → 括号嵌套:(A or B) or C——triton 3.2 报错信息里建议的 形式,3.2/3.5 都接受,AST 等价变换,零语义变化
  2. 指数统一升 fp32:fp16/bf16 分支 exponent.to(tl.float32)—— __hmf_powf 本来就是 fp32 入口,数值不变(int/半精度/浮点指数路径 实测逐位一致);查表命中 (fp32, fp32)
  3. spacemit/argmin.py 同模式一并修(AST 等价,无实机但纯语法变换风险极低)

不建议改 backends.yaml 升版本(issue 的另一方案):代码侧兼容是纯下行 安全的,升级 flagtree 是环境级变更影响面大,两者不冲突。

怎么用

前置:昇腾环境(triton-ascend 3.2 / 3.5 均可;本修复在 3.5.1 + CANN 8.5

  • flag_gems 5.3.5 验证,3.2 兼容性由 AST 层保证)。
# 部署 (两个文件)
cp src/pow.py             /path/to/site-packages/flag_gems/runtime/backend/_ascend/ops/pow.py
cp src/argmin_spacemit.py /path/to/site-packages/flag_gems/runtime/backend/_spacemit/ops/argmin.py
find /path/to/site-packages/flag_gems -name __pycache__ -exec rm -rf {} +

# 验证
ASCEND_LAUNCH_BLOCKING=1 python tests/test_pow.py    # 14/14 数值 + 0 链式残留
python tests/scan_chain.py                           # 全仓库 jit 内链式 or = 0 (门禁)

# 回归 (必须 --ref cpu)
cd /root/FlagGems && python -m pytest tests/test_pow.py -q --ref cpu
# 修复版: 11 failed, 241 passed (原版: 82 failed, 170 passed)

性能对比(修复前后实测,含严格 A/B)

严格 A/B(同机同进程口径,原版 pow.py 经 git stash 部署对照,中位数):

路径(4MB) 原版 gems 修复版 gems 判读
fp32 x.pow(2.0) 127.5 µs 111.6-133.3 µs 同噪声带(共享机器 ±15µs 波动),零回归
fp32 tensor 指数 112.9 µs 114.6-116.5 µs 同噪声带
int64 底数 pow(3) 109.5 µs 111.6-119.1 µs 同噪声带
fp16 x16.pow(2.0) 不可用(KeyError) 112.8 µs 修复净增可用路径
fp16 tensor 指数 不可用 115.6 µs 同上

修复的技术内容决定其零开销性质:括号化是 AST 层变换(编译期归一), exponent 升 fp32 与 __hmf_powf 的 fp32 入口一致(同一硬件指令)。 tests/bench_perf.py 可复现全部数据。

与原生的差距(gems ~113µs vs 原生 ~13µs,9x)是框架 pointwise host 链 的既有成本,见全局矩阵 NATIVE_VS_GEMS.md(五库通用),与本修复无关。

剩余 11 个失败的裁定(不装完美)

全部是 dtype1(fp16) × 指数 ±100.001/-111.999 的精度边界:Greatest relative difference 1.4e-6 对容差 1.3e-6——超 0.1e-6,属 fp16 大指数 (输出超 fp16 表示范围,走 fp32 中间量)的 ULP 级舍入差。原版这些用例 死在编译(从未执行),修复后第一次真正运行才暴露。非崩溃、非语义错误, 是测试容差与半精度大指数的边缘问题,建议上游单独评估容差或 kernel 内 精度策略。

目录结构

├── README.md              # 本文
├── arith.patch            # 97 行 diff (pow.py + spacemit/argmin.py), git apply 用
├── src/
│   ├── pow.py(.orig)      # _ascend/ops/pow.py 修复版 + 原版
│   └── argmin_spacemit.py(.orig)  # _spacemit/ops/argmin.py 同上
├── tests/
│   ├── test_pow.py        # 14 项数值矩阵 + 链式 or 门禁 (triton 3.2 兼容检查)
│   ├── scan_chain.py      # AST 全仓库扫描: jit 内链式 BoolOp 清单 (修复前 4 处)
│   ├── bench_perf.py      # 修复前后性能对比 (回归护栏 + 新增路径 + 原生对照)
│   └── ab_one.py          # 严格 A/B 单进程探针 (配合 git stash 部署原版对照)
└── docs/
    ├── ARITH_5732_REPORT.md
    └── NATIVE_VS_GEMS.md  # 六库通用: 原生 vs gems 全景矩阵

验证矩阵汇总

验证项 结果
AST 扫描定位 jit 内链式 or 共 4 处 (pow×3 + spacemit argmin×1), 修复后 0
数值矩阵 (3 dtype × 4 指数 × 3 入口 + int64 else 分支) 14/14 (含负底数**0.5=nan 语义)
仓库 pytest test_pow --ref cpu 原版 82f/170p → 修复版 11f/241p
剩余 11f 裁定 fp16 大指数 ±100/±111 精度边界 (rel diff 1.4e-6 vs 容差 1.3e-6), 原版从未执行到
性能 4MB pow: gems 133µs vs 原生 12µs (框架 pointwise host 既有差距, 与修复无关)
spacemit argmin AST 等价变换 + 解析通过 (无实机)

About

FlagGems pow chained-bool triton-3.2 compat + fp16 exponent dispatch KeyError fix (issue #5732)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages