已开启
[Feature]: 【社区任务】MaskedScatter算子贡献 #3310
Tream创建于  6月12日
Tream
Tream
6月12日 创建

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

Backgroud(背景信息)

需求背景

需求来源

基于 TBE MaskedScatter 算子历史版本,使用 Ascend C 进行改造与适配。改造目标是保持 TBE 语义和关键调度逻辑一致,并通过 aclnn 直调方式完成自定义 experimental 算子的编译、安装和 example 验证。

MaskedScatter 的语义为:按一维扁平顺序扫描 xmask,当 mask[i] 为 false 时输出 y[i] = x[i];当 mask[i] 为 true 时,按 true mask 出现顺序从 updates 中取值写入 y[i]

等价伪代码如下:

int64_t updateIndex = 0;
for (int64_t i = 0; i < numElemX; ++i) {
    if (mask[i] != 0) {
        y[i] = updates[updateIndex];
        updateIndex++;
    } else {
        y[i] = x[i];
    }
}

TBE 源码分析

通过对 TBE 内置 MaskedScatter 算子源码(masked_scatter.pyMaskedScatter 类、masked_scatter() 顶层函数、task_schedule()calc_updates_start()compute())进行逐行分析,当前支持的能力与核心逻辑如下。

TBE 算子源码路径:

${ASCEND_INSTALL_PATH}/opp/built-in/op_impl/ai_core/tbe/impl/ops_legacy/dynamic/masked_scatter.py

算子原型路径:

${ASCEND_INSTALL_PATH}/opp/built-in/op_proto/inc/

算子信息库路径:

${ASCEND_INSTALL_PATH}/opp/built-in/op_impl/ai_core/tbe/kernel/config/ascend910b/ops_legacy/masked_scatter.json

1. 支持的数据类型

TBE masked_scatter.py 的 dtype 字节表支持如下类型:

数据类型 字节大小 说明
float16 2 正常处理
float32 4 正常处理
int64 8 TBE 字节表包含
int32 4 正常处理
uint8 1 正常处理
int8 1 正常处理
bool 1 mask 使用 bool,TBE 内部按 int8 Tensor 读写
int16 2 正常处理
bfloat16 2 TBE 初始化时映射为 float16 处理

当前 Ascend C experimental 算子注册支持:

数据类型 字节大小 说明
float16 2 支持
float32 4 支持
uint8 1 支持
int8 1 支持
int16 2 支持
int32 4 支持
bool 1 支持
bfloat16 2 支持

TBE 整体流程图

流程图源码文件:01_tbe_overall.png

TBE task_schedule 流程图

流程图源码文件:02_tbe_task_schedule.png

TBE calc_updates_start 流程图

流程图源码文件:03_tbe_calc_updates_start.png

TBE compute 流程图

流程图源码文件:04_tbe_compute.png

需求分析

外部组件依赖

不涉及新增外部组件依赖。

本算子依赖 CANN 基础能力:

组件 作用
Ascend C kernel API GM/UB Tensor、DataCopyPad、Cast、WholeReduceSum、PipeBarrier 等 kernel 侧能力
op_host tiling API host 侧 dtype/shape 校验、tiling data 写入、block dim 设置
GE op 注册框架 算子原型、infer shape、AICore config 注册
aclnn 调用框架 example 中通过两段式 aclnn API 调用 custom 算子

需求模块设计

算子原型

名称 类别 dtype format shape 介绍
x 输入 float16 / float32 / uint8 / int8 / int16 / int32 / bool / bfloat16 ND all 原始输入张量
mask 输入 bool ND 与 x 相同 控制替换位置的 bool mask
updates 输入 与 x 相同 ND all 替换值来源
y 输出 与 x 相同 ND 与 x 相同 masked scatter 后的输出张量

属性

MaskedScatter 当前无属性。

算子支持型号

Atlas A2 训练系列产品 / Atlas 800I A2 推理产品(ascend910b)。

需求详细设计

使能方式

上层框架 涉及的框架勾选
TF训练/推理
Pytorch训练/推理
ATC推理
Aclnn直调
OPAT调优
SGAT子图切分

需求总体设计

host侧设计方案

MaskedScatter host 侧负责输出 shape 推导、输入合法性校验、tiling data 写入和 block dim 设置。

1) tiling 参数

Ascend C 使用与 TBE tiling_gm 一致的 4 个字段:

struct MaskedScatterTilingData {
    int64_t numElemX;
    int64_t numElemMask;
    int64_t numElemUpdates;
    int64_t tilingCoreNum;
};

host 侧写入规则:

numElemX = input x storage shape size
numElemMask = numElemX
numElemUpdates = input updates storage shape size
tilingCoreNum = PlatformAscendC.GetCoreNumAiv()

2) 分核策略

host 侧设置:

context->SetBlockDim(static_cast<uint32_t>(coreNum));

具体逻辑 task 长度、逻辑 task 数、每核 task 数在 kernel Init()Process() 中计算。该设计与 TBE 一致:host 只传入 tiling_core_num,kernel 内按 TASK_ALIGNMAX_VEC_PROCESS_NUM 计算分段。

Ascend C Host Tiling 流程图

流程图源码文件:05_ascend_host_tiling.png

kernel侧设计方案

Kernel 侧执行 InitProcess 两个阶段。

1) Kernel 入口

masked_scatter.cpp 中入口逻辑如下:

REGISTER_TILING_DEFAULT(MaskedScatterTilingData);
GET_TILING_DATA_WITH_STRUCT(MaskedScatterTilingData, tilingData, tiling);
NsMaskedScatter::MaskedScatter<DTYPE_X> op;
op.Init(x, mask, updates, y, &tilingData);
op.Process();

当前 schMode 仅作为模板参数声明存在,实际 kernel 使用单一路径,语义对应 TBE Python 版本的通用实现。

2) GetTaskBlockSize

Ascend C 与 TBE 保持一致:

int64_t blockSizeAligned = MAX_VEC_PROCESS_NUM;
if (tilingCoreNum_ != 0) {
    int64_t blockSize = numElemX_ / tilingCoreNum_;
    blockSizeAligned = AlignDiv(blockSize, MAX_VEC_PROCESS_NUM);
}
if (numElemX_ >= tilingCoreNum_ * TASK_ALIGN) {
    blockSizeAligned = TASK_ALIGN;
}
if (numElemX_ <= MAX_VEC_PROCESS_NUM) {
    blockSizeAligned = MAX_VEC_PROCESS_NUM;
}
return blockSizeAligned;

这里的 TASK_ALIGN = 4096MAX_VEC_PROCESS_NUM = 64 与 TBE 常量保持一致。

3) Process

每个 AI Core 只处理分配给自己的逻辑 task:

  1. aicoreIdx = GetBlockIdx()
  2. 计算本核 coreTaskNum
  3. 计算本核之前已经分配的 preCoreTaskNum
  4. coreTaskNum > 0,先调用:
taskUpdatesStart = CalcUpdatesStart(preCoreTaskNum * alignedElemPerCore_);
  1. 逐 task 调用:
taskUpdatesStart = Compute((preCoreTaskNum + taskIdx) * alignedElemPerCore_,
                           taskUpdatesStart);

该逻辑与 TBE task_schedule() 一致:每核只在开头统计一次前缀 mask true 数,本核内部 task 通过 Compute 返回值继续推进 updatesStart

4) CalcUpdatesStart

Ascend C 的 mask 前缀计数流程与 TBE 对齐,核心常量为:

constexpr int64_t TASK_ALIGN = 4096;
constexpr int64_t COUNT_REDUCE_LEN = TASK_ALIGN;

5) Compute

Compute(inputOffset, updatesStart) 处理一个逻辑 task。

Ascend C Kernel 流程图

1. Kernel 入口与 Init 流程图

流程图源码文件:06_ascend_kernel_entry_init.png

2. Process 分核调度流程图

流程图源码文件:07_ascend_process_schedule.png

3. CalcUpdatesStart 流程图

流程图源码文件:08_ascend_calc_updates_start.png

4. Compute 流程图

流程图源码文件:09_ascend_compute.png09_ascend_compute.png

支持的芯片版本 涉及勾选
香橙派OrangePi AIpro
Atlas 200I/500 A2推理产品
Atlas 800I/T A2
Atlas A2训练系列产品

算子约束限制

  • xmasky shape 必须完全相同。
  • mask dtype 必须为 bool。
  • xupdatesy dtype 必须完全相同。
  • 当前 experimental Ascend C 算子未注册 int64,即使 TBE dtype 字节表包含 int64
  • 当前 kernel 按扁平一维顺序处理所有元素,不区分原始 rank。
  • 当前实现没有显式校验 updates 元素数量必须大于等于 mask true 数;当当前 task 可读 updates 不足时,updatesOffsetUb < numRemainUpdates 会阻止继续写入,后续位置保留原始 x 值。

特性交叉分析可维可测分析

精度标准/性能标准

验收标准 描述(不涉及说明原因) 标准来源
精度标准 与 TBE 版本按元素比对一致 历史 TBE 对标
分核标准 task 长度、逻辑 task 分配、每核前缀 updates 起点与 TBE 一致 TBE 源码对齐
性能标准 大输入 task 长度为 4096,mask 前缀计数每 4096 个元素 reduce 一次 TBE 源码对齐

关联的 Issue

暂无。

文档更新

本文档。

类型标签

Origin(信息来源)

社区任务

Benefit / Necessity (价值/作用)

Design(设计方案)

likedislike
TreamTream
6月12日 关联了pull request:[社区任务]MaskedScatter算子
TreamTream
6月12日 关联了pull request:[社区任务]MaskedScatter算子
oscillatedoscillated成员
6月12日 将 TreamTik 设为负责人
tangweiwei2成员
7月27日 评论:

@Tream 您好,该社区任务自 6 月 12 日以来已有一段时间未更新,想跟进一下当前进展:

  1. 基于 Ascend C 改造 MaskedScatter 算子的开发进度如何?aclnn 直调方式的编译和 example 验证是否已完成?
  2. 开发过程中是否有阻塞问题需要协助?

欢迎同步最新状态,方便及时推进。感谢!

likedislike
Tream
Tream
7月27日 评论:

@Tream 您好,该社区任务自 6 月 12 日以来已有一段时间未更新,想跟进一下当前进展:

  1. 基于 Ascend C 改造 MaskedScatter 算子的开发进度如何?aclnn 直调方式的编译和 example 验证是否已完成?
  2. 开发过程中是否有阻塞问题需要协助?

欢迎同步最新状态,方便及时推进。感谢!

@tangweiwei2

https://gitcode.com/cann/ops-nn/merge_requests/6001 您好,已合并

likedislike
tangweiwei2成员
7月27日 评论:

@Tream 感谢确认!既然关联 MR !6001 已合并,该社区任务已完成。

麻烦您确认下是否可以关闭本 issue?如还有遗留事项可在关闭时补充说明。再次感谢贡献!

likedislike