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

SwigluMxQuantWithDualAxis

产品支持情况

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

功能说明

  • 算子功能:融合算子,实现SwiGLU激活函数与双轴动态块量化的组合计算。先对输入计算SwiGLU激活函数,然后对结果同时在-1轴和-2轴进行基于块的动态量化,输出低精度的FP4/FP8张量和对应的缩放因子。

  • 计算公式:

    阶段1:SwiGLU激活函数

    gate, hidden = split(x, axis=-1)

    swish = sigmoid(gate) * gate

    act = swish * hidden

    其中,当activate_left=True时,左半部分为hidden,右半部分为gate;当activate_left=False时,右半部分为hidden,左半部分为gate。

    阶段2:双轴动态块量化

    • -1轴量化(列方向):将SwiGLU结果在-1轴上按照32个数进行分组,一组32个数{{Vi}i=132}\{\{V_i\}_{i=1}^{32}\} 量化为 {mxscale1,{Pi}i=132}\{mxscale1, \{P_i\}_{i=1}^{32}\}

      shared_exp=floor(log2(maxi(∣Vi∣)))−emaxshared\_exp = floor(log_2(max_i(|V_i|))) - emax

      mxscale1=2shared_expmxscale1 = 2^{shared\_exp}

      Pi=cast_to_dst_type(Vi/mxscale1,round_mode), i from 1 to 32P_i = cast\_to\_dst\_type(V_i/mxscale1, round\_mode), \space i\space from\space 1\space to\space 32

    • -2轴量化(行方向):将SwiGLU结果在-2轴上按照32个数进行分组,一组32个数{{Vj}j=132}\{\{V_j\}_{j=1}^{32}\} 量化为 {mxscale2,{Pj}j=132}\{mxscale2, \{P_j\}_{j=1}^{32}\}

      shared_exp=floor(log2(maxj(∣Vj∣)))−emaxshared\_exp = floor(log_2(max_j(|V_j|))) - emax

      mxscale2=2shared_expmxscale2 = 2^{shared\_exp}

      Pj=cast_to_dst_type(Vj/mxscale2,round_mode), j from 1 to 32P_j = cast\_to\_dst\_type(V_j/mxscale2, round\_mode), \space j\space from\space 1\space to\space 32

    • emax: 对应数据类型的最大正则数的指数位。

      DataType emax
      FLOAT4_E2M1 2
      FLOAT4_E1M2 0
      FLOAT8_E4M3FN 8
      FLOAT8_E5M2 15

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入 输入张量,必须为2维,且最后一维必须为偶数(shape为[M, 2N]) FLOAT16、BFLOAT16 ND
group_index 可选输入 分组索引,用于控制分组量化边界。shape为[G],采用cumsum模式,表示每个group的行数累积值。传入空指针时表示不分组。 INT64 ND
activate_left 属性 SwiGLU激活侧选择。True表示左半部分为hidden,右半部分为gate;False表示右半部分为hidden,左半部分为gate。默认值为True。 BOOL -
round_mode 属性 舍入模式,用于量化时的类型转换。
当dst_type为FLOAT8_E5M2/FLOAT8_E4M3FN时,仅支持"rint";当dst_type为FLOAT4_E2M1/FLOAT4_E1M2时,支持"rint"、"floor"、"round"。默认值为"rint"。
STRING -
scale_alg 属性 缩放算法:取值为1时,表示使用cuBLAS算法;取值为0时,表示使用OCP算法。默认值为0。
当dst_type为FLOAT4_E2M1/FLOAT4_E1M2时仅支持scale_alg=0。
INT64 -
dst_type 属性 目标量化类型:35=FLOAT8_E5M2、36=FLOAT8_E4M3FN、40=FLOAT4_E2M1、41=FLOAT4_E1M2。默认值为35。 INT64 -
max_dtype_value 属性 预留参数,scale_alg=2且dst_type=FLOAT4_E1M2时生效。
当前仅支持取值为0.0。
FLOAT -
y1 输出 -1轴量化后的输出张量,形状为[M, N](SwiGLU输出的一半)。 FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT4_E2M1、FLOAT4_E1M2 ND
mx_scale1 输出 -1轴每个量化块对应的缩放因子。shape为[M, (ceil(N/32)+2-1)/2, 2],需进行偶数pad,pad填充值为0。 FLOAT8_E8M0 ND
y2 输出 -2轴量化后的输出张量,形状为[M, N]。 FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT4_E2M1、FLOAT4_E1M2 ND
mx_scale2 输出 -2轴每个量化块对应的缩放因子。当groupIndex存在时shape为[floor(M/64) + G, N, 2],当groupIndex不存在时shape为[ceil(M/64), N, 2],需进行偶数pad,pad填充值为0。输出需要对每两行数据进行交织处理。 FLOAT8_E8M0 ND

约束说明

  • 输入x必须为2维张量,最后一维必须能被2整除(shape为[M, 2N])。
  • FP8输出类型(FLOAT8_E5M2/FLOAT8_E4M3FN)仅支持“rint”舍入模式;FP4输出类型(FLOAT4_E2M1/FLOAT4_E1M2)支持“rint”、“floor”、“round”舍入模式。
  • 当dst_type为FLOAT4_E2M1/FLOAT4_E1M2时,仅支持scale_alg=0(OCP实现)。
  • 当group_index存在时,采用cumsum模式,每个值表示对应group的行数累积值,group_index的每个元素值需要大于0且最后一个元素值要等于M。

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_swiglu_mx_quant_with_dual_axis 通过aclnnSwigluMxQuantWithDualAxis接口方式调用SwigluMxQuantWithDualAxis算子。
图模式 - 通过算子IR构图方式调用SwigluMxQuantWithDualAxis算子。