已关闭
[Requirement|需求建议]: 950新增StridedSliceAssign算子 #2971
xuejinghui创建于  8 天前关闭于  8 天前
xuejinghui成员
8 天前 创建

一、背景信息 (必填)

在 Ascend 950(950PR/950DT)上新增 StridedSliceAssign 算子(非 D 版)。

功能:将输入张量 input_value 的内容,赋值给目标张量 var 中由 begin、end、strides 三个 1 维 int64 常量输入(配合 begin_mask/end_mask/ellipsis_mask/new_axis_mask/shrink_axis_mask 五个属性)指定的带步长切片位置,切片以外区域保持 var 原值不变(in-place / ref 语义,输出 var 与输入 var 同 shape 同 dtype)。

语义对齐 TensorFlow tf.raw_ops.StridedSliceAssign(稀疏切片 spec 形式:begin/end/strides 为 tensor 输入而非属性),补齐 950 平台该 TF 兼容算子的缺口,解决 TF 图迁移场景下该算子在 950 上无可调用实现的问题。

二、价值/作用 (必填)

  • TF 模型/图迁移到 950 平台时,StridedSliceAssign 是常用的切片写算子(变量按区间更新、序列/特征局部改写等场景),缺失会阻塞整图迁移
  • 与库上已有 StridedSliceAssignV2 / StridedSliceAssignD 形成完整系列:D 版(属性形式)覆盖老芯片,V2/本算子(tensor 输入形式)覆盖 TF 原生构图
  • Atlas A2/A3 已有支持通路(canndev const2attr 融合 pass 转 StridedSliceAssignD 走 built-in),本需求补齐 950 平台的确定性 SIMT 实现

三、设计方案 (必填)

3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)
  • GEIR 图模式调用:REG_OP(StridedSliceAssign) 在算子 op_graph/strided_slice_assign_proto.h 内注册(带 OPS_PROTO_DEF_STRIDEDSLICEASSIGN 宏隔离),TF 框架经适配器构图接入
  • 不提供 aclnn 直调接口
  • Atlas A2/A3:经 canndev const2attr pass 转为 StridedSliceAssignD 走 TBE built-in,无需本实现
3.2 总体设计
3.2.1 算子支持的数据类型
输入/输出 dtype
var(输入/输出,ref) float16 / float32 / bfloat16 / int32 / int16
input_value 与 var 相同 dtype(proto dtype 组合约束)
begin / end / strides int64(IndexNumberType,1 维常量输入)

属性:begin_mask / end_mask / ellipsis_mask / new_axis_mask / shrink_axis_mask(int,默认 0)

3.2.2 host侧设计
  • InferShape:输出 var shape 恒等于输入 var shape,直接透传;unknown rank / unknown shape 放行透传;静态校验 var dimNum∈[1,8]、begin/end/strides 为等长 1 维、input_value rank 合法性等(不读值);不注册 InferDataType(输出 dtype 由 proto dtype 组合推导,input_value==var 一致性兜底在 tiling)
  • Tiling
    • 值依赖:读取 begin/end/strides 常量值(TilingInputsDataDependency 声明)
    • mask 展开:稀疏 spec 展开为稠密 begin/end/strides(语义对齐 TF ValidateStridedSliceOp),校验 stride>0、最内维 stride==1、shrink 轴越界等
    • 切片形状计算 + input_value shape==切片 final shape 校验(无广播)
    • 核数两步法(按行/按元素取优,单核最少 1024 元素/32 行),DCACHE 固定预留 32KB
    • 外部输入校验日志统一使用库上 OP_LOGE_FOR_INVALID_*
  • Shape/Dtype 推导(graph):op_graph 内 graph_infer 独立通道
3.2.3 kernel侧设计
  • SIMT(arch35)实现,纯数据搬运:按 tiling 计算的切片区间,将 input_value 元素写入 var 对应位置,其余区域保留
  • 确定性实现,无原子操作/无并发写冲突(切片区间无重叠)
  • 位级精确(搬运算子,无算术,无精度损失,全 dtype binary_equal)
3.3 支持硬件
  • Ascend 950PR / 950DT(本实现,SIMT arch35)
  • Atlas A3 训练/推理系列、Atlas A2 训练/推理系列(经 canndev const2attr 转 StridedSliceAssignD,built-in 通路)

3.4 算子约束限制

  • 不支持广播:input_value shape 必须严格等于切片 final shape
  • strides 各维必须 > 0(不支持负步长反向切片);最内维 stride 必须为 1
  • shrink 轴 stride 必须为 1,且下标不得越界
  • ellipsis_mask 至多一个置位;稀疏 spec 长度 < 32
  • var dimNum ∈ [1, 8],不支持空 var
  • dtype 暂不支持 int8/int64/double/bool 等(proto 注册精确 dtype 集合:float16/float32/bfloat16/int32/int16)
  • 不提供 aclnn 接口,仅图模式调用

💡 备注(选填)

  • Golden 基准:TF tf.raw_ops.StridedSliceAssign(ref 语义经 tf.compat.v1.Variable(use_resource=False) 承载);TTK 黑盒 150 + 白盒 198 + 网络 13 全量通过(361/361)
  • Atlas A2/A3 支持机制:canndev const2attr_fusion_pass.ccREGISTER_CONST2ATTR("StridedSliceAssignD").OriginOpType("StridedSliceAssign")
likedislike
Xxuejinghui成员
8 天前 添加了label:requirement
Xxuejinghui成员
8 天前 修改了issue 的描述
Xxuejinghui成员
8 天前 修改标题为 “[Requirement|需求建议]: 950新增StridedSliceAssign算子”,原标题为“[Requirement|需求建议]: 950新增Col2ImV2算子”
xuejinghui成员
8 天前 评论:

/assign

likedislike
CANN-robotCANN-robot成员
8 天前 将 xuejinghui 设为负责人
CANN-robotCANN-robot成员
8 天前 关闭了 issue
CANN-robotCANN-robot成员
8 天前 添加了label:resolved