已关闭
[Requirement|需求建议]: 新增situ系列算子的torch_extension接口 #4749
guijianwei创建于  26 天前关闭于  21 天前
guijianwei
guijianwei
26 天前 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

situ系列算子为SiTU(Scalable Interpolated Tanh Units)相关的融合算子族,涵盖激活、反量化激活量化融合、激活MX量化融合等场景。本issue新增四个situ系列算子的torch_extension接口,各算子功能如下:

1. situ_glu

SiTU门控线性单元(SiTU Gated Linear Unit)激活函数。对输入张量 x 沿指定维度 dim 切分为门控(gate)与上路径(up)两半,按SiTU公式计算输出。

x 基于 dim 进行合轴,合轴后维度为 [pre, cut]cut 必须为偶数,令 h = cut // 2。根据 activate_left 切分:

  • activate_left=Truegate = x[..., :h]up = x[..., h:]
  • activate_left=Falsegate = x[..., h:]up = x[..., :h]

SiTU计算:

situ_a=βtanh(gateβ)σ(gate)situ\_a = \beta \cdot \tanh\left(\frac{gate}{\beta}\right) \cdot \sigma(gate)

linear_beta > 0 时:

up=linear_betatanh(uplinear_beta)up = linear\_beta \cdot \tanh\left(\frac{up}{linear\_beta}\right)

输出:

y=situ_aupy = situ\_a \cdot up

输出 ydim 轴大小为 x 的一半,其他维度与 x 相同。

2. situ_glu_grad

situ_glu 的反向梯度算子。给定前向输出梯度 grad_y 与前向输入 x,按SiTU门控线性单元公式对 x 求导,返回输入梯度 grad_x

记前向中间量:

tg=tanh(gateβ),sg=σ(gate),situ_a=βtgsgt_g = \tanh\left(\frac{gate}{\beta}\right), \quad s_g = \sigma(gate), \quad situ\_a = \beta \cdot t_g \cdot s_g

linear_beta > 0 时:up=linear_betatanh(up/linear_beta)up' = linear\_beta \cdot \tanh(up / linear\_beta),否则 up=upup' = up

乘积法则:

grad_situ_a=grad_yup,grad_up=grad_ysitu_agrad\_situ\_a = grad\_y \cdot up', \quad grad\_up' = grad\_y \cdot situ\_a

gate梯度:

grad_gate=grad_situ_asg[(1tg2)+βtg(1sg)]grad\_gate = grad\_situ\_a \cdot s_g \cdot \left[(1 - t_g^2) + \beta \cdot t_g \cdot (1 - s_g)\right]

up梯度:

  • linear_beta > 0 时:grad_up=grad_up(1tanh2(up/linear_beta))grad\_up = grad\_up' \cdot (1 - \tanh^2(up / linear\_beta))
  • linear_beta ≤ 0 时:grad_up=grad_upgrad\_up = grad\_up'

activate_left 拼接 grad_gategrad_upgrad_x(shape与 x 一致)。

3. dequant_situ_quant

将反量化(Dequant)、Situ激活函数和量化(Quant)三个操作融合为一个算子,减少中间数据的读写开销。支持INT8、INT32和BF16三种输入数据类型。

反量化:

  • INT8:dequantOut=Float(x)×weight_scale+biasdequantOut = Float(x) \times weight\_scale + bias
  • INT32:dequantOut=Float(x)×weight_scale×activation_scale+biasdequantOut = Float(x) \times weight\_scale \times activation\_scale + bias(可选按 group_index 分组选择对应行)
  • BF16:dequantOut=Float(x)dequantOut = Float(x)

Situ激活:同 situ_glu

量化:

  • static模式:y=Round(situOut/quant_scale+quant_offset)y = Round(situOut / quant\_scale + quant\_offset)scale 为无意义值
  • dynamic模式:scale=Max(situOut,axis=1)/127scale = Max(|situOut|, axis=-1) / 127y=Round(situOut/scale)y = Round(situOut / scale);当 quant_scale 不为空时先进行smooth scale加权

4. situ_mx_quant

将Situ激活函数与动态MX(Microscaling)量化融合为一个算子,输入为BF16,输出为FP8(E4M3FN/E5M2)+ E8M0 scale。

Situ激活:同 situ_glu

MX量化(OCP算法),按 blocksize = 32 分组:

shared_exp=log2(maxiVi)emaxshared\_exp = \lfloor \log_2(\max_i |V_i|) \rfloor - emax

mxscale=2shared_exp(E8M0格式)mxscale = 2^{shared\_exp} \quad (E8M0格式)

y=cast_to_fp8(Vi/mxscale)y = cast\_to\_fp8(V_i / mxscale)

其中 emax 为目标FP8数据类型的最大正则数指数位:E4M3FN为8,E5M2为15。

Origin(信息来源)

kimi k3大模型训练场景,SiTU激活路径涉及多个小算子串联(切分、tanh、sigmoid、乘法、量化等),存在大量小算子launch开销、GM读写开销和中间tensor materialization开销。通过算子融合降低上述开销,提升端到端训练性能。

Benefit / Necessity (价值/作用)

  • situ_glu / situ_glu_grad:将SiTU门控线性单元的切分、tanh、sigmoid、乘法等小算子融合为1个融合算子,减少算子launch和中间tensor读写开销。
  • dequant_situ_quant:将反量化、Situ激活、量化三段计算融合为1个融合算子,消除中间FP32 tensor的GM读写和materialization开销。
  • situ_mx_quant:将Situ激活和MX量化融合为1个融合算子,减少中间BF16/FP8数据的GM读写开销。

Design(设计方案)

为上述四个situ系列算子新增torch_extension接口,封装对应的aclnn API,支持单算子模式和TorchAir图模式调用。


1. situ_glu

torch.ops.cann_ops_nn.situ_glu(
    x: Tensor,
    *,
    dim: int = -1,
    beta: float = 1.0,
    linear_beta: float = 0.0,
    activate_left: bool = True
) -> Tensor

参数说明

参数名 类型 必选/可选 数据类型 数据格式 非连续Tensor支持 说明
x Tensor 必选 float16 / float32 / bfloat16 ND SiTU输入,dim 维度会被均分为gate与up两部分。支持1-8维,dim 维度为偶数。
dim int 可选 int64 - - 切分维度,取值范围 [-x.dim(), x.dim()-1],默认 -1
beta float 可选 float32 - - SiTU门控部分的缩放系数,控制tanh非线性强度。不能为0,默认 1.0
linear_beta float 可选 float32 - - up路径线性tanh的缩放系数。大于0时对up施加有界化变换;小于等于0时up直接透传。默认 0.0
activate_left bool 可选 bool - - 切分 x 时gate是否为前半部分。True 表示gate为前半、up为后半;False 表示gate为后半、up为前半。默认 True

返回值说明

返回值 类型 数据类型 数据格式 说明
y Tensor x 一致 ND SiTU激活结果。shape与 x 相同,但 dim 维度大小为 x.shape[dim] // 2

2. situ_glu_grad

torch.ops.cann_ops_nn.situ_glu_grad(
    grad_y: Tensor,
    x: Tensor,
    *,
    dim: int = -1,
    beta: float = 1.0,
    linear_beta: float = 0.0,
    activate_left: bool = True
) -> Tensor

参数说明

参数名 类型 必选/可选 数据类型 数据格式 非连续Tensor支持 说明
grad_y Tensor 必选 float16 / float32 / bfloat16 ND 前向输出 y 的梯度。数据类型需与 x 一致。shape与 x 相同,但 dim 维度大小为 x.shape[dim] // 2
x Tensor 必选 float16 / float32 / bfloat16 ND 前向输入,dim 维度会被均分为gate与up两部分。支持1-8维,dim 维度为偶数。
dim int 可选 int64 - - 切分维度,取值范围 [-x.dim(), x.dim()-1],默认 -1。需与前向 situ_glu 一致。
beta float 可选 float32 - - SiTU门控部分的缩放系数,需与前向 situ_glu 保持一致。不能为0,默认 1.0
linear_beta float 可选 float32 - - up路径线性tanh的缩放系数,需与前向 situ_glu 保持一致。默认 0.0
activate_left bool 可选 bool - - 切分 x 时gate是否为前半部分,需与前向 situ_glu 保持一致。默认 True

返回值说明

返回值 类型 数据类型 数据格式 说明
grad_x Tensor x 一致 ND 输入 x 的梯度。shape与 x 完全一致。

3. dequant_situ_quant

torch.ops.cann_ops_nn.dequant_situ_quant(
    x: Tensor,
    *,
    weight_scale: Tensor = None,
    activation_scale: Tensor = None,
    bias: Tensor = None,
    quant_scale: Tensor = None,
    quant_offset: Tensor = None,
    group_index: Tensor = None,
    beta: float = 4.0,
    linear_beta: float = 25.0,
    activate_left: bool = True,
    quant_type: str = "dynamic"
) -> (Tensor, Tensor)

参数说明

参数名 类型 必选/可选 数据类型 数据格式 非连续Tensor支持 说明
x Tensor 必选 int8 / int32 / bfloat16 ND 输入数据,最后一维 H 必须为偶数。INT8输入支持2-8维;INT32/BF16输入仅支持2维。
weight_scale Tensor 可选 float32 ND weight的反量化scale。INT8时必选,shape为 [H][1];INT32时必选,shape为 [H][1, H][groupNum, H]group_index 不为空时);BF16时不允许输入。
activation_scale Tensor 可选 float32 ND 激活函数的反量化scale。INT32时必选,shape为 [N][N, 1];INT8/BF16时不允许输入。
bias Tensor 可选 float32 ND 反量化的bias。INT8时shape为 [H][1];INT32时shape与 weight_scale 一致;BF16时不允许输入。
quant_scale Tensor 可选 float32 ND 量化的scale。quant_mode="static" 时必选,shape为 [H/2][1]quant_mode="dynamic" 时可选,作为smooth scale加权;INT32/BF16时不允许输入。
quant_offset Tensor 可选 float32 ND 量化的offset。quant_mode="static" 时可选,shape为 [H/2][1]quant_mode="dynamic" 时不允许输入;INT32/BF16时不允许输入。
group_index Tensor 可选 int64 ND MoE分组的group_index。INT32时可选,shape为 [groupNum],每个元素表示对应专家分组处理的行数;INT8/BF16时不允许输入。
beta float 可选 float32 - - Situ激活的beta参数。不能为0,默认 4.0
linear_beta float 可选 float32 - - Situ激活的linear_beta参数。大于0时启用up分支的tanh加权,小于等于0时不启用。默认 25.0
activate_left bool 可选 bool - - gate在 x 的前半部分(True)还是后半部分(False)。默认 True
quant_type str 可选 str - - 量化模式,支持 "static""dynamic"。INT32/BF16输入时仅支持 "dynamic"。默认 "dynamic"

返回值说明

返回值 类型 数据类型 数据格式 说明
y Tensor int8 ND 量化输出。shape为 x.shape[:-1] + [H/2],其中 H = x.shape[-1]
y_scale Tensor float32 ND 动态量化scale。INT8输入时shape为 x.shape[:-1];INT32/BF16输入时shape为 [N]quant_mode="static" 时为无意义值。

4. situ_mx_quant

torch.ops.cann_ops_nn.situ_mx_quant(
    x: Tensor,
    beta: float = 1.0,
    linear_beta: float = 0.0,
    activate_left: bool = False,
    dst_type: int = 36
) -> (Tensor, Tensor)

参数说明

参数名 类型 必选/可选 数据类型 数据格式 非连续Tensor支持 说明
x Tensor 必选 bfloat16 ND 输入数据,最后一维 2H 必须为偶数。支持1-7维。
beta float 可选 float32 - - Situ激活的beta参数。必须大于0,默认 1.0
linear_beta float 可选 float32 - - Situ激活的linear_beta参数。大于0时启用up分支的tanh加权,小于等于0时不启用。默认 0.0
activate_left bool 可选 bool - - gate在 x 的前半部分(True)还是后半部分(False)。默认 False
dst_type int 可选 int64 - - 目标FP8类型:36=FP8_E4M3FN(4位指数+3位尾数,emax=8),35=FP8_E5M2(5位指数+2位尾数,emax=15)。默认 36

返回值说明

返回值 类型 数据类型 数据格式 说明
y Tensor float8_e4m3fn / float8_e5m2 ND FP8量化输出。shape为 x.shape[:-1] + [H],其中 H = x.shape[-1] / 2。数据类型由 dst_type 决定。
y_scale Tensor float8_e8m0 ND MX量化的E8M0 scale。shape为 x.shape[:-1] + [ceil(H/64), 2]

产品支持情况汇总

算子 Ascend 950PR/950DT Atlas A3 训练/推理 Atlas A2 训练/推理
situ_glu
situ_glu_grad
dequant_situ_quant ×
situ_mx_quant × ×

约束说明

  • 四个算子的torch_extension接口均支持单算子模式和TorchAir图模式调用。
  • 输入Tensor均需为NPU Tensor,数据格式仅支持ND。
  • 四个算子均不支持非连续Tensor和空Tensor。
  • situ_glu / situ_glu_grad:dim 维度必须为偶数;beta 不能为0;dim 非 last 维时若尾部维度组合不满足半行32B对齐会回退到Long-H路径,功能正确但性能略降。
  • dequant_situ_quant:beta 不能为0;quant_mode="static"quant_scale 必选;INT32/BF16输入时 quant_mode 必须为 "dynamic"
  • situ_mx_quant:beta 必须大于0;dst_type 只支持36(E4M3FN)或35(E5M2);axis 固定为 -1
  • 四个算子均默认支持确定性计算。
likedislike
guijianweiguijianwei
26 天前 修改了issue 的描述
guijianweiguijianwei
26 天前 修改了issue 的描述
guijianweiguijianwei
26 天前 修改标题为 “[Requirement|需求建议]: 新增situ系列算子的torch_extension接口”,原标题为“[Requirement|需求建议]: ”
yuning_chenyuning_chen成员
26 天前 将 guijianwei 设为负责人
guijianweiguijianwei
25 天前 修改了issue 的描述
guijianwei
guijianwei
24 天前 评论:

torch.ops去掉

likedislike
guijianwei
guijianwei
24 天前 评论:

A5 : cann_ops_nn.situ_quant

likedislike
guijianwei
guijianwei
24 天前 评论:

mxFP4原型增加

likedislike
CANN-robotCANN-robot成员
21 天前 关闭了 issue
CANN-robotCANN-robot成员
21 天前 添加了label:resolved