已关闭
[Requirement|需求建议]: 新增situ系列算子的torch_extension接口 #4749
guijianwei创建于 26 天前关闭于 21 天前
26 天前 修改了issue 的描述
26 天前 修改了issue 的描述
26 天前 修改标题为 “[Requirement|需求建议]: 新增situ系列算子的torch_extension接口”,原标题为“[Requirement|需求建议]: ”
26 天前 将 guijianwei 设为负责人
25 天前 修改了issue 的描述
guijianwei
24 天前 评论:
24 天前 评论:
torch.ops去掉


guijianwei
24 天前 评论:
24 天前 评论:
A5 : cann_ops_nn.situ_quant


guijianwei
24 天前 评论:
24 天前 评论:
mxFP4原型增加


21 天前 关闭了 issue
21 天前 添加了label:resolved
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=True:gate = x[..., :h],up = x[..., h:]activate_left=False:gate = x[..., h:],up = x[..., :h]SiTU计算:
situ_a=β⋅tanh(βgate)⋅σ(gate)
当
linear_beta > 0时:up=linear_beta⋅tanh(linear_betaup)
输出:
y=situ_a⋅up
输出
y的dim轴大小为x的一半,其他维度与x相同。2. situ_glu_grad
situ_glu的反向梯度算子。给定前向输出梯度grad_y与前向输入x,按SiTU门控线性单元公式对x求导,返回输入梯度grad_x。记前向中间量:
tg=tanh(βgate),sg=σ(gate),situ_a=β⋅tg⋅sg
当
linear_beta > 0时:up′=linear_beta⋅tanh(up/linear_beta),否则 up′=up。乘积法则:
grad_situ_a=grad_y⋅up′,grad_up′=grad_y⋅situ_a
gate梯度:
grad_gate=grad_situ_a⋅sg⋅[(1−tg2)+β⋅tg⋅(1−sg)]
up梯度:
linear_beta > 0时:grad_up=grad_up′⋅(1−tanh2(up/linear_beta))linear_beta ≤ 0时:grad_up=grad_up′按
activate_left拼接grad_gate、grad_up为grad_x(shape与x一致)。3. dequant_situ_quant
将反量化(Dequant)、Situ激活函数和量化(Quant)三个操作融合为一个算子,减少中间数据的读写开销。支持INT8、INT32和BF16三种输入数据类型。
反量化:
group_index分组选择对应行)Situ激活:同
situ_glu。量化:
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(imax∣Vi∣)⌋−emax
mxscale=2shared_exp(E8M0格式)
y=cast_to_fp8(Vi/mxscale)
其中
emax为目标FP8数据类型的最大正则数指数位:E4M3FN为8,E5M2为15。Origin(信息来源)
kimi k3大模型训练场景,SiTU激活路径涉及多个小算子串联(切分、tanh、sigmoid、乘法、量化等),存在大量小算子launch开销、GM读写开销和中间tensor materialization开销。通过算子融合降低上述开销,提升端到端训练性能。
Benefit / Necessity (价值/作用)
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参数说明
dim维度会被均分为gate与up两部分。支持1-8维,dim维度为偶数。[-x.dim(), x.dim()-1],默认-1。1.0。0.0。x时gate是否为前半部分。True表示gate为前半、up为后半;False表示gate为后半、up为前半。默认True。返回值说明
x一致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参数说明
y的梯度。数据类型需与x一致。shape与x相同,但dim维度大小为x.shape[dim] // 2。dim维度会被均分为gate与up两部分。支持1-8维,dim维度为偶数。[-x.dim(), x.dim()-1],默认-1。需与前向situ_glu一致。situ_glu保持一致。不能为0,默认1.0。situ_glu保持一致。默认0.0。x时gate是否为前半部分,需与前向situ_glu保持一致。默认True。返回值说明
x一致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)参数说明
H必须为偶数。INT8输入支持2-8维;INT32/BF16输入仅支持2维。[H]或[1];INT32时必选,shape为[H]、[1, H]或[groupNum, H](group_index不为空时);BF16时不允许输入。[N]或[N, 1];INT8/BF16时不允许输入。[H]或[1];INT32时shape与weight_scale一致;BF16时不允许输入。quant_mode="static"时必选,shape为[H/2]或[1];quant_mode="dynamic"时可选,作为smooth scale加权;INT32/BF16时不允许输入。quant_mode="static"时可选,shape为[H/2]或[1];quant_mode="dynamic"时不允许输入;INT32/BF16时不允许输入。[groupNum],每个元素表示对应专家分组处理的行数;INT8/BF16时不允许输入。4.0。25.0。x的前半部分(True)还是后半部分(False)。默认True。"static"和"dynamic"。INT32/BF16输入时仅支持"dynamic"。默认"dynamic"。返回值说明
x.shape[:-1] + [H/2],其中H = x.shape[-1]。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)参数说明
2H必须为偶数。支持1-7维。1.0。0.0。x的前半部分(True)还是后半部分(False)。默认False。36=FP8_E4M3FN(4位指数+3位尾数,emax=8),35=FP8_E5M2(5位指数+2位尾数,emax=15)。默认36。返回值说明
x.shape[:-1] + [H],其中H = x.shape[-1] / 2。数据类型由dst_type决定。x.shape[:-1] + [ceil(H/64), 2]。产品支持情况汇总
约束说明
dim维度必须为偶数;beta不能为0;dim非 last 维时若尾部维度组合不满足半行32B对齐会回退到Long-H路径,功能正确但性能略降。beta不能为0;quant_mode="static"时quant_scale必选;INT32/BF16输入时quant_mode必须为"dynamic"。beta必须大于0;dst_type只支持36(E4M3FN)或35(E5M2);axis固定为-1。