已关闭
[Requirement|需求建议]: ops-nn仓新增ClippedSwigluGrad的torch_npu接口 #4754
shilulu创建于  25 天前关闭于  21 天前
shilulu
25 天前 创建

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

Backgroud(背景信息)

ClippedSwiglu为带截断的SwiGLU(Swish Gated Linear Unit)激活函数,正向公式为 y=(clamp(x1,l,l)+bias)σ(αmin(x0,l))y = (\text{clamp}(x_1, -l, l) + \text{bias}) \cdot \sigma(\alpha \cdot \min(x_0, l))。本issue新增ClippedSwiglu反向算子的torch_extension接口,用于计算输入x的梯度。

正向计算流程:

  1. 将x基于dim进行合轴,合轴后维度为[pre, cut],cut必须为偶数,令 H=cut/2H = \text{cut} / 2
  2. 根据interleaved切分x为 x0x_0x1x_1 两部分:
    • interleaved=True(奇偶切分):x0=x[:,::2]x_0 = x[:, ::2]x1=x[:,1::2]x_1 = x[:, 1::2]
    • interleaved=False(前后切分):x0=x[:,:H]x_0 = x[:, :H]x1=x[:,H:]x_1 = x[:, H:]
  3. 重算正向中间量:
    • A=min(x0,l)A = \min(x_0, l)(上界截断)
    • B=clamp(x1,l,l)B = \text{clamp}(x_1, -l, l)(双向截断)
    • s=σ(αA)=11+eαAs = \sigma(\alpha \cdot A) = \frac{1}{1 + e^{-\alpha \cdot A}}
  4. 正向输出:y=(B+bias)sy = (B + \text{bias}) \cdot s

反向计算流程:

  1. 合轴切分(同正向)、GroupIndex处理(可选)。
  2. Clamp反向传播掩码计算:
    • maskA=I(x0l)\text{maskA} = \mathbb{I}(x_0 \leq l)
    • maskB=I(lx1l)\text{maskB} = \mathbb{I}(-l \leq x_1 \leq l)
  3. SwiGLU反向梯度计算:
    • gradx0=grady(B+bias)s(1+αA(1s))maskA\text{grad}_{x_0} = \text{grad}_y \cdot (B + \text{bias}) \cdot s \cdot (1 + \alpha \cdot A \cdot (1 - s)) \cdot \text{maskA}
    • gradx1=gradyAsmaskB\text{grad}_{x_1} = \text{grad}_y \cdot A \cdot s \cdot \text{maskB}
  4. 梯度拼接输出(按interleaved散回,与正向切分方式对应)。
  5. GroupIndex置零处理(可选):超出validRows的行梯度置零。

Origin(信息来源)

大规模MOE场景容易出现激活异常值,导致MOE训练不稳定。因此DeepSeekV4、MiniMax M3等模型中MOE部分的Swiglu额外做了clamp操作。带clamp操作的Swiglu操作由小算子拼接,进行融合后预计能在整网带来5%+的性能提升,目前ClippedSwiglu正向算子已经实现,需要补齐训练场景下的反向算子实现。

Benefit / Necessity(价值/作用)

  • 将ClippedSwiglu反向传播中的clamp掩码计算、sigmoid重算、梯度乘法、散回拼接等多个小算子融合为1个融合算子,减少算子launch和中间tensor读写开销。
  • 支持group_index分组(MoE场景),仅处理有效行,无效行梯度自动置零。
  • 支持950PR/950DT、A3、A2三平台,A5采用RegBase VF(MicroAPI)实现。
  • 无workspace(纯向量算子),减少内存分配开销。

Design(设计方案)

为ClippedSwigluGrad算子新增torch_extension接口,封装对应的aclnn API,支持单算子模式和TorchAir图模式调用。

Python接口

cann_ops_nn.clipped_swiglu_grad(
    grad_y: Tensor,
    x: Tensor,
    *,
    group_index: Tensor = None,
    dim: int = -1,
    alpha: float = 1.702,
    limit: float = 7.0,
    bias: float = 1.0,
    interleaved: bool = True,
    clamp_mode: int = 0
) -> Tensor

参数说明

参数名 类型 必选/可选 数据类型 数据格式 非连续Tensor支持 说明
grad_y Tensor 必选 float16 / float32 / bfloat16 ND 正向输出y的梯度。dim维度为x的一半,其他维度与x一致。数据类型需与x一致。
x Tensor 必选 float16 / float32 / bfloat16 ND 正向输入x。dim维度会被均分为x_0和x_1两部分。支持1-8维,dim维度为偶数。
group_index Tensor 可选 int64 ND 分组索引。1维,元素个数不超过8192。第i个元素表示第i组处理的batch数。不提供时全部行参与计算。
dim int 可选 int64 - - 合轴和切分的维度序号,取值范围[-x.dim(), x.dim()-1],默认-1。
alpha float 可选 float32 - - SwiGLU激活系数,控制sigmoid非线性强度,默认1.702。
limit float 可选 float32 - - 截断门限值,必须大于0,默认7.0。
bias float 可选 float32 - - 线性计算偏差,默认1.0。
interleaved bool 可选 bool - - 切分方式。True=奇偶切分,False=前后切分。默认True。
clamp_mode int 可选 int64 - - clamp位置控制。0=clamp在silu之前,1=clamp移至silu之后(忽略alpha和bias)。当前只支持0,默认0。

返回值说明

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

产品支持情况

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

约束说明

  • 接口支持单算子模式和TorchAir图模式调用。
  • 输入Tensor均需为NPU Tensor,数据格式仅支持ND。
  • 不支持非连续Tensor和空Tensor。
  • x的dim维度必须能被2整除。
  • grad_y的dim维度必须为x的dim维度的一半,其他维度必须一致。
  • grad_y的数据类型必须与x一致。
  • limit必须大于0。
  • clamp_mode当前只支持取值0。
  • 无workspace(纯向量算子)。
  • 支持确定性计算。
likedislike
yuning_chenyuning_chen成员
25 天前 将 shilulu 设为负责人
Sshilulu
25 天前 修改了issue 的描述
Sshilulu
25 天前 修改了issue 的描述
shilulu
24 天前 评论:

修改意见:
torch.ops. -> cann_ops_nn.clipped_swiglu_grad
公式描述需要与参数一致

likedislike
Sshilulu
21 天前 修改了issue 的描述
CANN-robotCANN-robot成员
21 天前 关闭了 issue
CANN-robotCANN-robot成员
21 天前 添加了label:resolved
Sshilulu
21 天前 关联了pull request:clipped_swiglu_grad torch extension bug fix