已合并
【社区任务】AsinGrad算子设计文档 #547
寧懿.创建于 7月9日
【社区任务】AsinGrad算子设计文档 #547
已合并
寧懿.创建于 7月9日
从已删除 :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+新算子,不涉及历史二进制兼容。