Power

📄 查看源码

本算子仅提供 GE IR 通路,不提供aclnn接口。在计算图中以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_ZEROS scale·power == 0shift==0power>0 y = 0
    BROADCAST_SCALAR power==0,或scale==0power≠0 y = bcastVal(host端预算pow(shift,power)1.0NaN+inf之一)
    LINEAR scale·power≠0power==1 y = x·scale + shift
    SQUARE scale·power≠0power==2 y = (x·scale + shift)^2
    CUBE scale·power≠0power==3 y = (x·scale + shift)^3
    GENERIC_POW_POS scale·power≠0power>0power∉{1,2,3} 通用幂运算,base==0时输出0
    GENERIC_POW_NEG scale·power≠0power<0 通用幂运算,base==0时输出+inf

    其中power∈{0,1,2,3}IsClose判等容差为atol=1e-8rtol=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 < 0power为整数 `(-1)^power · exp(power · log( base
base < 0power非整数 NaN 实数域未定义,host预置negScalar = NaN,乘加传递得NaN
base == 0power > 0 0 GENERIC_POW_POS / ALL_ZEROS
base == 0power < 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
  • 形状一致:不支持广播,xy的shape必须完全一致;0 维标量在host端被视为[1]
  • 平台:当前仅生成ascend950平台的kernel二进制;其它平台编译时不下发本算子。
  • 接口形式:本算子提供aclnn单算子接口,仅作为图算子节点存在,需通过GE IR或图编译器接入;如需host直调(kernel-launch)场景请直接复用本仓的kernel源码而非aclnn API。
  • 属性默认值:缺省时power=1.0scale=1.0shift=0.0,等价于identity映射(y = x),此时走LINEAR路径。
  • 确定性计算:本算子默认确定性实现,相同输入与属性下每次执行输出一致。

调用方式

暂不支持aclnn接口方式调用本算子,推荐接入方式:

  1. 图算子方式:在通过GE IR / ATC构图时,将Power作为节点加入计算图,由图编译器自动完成tiling、kernel选择与下发。属性power / scale / shift通过节点attr注入。
  2. Kernel直调方式:参考op_kernel/power_apt.cpppower(...)入口与op_host/arch35/power_tiling_arch35.cppTiling4Power(...)回调,自行构造PowerTilingData并通过<<<>>>直接发射kernel。tilingKey由GET_TPL_TILING_KEY(schMode, culType, dType)编码,三段含义见op_kernel/arch35/power_struct.h

相关文档

  • Power README:文件清单、目录结构与对外承诺。