已关闭
[Requirement|需求建议]: 【社区任务】AssignSub算子AscendC实现贡献 #2277
ิีิีีึีึีึีึ创建于  7月21日关闭于  7月27日
ิีิีีึีึีึีึ
7月21日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

一、背景信息

AssignSub算子实现变量减法赋值操作,对应TensorFlow中的tf.assign_sub接口,计算var = var - value并将结果写入输出。该算子是训练过程中参数更新的基础原子操作,广泛用于优化器(如SGD、Adam)的权重更新步骤。

当前CANN内置的AssignSub仅有ascend950平台的DAG实现,本需求为Atlas A2/A3训练系列产品提供高性能的AscendC实现。

二、价值/作用

  • 优化器基础算子:SGD优化器的核心更新步骤weight -= lr * grad直接依赖AssignSub,缺少该算子会导致训练流程中断或性能劣化
  • 多数据类型支持:覆盖训练场景中常见的7种数据类型(float32/float16/bf16/int8/uint8/int32/int64),满足混合精度训练和量化训练需求
  • 高带宽利用率:作为纯访存瓶颈算子,实现了97%的HBM带宽利用率,接近硬件理论极限

三、设计方案

3.1 使能方式
  • 图模式:通过GE IR构图调用(OpType: AssignSub)
  • Cannjudge框架下支持aclnn直调(框架自动生成aclnnAssignSub接口)
3.2 总体设计

计算公式:var_out = var - value

对于int8/uint8类型,减法结果按模256环绕(与TBE行为对齐):

  • int8: 结果范围[-128, 127]
  • uint8: 结果范围[0, 255]
3.2.1 算子支持的数据类型
输入/输出 数据类型 数据格式
var float16, int8, float32, int32, uint8, bfloat16, int64 ND
value float16, int8, float32, int32, uint8, bfloat16, int64 ND
var_out float16, int8, float32, int32, uint8, bfloat16, int64 ND

var和value的数据类型必须一致,输出类型与输入一致。

3.2.2 host侧设计

InferShape:输出shape与输入var完全一致。

Tiling策略

  • 通过platform_ascendc::PlatformAscendC动态获取核数和UB大小
  • 核间切分:按元素总数均分到各核,每核处理的元素数向上对齐到32B block边界
  • 大shape优化:当数据量较大时,将核间边界进一步对齐到512B(HBM burst边界),提升带宽利用率
  • UB分块:根据数据类型计算每元素的UB占用量(含双缓冲和类型转换临时空间),向下对齐到block边界
    • float16/float32/int32:3队列 × 2缓冲 × dtype_size
    • int8/uint8:3队列 × 2缓冲 × 1B + 2 × half临时缓存
    • bfloat16:3队列 × 2缓冲 × 2B + 2 × float临时缓存
    • int64:3队列 × 2缓冲 × 8B + 2 × int32临时缓存
  • 模板参数(tiling key):按数据类型分7个模板实例
3.2.3 kernel侧设计

整体流水:CopyIn → Compute → CopyOut(三级流水,双缓冲)

分类型计算策略

数据类型 计算方式 中间类型
float16, float32, int32 直接Sub 无需转换
int8, uint8 Cast→half, Sub, 模256环绕, Cast回 half
bfloat16 Cast→float32, Sub, Cast回 float32
int64 Cast→int32, Sub, Cast回 int32

int8/uint8模256环绕实现

  • Cast到half后相减
  • 结果Cast到int16
  • int8: 算术左移8位+算术右移8位(保留低8位并符号扩展)
  • uint8: 逻辑左移8位+逻辑右移8位(保留低8位无符号)
  • Cast回目标类型

数据搬运优化

  • 对齐块(currentNum % alignElem == 0):使用轻量DataCopy
  • 非对齐尾块:使用DataCopyPad处理
3.3 支持硬件
硬件平台 SoC版本 支持状态
Atlas A2训练系列产品 ascend910b
Atlas A3训练系列产品 ascend910_93

3.4 算子约束限制

  • var和value的shape必须完全一致(不支持broadcast)
  • var和value的数据类型必须一致
  • 数据格式仅支持ND
  • 数据类型支持:float16、int8、float32、int32、uint8、bfloat16、int64
  • int8/uint8的减法溢出按模256环绕处理(与TBE实现行为一致)

💡 备注

likedislike
ิีิีีึีึีึีึ
7月21日 评论:

/assign

likedislike
CANN-robotCANN-robot成员
7月21日 将 qq_64858158 设为负责人
ิีิีีึีึีึีึิีิีีึีึีึีึ
7月21日 关联了pull request:docs: 新增assignsub算子实现
Ffulltower成员
7月22日 将 fullt 设为负责人
fulltower成员
7月22日 评论:

我们会进行评审

likedislike
CANN-robotCANN-robot成员
7月27日 关闭了 issue
CANN-robotCANN-robot成员
7月27日 添加了label:resolved