Power
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | √ |
| Atlas 推理系列产品 | √ |
| Atlas 训练系列产品 | × |
功能说明
-
算子功能:对输入张量
x逐元素执行线性变换后再做幂运算。 -
计算公式:
yi=(scale⋅xi+shift)powery_i = \big(\text{scale} \cdot x_i + \text{shift}\big)^{\text{power}}
其中power / scale / shift均为标量属性,在host端tiling阶段完成全部分支决策与可预计算的常数折叠,kernel端按tilingKey路由到对应的DAG,无运行时分支。
-
等价PyTorch表达:
base = scale * x + shift y = torch.pow(base, power) -
实现要点:
类别 内容 计算模板 Elementwise( ElewiseBaseTiling+ElementwiseSch+DAGSch)计算精度 fp16/bf16在kernel内cast到fp32计算后回写到原dtype分支前移 culType×dtype×schMode三段编码到tilingKey,kernel模板实例化时已选定DAG性能优化 power ∈ {1,2,3}走乘法展开;power ∉ {0,1,2,3}走`exp(power·log( -
算子内部分支(host端
culTypeEnum):culType 触发条件 计算 ALL_ZEROSscale·power == 0且shift==0且power>0y = 0BROADCAST_SCALARpower==0,或scale==0且power≠0y = bcastVal(host端预算pow(shift,power)、1.0、NaN、+inf之一)LINEARscale·power≠0且power==1y = x·scale + shiftSQUAREscale·power≠0且power==2y = (x·scale + shift)^2CUBEscale·power≠0且power==3y = (x·scale + shift)^3GENERIC_POW_POSscale·power≠0且power>0且power∉{1,2,3}通用幂运算, base==0时输出0GENERIC_POW_NEGscale·power≠0且power<0通用幂运算, base==0时输出+inf其中
power∈{0,1,2,3}的IsClose判等容差为atol=1e-8、rtol=1e-5,与math/is_close对齐。
算子参数说明
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| x(Tensor) | 输入 | 幂运算的底数前置量,公式中的x。 | shape与y完全一致;不支持广播。 | FLOAT16、BFLOAT16、FLOAT | ND | 不限制(含0 维标量,标量按 [1] 处理) | × |
| y(Tensor) | 输出 | 幂运算结果,公式中的y。 | dtype与x一致;shape与x一致。 | FLOAT16、BFLOAT16、FLOAT | ND | 与x保持一致 | × |
| power(attr,float) | 属性 | 幂指数,公式中的power。 | 可选,默认 1.0。 |
FLOAT | - | 标量 | - |
| scale(attr,float) | 属性 | 输入线性缩放因子,公式中的scale。 | 可选,默认 1.0。 |
FLOAT | - | 标量 | - |
| shift(attr,float) | 属性 | 输入线性平移量,公式中的shift。 | 可选,默认 0.0。 |
FLOAT | - | 标量 | - |
异常值约定
逐元素遵循下表(base = scale·x + shift,整数power指power == floor(power)且有限):
| 场景 | 输出值 | 说明 |
|---|---|---|
base > 0 |
exp(power · log(base)) |
通用主路径 |
base < 0且power为整数 |
`(-1)^power · exp(power · log( | base |
base < 0且power非整数 |
NaN |
实数域未定义,host预置negScalar = NaN,乘加传递得NaN |
base == 0且power > 0 |
0 |
GENERIC_POW_POS / ALL_ZEROS |
base == 0且power < 0 |
+inf |
GENERIC_POW_NEG / BROADCAST_SCALAR,IEEE 754 语义 |
power == 0(含0^0) |
1 |
按约定,走BROADCAST_SCALAR,host预置bcastVal = 1.0 |
约束说明
- 数据类型:输入
x与输出y必须为FLOAT16/BFLOAT16/FLOAT三者之一,且二者完全相同;tiling在host端会拒绝其它dtype并返回GRAPH_FAILED。 - 形状一致:不支持广播,
x与y的shape必须完全一致;0 维标量在host端被视为[1]。 - 平台:当前仅生成
ascend950平台的kernel二进制;其它平台编译时不下发本算子。 - 接口形式:本算子不提供aclnn单算子接口,仅作为图算子节点存在,需通过GE IR或图编译器接入;如需host直调(kernel-launch)场景请直接复用本仓的kernel源码而非aclnn API。
- 属性默认值:缺省时
power=1.0、scale=1.0、shift=0.0,等价于identity映射(y = x),此时走LINEAR路径。 - 确定性计算:本算子默认确定性实现,相同输入与属性下每次执行输出一致。
调用方式
暂不支持aclnn接口方式调用本算子,推荐接入方式:
图算子方式:在通过GE IR / ATC构图时,将Power作为节点加入计算图,由图编译器自动完成tiling、kernel选择与下发。属性power / scale / shift通过节点attr注入。