已合并
fixed README.md #4211
littlecc创建于 7月23日
fixed README.md #4211
已合并
littlecc创建于 7月23日
2 个文件变更+142-227
@@ -1,76 +1,158 @@
1-# Power算子1+# Power
2 2 
3-## 1.功能描述3+## 产品支持情况
4 4 
5-逐元素计算:5+| 产品 | 是否支持 |
6+| :----------------------------------------------------------- | :------: |
7+| <term>Ascend 950PR/Ascend 950DT</term> | √ |
8+| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
9+| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
10+| <term>Atlas 200I/500 A2 推理产品</term> | √ |
11+| <term>Atlas 推理系列产品</term> | √ |
12+| <term>Atlas 训练系列产品</term> | × |
6 13 
7-```text14+## 功能说明
8-y = exp(power * log(x * scale + shift))
9-```
10 15 
11-支持的输入数据类型:`fp16` / `bf16` / `fp32`(fp16、bf16在kernel内cast为fp32计算回写)16+- 算子功能:对输入张量`x`逐元素执行线性变换再做幂运算
12 17 
13-## 2.子属性18+- 公式:
14 19 
15-| 属性 | 类型 | 必选 | 默认值 | 含义 |20+$$
16-|---|---|---|---|---|21+y_i = \big(\text{scale} \cdot x_i + \text{shift}\big)^{\text{power}}
17-| `power` | float | 是 | - | 幂指数 |22+$$
18-| `scale` | float | 否 | `1.0` | 输入线性缩放因子 |
19-| `shift` | float | 否 | `0.0` | 输入线性平移量 |
20 23 
21-## 3.实现要点24+ 其中`power` / `scale` / `shift`均为标量属性,在host端tiling阶段完成全部分支决策与可预计算的常数折叠,kernel端按`tilingKey`路由到对应的DAG,无运行时分支。
22 25 
23-- 使用 **Elementwise模板**(`ElewiseBaseTiling` + `ElementwiseSchWithScalar` + `DAGSch`)。26+- 等价PyTorch表达:
24-- 在tiling层完成所有`scale / power / shift`标量计算与分支决策,最终通过`culType`枚举值
25- - dtype编码出`tilingKey`,在kernel中实例化对应的DAG,避免kernel内运行时分支:
26 27 
27-| culType | 计算 |28+ ```python
28-|---|---|29+ base = scale * x + shift
29-| `ALL_ZEROS` | `y = 0` |30+ y = torch.pow(base, power)
30-| `BROADCAST_SCALAR` | `y = bcastVal` (host端预计算`pow(shift, power)`、异常值`+inf`/`NaN`,或`power=0`时的`1.0`) |31+ ```
31-| `LINEAR` (`power=1`) | `y = x*scale + shift` (fused MulAdd语义:`Duplicate(shift) + Axpy(scale)`) |
32-| `SQUARE` (`power=2`) | `base = x*scale + shift; y = base * base` |
33-| `CUBE` (`power=3`) | `base = x*scale + shift; y = base^3` |
34-| `GENERIC_POW_POS` | 通用,`power>0`,零底数映射为0 |
35-| `GENERIC_POW_NEG` | 通用,`power<0`,零底数映射为`+inf` |
36 32 
37-通用分支在kernel内通过`Compare + Select`合并三种取值(正/负/零底数)33+- 实现要点
38 34 
39-```text35+ | 类别 | 内容 |
40-absBase = |base|36+ |---|---|
41-logAbs = log(absBase)37+ | 计算模板 | Elementwise(`ElewiseBaseTiling` + `ElementwiseSch` + `DAGSch`) |
42-rawExp = exp(power * logAbs)38+ | 计算精度 | `fp16`/`bf16`在kernel内cast到`fp32`计算后回写到原dtype |
43-posVal = rawExp // base > 039+ | 分支前移 | `culType` × `dtype` × `schMode`三段编码到`tilingKey`,kernel模板实例化时已选定DAG |
44-negVal = rawExp * negScalar // negScalar = ±1 (整数power)或NaN (非整数power)40+ | 性能优化 | `power {1,2,3}`走乘法展开;`power ∉ {0,1,2,3}`走`exp(power·log(|base|))`通用路径 |
45-zeroVal = 0 / +inf // 由tilingKey区分
46-tmp = base > 0 ? posVal : negVal
47-y = base == 0 ? zeroVal : tmp
48-```
49 41 
50-`isclose`判等方法参考`math/is_close`:`|a-b| <= atol + rtol*|b|``atol=1e-8`,`rtol=1e-5`。42+- 算子内部分支(host端`culTypeEnum`):
51 43 
52-## 4.不支持范围44+ | culType | 触发条件 | 计算 |
45+ |---|---|---|
46+ | `ALL_ZEROS` | `scale·power == 0``shift==0``power>0` | `y = 0` |
47+ | `BROADCAST_SCALAR` | `power==0`,或`scale==0``power≠0` | `y = bcastVal`(host端预算`pow(shift,power)``1.0``NaN``+inf`之一) |
48+ | `LINEAR` | `scale·power≠0``power==1` | `y = x·scale + shift` |
49+ | `SQUARE` | `scale·power≠0``power==2` | `y = (x·scale + shift)^2` |
50+ | `CUBE` | `scale·power≠0``power==3` | `y = (x·scale + shift)^3` |
51+ | `GENERIC_POW_POS` | `scale·power≠0``power>0``power∉{1,2,3}` | 通用幂运算,`base==0`时输出`0` |
52+ | `GENERIC_POW_NEG` | `scale·power≠0``power<0` | 通用幂运算,`base==0`时输出`+inf` |
53 53 
54-- 不支持广播;输入输出shape一致54+ 其中`power∈{0,1,2,3}`的`IsClose`判等容差为`atol=1e-8`、`rtol=1e-5`,`math/is_close`对齐
55-- 输入dtype仅支持`fp16/bf16/fp32`;其它类型在tiling层会直接报错。
56-- 本算子目前仅生成ascend950平台的二进制,不包含aclnn接口模块。
57 55 
58-## 5.目录结构56+## 算子参数说明
59 57 
60-```text58+ <table style="undefined;table-layout: fixed; width: 1280px"><colgroup>
61-power/59+ <col style="width: 140px">
62-├── CMakeLists.txt60+ <col style="width: 100px">
63-├── README.md61+ <col style="width: 280px">
64-├── docs/DESIGN.md62+ <col style="width: 200px">
65-├── op_graph/power_proto.h63+ <col style="width: 220px">
66-├── op_host/64+ <col style="width: 100px">
67-│ ├── power_def.cpp65+ <col style="width: 140px">
68-│ ├── power_infershape.cpp66+ <col style="width: 100px">
69-│ ├── arch35/67+ </colgroup>
70-│ │ ├── power_tiling_arch35.h68+ <thead>
71-│ │ └── power_tiling_arch35.cpp69+ <tr>
72-│ └── config/ascend950/{power_binary.json, power_simplified_key.ini}70+ <th>参数名</th>
73-└── op_kernel/71+ <th>输入/输出</th>
74- ├── power_apt.cpp72+ <th>描述</th>
75- └── arch35/{power_struct.h, power_dag.h}73+ <th>使用说明</th>
76-```74+ <th>数据类型</th>
75+ <th>数据格式</th>
76+ <th>维度(shape)</th>
77+ <th>非连续Tensor</th>
78+ </tr></thead>
79+ <tbody>
80+ <tr>
81+ <td>x(Tensor)</td>
82+ <td>输入</td>
83+ <td>幂运算的底数前置量,公式中的x。</td>
84+ <td>shape与y完全一致;不支持广播。</td>
85+ <td>FLOAT16、BFLOAT16、FLOAT</td>
86+ <td>ND</td>
87+ <td>不限制(含0 维标量,标量按 [1] 处理)</td>
88+ <td>×</td>
89+ </tr>
90+ <tr>
91+ <td>y(Tensor)</td>
92+ <td>输出</td>
93+ <td>幂运算结果,公式中的y。</td>
94+ <td>dtype与x一致;shape与x一致。</td>
95+ <td>FLOAT16、BFLOAT16、FLOAT</td>
96+ <td>ND</td>
97+ <td>与x保持一致</td>
98+ <td>×</td>
99+ </tr>
100+ <tr>
101+ <td>power(attr,float)</td>
102+ <td>属性</td>
103+ <td>幂指数,公式中的power。</td>
104+ <td>可选,默认 <code>1.0</code>。</td>
105+ <td>FLOAT</td>
106+ <td>-</td>
107+ <td>标量</td>
108+ <td>-</td>
109+ </tr>
110+ <tr>
111+ <td>scale(attr,float)</td>
112+ <td>属性</td>
113+ <td>输入线性缩放因子,公式中的scale。</td>
114+ <td>可选,默认 <code>1.0</code>。</td>
115+ <td>FLOAT</td>
116+ <td>-</td>
117+ <td>标量</td>
118+ <td>-</td>
119+ </tr>
120+ <tr>
121+ <td>shift(attr,float)</td>
122+ <td>属性</td>
123+ <td>输入线性平移量,公式中的shift。</td>
124+ <td>可选,默认 <code>0.0</code>。</td>
125+ <td>FLOAT</td>
126+ <td>-</td>
127+ <td>标量</td>
128+ <td>-</td>
129+ </tr>
130+ </tbody></table>
131+ 
132+## 异常值约定
133+ 
134+逐元素遵循下表(`base = scale·x + shift`,整数power指`power == floor(power)`且有限):
135+ 
136+ | 场景 | 输出值 | 说明 |
137+ |---|---|---|
138+ | `base > 0` | `exp(power · log(base))` | 通用主路径 |
139+ | `base < 0``power`为整数 | `(-1)^power · exp(power · log(|base|))` | 由host预置`negScalar = ±1`,kernel端用`Compare + Select`合并 |
140+ | `base < 0``power`非整数 | `NaN` | 实数域未定义,host预置`negScalar = NaN`,乘加传递得NaN |
141+ | `base == 0``power > 0` | `0` | GENERIC_POW_POS / ALL_ZEROS |
142+ | `base == 0``power < 0` | `+inf` | GENERIC_POW_NEG / BROADCAST_SCALAR,IEEE 754 语义 |
143+ | `power == 0`(含`0^0`) | `1` | 按约定,走BROADCAST_SCALAR,host预置`bcastVal = 1.0` |
144+ 
145+## 约束说明
146+ 
147+- **数据类型**:输入`x`与输出`y`必须为`FLOAT16` / `BFLOAT16` / `FLOAT`三者之一,且二者完全相同;tiling在host端会拒绝其它dtype并返回`GRAPH_FAILED`
148+- **形状一致**:不支持广播,`x``y`的shape必须完全一致;0 维标量在host端被视为`[1]`
149+- **平台**:当前仅生成`ascend950`平台的kernel二进制;其它平台编译时不下发本算子。
150+- **接口形式**:本算子**不**提供aclnn单算子接口,仅作为图算子节点存在,需通过GE IR或图编译器接入;如需host直调(kernel-launch)场景请直接复用本仓的kernel源码而非aclnn API。
151+- **属性默认值**:缺省时`power=1.0``scale=1.0``shift=0.0`,等价于identity映射(`y = x`),此时走`LINEAR`路径。
152+- **确定性计算**:本算子默认确定性实现,相同输入与属性下每次执行输出一致。
153+ 
154+## 调用方式
155+ 
156+暂不支持aclnn接口方式调用本算子,推荐接入方式:
157+ 
158+ **图算子方式**:在通过GE IR / ATC构图时,将`Power`作为节点加入计算图,由图编译器自动完成tiling、kernel选择与下发。属性`power` / `scale` / `shift`通过节点attr注入。
@@ -1,167 +0,0 @@
1-# Power
2- 
3-[📄 查看源码](https://gitcode.com/cann/ops-math/tree/master/math/power)
4- 
5-> 本算子仅提供 **GE IR** 通路,**不提供aclnn接口**。在计算图中以`Power`算子节点的形式接入,由图编译器调度执行。
6- 
7-## 产品支持情况
8- 
9-| 产品 | 是否支持 |
10-| :----------------------------------------------------------- | :------: |
11-| <term>Ascend 950PR/Ascend 950DT</term> | √ |
12-| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
13-| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
14-| <term>Atlas 200I/500 A2 推理产品</term> | √ |
15-| <term>Atlas 推理系列产品</term> | √ |
16-| <term>Atlas 训练系列产品</term> | × |
17- 
18-## 功能说明
19- 
20-- 算子功能:对输入张量`x`逐元素执行线性变换后再做幂运算。
21- 
22-- 计算公式:
23- 
24-$$
25-y_i = \big(\text{scale} \cdot x_i + \text{shift}\big)^{\text{power}}
26-$$
27- 
28- 其中`power` / `scale` / `shift`均为标量属性,在host端tiling阶段完成全部分支决策与可预计算的常数折叠,kernel端按`tilingKey`路由到对应的DAG,无运行时分支。
29- 
30-- 等价PyTorch表达:
31- 
32- ```python
33- base = scale * x + shift
34- y = torch.pow(base, power)
35- ```
36- 
37-- 实现要点:
38- 
39- | 类别 | 内容 |
40- |---|---|
41- | 计算模板 | Elementwise(`ElewiseBaseTiling` + `ElementwiseSch` + `DAGSch`) |
42- | 计算精度 | `fp16`/`bf16`在kernel内cast到`fp32`计算后回写到原dtype |
43- | 分支前移 | `culType` × `dtype` × `schMode`三段编码到`tilingKey`,kernel模板实例化时已选定DAG |
44- | 性能优化 | `power ∈ {1,2,3}`走乘法展开;`power ∉ {0,1,2,3}``exp(power·log(|base|))`通用路径 |
45- 
46-- 算子内部分支(host端`culTypeEnum`):
47- 
48- | culType | 触发条件 | 计算 |
49- |---|---|---|
50- | `ALL_ZEROS` | `scale·power == 0``shift==0``power>0` | `y = 0` |
51- | `BROADCAST_SCALAR` | `power==0`,或`scale==0``power≠0` | `y = bcastVal`(host端预算`pow(shift,power)``1.0``NaN``+inf`之一) |
52- | `LINEAR` | `scale·power≠0``power==1` | `y = x·scale + shift` |
53- | `SQUARE` | `scale·power≠0``power==2` | `y = (x·scale + shift)^2` |
54- | `CUBE` | `scale·power≠0``power==3` | `y = (x·scale + shift)^3` |
55- | `GENERIC_POW_POS` | `scale·power≠0``power>0``power∉{1,2,3}` | 通用幂运算,`base==0`时输出`0` |
56- | `GENERIC_POW_NEG` | `scale·power≠0``power<0` | 通用幂运算,`base==0`时输出`+inf` |
57- 
58- 其中`power∈{0,1,2,3}``IsClose`判等容差为`atol=1e-8``rtol=1e-5`,与`math/is_close`对齐。
59- 
60-## 算子参数说明
61- 
62- <table style="undefined;table-layout: fixed; width: 1280px"><colgroup>
63- <col style="width: 140px">
64- <col style="width: 100px">
65- <col style="width: 280px">
66- <col style="width: 200px">
67- <col style="width: 220px">
68- <col style="width: 100px">
69- <col style="width: 140px">
70- <col style="width: 100px">
71- </colgroup>
72- <thead>
73- <tr>
74- <th>参数名</th>
75- <th>输入/输出</th>
76- <th>描述</th>
77- <th>使用说明</th>
78- <th>数据类型</th>
79- <th>数据格式</th>
80- <th>维度(shape)</th>
81- <th>非连续Tensor</th>
82- </tr></thead>
83- <tbody>
84- <tr>
85- <td>x(Tensor)</td>
86- <td>输入</td>
87- <td>幂运算的底数前置量,公式中的x。</td>
88- <td>shape与y完全一致;不支持广播。</td>
89- <td>FLOAT16、BFLOAT16、FLOAT</td>
90- <td>ND</td>
91- <td>不限制(含0 维标量,标量按 [1] 处理)</td>
92- <td>×</td>
93- </tr>
94- <tr>
95- <td>y(Tensor)</td>
96- <td>输出</td>
97- <td>幂运算结果,公式中的y。</td>
98- <td>dtype与x一致;shape与x一致。</td>
99- <td>FLOAT16、BFLOAT16、FLOAT</td>
100- <td>ND</td>
101- <td>与x保持一致</td>
102- <td>×</td>
103- </tr>
104- <tr>
105- <td>power(attr,float)</td>
106- <td>属性</td>
107- <td>幂指数,公式中的power。</td>
108- <td>可选,默认 <code>1.0</code>。</td>
109- <td>FLOAT</td>
110- <td>-</td>
111- <td>标量</td>
112- <td>-</td>
113- </tr>
114- <tr>
115- <td>scale(attr,float)</td>
116- <td>属性</td>
117- <td>输入线性缩放因子,公式中的scale。</td>
118- <td>可选,默认 <code>1.0</code>。</td>
119- <td>FLOAT</td>
120- <td>-</td>
121- <td>标量</td>
122- <td>-</td>
123- </tr>
124- <tr>
125- <td>shift(attr,float)</td>
126- <td>属性</td>
127- <td>输入线性平移量,公式中的shift。</td>
128- <td>可选,默认 <code>0.0</code>。</td>
129- <td>FLOAT</td>
130- <td>-</td>
131- <td>标量</td>
132- <td>-</td>
133- </tr>
134- </tbody></table>
135- 
136-## 异常值约定
137- 
138-逐元素遵循下表(`base = scale·x + shift`,整数power指`power == floor(power)`且有限):
139- 
140- | 场景 | 输出值 | 说明 |
141- |---|---|---|
142- | `base > 0` | `exp(power · log(base))` | 通用主路径 |
143- | `base < 0``power`为整数 | `(-1)^power · exp(power · log(|base|))` | 由host预置`negScalar = ±1`,kernel端用`Compare + Select`合并 |
144- | `base < 0``power`非整数 | `NaN` | 实数域未定义,host预置`negScalar = NaN`,乘加传递得NaN |
145- | `base == 0``power > 0` | `0` | GENERIC_POW_POS / ALL_ZEROS |
146- | `base == 0``power < 0` | `+inf` | GENERIC_POW_NEG / BROADCAST_SCALAR,IEEE 754 语义 |
147- | `power == 0`(含`0^0`) | `1` | 按约定,走BROADCAST_SCALAR,host预置`bcastVal = 1.0` |
148- 
149-## 约束说明
150- 
151-- **数据类型**:输入`x`与输出`y`必须为`FLOAT16` / `BFLOAT16` / `FLOAT`三者之一,且二者完全相同;tiling在host端会拒绝其它dtype并返回`GRAPH_FAILED`
152-- **形状一致**:不支持广播,`x``y`的shape必须完全一致;0 维标量在host端被视为`[1]`
153-- **平台**:当前仅生成`ascend950`平台的kernel二进制;其它平台编译时不下发本算子。
154-- **接口形式**:本算子**不**提供aclnn单算子接口,仅作为图算子节点存在,需通过GE IR或图编译器接入;如需host直调(kernel-launch)场景请直接复用本仓的kernel源码而非aclnn API。
155-- **属性默认值**:缺省时`power=1.0``scale=1.0``shift=0.0`,等价于identity映射(`y = x`),此时走`LINEAR`路径。
156-- **确定性计算**:本算子默认确定性实现,相同输入与属性下每次执行输出一致。
157- 
158-## 调用方式
159- 
160-暂不支持aclnn接口方式调用本算子,推荐接入方式:
161- 
162-1. **图算子方式**:在通过GE IR / ATC构图时,将`Power`作为节点加入计算图,由图编译器自动完成tiling、kernel选择与下发。属性`power` / `scale` / `shift`通过节点attr注入。
163-2. **Kernel直调方式**:参考`op_kernel/power_apt.cpp``power(...)`入口与`op_host/arch35/power_tiling_arch35.cpp``Tiling4Power(...)`回调,自行构造`PowerTilingData`并通过`<<<>>>`直接发射kernel。tilingKey由`GET_TPL_TILING_KEY(schMode, culType, dType)`编码,三段含义见`op_kernel/arch35/power_struct.h`
164- 
165-## 相关文档
166- 
167-- [Power README](../README.md):文件清单、目录结构与对外承诺。