文件最后提交记录最后更新时间
27 天前
27 天前
27 天前
26 天前
26 天前
27 天前
2 个月前
10 天前
README

Celu

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品
Atlas 训练系列产品

功能说明

  • 算子功能:Celu(Continuously Differentiable Exponential Linear Unit)是一种连续可微的激活函数,对正输入按系数线性缩放,对负输入返回指数衰减曲线,兼容ONNX的Celu算子。

  • 计算公式

    y={α3×xif x≥0α1×(exp⁡(xα2)−1)if x<0y = \begin{cases} \alpha_3 \times x & \text{if } x \geq 0 \\ \alpha_1 \times \left(\exp\left(\frac{x}{\alpha_2}\right) - 1\right) & \text{if } x < 0 \end{cases}

    其中:

    • x >= 0时,输出为线性缩放alpha3 * x(默认alpha3=1时即恒等y=x,与ONNX Celu一致)
    • x < 0时,输出为指数衰减形式alpha1 * (exp(x / alpha2) - 1)
  • 精度保护:针对exp操作的溢出风险,对负区输入执行Mins截断保护(FP16截断阈值为11.0,FP32截断阈值为87.0),防止指数运算溢出。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入 公式中的输入x,任意形状的张量。 FLOAT16、FLOAT ND
alpha1 属性
  • 负区指数曲线的缩放系数,对应公式中的alpha1。
  • 默认值为1.0。
FLOAT -
alpha2 属性
  • 负区指数曲线的斜率系数,对应公式中的alpha2,必须大于0。
  • 默认值为1.0。
FLOAT -
alpha3 属性
  • 正区的固定输出值,对应公式中的alpha3。
  • 默认值为1.0。
FLOAT -
y 输出 公式中的输出y,与输入x形状相同的张量。 FLOAT16、FLOAT ND

约束说明

  • 本算子仅支持Ascend 950PR/Ascend 950DT(DAV_3510/arch35),不支持其他芯片。
  • 输入tensor支持的数据类型为FLOAT16和FLOAT,输出与输入数据类型一致。
  • 输入tensor的格式必须为ND。
  • alpha2属性必须大于0,否则可能导致除零或数值不稳定。
  • 输入支持任意shape(包括0元素空tensor),算子内部已处理空tensor边界情况。

调用说明

调用方式 调用样例 说明
图模式调用 test_geir_celu 通过算子IR构图方式调用Celu算子。