文件最后提交记录最后更新时间
9 小时前
17 天前
17 天前
17 天前
17 天前
17 天前
7 个月前
26 天前
README

DequantSwigluQuant

产品支持情况

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

功能说明

  • 算子功能:在Swish门控线性单元激活函数前后添加dequant和quant操作,实现x的DequantSwigluQuant计算。

  • swiglu_mode为0时的计算公式:

    dequantOuti=Dequant(xi)dequantOut_i = Dequant(x_i)

    swigluOuti=Swiglu(dequantOuti)=Swish(Ai)∗BiswigluOut_i = Swiglu(dequantOut_i)=Swish(A_i)*B_i

    outi=Quant(swigluOuti)out_i = Quant(swigluOut_i)

    其中,Ai表示dequantOuti的前半部分,Bi表示dequantOuti的后半部分。

  • swiglu_mode为1时的计算公式:

    dequantOuti=Dequant(xi)dequantOut_i = Dequant(x_i)

    x_glu=x_glu.clamp(min=None,max=clamp_limit)x\_glu = x\_glu.clamp(min=None, max=clamp\_limit)

    x_linear=x_linear.clamp(min=−clamp_limit,max=clamp_limit)x\_linear = x\_linear.clamp(min=-clamp\_limit, max=clamp\_limit)

    out_glu=x_glu∗sigmoid(glu_alpha∗x_glu)out\_glu = x\_glu * sigmoid(glu\_alpha * x\_glu)

    swigluOuti=out_glu∗(x_linear+glu_bias)swigluOut_i = out\_glu * (x\_linear + glu\_bias)

    outi=Quant(swigluOuti)out_i = Quant(swigluOut_i)

    其中,x_glu表示dequantOuti的偶数索引部分,x_linear表示dequantOuti的奇数索引部分。

  • swiglu_mode为2时的计算公式:

    计算逻辑与swiglu_mode为1时相同(同为GPT-OSS变体SwiGLU,使用clamp_limit、glu_alpha和glu_bias):

    dequantOuti=Dequant(xi)dequantOut_i = Dequant(x_i)

    x_glu=x_glu.clamp(min=None,max=clamp_limit)x\_glu = x\_glu.clamp(min=None, max=clamp\_limit)

    x_linear=x_linear.clamp(min=−clamp_limit,max=clamp_limit)x\_linear = x\_linear.clamp(min=-clamp\_limit, max=clamp\_limit)

    out_glu=x_glu∗sigmoid(glu_alpha∗x_glu)out\_glu = x\_glu * sigmoid(glu\_alpha * x\_glu)

    swigluOuti=out_glu∗(x_linear+glu_bias)swigluOut_i = out\_glu * (x\_linear + glu\_bias)

    outi=Quant(swigluOuti)out_i = Quant(swigluOut_i)

    与swiglu_mode为1的区别在于x_glu与x_linear的切分方式:swiglu_mode为2时,x_glu表示dequantOuti的前半部分,x_linear表示dequantOuti的后半部分(与swiglu_mode为0的切分方式一致);swiglu_mode为1时为奇偶索引交错切分。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入 输入待处理的数据,公式中的x。输入不支持包含±inf或nan。 FLOAT16、BFLOAT16、INT32 ND
weight_scale 输入 输入不支持包含±inf或nan。 FLOAT ND
activation_scale 输入 激活函数的反量化scale。输入不支持包含±inf或nan。 FLOAT ND
bias 输入 Matmul的bias,公式中的bias。输入不支持包含±inf或nan。 FLOAT、FLOAT16、BFLOAT16、INT32 ND
quant_scale 输入 量化的scale,公式中的quant_scale。输入不支持包含±inf或nan。 FLOAT、FLOAT16 ND
quant_offset 输入 量化的offset。输入不支持包含±inf或nan。 FLOAT ND
group_index 输入 MoE分组需要的group_index。输入不支持包含±inf或nan。 INT64 ND
activate_left 属性 表示是否对输入的左半部分做swiglu激活。 BOOL -
quant_mode 属性 表示使用动态量化。 STRING -
dst_type 属性 表示指定输出y的数据类型。 INT64 -
round_mode 属性 表示对输出y结果的舍入模式。 STRING -
activate_dim 属性 表示进行swish计算时,选择的指定切分轴。 INT64 -
swiglu_mode 属性 表示swiglu的计算模式,取值0/1/2:0为传统swiglu;1为变体swiglu(奇偶分块);2为变体swiglu(连续前后半分块)。 INT64 -
clamp_limit 属性 表示变体swiglu使用的门限值。 FLOAT -
glu_alpha 属性 表示变体swiglu使用的参数。 FLOAT -
glu_bias 属性 表示变体swiglu使用的偏差参数。 FLOAT -
y 输出 - INT8、FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT4_E2M1、FLOAT4_E1M2 ND
scale 输出 - FLOAT ND
  • Kirin X90/Kirin 9030 处理器系列产品:
    • 输入x:数据类型不支持BFLOAT16。
    • 输入bias:数据类型不支持BFLOAT16。
    • 输入quant_scale:数据类型不支持FLOAT16。
    • 输出y:数据类型不支持FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT4_E2M1、FLOAT4_E1M2。

约束说明

  • Ascend 950PR/Ascend 950DT:

    • 输入x对应activate_dim的维度需要是2的倍数,且x的维数必须大于1维。
    • 当输入x的数据类型为INT32时,weight_scale不能为空;当输入x的数据类型不为INT32时,weight_scale不允许输入,传入空指针。
    • 当输入x的数据类型不为INT32时,activation_scale不允许输入,参数置为空指针。
    • 当输入x的数据类型不为INT32时,bias不允许输入,参数置为空指针。
    • 当输出y的数据类型为FLOAT4_E2M1、FLOAT4_E1M2时,y的最后一维需要是2的倍数。
    • 输出y的尾轴不超过5120.
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:

    • swiglu_mode、clamp_limit、glu_alpha和glu_bias四个参数用于GPT-OSS变体SwiGLU的使用。
    • x的最后一维需要是2的倍数,且x的维数必须大于1维。
    • 当quant_mode为static时,quant_scale和quant_offset为1维,值为1;quant_mode为dynamic时,quant_scale和quant_offset
    • 算子支持的输入张量的内存大小有上限,校验公式:weight_scale张量内存大小+bias张量内存大小+quant_scale张量内存大小+quant_offset张量内存大小 + (activation_scale张量内存大小 + scale张量内存大小)/40 + x张量最后一维H内存大小 * 10 < 192KB。

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_dequant_swiglu_quant 通过aclnnDequantSwigluQuant接口方式调用DequantSwigluQuant算子。
aclnn调用 test_aclnn_dequant_swiglu_quant_v2 通过aclnnDequantSwigluQuantV2接口方式调用DequantSwigluQuant算子。
图模式调用 - 通过算子IR构图方式调用DequantSwigluQuant算子。