已合并
【社区任务】AsinGrad算子设计文档 #547
【社区任务】AsinGrad算子设计文档 #547
已合并
从已删除 :asin-grad-design-doc合入到cann/cann-ops-competitionsmaster
共 1 个文件变更+353-0
| @@ -0,0 +1,353 @@ | |||
| 1 | +# AsinGrad算子AscendC设计方案 | ||
| 2 | + | ||
| 3 | +## 需求背景 | ||
| 4 | + | ||
| 5 | +### 需求来源 | ||
| 6 | + | ||
| 7 | +依据CANN社区算子迁移任务要求,对历史TBE实现的AsinGrad算子进行AscendC重写,使其能够在Atlas A2/Ascend 910B环境下通过CANNJudge功能、精度和性能验证。 | ||
| 8 | + | ||
| 9 | +### 背景介绍 | ||
| 10 | + | ||
| 11 | +AsinGrad用于计算反正弦函数Asin的反向梯度,数学表达式为: | ||
| 12 | + | ||
| 13 | +```text | ||
| 14 | +z = dy / sqrt(1 - y * y) | ||
| 15 | +``` | ||
| 16 | + | ||
| 17 | +其中`y`为前向输入,`dy`为上游梯度,`z`为输出梯度。原TBE实现采用DSL向量接口完成`vmul`、`vmuls`、`vadds`、`vsqrt`、`vdiv`等逐元素计算;本方案使用AscendC在AICore侧显式实现数据搬运、向量计算和结果写回。 | ||
| 18 | + | ||
| 19 | +AsinGrad算子TBE实现路径和相关API路径如下: | ||
| 20 | + | ||
| 21 | +- AsinGrad算子实现路径为:`/usr/local/Ascend/ascend-toolkit/latest/opp/built-in/op_impl/ai_core/tbe/impl` | ||
| 22 | +- AsinGrad算子实现中的API路径:`/usr/local/Ascend/ascend-toolkit/latest/python/site-packages/tbe/dsl` | ||
| 23 | + | ||
| 24 | + | ||
| 25 | +### AsinGrad算子现状分析 | ||
| 26 | + | ||
| 27 | +通过对AsinGrad历史TBE版本的功能分析,当前核心能力如下: | ||
| 28 | + | ||
| 29 | +- 输入`y`、`dy`形状必须一致,输出`z`与输入保持相同shape。 | ||
| 30 | +- TBE历史版本主要支持`float16`、`float32`;CANNJudge任务模板扩展到`float16`、`float32`、`bfloat16`。 | ||
| 31 | +- `float16`路径在硬件支持时先转换为`float32`计算,再转换回`float16`,以降低`1 - y * y`接近0时的精度损失。 | ||
| 32 | +- 计算过程为纯逐元素操作,不涉及reduce、broadcast、shape变换或workspace依赖。 | ||
| 33 | + | ||
| 34 | +| 名称 | 类别 | dtype | format | shape | 介绍 | | ||
| 35 | +| --- | --- | --- | --- | --- | --- | | ||
| 36 | +| y | 输入 | fp16/fp32/bf16 | ND | all | Asin前向输入 | | ||
| 37 | +| dy | 输入 | fp16/fp32/bf16 | ND | 同y | 上游梯度 | | ||
| 38 | +| z | 输出 | fp16/fp32/bf16 | ND | 同y | 输出梯度 | | ||
| 39 | + | ||
| 40 | +AsinGrad算子TBE版本的整体流程图如下图所示: | ||
| 41 | +```mermaid | ||
| 42 | +flowchart TD | ||
| 43 | + A["AsinGrad算子入口 asin_grad(y, dy, z, kernel_name)"] --> B["获取输入shape和dtype"] | ||
| 44 | + B --> C["check_shape(y), check_shape(dy)"] | ||
| 45 | + C --> D{"shape_y == shape_dy ?"} | ||
| 46 | + D -- 否 --> E["抛出shape不一致错误"] | ||
| 47 | + D -- 是 --> F["refine_shape_axes(shape_y / shape_dy)"] | ||
| 48 | + | ||
| 49 | + F --> G["check_dtype: float16 / float32"] | ||
| 50 | + G --> H{"dtype_y == dtype_dy ?"} | ||
| 51 | + H -- 否 --> I["抛出dtype不一致错误"] | ||
| 52 | + H -- 是 --> J["创建TVM placeholder: data_y, data_dy"] | ||
| 53 | + | ||
| 54 | + J --> K["调用 asin_grad_compute(data_y, data_dy, z)"] | ||
| 55 | + K --> L["生成计算表达式 res"] | ||
| 56 | + L --> M["with tvm.target.cce()"] | ||
| 57 | + M --> N["auto_schedule(res)"] | ||
| 58 | + N --> O["build(schedule, config)"] | ||
| 59 | + O --> P["生成TBE算子二进制"] | ||
| 60 | +``` | ||
| 61 | + | ||
| 62 | +TBE Compute计算流程图 | ||
| 63 | +```mermaid | ||
| 64 | +flowchart TD | ||
| 65 | + A["asin_grad_compute(y, dy, z)"] --> B["记录原始dtype = y.dtype"] | ||
| 66 | + B --> C{"dtype == float16 且支持float32向量API ?"} | ||
| 67 | + | ||
| 68 | + C -- 是 --> D["cast_to(y, float32)"] | ||
| 69 | + D --> E["cast_to(dy, float32)"] | ||
| 70 | + C -- 否 --> F["保持原dtype计算"] | ||
| 71 | + | ||
| 72 | + E --> G["vmul(y, y)"] | ||
| 73 | + F --> G | ||
| 74 | + | ||
| 75 | + G --> H["vmuls(data, -1)"] | ||
| 76 | + H --> I["vadds(data, 1)"] | ||
| 77 | + I --> J["vsqrt(1 - y*y)"] | ||
| 78 | + J --> K["vdiv(dy, sqrt_res)"] | ||
| 79 | + K --> L{"原始dtype == float16 ?"} | ||
| 80 | + L -- 是 --> M["cast_to(res, float16)"] | ||
| 81 | + L -- 否 --> N["直接输出res"] | ||
| 82 | + M --> O["返回结果z"] | ||
| 83 | + N --> O | ||
| 84 | +``` | ||
| 85 | + | ||
| 86 | +接口具体实现图: | ||
| 87 | +```mermaid | ||
| 88 | +flowchart LR | ||
| 89 | + Y["输入 y"] --> C1{"float16路径?"} | ||
| 90 | + DY["输入 dy"] --> C1 | ||
| 91 | + | ||
| 92 | + C1 -- 是 --> CY["cast_to(y, float32)"] | ||
| 93 | + C1 -- 是 --> CDY["cast_to(dy, float32)"] | ||
| 94 | + C1 -- 否 --> Y2["y保持原dtype"] | ||
| 95 | + C1 -- 否 --> DY2["dy保持原dtype"] | ||
| 96 | + | ||
| 97 | + CY --> MUL["vmul(y, y)"] | ||
| 98 | + Y2 --> MUL | ||
| 99 | + | ||
| 100 | + MUL --> MULS["vmuls(data, -1)"] | ||
| 101 | + MULS --> ADDS["vadds(data, 1)"] | ||
| 102 | + ADDS --> SQRT["vsqrt(1 - y*y)"] | ||
| 103 | + | ||
| 104 | + CDY --> DIV["vdiv(dy, sqrt_res)"] | ||
| 105 | + DY2 --> DIV | ||
| 106 | + SQRT --> DIV | ||
| 107 | + | ||
| 108 | + DIV --> C2{"原始dtype为float16?"} | ||
| 109 | + C2 -- 是 --> OUT16["cast_to(res, float16)"] | ||
| 110 | + C2 -- 否 --> OUT32["res保持float32"] | ||
| 111 | + OUT16 --> Z["输出 z"] | ||
| 112 | + OUT32 --> Z | ||
| 113 | +``` | ||
| 114 | + | ||
| 115 | +数学表达式与TBE API对应关系: | ||
| 116 | +```mermaid | ||
| 117 | +flowchart TD | ||
| 118 | + A["目标公式: z = dy / sqrt(1 - y*y)"] --> B["y*y"] | ||
| 119 | + B --> B1["tbe.vmul(y, y)"] | ||
| 120 | + | ||
| 121 | + A --> C["-(y*y)"] | ||
| 122 | + C --> C1["tbe.vmuls(data, -1)"] | ||
| 123 | + | ||
| 124 | + A --> D["1 - y*y"] | ||
| 125 | + D --> D1["tbe.vadds(data, 1)"] | ||
| 126 | + | ||
| 127 | + A --> E["sqrt(1 - y*y)"] | ||
| 128 | + E --> E1["tbe.vsqrt(num_to_vrsqrt, 1)"] | ||
| 129 | + | ||
| 130 | + A --> F["dy / sqrt(...)"] | ||
| 131 | + F --> F1["tbe.vdiv(dy, vsqrt_res)"] | ||
| 132 | + | ||
| 133 | + A --> G["float16精度策略"] | ||
| 134 | + G --> G1["输入float16先cast_to float32计算"] | ||
| 135 | + G1 --> G2["最终cast_to float16输出"] | ||
| 136 | +``` | ||
| 137 | + | ||
| 138 | +dtype分支逻辑图: | ||
| 139 | +```mermaid | ||
| 140 | +flowchart TD | ||
| 141 | + A["输入dtype"] --> B{"dtype == float16 ?"} | ||
| 142 | + B -- 是 --> C{"平台支持float32向量API ?"} | ||
| 143 | + C -- 是 --> D["y, dy 转float32"] | ||
| 144 | + D --> E["按float32执行 vmul/vmuls/vadds/vsqrt/vdiv"] | ||
| 145 | + E --> F["结果转回float16"] | ||
| 146 | + F --> G["输出z"] | ||
| 147 | + | ||
| 148 | + C -- 否 --> H["按float16执行计算"] | ||
| 149 | + H --> G | ||
| 150 | + | ||
| 151 | + B -- 否: float32 --> I["按float32执行计算"] | ||
| 152 | + I --> G | ||
| 153 | +``` | ||
| 154 | + | ||
| 155 | + | ||
| 156 | +## 需求分析 | ||
| 157 | + | ||
| 158 | +### 外部组件依赖 | ||
| 159 | + | ||
| 160 | +不涉及外部组件依赖。 | ||
| 161 | + | ||
| 162 | +### 内部适配模块 | ||
| 163 | + | ||
| 164 | +适配Aclnn接口和图模式调用。 | ||
| 165 | + | ||
| 166 | + | ||
| 167 | +### 需求模块设计 | ||
| 168 | + | ||
| 169 | +#### 算子原型 | ||
| 170 | + | ||
| 171 | +- 原型设计:`AsinGrad(y, dy) -> z`。 | ||
| 172 | +- 相关约束:`y`、`dy`、`z` dtype保持一致;shape保持一致;format支持ND。 | ||
| 173 | +- Atlas A2训练系列产品/Atlas 800I A2推理产品支持`float16`、`float32`、`bfloat16`。 | ||
| 174 | + | ||
| 175 | +## 需求详细设计 | ||
| 176 | + | ||
| 177 | +### 使能方式 | ||
| 178 | + | ||
| 179 | +| 上层框架 | 涉及的框架勾选 | | ||
| 180 | +| --- | --- | | ||
| 181 | +| TF训练/推理 | | | ||
| 182 | +| Pytorch训练/推理 | | | ||
| 183 | +| ATC推理 | √ | | ||
| 184 | +| Aclnn直调 | √ | | ||
| 185 | +| OPAT调优 | | | ||
| 186 | + | ||
| 187 | +### 需求总体设计 | ||
| 188 | + | ||
| 189 | +#### host侧设计 | ||
| 190 | + | ||
| 191 | +AsinGrad计算过程不依赖原始维度的逐维信息,host侧将输入视为一维连续向量,仅关注总元素个数、数据类型、UB大小和AIV核数。Host侧通过`GetInputShape`获取`y`、`dy`的shape,检查输入输出shape一致;通过`GetInputDesc`/`GetOutputDesc`获取dtype并检查一致性。 | ||
| 192 | + | ||
| 193 | +**分核策略:** | ||
| 194 | + | ||
| 195 | +- 优先使用可用AIV核,将`totalNum`按`coreNum`近似均分。 | ||
| 196 | +- 单核长度按cache line对应元素数向上对齐得到`blockFactor`。 | ||
| 197 | +- `usedCoreNum`由`totalNum`和`blockFactor`反推,避免启动空核。 | ||
| 198 | + | ||
| 199 | +**数据分块和内存优化策略:** | ||
| 200 | + | ||
| 201 | +- 根据UB空间和kernel侧实际buffer数量计算`ubFactor`。 | ||
| 202 | +- FP32路径使用`y`、`dy`、`z`双缓冲队列以及`tmp`临时buffer。 | ||
| 203 | +- FP16/BF16路径使用低精度输入输出队列,并额外申请float32中间buffer完成升精度计算。 | ||
| 204 | +- 单次UB循环仅处理真实`currentNum`元素,尾块不对无效padding区域执行计算。 | ||
| 205 | + | ||
| 206 | +**tilingData规划:** | ||
| 207 | + | ||
| 208 | +| 字段 | 含义 | | ||
| 209 | +| --- | --- | | ||
| 210 | +| totalNum | 展平后的总元素数 | | ||
| 211 | +| blockFactor | 每个核处理的最大元素数 | | ||
| 212 | +| ubFactor | 单次UB循环处理的元素数 | | ||
| 213 | + | ||
| 214 | +**tilingKey规划:** | ||
| 215 | + | ||
| 216 | +| schMode | 数据类型 | kernel分支 | | ||
| 217 | +| --- | --- | --- | | ||
| 218 | +| 0 | FP16 | `AsinGrad<half>` | | ||
| 219 | +| 1 | FP32 | `AsinGrad<float>` | | ||
| 220 | +| 2 | BF16 | `AsinGrad<bfloat16_t>` | | ||
| 221 | + | ||
| 222 | +`schMode`使用2 bit声明,保证三种取值均可表达。 | ||
| 223 | + | ||
| 224 | +**数据检测:** | ||
| 225 | + | ||
| 226 | +- shape约束:`y`、`dy`、`z`必须同shape。 | ||
| 227 | +- dtype约束:`y`、`dy`、`z`必须同dtype,支持`DT_FLOAT16`、`DT_FLOAT`、`DT_BF16`。 | ||
| 228 | +- format约束:支持ND格式。 | ||
| 229 | + | ||
| 230 | +#### kernel侧设计 | ||
| 231 | + | ||
| 232 | +Kernel侧分为`Init`和`Process`两个阶段,其中`Process`包括数据搬入`CopyIn`、计算`Compute`、搬出`CopyOut`三个阶段。 | ||
| 233 | + | ||
| 234 | +- `Init`阶段:根据`GetBlockIdx`获取当前核编号,结合`blockFactor`计算本核GM起始偏移和`blockLength`;初始化输入输出`GlobalTensor`以及UB队列/临时buffer。 | ||
| 235 | +- `CopyIn`阶段:使用`DataCopyPad`将`y`、`dy`从GM搬入UB队列。 | ||
| 236 | +- `Compute`阶段:FP32路径直接在float上计算;FP16/BF16路径先Cast到float32,完成`Mul`、`Muls`、`Adds`、`Sqrt`、`Div`后再Cast回原dtype。 | ||
| 237 | +- `CopyOut`阶段:将`z`从UB搬回GM,只写回当前tile真实元素数。 | ||
| 238 | +- 循环策略:每个核内部按`ubFactor`循环处理,尾块`currentNum`小于`ubFactor`时仅对真实元素执行向量计算,避免无效padding区域影响结果。 | ||
| 239 | + | ||
| 240 | +AscendC核心计算流程如下: | ||
| 241 | + | ||
| 242 | +```text | ||
| 243 | +CopyIn: yGM, dyGM -> yLocal, dyLocal | ||
| 244 | +Compute: tmp = yLocal * yLocal | ||
| 245 | + tmp = 1.0 - tmp | ||
| 246 | + tmp = sqrt(tmp) | ||
| 247 | + zLocal = dyLocal / tmp | ||
| 248 | +CopyOut: zLocal -> zGM | ||
| 249 | +``` | ||
| 250 | +Ascend C整体执行流程图: | ||
| 251 | +```mermaid | ||
| 252 | +flowchart TD | ||
| 253 | + A["调用 AsinGrad(y, dy)"] --> B["Host侧算子注册与校验"] | ||
| 254 | + B --> C["InferShape: z.shape = y.shape"] | ||
| 255 | + C --> D["TilingFunc"] | ||
| 256 | + D --> E["获取输入shape / dtype / format"] | ||
| 257 | + E --> F{"shape和dtype是否合法?"} | ||
| 258 | + F -- 否 --> G["返回GRAPH_FAILED"] | ||
| 259 | + F -- 是 --> H["获取平台信息: AIV核数 / UB大小"] | ||
| 260 | + | ||
| 261 | + H --> I["计算totalNum"] | ||
| 262 | + I --> J["计算blockFactor: 单核处理元素数"] | ||
| 263 | + J --> K["计算usedCoreNum并设置BlockDim"] | ||
| 264 | + K --> L["计算ubFactor: 单次UB处理元素数"] | ||
| 265 | + L --> M["设置TilingData: totalNum, blockFactor, ubFactor"] | ||
| 266 | + M --> N["根据dtype设置TilingKey"] | ||
| 267 | + | ||
| 268 | + N --> O["启动AICore Kernel asin_grad<schMode>"] | ||
| 269 | + O --> P["Kernel Init"] | ||
| 270 | + P --> Q["Kernel Process循环"] | ||
| 271 | + Q --> R["CopyIn: GM -> UB"] | ||
| 272 | + R --> S["Compute: dy / sqrt(1 - y*y)"] | ||
| 273 | + S --> T["CopyOut: UB -> GM"] | ||
| 274 | + T --> U{"本核数据处理完成?"} | ||
| 275 | + U -- 否 --> Q | ||
| 276 | + U -- 是 --> V["Kernel结束"] | ||
| 277 | +``` | ||
| 278 | + | ||
| 279 | +Kernel侧Init流程图: | ||
| 280 | +```mermaid | ||
| 281 | +flowchart TD | ||
| 282 | + A["Kernel入口 asin_grad<schMode>"] --> B["读取TilingData"] | ||
| 283 | + B --> C["根据schMode选择模板类型"] | ||
| 284 | + C --> D["创建AsinGrad<T>对象"] | ||
| 285 | + D --> E["Init(y, dy, z, tilingData)"] | ||
| 286 | + | ||
| 287 | + E --> F["blockIdx = GetBlockIdx()"] | ||
| 288 | + F --> G["offset = blockFactor * blockIdx"] | ||
| 289 | + G --> H["remain = totalNum - offset"] | ||
| 290 | + H --> I{"remain <= 0 或 ubFactor <= 0 ?"} | ||
| 291 | + I -- 是 --> J["blockLength = 0, 返回"] | ||
| 292 | + I -- 否 --> K["blockLength = min(remain, blockFactor)"] | ||
| 293 | + K --> L["ubLength = ubFactor"] | ||
| 294 | + L --> M["设置yGM / dyGM / zGM GlobalTensor"] | ||
| 295 | + M --> N["初始化yQueue / dyQueue / zQueue"] | ||
| 296 | + N --> O{"T == float ?"} | ||
| 297 | + O -- 是 --> P["初始化tmpBuf"] | ||
| 298 | + O -- 否 --> Q["初始化yFp32Buf / dyFp32Buf / tmpBuf / zFp32Buf"] | ||
| 299 | +``` | ||
| 300 | + | ||
| 301 | +Kernel侧Process流程图: | ||
| 302 | +```mermaid | ||
| 303 | +flowchart TD | ||
| 304 | + A["Process()"] --> B{"blockLength <= 0 ?"} | ||
| 305 | + B -- 是 --> C["直接返回"] | ||
| 306 | + B -- 否 --> D["loopCount = ceil(blockLength / ubLength)"] | ||
| 307 | + | ||
| 308 | + D --> E["for i in loopCount"] | ||
| 309 | + E --> F["currentNum = 当前tile真实元素数"] | ||
| 310 | + F --> G["CopyIn(i, currentNum)"] | ||
| 311 | + G --> H["Compute(currentNum)"] | ||
| 312 | + H --> I["CopyOut(i, currentNum)"] | ||
| 313 | + I --> J{"所有tile完成?"} | ||
| 314 | + J -- 否 --> E | ||
| 315 | + J -- 是 --> K["Process结束"] | ||
| 316 | +``` | ||
| 317 | + | ||
| 318 | + | ||
| 319 | +### 支持硬件 | ||
| 320 | + | ||
| 321 | +| 支持的芯片版本 | 涉及勾选 | | ||
| 322 | +| --- | --- | | ||
| 323 | +| 香橙派OrangePi AIpro | | | ||
| 324 | +| Atlas 200I/500 A2推理产品 | | | ||
| 325 | +| Atlas 800I/T A2 | √ | | ||
| 326 | + | ||
| 327 | +### 算子约束限制 | ||
| 328 | + | ||
| 329 | +不支持广播;不支持`y`、`dy` dtype不一致;不支持非ND格式;不支持复数、整数、bool等非浮点输入。输入值通常应位于`[-1, 1]`范围内,超出定义域时结果遵循硬件`sqrt/div`对NaN或Inf的处理行为。 | ||
| 330 | + | ||
| 331 | +## 特性交叉分析 | ||
| 332 | + | ||
| 333 | +AsinGrad为逐元素数学算子,与reduce、broadcast、layout转换、原子写、多输出等特性无交叉依赖。算子不使用workspace,不依赖外部状态;可作为图模式或aclnn直调路径中的基础反向算子。 | ||
| 334 | + | ||
| 335 | +## 可维可测分析 | ||
| 336 | + | ||
| 337 | +### 精度标准/性能标准 | ||
| 338 | + | ||
| 339 | +| 验收标准 | 描述(不涉及说明原因) | 标准来源 | | ||
| 340 | +| --- | --- | --- | | ||
| 341 | +| 精度标准 | 不低于TBE版本。 | | | ||
| 342 | +| 性能标准 | 不低于TBE版本。 | | | ||
| 343 | + | ||
| 344 | +测试建议覆盖以下场景: | ||
| 345 | + | ||
| 346 | +- 基础shape:一维、二维、多维、标量shape。 | ||
| 347 | +- 数据类型:`float16`、`float32`、`bfloat16`。 | ||
| 348 | +- 规模场景:小shape、单核、满核、多tile、大shape和尾块。 | ||
| 349 | +- 数值场景:普通随机值、接近0、接近±1、`dy`包含正负值。 | ||
| 350 | + | ||
| 351 | +### 兼容性分析 | ||
| 352 | + | ||
| 353 | +新算子,不涉及历史二进制兼容。 | ||