已开启
feat(grad_minmax_bwd_cmp): adapt UpdateGradMinMax_hetero and BackwardSegmentCmp for Ascend NPU #23
feat(grad_minmax_bwd_cmp): adapt UpdateGradMinMax_hetero and BackwardSegmentCmp for Ascend NPU #23
已开启
lhp_lhp创建于 4 天前
lhp_lhp成员
4 天前

Description

为 Ascend 910B3 (dav-2201) 适配 2 个 segment reduce 反向算子的 NPU 原生路径,替换 CPU 回退和 LOG(FATAL)。

适配算子

算子 数学语义 实现方式 操作类型
BackwardSegmentCmp out[arg[i,k], k] = feat[i,k] (arg>=0) 直接赋值(无原子,arg 列内唯一无写冲突) scatter write
UpdateGradMinMax_hetero if (type==idx_type[r,c]) out[idx[r,c], c] += feat[r,c] 条件性原子累加(SetAtomicAdd) conditional scatter add

实现方案

  • 架构:SIMD / MemBase,参照 scatter_add_kernel.cpp
  • 核心设计:单元素散列写(scratchBuf + GetValue/SetValue + DataCopyPad blockLen=4)
  • 同步机制:复用 scatter_add 验证过的范式(PipeBarrier + SetFlag/WaitFlag + SetAtomicNone)
  • 多核切分:按 N 切分,blockDim=min(N,40)
  • host launcher:hetero 遍历 etype 串行 launch(与 CUDA/CPU 一致)

支持的数据类型

IdType Dtype 状态
int32 float32 ✅ 原生支持
int64 float32 ✅ 支持(int64 idx/arg 转 int32,同 segment_reduce.cc 范式)
int32 half (float16) ❌ LOG(FATAL) 占位
int64 half (float16) ❌ LOG(FATAL) 占位
int32 bfloat16 ❌ LOG(FATAL) 占位
int64 bfloat16 ❌ LOG(FATAL) 占位
int32 double (float64) ❌ LOG(FATAL) 占位
int64 double (float64) ❌ LOG(FATAL) 占位

两个算子同时涉及 IdType(索引类型 int32/int64)和 DType(特征数据类型 float32)。IdType 用于 arg/idx/idx_etype 数组,DType 用于 feat/out 张量。

Checklist

Changes

Test Results

  • 58/59 passed, 1 xfail(DGL 绑定层 use-after-free,非算子缺陷)
  • QA 独立探针 43/43 pass
  • BackwardSegmentCmp 全部 bit-exact
  • UpdateGradMinMax_hetero max_abs_error=3.81e-6(远低于 1e-2 阈值)
likedislike
合并受阻
Llhp_lhp成员
4 天前 修改了pull request 的描述
Llhp_lhp成员
4 天前 修改了pull request 的描述
Llhp_lhp成员
4 天前 修改了pull request 的描述
Llhp_lhp成员
4 天前 强制推送  1 个提交:e7482e6f-feat(grad_minmax_bwd_cmp): adapt UpdateGradMinMax_hetero and BackwardSegmentCmp for Ascend NPU