文件最后提交记录最后更新时间
9 天前
16 天前
26 天前
16 天前
26 天前
16 天前
1 个月前
10 天前
README

SwigluGroupQuantGrad

产品支持情况

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

功能说明

  • 算子功能:SwigluGroupQuantGrad算子实现SwiGLU激活函数分组量化的反向梯度计算,用于计算输入梯度grad_x和权重梯度grad_weight

  • 算子支持范围:支持MoE场景(传入groupIndex)和非MoE场景(groupIndex传空),支持可选的Clamp反向传播掩码,支持可选的Weight梯度计算。

  • 计算流程:

    • 步骤〇:GroupIndex处理(可选)→ 计算trunc
    • 步骤一:输入切分(将x切分为x0和x1)
    • 步骤二:Clamp处理(可选)
    • 步骤三:SwiGLU反向传播计算
    • 步骤四:Weight梯度计算(可选)
    • 步骤五:梯度拼接输出
  • MoE场景GroupIndex处理公式:

    trunc=∑g=0G−1groupIndex[g] \text{trunc} = \sum_{g=0}^{G-1} \text{groupIndex}[g]

    其中:GG 为MoE专家分组数,后续所有步骤仅处理前 trunc\text{trunc} 行数据。

  • 输入切分公式:

    x0[t,h]=x[t,h],h∈[0,H) \mathbf{x}_0[t, h] = \mathbf{x}[t, h], \quad h \in [0, H)

    x1[t,h]=x[t,h+H],h∈[0,H) \mathbf{x}_1[t, h] = \mathbf{x}[t, h + H], \quad h \in [0, H)

  • Clamp处理公式(当clamp_limit > 0时):

    x0′[t,h]=min⁡(x0[t,h],c) \mathbf{x}_0'[t, h] = \min(\mathbf{x}_0[t, h], c)

    x1′[t,h]=min⁡(max⁡(x1[t,h],−c),c) \mathbf{x}_1'[t, h] = \min(\max(\mathbf{x}_1[t, h], -c), c)

    其中 ccclamp_limit

  • SiLU梯度公式:

    dSiLUdx0′=σ(x0′)⋅(1+x0′⋅(1−σ(x0′))) \frac{d\text{SiLU}}{d\mathbf{x}_0'} = \sigma(\mathbf{x}_0') \cdot \left(1 + \mathbf{x}_0' \cdot (1 - \sigma(\mathbf{x}_0'))\right)

    其中:σ(x0′)=11+e−x0′\sigma(\mathbf{x}_0') = \frac{1}{1 + e^{-\mathbf{x}_0'}}

  • 输入梯度计算公式:

    gradx0[t,h]=grady0[t,h]⋅x1′[t,h]⋅dSiLUdx0′[t,h] \mathbf{grad}_{x_0}[t, h] = \mathbf{grad}_{y_0}[t, h] \cdot \mathbf{x}_1'[t, h] \cdot \frac{d\text{SiLU}}{d\mathbf{x}_0'}[t, h]

    gradx1[t,h]=grady0[t,h]⋅SiLU(x0′[t,h]) \mathbf{grad}_{x_1}[t, h] = \mathbf{grad}_{y_0}[t, h] \cdot \text{SiLU}(\mathbf{x}_0'[t, h])

    其中:如果提供了weight,则 grady0=grady⋅weight\mathbf{grad}_{y_0} = \mathbf{grad}_{y} \cdot \mathbf{weight};如果未提供weight,则 grady0=grady\mathbf{grad}_{y_0} = \mathbf{grad}_{y}

  • Weight梯度计算公式(可选):

    gradweight[t]=∑h=0H−1grady[t,h]⋅yorigin[t,h] \mathbf{grad}_{\text{weight}}[t] = \sum_{h=0}^{H-1} \mathbf{grad}_{y}[t, h] \cdot \mathbf{y}_{\text{origin}}[t, h]

    其中:yorigin\mathbf{y}_{\text{origin}} 为SwiGLU前向传播的原始激活值输出,沿最后一维(H维度)求和。

    gradweight[t]=gradweight[t]⋅I(t<trunc) \mathbf{grad}_{\text{weight}}[t] = \mathbf{grad}_{\text{weight}}[t] \cdot \mathbb{I}(t < \text{trunc})

  • Clamp反向传播掩码公式(当clamp_limit > 0时):

    gradx0[t,h]=gradx0[t,h]⋅I(x0[t,h]<c) \mathbf{grad}_{x_0}[t, h] = \mathbf{grad}_{x_0}[t, h] \cdot \mathbb{I}(\mathbf{x}_0[t, h] < c)

    gradx1[t,h]=gradx1[t,h]⋅I(−c<x1[t,h]<c) \mathbf{grad}_{x_1}[t, h] = \mathbf{grad}_{x_1}[t, h] \cdot \mathbb{I}(-c < \mathbf{x}_1[t, h] < c)

    其中I\mathbb{I}为指示函数。

  • 梯度拼接与GroupIndex处理公式:

    gradx[t,h]={gradx0[t,h]h∈[0,H)gradx1[t,h−H]h∈[H,2H) \mathbf{grad}_x[t, h] = \begin{cases} \mathbf{grad}_{x_0}[t, h] & h \in [0, H) \\ \mathbf{grad}_{x_1}[t, h-H] & h \in [H, 2H) \end{cases}

    gradx[t,:]=gradx[t,:]⋅I(t<trunc) \mathbf{grad}_x[t, :] = \mathbf{grad}_x[t, :] \cdot \mathbb{I}(t < \text{trunc})

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
gradY 输入 梯度输出张量,来自下游层的梯度。 BFLOAT16、FLOAT16、FLOAT ND
x 输入 前向传播的输入张量。 BFLOAT16、FLOAT16、FLOAT ND
weightOptional 输入 MoE权重张量。 FLOAT ND
yOriginOptional 输入 SwiGLU前向传播的原始激活值输出。 BFLOAT16、FLOAT16、FLOAT ND
groupIndexOptional 输入 GroupIndex张量,动态核分配。 INT64 ND
clampLimit 属性 Clamp阈值。 FLOAT -
gradXOut 输出 输入梯度张量。 BFLOAT16、FLOAT16、FLOAT ND
gradWeightOutOptional 输出 权重梯度张量。 FLOAT ND

约束说明

  • 确定性计算:默认确定性实现。

  • 输入shape约束:

    • x最后一维必须为偶数(2H2H
    • gradY最后一维为 HH,与x最后一维的一半对应
    • gradY与x的前n-1维shape必须一致
  • 可选参数约束:

    • weight提供时,必须同时提供yOrigin才能计算gradWeight
    • weight元素个数需等于x或gradY除最后一维外的元素个数之积
    • yOrigin的shape需与gradY一致
  • 数据类型约束:

    • gradY、x、yOrigin、gradXOut数据类型必须一致(FLOAT、FLOAT16或BFLOAT16)
    • weight、gradWeightOutOptional必须为FLOAT类型
    • groupIndex必须为INT64类型
  • Clamp约束:

    • clampLimit取值范围为-1.0或>0.0
    • clampLimit=-1.0表示不启用Clamp反向传播掩码,启用时clampLimit必须>0.0
  • 规格约束:

    规格项 规格 规格说明
    B 1~31 -
    S 0~128K -
    H 512, 768, 1024, 1536, 1792, 2048, 2560, 4096 -

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_swiglu_group_quant_grad 通过aclnnSwigluGroupQuantGrad接口方式调用SwigluGroupQuantGrad算子。
图模式调用 - 通过算子IR构图方式调用SwigluGroupQuantGrad算子。