文件最后提交记录最后更新时间
2 个月前
2 个月前
29 天前
1 个月前
2 个月前
28 天前
README

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接口方式调用本算子,推荐接入方式:

图算子方式:在通过GE IR / ATC构图时,将Power作为节点加入计算图,由图编译器自动完成tiling、kernel选择与下发。属性power / scale / shift通过节点attr注入。