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 | 属性 |
|
FLOAT | - |
| alpha2 | 属性 |
|
FLOAT | - |
| alpha3 | 属性 |
|
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算子。 |