graph TD
A["开始"] --> B["Host 校验 shape / dtype 并生成 tiling"]
B --> C["Kernel Init: 读取 GM 地址与 tiling"]
C --> D["CopyIn: 搬入 x tile 到 UB"]
D --> E{"dtype 是否 bfloat16?"}
E -->|是| F["Cast 到 float32"]
E -->|否| G["保持原 dtype"]
F --> H["执行向量取负"]
G --> H
H --> I{"dtype 是否 int8/uint8?"}
I -->|是| J["按 8 bit 回绕语义修正"]
I -->|否| K["保持计算结果"]
J --> L["CopyOut: 写回 y"]
K --> L
L --> M["结束"]
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
一、背景信息 (必填)
参考昇腾版本内置neg算子的 TBE 实现,在昇腾 NPU 上基于 Ascend C 编程语言实现功能一致的算子,完成算子设计、开发、测试全流程工作,验收通过后将算子提交至昇腾算子开源仓。
二、价值/作用 (必填)
与原 TBE 算子核心功能完全对齐,支持原算子对应的所有数据类型、数据格式;特别地需要支持int16、uint8、int64类型输入。说明:当输入类型为 uint8时,其行为和torch.neg一致,torch.neg(uint8) 不会返回负数,而是返回一个 uint8 类型的张量,其值等于 256 - x(对于非零值)或 0(对于零值),即发生“回绕”(wrap-around)或“截断”效果,而不是数学上的取负。
必须实现算子泛化功能,满足各类合法输入场景的计算需求,验收阶段将采用泛化数据进行验收。
三、设计方案 (必填)
host 侧设计
参数解析与校验
x,必选张量。y,必选张量。x/y非空。x.dtype == y.dtype。x.shape == y.shape。BFLOAT16/FLOAT16/FLOAT32/INT8/UINT8/INT16/INT32/INT64,其中INT16/UINT8为本次新增支持类型。tiling 策略
Neg是无属性、无广播的一元逐元素算子,可复用通用逐元素 tiling。tilingKey 规划策略
tilingKey用于区分调度模式和 dtype 模板。bfloat16可使用float32中间计算后转回。int8/uint8需要保留与 TBE 一致的低 bit 整数回绕语义。kernel 侧设计
kernel 侧实现描述
Init和Process两个阶段,其中Process包括搬入、计算、搬出。x到 UB。float16/float32/int16/int32/int64:直接执行向量取负。bfloat16:转换为float32取负,再转回bfloat16。int8:按有符号 8 bit 回绕语义处理,重点验证-128。uint8:按无符号 8 bit 回绕语义处理,等价于(-x) mod 256。AscendC 实现流程图
graph TD A["开始"] --> B["Host 校验 shape / dtype 并生成 tiling"] B --> C["Kernel Init: 读取 GM 地址与 tiling"] C --> D["CopyIn: 搬入 x tile 到 UB"] D --> E{"dtype 是否 bfloat16?"} E -->|是| F["Cast 到 float32"] E -->|否| G["保持原 dtype"] F --> H["执行向量取负"] G --> H H --> I{"dtype 是否 int8/uint8?"} I -->|是| J["按 8 bit 回绕语义修正"] I -->|否| K["保持计算结果"] J --> L["CopyOut: 写回 y"] K --> L L --> M["结束"]AscendC 实现与 TBE 实现差异点和原因
classify + variable_shape + auto_schedule自动处理动态 shape;AscendC 需要在 host tiling 中显式完成元素数统计、分核和 tile 大小规划。int32/int64通过乘以-1实现;扩展int16时也应保持整数回绕语义。int8经过float16中间计算和溢出修正;扩展uint8时同样需要按uint8取模语义修正,避免饱和转换导致结果不一致。bfloat16建议在 AscendC 侧使用float32中间计算再转回,避免不同后端对直接 bf16 取负的支持差异。支持硬件
算子约束限制
Neg没有广播语义,不需要处理多输入 shape 对齐。