| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 1 个月前 | ||
| 3 个月前 | ||
| 3 个月前 | ||
| 3 个月前 | ||
| 1 个月前 | ||
| 2 个月前 | ||
| 2 个月前 | ||
| 3 个月前 | ||
| 1 个月前 |
GroupedDynamicMxQuant
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | × |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | × |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
功能说明
-
算子功能:根据传入的分组索引的起始值,对传入的数据进行分组的float8的动态量化。
-
计算公式:
-
场景1,当scaleAlg为0时:
- 将输入x在第0维上先按照groupIndex进行分组,每个group内按k = blocksize个数分组,一组k个数 {{xi}i=1k} 计算出这组数对应的量化尺度mxscale_pre {mxscale_pre, {Pi}i=1k},计算公式为下面公式(1)(2)。
shared_exp=floor(log2(maxi(∣Vi∣)))−emax(1)shared\_exp = floor(log_2(max_i(|V_i|))) - emax \tag{1}
mxscale_pre=2shared_exp(2)mxscale\_pre = 2^{shared\_exp} \tag{2}
- 这组数每个数都除以mxscale,根据round_mode转换到对应的dst_type,得到量化结果y,计算公式为下面公式(3)。
Pi=cast_to_dst_type(Vi/mxscale,round_mode), i from 1 to blocksize(3)P_i = cast\_to\_dst\_type(V_i/mxscale, round\_mode), \space i\space from\space 1\space to\space blocksize \tag{3}
-
量化后的 PiP_{i} 按对应的 ViV_{i} 的位置组成输出y,mxscale_pre按对应的groupIndex分组,分组内第一个维度pad为偶数,组成输出mxscale。
-
emax:对应数据类型的最大正则数的指数位。
DataType emax FLOAT8_E4M3FN 8 FLOAT8_E5M2 15
-
场景2,当scaleAlg为1时:
-
将长向量按块分,每块长度为k,对每块单独计算一个块缩放因子Sfp32bS_{fp32}^b,再把块内所有元素用同一个Sfp32bS_{fp32}^b映射到目标低精度类型FP8。如果最后一块不足k个元素,把缺失值视为0,按照完整块处理。
-
找到该块中数值的最大绝对值:
Amax(Dfp32b)=max({∣di∣}i=1k)Amax(D_{fp32}^b)=max(\{|d_{i}|\}_{i=1}^{k})
-
将FP32映射到目标数据类型FP8可表示的范围内,其中Amax(DType)Amax(DType)是目标精度能表示的最大值。
Sfp32b=Amax(Dfp32b)Amax(DType)S_{fp32}^b = \frac{Amax(D_{fp32}^b)}{Amax(DType)}
-
将块缩放因子Sfp32bS_{fp32}^b转换为FP8格式下可表示的缩放值Sue8m0bS_{ue8m0}^b
-
从块的浮点缩放因子Sfp32bS_{fp32}^b中提取无偏指数EintbE_{int}^b和尾数MfixpbM_{fixp}^b
-
为保证量化时不溢出,对指数进行向上取整,且在FP8可表示的范围内:
Eintb={Eintb+1,如果Sfp32b为正规数,且Eintb<254且Mfixpb>0Eintb+1,如果Sfp32b为非正规数,且Mfixpb>0.5Eintb,否则E_{int}^b = \begin{cases} E_{int}^b + 1, & \text{如果} S_{fp32}^b \text{为正规数,且} E_{int}^b < 254 \text{且} M_{fixp}^b > 0 \\ E_{int}^b + 1, & \text{如果} S_{fp32}^b \text{为非正规数,且} M_{fixp}^b > 0.5 \\ E_{int}^b, & \text{否则} \end{cases}
-
计算块缩放因子:Sue8m0b=2EintbS_{ue8m0}^b=2^{E_{int}^b}
-
计算块转换因子:Rfp32b=1fp32(Sue8m0b)R_{fp32}^b=\frac{1}{fp32(S_{ue8m0}^b)}
-
应用到量化的最终步骤,对于每个块内元素,di=DType(dfp32i⋅Rfp32n)d^i = DType(d_{fp32}^i \cdot R_{fp32}^n),最终输出的量化结果是(Sb,[di]i=1k)\left(S^b, [d^i]_{i=1}^k\right),其中SbS^b代表块的缩放因子,这里指Sue8m0bS_{ue8m0}^b,[di]i=1k[d^i]_{i=1}^k代表块内量化后的数据。
-
-
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x | 输入 | Device侧的aclTensor,计算公式中的输入x。shape仅支持2维。支持非连续的Tensor,支持空Tensor。 | FLOAT16、BFLOAT16 | ND |
| groupIndex | 输入 | Device侧的aclTensor,量化分组的起始索引。shape仅支持1维。支持非连续的Tensor,不支持空Tensor。 | INT32 | ND |
| roundMode | 属性 | host侧的string,公式中的round_mode,数据转换的模式,仅支持"rint"模式。 | STRING | - |
| dstType | 属性 | host侧的int64_t,公式中的dst_type,指定数据转换后y的类型,输入范围为{35, 36},分别对应输出y的数据类型为{35: FLOAT8_E5M2, 36: FLOAT8_E4M3FN}。 | INT64 | - |
| blocksize | 属性 | host侧的int64_t,公式中的blocksize,指定每次量化的元素个数,仅支持32。 | INT64 | - |
| scaleAlg | 属性 | host侧的int64_t,指定mxscale计算时采用的算法,仅支持0和1。 | INT64 | - |
| dstTypeMax | 属性 | host侧的float32,在scale_alg=2时生效。默认值0.0表示max_type为目标数据类型的最大值,若传入其它数值,则需要按照传入的数值计算mxscale。当前支持取值为0.0/6.0-12.0,只支持在FLOAT4_E2M1场景设置该值。 | FLOAT | - |
| y | 输出 | Device侧的aclTensor,公式中的输出y,输入x量化后的对应结果。需与dstType对应,shape仅支持2维,支持空Tensor,shape和输入x一致。 | FLOAT8_E4M3FN、FLOAT8_E5M2 | ND |
| mxscale | 输出 | Device侧的aclTensor,公式中的mxscale_pre组成的输出mxscale,每个分组对应的量化尺度。shape仅支持3维,支持空Tensor。假设x的shape为 [m,n],groupedIndex的shape为 [g],则mxscale的shape为 [(m/(blocksize * 2)+g), n, 2]。 | FLOAT8_E8M0 | ND |
约束说明
- 关于x、groupIndex、y、mxscale的约束说明如下:
- groupIndex中的值必须非递减,且不能小于0,最后一个元素必须为x第一个维度的长度。
- rank(mxscale)=rank(x)+1rank(mxscale) = rank(x) + 1。
- 假设x的shape为 [m,n][m,n],groupedIndex的shape为 [g][g],则mxscale的shape为 [(m/(blocksize∗2)+g),n,2][(m/(blocksize * 2)+g), n, 2]。
- mxscale.shape[−1]=2mxscale.shape[-1] = 2。
- 输出y的shape与输入x一致。
调用说明
| 调用方式 | 调用样例 | 说明 |
|---|---|---|
| aclnn调用 | test_aclnn_grouped_dynamic_mx_quant | 通过aclnnGroupedDynamicMxQuant接口方式调用GroupedDynamicMxQuant算子。 |
| aclnn调用 | test_aclnn_grouped_dynamic_mx_quant_v2 | 通过aclnnGroupedDynamicMxQuantV2接口方式调用GroupedDynamicMxQuant算子。 |