文件最后提交记录最后更新时间
1 个月前
2 个月前
1 个月前
2 个月前
1 个月前
1 个月前
1 个月前
1 个月前
2 个月前
2 个月前
README

SwigluGroupGrad

产品支持情况

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

功能说明

  • 算子功能:完成ClampedSwiglu激活函数的反向梯度计算。从上游梯度grad_y和前向输入x重算clamp mask与sigmoid,输出grad_x与可选grad_weight。

  • 计算公式:

    前向分解:x按hidden维劈半得到gate(g)和up(u);可选clamp产生g̃ = min(c, g)、ũ = clip(u, −c, c);SiLU(g̃) = g̃·σ(g̃);y = SiLU(g̃)·ũ·w_t。

    silu′(g~)=s+f−f⋅ssilu'(g̃) = s + f − f·s

    dg=grad_y⋅silu′(g~)⋅u~⋅wt⋅I(g<c)⋅mrdg = grad\_y \cdot silu'(g̃) \cdot ũ \cdot w_t \cdot I(g < c) \cdot m_r

    du=grad_y⋅f⋅wt⋅I(−c<u<c)⋅mrdu = grad\_y \cdot f \cdot w_t \cdot I(−c < u < c) \cdot m_r

    grad_weight=Σ(grad_y⋅y_origin) along hidden dimgrad\_weight = \Sigma(grad\_y \cdot y\_origin) \text{ along hidden dim}

    其中I为开区间指示函数(边界值时mask=0),m_r为group_index mask,w_t为weight的broadcast。

    约束:weight和y_origin必须同时提供或同时为空;成对提供时计算grad_weight。

    grad_x拼回:gradX[..., :H] = dg,gradX[..., H:] = du。

  • 关键特性:支持MoE场景的group_index动态分组和weight权重梯度计算。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
grad_y 输入 上游梯度,shape (T, H) 或 (B, S, H)。 BFLOAT16、FLOAT16、FLOAT32 ND
x 输入 前向输入,shape (T, 2H) 或 (B, S, 2H),包含 gate 和 up 分支。 BFLOAT16、FLOAT16、FLOAT32 ND
weight 可选输入 MoE top-k 路由权重,shape (T, 1) 或 (B, S, 1),dtype FP32。缺省视作全1。 FLOAT32 ND
y_origin 可选输入 前向输出 y,shape (T, H) 或 (B, S, H),dtype 同 grad_y;weight 存在时 y 已乘该权重。weight 提供时必须同时提供。 BFLOAT16、FLOAT16、FLOAT32 ND
group_index 可选输入 各分组 token/batch 数量索引,shape (G,),G > 0,dtype INT64。缺省视作全部行有效。 INT64 ND
clamp_limit 属性 截断门限标量 c;缺省 0 表示不 clamp(等价 c=+∞)。 FLOAT -
grad_x 输出 x 的梯度,shape (T, 2H) 或 (B, S, 2H),dtype 同 grad_y。 BFLOAT16、FLOAT16、FLOAT32 ND
grad_weight 可选输出 weight 的梯度,shape (T, 1) 或 (B, S, 1),dtype FP32。仅 weight 和 y_origin 同时提供时计算。 FLOAT32 ND

约束说明

  • H > 0
  • x.shape[-1] = 2 × H(grad_y.shape[-1])
  • grad_y 与 x 的前导维度必须一致,且二者均为 2D 或 3D Tensor
  • weight 和 y_origin 必须同时提供才能计算 grad_weight
  • clamp_limit 缺省时禁用 clamp(等价 c = +∞)
  • group_index 非空时必须是一维非空 Tensor(G > 0)
  • group_index 缺省时所有前导维度展平后的行均为有效行

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_swiglu_group_grad 通过aclnnSwigluGroupGrad调用SwigluGroupGrad算子
图模式 test_geir_swiglu_group_grad 通过算子IR调用SwigluGroupGrad算子
torch接口 test_torch_swiglu_group_grad 通过torch.ops.cann_ops_nn.swiglu_group_backward调用SwigluGroupGrad算子