文件最后提交记录最后更新时间
6 天前
20 天前
1 个月前
17 天前
20 天前
1 个月前
19 天前
1 个月前
6 天前
README

SituMxQuant

产品支持情况

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

功能说明

  • 算子功能:将Situ激活函数与动态MX(Microscaling)量化融合为一个算子。

  • 计算公式:

    1. Situ激活:

      situa=β×tanh⁡(gate/β)×sigmoid(gate)situ_a = \beta \times \tanh(gate / \beta) \times sigmoid(gate)

      当linear_beta > 0时:

      up=linear_beta×tanh⁡(up/linear_beta)up = linear\_beta \times \tanh(up / linear\_beta)

      situOut=situa×upsituOut = situ_a \times up

      其中,当activate_left为true时,gate取x的前半部分,up取后半部分;当activate_left为false时,gate取x的后半部分,up取前半部分。

    2. MX量化(OCP算法):

      shared_exp=floor(log2(max(∣situOuti∣)))−emaxshared\_exp = floor(log2(max(|situOut_i|))) - emax

      y_scale=2shared_exp(E8M0)y\_scale = 2^{shared\_exp} (E8M0)

      y=cast_to_fp8(situOut/y_scale)y = cast\_to\_fp8(situOut / y\_scale)

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x 输入 输入待处理的数据,公式中的x。最后一维需要是2的倍数。输入不支持包含±inf或nan。 FLOAT16、BFLOAT16 ND
beta 属性 Situ激活的beta参数,公式中的β。必须大于0。默认1.0。 FLOAT -
linear_beta 属性 Situ激活的linear_beta参数,公式中的linear_beta。当值≤0时不启用。默认0.0。 FLOAT -
activate_left 属性 表示gate取x的前半部分还是后半部分,公式中的activate_left。默认false。 BOOL -
axis 属性 表示量化轴,公式中max和scale的计算轴。当前仅支持-1。默认-1。 INT -
dst_type 属性 表示输出y的数据类型,对应公式中cast_to_fp8的目标类型:36=FLOAT8_E4M3FN,35=FLOAT8_E5M2。默认36。 INT -
round_mode 属性 表示量化舍入模式,公式中cast的舍入方式。支持"rint"、"round"、"floor"。FP8输出仅支持"rint"。默认"rint"。 STRING -
y 输出 量化后的输出,公式中的y。 FLOAT8_E4M3FN、FLOAT8_E5M2 ND
y_scale 输出 MX量化的scale(E8M0格式),公式中的y_scale。 FLOAT8_E8M0 ND

约束说明

  • x的最后一维需要是2的倍数。
  • x的维数必须大于等于1维。
  • axis当前仅支持-1(尾轴量化)。
  • dst_type支持36(FLOAT8_E4M3FN)或35(FLOAT8_E5M2)。
  • round_mode必须为"rint"。
  • y的数据类型必须与dst_type匹配,y_scale的数据类型必须为FLOAT8_E8M0。
  • y、y_scale的shape需要与推导结果一致(见参数说明)。
  • 关于y_scale的shape约束说明如下:
    • H = x.shape[-1] / 2
    • scaleNum = ceil(H / 64)
    • y_scale.shape = x.shape[:-1] + [scaleNum, 2]

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_situ_mx_quant 通过aclnnSituMxQuant接口方式调用SituMxQuant算子。
图模式调用 - 通过算子IR构图方式调用SituMxQuant算子。