已合并
add dynamic_mx_quant dst_type_value && rewrite tail_axis #3145
add dynamic_mx_quant dst_type_value && rewrite tail_axis #3145
已合并
pyongq创建于 3月25日
共 21 个文件变更+4270-4397
@@ -55,7 +55,37 @@
55 $$55 $$
56 - 计算块缩放因子:$S_{ue8m0}^b=2^{E_{int}^b}$56 - 计算块缩放因子:$S_{ue8m0}^b=2^{E_{int}^b}$
57 - 计算块转换因子:$R_{fp32}^b=\frac{1}{fp32(S_{ue8m0}^b)}$57 - 计算块转换因子:$R_{fp32}^b=\frac{1}{fp32(S_{ue8m0}^b)}$
58- - 应用到量化的最终步骤,对于每个块内元素,$d^i = DType(d_{fp32}^i \cdot R_{fp32}^b)$,最终输出的量化结果是$\left(S^b, [d^i]_{i=1}^k\right)$,其中$S^b$代表块的缩放因子,这里指$S_{ue8m0}^b$,$[d^i]_{i=1}^k$代表块内量化后的数据。58+ - 应用到量化的最终步骤,对于每个块内元素,$d^i = DType(d_{fp32}^i \cdot R_{fp32}^n)$,最终输出的量化结果是$\left(S^b, [d^i]_{i=1}^k\right)$,其中$S^b$代表块的缩放因子,这里指$S_{ue8m0}^b$,$[d^i]_{i=1}^k$代表块内量化后的数据。
59+ - 场景3,当scaleAlg为2时,只涉及FP4_E2M1类型:
60+ - 当dst_max_value = 0.0/6.0/7.0时:
61+ - 将输入x在axis维度上按k = blocksize个数分组,一组k个数 $\{\{V_i\}_{i=1}^{k}\}$ 动态量化为 $\{mxscale1, \{P_i\}_{i=1}^{k}\}$, k = blocksize:
62+ $$
63+ shared\_exp = \begin{cases} ceil(log_2(max_i(|V_i|))) - emax, & \text{如果} 尾数位的高比特前一/两位 \text{为1,且尾数不全为0} \\ floor(log_2(max_i(|V_i|))) - emax, & \text{其它} \end{cases} \\
64+ $$
65+ $$
66+ P_i = cast\_to\_dst\_type(V_i/mxscale, round\_mode), \space i\space from\space 1\space to\space blocksize\\
67+ $$
68+ - ​量化后的 $P_{i}$ 按对应的 $V_{i}$ 的位置组成输出yOut,mxscale按对应的axis维度上的分组组成输出mxscaleOut。
69+ - 当dst_max_value != 0.0/6.0/7.0时:
70+ - 将长向量按块分,每块长度为k,对每块单独计算一个块缩放因子$S_{fp32}^b$,再把块内所有元素用同一个$S_{fp32}^b$映射到目标低精度类型FP8。如果最后一块不足k个元素,把缺失值视为0,按照完整块处理。
71+ - 找到该块中数值的最大绝对值:
72+ $$
73+ Amax(D_{fp32}^b)=max(\{|d_{i}|\}_{i=1}^{k})
74+ $$
75+ - 将FP32映射到目标数据类型FP8可表示的范围内,其中当dst_max_value=0时,$Amax(DType)$是目标精度能表示的最大值;当dst_max_value!=0时,$Amax(DType)$是dst_max_value传入值。
76+ $$
77+ S_{fp32}^b = \frac{Amax(D_{fp32}^b)}{Amax(DType)}
78+ $$
79+ - 将块缩放因子$S_{fp32}^b$转换为FP8格式下可表示的缩放值$S_{ue8m0}^b$。
80+ - 从块的浮点缩放因子$S_{fp32}^b$中提取无偏指数$E_{int}^b$和尾数$M_{fixp}^b$。
81+ - 为保证量化时不溢出,对指数进行向上取整,且在FP8可表示的范围内:
82+ $$
83+ 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, & \text{否则} \end{cases}
84+ $$
85+ - 计算块缩放因子:$S_{ue8m0}^b=2^{E_{int}^b}$
86+ - 计算块转换因子:$R_{fp32}^b=\frac{1}{fp32(S_{ue8m0}^b)}$
87+ - 应用到量化的最终步骤,对于每个块内元素,$d^i = DType(d_{fp32}^i \cdot R_{fp32}^n)$,最终输出的量化结果是$\left(S^b, [d^i]_{i=1}^k\right)$,其中$S^b$代表块的缩放因子,这里指$S_{ue8m0}^b$,$[d^i]_{i=1}^k$代表块内量化后的数据。
88+ - ​量化后的 $P_{i}$ 按对应的 $V_{i}$ 的位置组成输出yOut,mxscale按对应的axis维度上的分组组成输出mxscaleOut。
59 89 
60## 函数原型90## 函数原型
61 91 
@@ -69,6 +99,7 @@ aclnnStatus aclnnDynamicMxQuantGetWorkspaceSize(
69 int64_t dstType,99 int64_t dstType,
70 int64_t blocksize,100 int64_t blocksize,
71 int64_t scaleAlg,101 int64_t scaleAlg,
102+ float dstMaxValue,
72 aclTensor *yOut,103 aclTensor *yOut,
73 aclTensor *mxscaleOut,104 aclTensor *mxscaleOut,
74 uint64_t *workspaceSize,105 uint64_t *workspaceSize,
@@ -163,12 +194,22 @@ aclnnStatus aclnnDynamicMxQuant(
163 <td>scaleAlg</td>194 <td>scaleAlg</td>
164 <td>输入</td>195 <td>输入</td>
165 <td>表示mxscaleOut的计算方法,对应公式中的scaleAlg。</td>196 <td>表示mxscaleOut的计算方法,对应公式中的scaleAlg。</td>
166- <td><ul><li>支持取值0和1,取值为0代表场景1,为1代表场景2。</li><li>当dstType为FLOAT4_E2M1/FLOAT4_E1M2时仅支持取值为0。</li></ul></td>197+ <td><ul><li>支持取值0、1和2,取值为0代表场景1,为1代表场景2,为2代表场景3。</li><li>当dstType为FLOAT4_E1M2时仅支持取值为0;当dstType为FLOAT4_E2M1时仅支持取值为0和2;当dstType为FLOAT8时仅支持取值为0和1。</li></ul></td>
167 <td>INT64</td>198 <td>INT64</td>
168 <td>ND</td>199 <td>ND</td>
169 <td>-</td>200 <td>-</td>
170 <td>-</td>201 <td>-</td>
171 </tr>202 </tr>
203+ <tr>
204+ <td>dstMaxValue</td>
205+ <td>输入</td>
206+ <td>表示maxType的取值,对应公式中的Amax(DType)。</td>
207+ <td><ul><li>支持取值0.0和6.0-12.0,取值为0.0代表Amax(DType)为量化结果数据类型的最大值;取值为6.0-12.0代表Amax(DType)为传入值。</li></ul></td>
208+ <td>FLOAT</td>
209+ <td>ND</td>
210+ <td>-</td>
211+ <td>-</td>
212+ </tr>
172 <tr>213 <tr>
173 <td>yOut</td>214 <td>yOut</td>
174 <td>输出</td>215 <td>输出</td>
@@ -307,6 +348,13 @@ aclnnStatus aclnnDynamicMxQuant(
307 - mxscaleOut.shape[axis_change] = (ceil(x.shape[axis] / blocksize) + 2 - 1) / 2。348 - mxscaleOut.shape[axis_change] = (ceil(x.shape[axis] / blocksize) + 2 - 1) / 2。
308 - mxscaleOut.shape[-1] = 2。349 - mxscaleOut.shape[-1] = 2。
309 - 其他维度与输入x一致。350 - 其他维度与输入x一致。
351+- 关于参数的约束说明如下:
352+ - x/yOut:rank(x) = rank(yOut) = 1-7,输入输出数据类型保持一致。
353+ - round_mode:量化结果数据类型为FLOAT4时,支持"rint"、"round"、"floor";量化结果数据类型为FLOAT8时,仅支持"rint"。
354+ - dst_type:若量化结果数据类型为FLOAT4_E2M1,取值为40;若量化结果数据类型为FLOAT4_E1M2,取值为41;若量化结果数据类型为FLOAT8_E5M2,取值为35;若量化结果数据类型为FLOAT8_E4M3FN,取值为36。
355+ - scale_alg:若量化结果数据类型为FLOAT4_E1M2,取值仅支持0(OCP Microscaling Formats (Mx) Specification 实现);若量化结果数据类型为FLOAT4kt_E2M1,取值仅支持0或2(Dynamic dtype Range 实现);若量化结果数据类型为FLOAT8,取值仅支持0或1(cuBLAS 实现)。
356+ - blocksize:scale_alg为2时,blocksize必须为32。
357+ - dst_max_value:取值仅支持0.0或6.0-12.0,在scale_alg=2时生效。默认值0.0代表maxType为目标数据类型的最大值,若传入其它数值则按照传入的数值计算mxscale。仅支持在FLOAT4_E2M1场景设置该值。
310 358 
311 359 
312## 调用示例360## 调用示例
@@ -420,6 +468,7 @@ int aclnnDynamicMxQuantTest(int32_t deviceId, aclrtStream& stream)
420 int64_t dstType = 36;468 int64_t dstType = 36;
421 int64_t blocksize = 32;469 int64_t blocksize = 32;
422 int64_t scaleAlg = 0;470 int64_t scaleAlg = 0;
471+ float dstMaxValue = 0.0;
423 // 创建x aclTensor472 // 创建x aclTensor
424 ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_BF16, &x);473 ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_BF16, &x);
425 std::unique_ptr<aclTensor, aclnnStatus (*)(const aclTensor*)> xTensorPtr(x, aclDestroyTensor);474 std::unique_ptr<aclTensor, aclnnStatus (*)(const aclTensor*)> xTensorPtr(x, aclDestroyTensor);
@@ -441,7 +490,7 @@ int aclnnDynamicMxQuantTest(int32_t deviceId, aclrtStream& stream)
441 aclOpExecutor* executor;490 aclOpExecutor* executor;
442 491 
443 // 调用aclnnDynamicMxQuant第一段接口492 // 调用aclnnDynamicMxQuant第一段接口
444- ret = aclnnDynamicMxQuantGetWorkspaceSize(x, axis, roundModeOptional, dstType, blocksize, scaleAlg, yOut, mxscaleOut, &workspaceSize, &executor);493+ ret = aclnnDynamicMxQuantGetWorkspaceSize(x, axis, roundModeOptional, dstType, blocksize, scaleAlg, dstMaxValue, yOut, mxscaleOut, &workspaceSize, &executor);
445 CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnDynamicMxQuantGetWorkspaceSize failed. ERROR: %d\n", ret);494 CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnDynamicMxQuantGetWorkspaceSize failed. ERROR: %d\n", ret);
446 return ret);495 return ret);
447 // 根据第一段接口计算出的workspaceSize申请device内存496 // 根据第一段接口计算出的workspaceSize申请device内存
@@ -36,8 +36,11 @@ namespace ge {
36* @li dst_type: An optional int. Declare the output y dtype. Support FLOAT4_E2M1, FLOAT4_E1M2,36* @li dst_type: An optional int. Declare the output y dtype. Support FLOAT4_E2M1, FLOAT4_E1M2,
37* FLOAT8_E4M3FN or FLOAT8_E5M2. Defaults to FLOAT4_E2M1.37* FLOAT8_E4M3FN or FLOAT8_E5M2. Defaults to FLOAT4_E2M1.
38* @li blocksize: An optional int. Block size for quantization scaling factors.Defaults to 32.38* @li blocksize: An optional int. Block size for quantization scaling factors.Defaults to 32.
39+* When scale_alg is 2, blocksize must be 32.
39* @li scale_alg: An optional int.The algorithm for the scale in quantization.Default to 0.40* @li scale_alg: An optional int.The algorithm for the scale in quantization.Default to 0.
40-* Support MxFP8(OCP Microscaling Formats (Mx) Specification, count 0) or MxFP8(nvidia-cuBLAS , count 1).41+* Support MxFP8/MxFP4(OCP Microscaling Formats (Mx) Specification , count 0) or MxFP8(nvidia-cuBLAS , count 1) or MxFP4(Dynamic Dtype Range , count 2).
42+* @li dst_type_max: An optional Float.Max_dtype takes the maximum value of the quant_data_type, or the provided value.Defaults to 0.
43+* Only support in FP4_E2M1 mode, with a valid range of 6.0 to 12.0.
41 44 
42* @par Outputs:45* @par Outputs:
43* @li y: Quantized output tensor. It has the same shape and rank as input x.46* @li y: Quantized output tensor. It has the same shape and rank as input x.
@@ -54,23 +57,25 @@ namespace ge {
54* @li When dst_type is DT_FLOAT4_E2M1 or DT_FLOAT4_E1M2, round_mode supports "rint", "floor" and "round".57* @li When dst_type is DT_FLOAT4_E2M1 or DT_FLOAT4_E1M2, round_mode supports "rint", "floor" and "round".
55* @li If dst_type is DT_FLOAT4_E2M1 or DT_FLOAT4_E1M2, the input x last dimension of the shape must be divisible by 2.58* @li If dst_type is DT_FLOAT4_E2M1 or DT_FLOAT4_E1M2, the input x last dimension of the shape must be divisible by 2.
56* @li The blocksize must be a multiple of 32 (non-zero) and ≤ 1024.59* @li The blocksize must be a multiple of 32 (non-zero) and ≤ 1024.
60+* When scale_alg is 2, blocksize must be 32.
61+* @li The value of dst_max_value only supports 0.0 or 6.0-12.0 and is effective when scale_alg=2.
62+* The default value 0.0 means that maxType corresponds to the maximum value of the target data type.
63+* If other values are provided, mxscale is calculated based on the provided value.
57 64 
58* @par Third-party framework compatibility65* @par Third-party framework compatibility
59* It is a custom operator. It has no corresponding operator in Caffe, ONNX, TensorFlow, or PyTorch.66* It is a custom operator. It has no corresponding operator in Caffe, ONNX, TensorFlow, or PyTorch.
60*/67*/
61REG_OP(DynamicMxQuant)68REG_OP(DynamicMxQuant)
62 .INPUT(x, TensorType({DT_FLOAT16, DT_BF16}))69 .INPUT(x, TensorType({DT_FLOAT16, DT_BF16}))
63- .OUTPUT(70+ .OUTPUT(y, TensorType({DT_FLOAT4_E2M1, DT_FLOAT4_E1M2, DT_FLOAT6_E3M2, DT_FLOAT6_E2M3, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2}))
64- y,
65- TensorType({DT_FLOAT4_E2M1, DT_FLOAT4_E1M2, DT_FLOAT6_E3M2, DT_FLOAT6_E2M3, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2}))
66 .OUTPUT(mxscale, TensorType({DT_FLOAT8_E8M0}))71 .OUTPUT(mxscale, TensorType({DT_FLOAT8_E8M0}))
67 .ATTR(axis, Int, -1)72 .ATTR(axis, Int, -1)
68 .ATTR(round_mode, String, "rint")73 .ATTR(round_mode, String, "rint")
69 .ATTR(dst_type, Int, DT_FLOAT4_E2M1)74 .ATTR(dst_type, Int, DT_FLOAT4_E2M1)
70 .ATTR(blocksize, Int, 32)75 .ATTR(blocksize, Int, 32)
71 .ATTR(scale_alg, Int, 0)76 .ATTR(scale_alg, Int, 0)
77+ .ATTR(dst_type_max, Float, 0.0)
72 .OP_END_FACTORY_REG(DynamicMxQuant)78 .OP_END_FACTORY_REG(DynamicMxQuant)
73- 
74} // namespace ge79} // namespace ge
75 80 
76#endif // QUANT_DYNAMIC_MX_QUANT_PROTO_H_81#endif // QUANT_DYNAMIC_MX_QUANT_PROTO_H_
Mquant/dynamic_mx_quant/op_host/arch35/dynamic_mx_quant_tiling_arch35.cpp+144-349文件内容审核中,请稍后刷新重试
Mquant/dynamic_mx_quant/op_host/arch35/dynamic_mx_quant_tiling_arch35.h+64-119文件内容审核中,请稍后刷新重试
@@ -20,6 +20,7 @@ namespace ops {
20static constexpr int32_t DEFAULT_BLOCK_SIZE = 32;20static constexpr int32_t DEFAULT_BLOCK_SIZE = 32;
21static constexpr int32_t DEFAULT_DST_TYPE = 40;21static constexpr int32_t DEFAULT_DST_TYPE = 40;
22static constexpr int32_t DEFAULT_SCALE_ALG = 0;22static constexpr int32_t DEFAULT_SCALE_ALG = 0;
23+static constexpr float DEFAULT_DST_TYPE_MAX = 0.0;
23class DynamicMxQuant : public OpDef {24class DynamicMxQuant : public OpDef {
24public:25public:
25 explicit DynamicMxQuant(const char* name) : OpDef(name)26 explicit DynamicMxQuant(const char* name) : OpDef(name)
@@ -63,6 +64,7 @@ public:
63 this->Attr("dst_type").AttrType(OPTIONAL).Int(DEFAULT_DST_TYPE);64 this->Attr("dst_type").AttrType(OPTIONAL).Int(DEFAULT_DST_TYPE);
64 this->Attr("blocksize").AttrType(OPTIONAL).Int(DEFAULT_BLOCK_SIZE);65 this->Attr("blocksize").AttrType(OPTIONAL).Int(DEFAULT_BLOCK_SIZE);
65 this->Attr("scale_alg").AttrType(OPTIONAL).Int(DEFAULT_SCALE_ALG);66 this->Attr("scale_alg").AttrType(OPTIONAL).Int(DEFAULT_SCALE_ALG);
67+ this->Attr("dst_type_max").AttrType(OPTIONAL).Version(2).Float(DEFAULT_DST_TYPE_MAX);
66 68 
67 OpAICoreConfig aicoreConfig;69 OpAICoreConfig aicoreConfig;
68 aicoreConfig.DynamicCompileStaticFlag(true)70 aicoreConfig.DynamicCompileStaticFlag(true)
@@ -17,6 +17,8 @@
17#define DYNAMIC_MX_QUANT_COMMON_H17#define DYNAMIC_MX_QUANT_COMMON_H
18 18 
19#include "kernel_operator.h"19#include "kernel_operator.h"
20+#include "../inc/platform.h"
21+#include "dynamic_mx_quant_tilingdata.h"
20namespace DynamicMxQuant {22namespace DynamicMxQuant {
21 23 
22template <typename Tp, Tp v>24template <typename Tp, Tp v>
@@ -32,12 +34,30 @@ struct IsSame<Tp, Tp> : public trueType {};
32 34 
33constexpr int64_t DB_BUFFER = 2;35constexpr int64_t DB_BUFFER = 2;
34constexpr int64_t DIM2 = 2;36constexpr int64_t DIM2 = 2;
37+constexpr int64_t DIGIT_ZERO = 0;
38+constexpr int64_t DIGIT_ONE = 1;
35constexpr int64_t DIGIT_TWO = 2;39constexpr int64_t DIGIT_TWO = 2;
36constexpr int64_t DIGIT_FOUR = 4;40constexpr int64_t DIGIT_FOUR = 4;
41+constexpr int64_t DIGIT_EIGHT = 8;
37constexpr int64_t DIGIT_SIXTY_THREE = 63;42constexpr int64_t DIGIT_SIXTY_THREE = 63;
43+constexpr float DIGIT_ZERO_FLOAT = 0.0;
44+constexpr float DIGIT_SIX_FLOAT = 6.0;
45+constexpr float DIGIT_SEVEN_FLOAT = 7.0;
46+constexpr int64_t ModeZero = 0;
47+constexpr int64_t ModeOne = 1;
48+constexpr int64_t ModeTwo = 2;
49+constexpr int64_t ModeThree = 3;
50+ 
51+constexpr uint32_t vfLen16 = platform::GetVRegSize() / sizeof(uint16_t);
52+constexpr uint32_t vfLen16Double = vfLen16 * 2;
53+constexpr uint32_t vfLen32 = platform::GetVRegSize() / sizeof(uint32_t);
54+constexpr int64_t UBBlockSize_ = platform::GetUbBlockSize();
55+constexpr uint16_t elementAfterReduce_ = platform::GetVRegSize() / UBBlockSize_;
56+ 
57+constexpr uint16_t ADD_VALUE_FOR_BF16_MAN1 = 0x003f;
58+constexpr uint16_t ADD_VALUE_FOR_BF16_MAN2 = 0x001f;
38constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64;59constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64;
39constexpr int64_t OUT_ELE_NUM_ONE_BLK_FP8 = 32;60constexpr int64_t OUT_ELE_NUM_ONE_BLK_FP8 = 32;
40-constexpr int64_t OUT_ALL = 256;
41constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81;61constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81;
42constexpr uint32_t NAN_CUSTOMIZATION_FP32 = 0x7f810000;62constexpr uint32_t NAN_CUSTOMIZATION_FP32 = 0x7f810000;
43constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80;63constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80;
@@ -73,6 +93,9 @@ constexpr int32_t SCALE_BUFFER_SIZE = 16 * 1024;
73constexpr int32_t MAX_MTE_BLOCK_COUNT = 4095;93constexpr int32_t MAX_MTE_BLOCK_COUNT = 4095;
74constexpr uint16_t NAN_CUSTOMIZATION_PACK = 0x00007f81;94constexpr uint16_t NAN_CUSTOMIZATION_PACK = 0x00007f81;
75constexpr uint16_t ABS_MASK_FOR_16BIT = 0x7fff;95constexpr uint16_t ABS_MASK_FOR_16BIT = 0x7fff;
96+constexpr uint32_t ABS_MASK_FOR_32BIT = 0x7fffffff;
97+constexpr uint32_t SUB_NUM_FOR_SCALE_32BIT = 0x000000e1;
98+constexpr uint16_t SUB_NUM_FOR_SCALE_16BIT = 0x00e1;
76constexpr uint32_t MAN_MASK_FLOAT = 0x007fffff;99constexpr uint32_t MAN_MASK_FLOAT = 0x007fffff;
77constexpr uint32_t FP32_EXP_BIAS_CUBLAS = 0x00007f00;100constexpr uint32_t FP32_EXP_BIAS_CUBLAS = 0x00007f00;
78constexpr uint32_t FP8_E5M2_MAX = 0x37924925; // 1/57344的float32表示 57334是E5M2所能表示的最大值101constexpr uint32_t FP8_E5M2_MAX = 0x37924925; // 1/57344的float32表示 57334是E5M2所能表示的最大值
@@ -114,6 +137,26 @@ __aicore__ inline constexpr T GetMaxExp()
114 }137 }
115}138}
116 139 
140+template <typename T>
141+__aicore__ inline constexpr T GetabsForX()
142+{
143+ if constexpr (IsSame<T, uint16_t>::value) {
144+ return ABS_MASK_FOR_16BIT;
145+ } else {
146+ return ABS_MASK_FOR_32BIT;
147+ }
148+}
149+ 
150+template <typename T>
151+__aicore__ inline constexpr T GetSubNumForScale()
152+{
153+ if constexpr (IsSame<T, uint16_t>::value) {
154+ return SUB_NUM_FOR_SCALE_32BIT;
155+ } else {
156+ return SUB_NUM_FOR_SCALE_16BIT;
157+ }
158+}
159+ 
117template <typename T>160template <typename T>
118__aicore__ inline constexpr T GetFp4MaxExp()161__aicore__ inline constexpr T GetFp4MaxExp()
119{162{
@@ -166,77 +209,72 @@ __aicore__ inline constexpr T GetSpecialExp()
166 209 
167template <AscendC::RoundMode roundMode, typename outType, typename inType>210template <AscendC::RoundMode roundMode, typename outType, typename inType>
168__aicore__ inline void CalcElement(211__aicore__ inline void CalcElement(
169- AscendC::MicroAPI::RegTensor<inType>& in, AscendC::MicroAPI::RegTensor<int32_t>& maxEle,212+ AscendC::Reg::RegTensor<inType>& in, AscendC::Reg::RegTensor<int32_t>& maxEle, AscendC::Reg::MaskReg mask)
170- AscendC::MicroAPI::MaskReg mask)
171{213{
172- AscendC::MicroAPI::RegTensor<float> y1;214+ AscendC::Reg::RegTensor<float> y1;
173- AscendC::MicroAPI::MaskReg negValueMask;215+ AscendC::Reg::MaskReg negValueMask;
174- AscendC::MicroAPI::MaskReg zeroMask;216+ AscendC::Reg::MaskReg zeroMask;
175- AscendC::MicroAPI::MaskReg negZeroMask;217+ AscendC::Reg::MaskReg negZeroMask;
176- AscendC::MicroAPI::MaskReg zeroNegMask;218+ AscendC::Reg::MaskReg zeroNegMask;
177- AscendC::MicroAPI::RegTensor<int32_t> negZero;219+ AscendC::Reg::RegTensor<int32_t> negZero;
178- AscendC::MicroAPI::Duplicate(negZero, NEG_ZERO);220+ AscendC::Reg::Duplicate(negZero, NEG_ZERO);
179- AscendC::MicroAPI::CompareScalar<int32_t, AscendC::CMPMODE::EQ>(221+ AscendC::Reg::CompareScalar<int32_t, AscendC::CMPMODE::EQ>(
180- zeroNegMask, (AscendC::MicroAPI::RegTensor<int32_t>&)in, NEG_ZERO, mask);222+ zeroNegMask, (AscendC::Reg::RegTensor<int32_t>&)in, NEG_ZERO, mask);
181 if constexpr (IsSame<outType, fp4x2_e2m1_t>::value) {223 if constexpr (IsSame<outType, fp4x2_e2m1_t>::value) {
182- AscendC::MicroAPI::RegTensor<int32_t> exp1;224+ AscendC::Reg::RegTensor<int32_t> exp1;
183- AscendC::MicroAPI::RegTensor<int32_t> exp2;225+ AscendC::Reg::RegTensor<int32_t> exp2;
184- AscendC::MicroAPI::And(exp1, (AscendC::MicroAPI::RegTensor<int32_t>&)in, maxEle, mask);226+ AscendC::Reg::And(exp1, (AscendC::Reg::RegTensor<int32_t>&)in, maxEle, mask);
185- AscendC::MicroAPI::ShiftRights(exp1, exp1, SHR_NUM_FOR_FP32, mask);227+ AscendC::Reg::ShiftRights(exp1, exp1, SHR_NUM_FOR_FP32, mask);
186- AscendC::MicroAPI::Adds(exp1, exp1, FP32_BIAS_NEG, mask);228+ AscendC::Reg::Adds(exp1, exp1, FP32_BIAS_NEG, mask);
187- AscendC::MicroAPI::Maxs(exp1, exp1, 0, mask);229+ AscendC::Reg::Maxs(exp1, exp1, 0, mask);
188- AscendC::MicroAPI::Adds(exp1, exp1, NEG_ONE, mask);230+ AscendC::Reg::Adds(exp1, exp1, NEG_ONE, mask);
189- AscendC::MicroAPI::Muls(exp2, exp1, NEG_ONE, mask);231+ AscendC::Reg::Muls(exp2, exp1, NEG_ONE, mask);
190- AscendC::MicroAPI::Adds(exp2, exp2, FP32_BIAS, mask);232+ AscendC::Reg::Adds(exp2, exp2, FP32_BIAS, mask);
191- AscendC::MicroAPI::ShiftLefts(exp2, exp2, SHR_NUM_FOR_FP32, mask);233+ AscendC::Reg::ShiftLefts(exp2, exp2, SHR_NUM_FOR_FP32, mask);
192 234 
193- AscendC::MicroAPI::Mul(y1, in, (AscendC::MicroAPI::RegTensor<float>&)exp2, mask);235+ AscendC::Reg::Mul(y1, in, (AscendC::Reg::RegTensor<float>&)exp2, mask);
194- AscendC::MicroAPI::Adds(exp1, exp1, FP32_BIAS, mask);236+ AscendC::Reg::Adds(exp1, exp1, FP32_BIAS, mask);
195- AscendC::MicroAPI::ShiftLefts(exp1, exp1, SHR_NUM_FOR_FP32, mask);237+ AscendC::Reg::ShiftLefts(exp1, exp1, SHR_NUM_FOR_FP32, mask);
196- AscendC::MicroAPI::CompareScalar<float, AscendC::CMPMODE::LT>(negValueMask, y1, 0, mask);238+ AscendC::Reg::CompareScalar<float, AscendC::CMPMODE::LT>(negValueMask, y1, 0, mask);
197- AscendC::MicroAPI::Truncate<float, roundMode>(y1, y1, mask);239+ AscendC::Reg::Truncate<float, roundMode>(y1, y1, mask);
198- AscendC::MicroAPI::Mul(in, y1, (AscendC::MicroAPI::RegTensor<float>&)exp1, mask);240+ AscendC::Reg::Mul(in, y1, (AscendC::Reg::RegTensor<float>&)exp1, mask);
199 } else {241 } else {
200- AscendC::MicroAPI::Muls(y1, in, FOUR, mask);242+ AscendC::Reg::Muls(y1, in, FOUR, mask);
201- AscendC::MicroAPI::CompareScalar<float, AscendC::CMPMODE::LT>(negValueMask, y1, 0, mask);243+ AscendC::Reg::CompareScalar<float, AscendC::CMPMODE::LT>(negValueMask, y1, 0, mask);
202- AscendC::MicroAPI::Truncate<float, roundMode>(y1, y1, mask);244+ AscendC::Reg::Truncate<float, roundMode>(y1, y1, mask);
203- AscendC::MicroAPI::Muls(in, y1, ONE_FOURTH, mask);245+ AscendC::Reg::Muls(in, y1, ONE_FOURTH, mask);
204 }246 }
205- AscendC::MicroAPI::CompareScalar<float, AscendC::CMPMODE::EQ>(zeroMask, in, 0, mask);247+ AscendC::Reg::CompareScalar<float, AscendC::CMPMODE::EQ>(zeroMask, in, 0, mask);
206- AscendC::MicroAPI::MaskAnd(negZeroMask, zeroMask, negValueMask, mask);248+ AscendC::Reg::MaskAnd(negZeroMask, zeroMask, negValueMask, mask);
207- AscendC::MicroAPI::MaskOr(zeroMask, negZeroMask, zeroNegMask, mask);249+ AscendC::Reg::MaskOr(zeroMask, negZeroMask, zeroNegMask, mask);
208- AscendC::MicroAPI::Copy((AscendC::MicroAPI::RegTensor<int32_t>&)in, negZero, zeroMask);250+ AscendC::Reg::Copy((AscendC::Reg::RegTensor<int32_t>&)in, negZero, zeroMask);
209}251}
210 252 
211template <AscendC::RoundMode roundMode, typename outType, typename inType, typename calcTypeInt>253template <AscendC::RoundMode roundMode, typename outType, typename inType, typename calcTypeInt>
212__aicore__ inline void CalcElement(254__aicore__ inline void CalcElement(
213- AscendC::MicroAPI::RegTensor<inType>& in, AscendC::MicroAPI::RegTensor<calcTypeInt>& scaleReprocal,255+ AscendC::Reg::RegTensor<inType>& in, AscendC::Reg::RegTensor<calcTypeInt>& scaleReprocal,
214- AscendC::MicroAPI::RegTensor<calcTypeInt>& maxEle, AscendC::MicroAPI::RegTensor<uint8_t>& out,256+ AscendC::Reg::RegTensor<calcTypeInt>& maxEle, AscendC::Reg::RegTensor<uint8_t>& out, AscendC::Reg::MaskReg mask)
215- AscendC::MicroAPI::MaskReg mask)
216{257{
217- static constexpr AscendC::MicroAPI::CastTrait castTrait = {258+ static constexpr AscendC::Reg::CastTrait castTrait = {
218- AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN,259+ AscendC::Reg::RegLayout::ZERO, AscendC::Reg::SatMode::UNKNOWN, AscendC::Reg::MaskMergeMode::ZEROING, roundMode};
219- AscendC::MicroAPI::MaskMergeMode::ZEROING, roundMode};260+ static constexpr AscendC::Reg::CastTrait castTraitFp32ToBf16 = {
220- static constexpr AscendC::MicroAPI::CastTrait castTraitFp32ToBf16 = {261+ AscendC::Reg::RegLayout::ZERO, AscendC::Reg::SatMode::NO_SAT, AscendC::Reg::MaskMergeMode::ZEROING, roundMode};
221- AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::NO_SAT,262+ AscendC::Reg::RegTensor<bfloat16_t> valueRegTensor;
222- AscendC::MicroAPI::MaskMergeMode::ZEROING, roundMode};263+ AscendC::Reg::RegTensor<outType> y;
223- AscendC::MicroAPI::RegTensor<bfloat16_t> valueRegTensor;264+ AscendC::Reg::RegTensor<uint16_t> yRegTensor;
224- AscendC::MicroAPI::RegTensor<outType> y;265+ AscendC::Reg::Mul(in, in, (AscendC::Reg::RegTensor<inType>&)scaleReprocal, mask);
225- AscendC::MicroAPI::RegTensor<uint16_t> yRegTensor;
226- AscendC::MicroAPI::Mul(in, in, (AscendC::MicroAPI::RegTensor<inType>&)scaleReprocal, mask);
227 if constexpr (IsSame<inType, float>::value) {266 if constexpr (IsSame<inType, float>::value) {
228- CalcElement<roundMode, outType, inType>(in, (AscendC::MicroAPI::RegTensor<int32_t>&)maxEle, mask);267+ CalcElement<roundMode, outType, inType>(in, (AscendC::Reg::RegTensor<int32_t>&)maxEle, mask);
229- AscendC::MicroAPI::Cast<bfloat16_t, inType, castTraitFp32ToBf16>(valueRegTensor, in, mask);268+ AscendC::Reg::Cast<bfloat16_t, inType, castTraitFp32ToBf16>(valueRegTensor, in, mask);
230- AscendC::MicroAPI::Pack(269+ AscendC::Reg::Pack(
231- (AscendC::MicroAPI::RegTensor<uint16_t>&)valueRegTensor,270+ (AscendC::Reg::RegTensor<uint16_t>&)valueRegTensor, (AscendC::Reg::RegTensor<uint32_t>&)valueRegTensor);
232- (AscendC::MicroAPI::RegTensor<uint32_t>&)valueRegTensor);271+ AscendC::Reg::Cast<outType, bfloat16_t, castTrait>(y, valueRegTensor, mask);
233- AscendC::MicroAPI::Cast<outType, bfloat16_t, castTrait>(y, valueRegTensor, mask);
234 } else {272 } else {
235- AscendC::MicroAPI::Cast<outType, inType, castTrait>(y, in, mask);273+ AscendC::Reg::Cast<outType, inType, castTrait>(y, in, mask);
236 }274 }
237 275 
238- AscendC::MicroAPI::Pack(yRegTensor, (AscendC::MicroAPI::RegTensor<uint32_t>&)y);276+ AscendC::Reg::Pack(yRegTensor, (AscendC::Reg::RegTensor<uint32_t>&)y);
239- AscendC::MicroAPI::Pack(out, yRegTensor);277+ AscendC::Reg::Pack(out, yRegTensor);
240}278}
241 279 
242} // namespace DynamicMxQuant280} // namespace DynamicMxQuant
@@ -20,6 +20,7 @@
20#include "dynamic_mx_quant_common.h"20#include "dynamic_mx_quant_common.h"
21#include "kernel_tiling/kernel_tiling.h"21#include "kernel_tiling/kernel_tiling.h"
22#include "op_kernel/platform_util.h"22#include "op_kernel/platform_util.h"
23+#include "dynamic_mx_quant_tilingdata.h"
23 24 
24namespace DynamicMxQuant {25namespace DynamicMxQuant {
25using namespace AscendC;26using namespace AscendC;
@@ -72,11 +73,15 @@ protected:
72 int64_t tailBlockSize_ = 0;73 int64_t tailBlockSize_ = 0;
73 int64_t postAxisSize_ = 0;74 int64_t postAxisSize_ = 0;
74 int64_t mxScaleSize_ = 0;75 int64_t mxScaleSize_ = 0;
76+ int64_t scaleAlg_ = 0;
77+ float dstTypeMax_ = 0;
78+ float invDstTypeMax_ = 0;
75 bool isPad_ = false;79 bool isPad_ = false;
76 bool isTailBlock_ = false;80 bool isTailBlock_ = false;
77 using intCalcType = typename std::conditional<IsSame<T, half>::value, uint32_t, uint16_t>::type;81 using intCalcType = typename std::conditional<IsSame<T, half>::value, uint32_t, uint16_t>::type;
78 constexpr static int16_t shrNum_ = GetShrNum<T>();82 constexpr static int16_t shrNum_ = GetShrNum<T>();
79 constexpr static intCalcType maxExp_ = GetMaxExp<intCalcType>();83 constexpr static intCalcType maxExp_ = GetMaxExp<intCalcType>();
84+ constexpr static intCalcType absForX_ = GetabsForX<intCalcType>();
80 constexpr static intCalcType f4Emax_ = GetFp4MaxExp<intCalcType>();85 constexpr static intCalcType f4Emax_ = GetFp4MaxExp<intCalcType>();
81 constexpr static intCalcType f8Emax_ = GetFp8MaxExp<intCalcType>();86 constexpr static intCalcType f8Emax_ = GetFp8MaxExp<intCalcType>();
82 constexpr static intCalcType maxBias_ = GetMaxBias<intCalcType>();87 constexpr static intCalcType maxBias_ = GetMaxBias<intCalcType>();
@@ -116,6 +121,9 @@ __aicore__ inline void DynamicMxQuantBase<T, U, ISTAIL>::ParseTilingData(const D
116 postAxisSize_ = tilingData->postAxisSize;121 postAxisSize_ = tilingData->postAxisSize;
117 isPad_ = tilingData->isPad == 1;122 isPad_ = tilingData->isPad == 1;
118 mxScaleSize_ = tilingData->mxScaleSize;123 mxScaleSize_ = tilingData->mxScaleSize;
124+ scaleAlg_ = tilingData->scaleAlg;
125+ dstTypeMax_ = tilingData->dstTypeMax;
126+ invDstTypeMax_ = tilingData->invDstTypeMax;
119}127}
120 128 
121template <typename T, typename U, const bool ISTAIL>129template <typename T, typename U, const bool ISTAIL>
Mquant/dynamic_mx_quant/op_kernel/arch35/dynamic_mx_quant_tail_axis.h+927-673文件内容审核中,请稍后刷新重试
@@ -0,0 +1,99 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/* !
12+ * \file dynamic_mx_quant_with_dual_axis_tilingdata.h
13+ * \brief
14+ */
15+ 
16+#ifndef OPS_NN_DYNAMIC_MX_QUANT_H
17+#define OPS_NN_DYNAMIC_MX_QUANT_H
18+ 
19+#include <cstdint>
20+ 
21+struct DynamicMxQuantTilingData {
22+ int64_t totalCoreNum;
23+ int64_t usedCoreNum; // 实际使用的核数
24+ int64_t blockFactor; // 单核循环次数
25+ int64_t tailBlockFactor; // 尾核循环次数
26+ int64_t ubDim; // 合轴后,ubfactor所切的轴
27+ int64_t uo; // 切分轴上的循环次数
28+ int64_t ubFactor; // 单次循环要处理的数据大小
29+ int64_t tailUbFactor; // 尾循环要处理的数据大小
30+ int64_t roundMode; // 数据类型转换的模式
31+ int64_t dstType; // 输出y的数据类型
32+ int64_t blockSize; // 进行微缩的数据块大小
33+ int64_t scaleAlg; // scale计算方法
34+ int64_t blockSizeNumInAxis; // 在axis轴上有多少个blocksize
35+ int64_t tailBlockSize; // 指定轴要进行微缩的最后一个数据块大小
36+ int64_t isPad; // axis指定的轴是否需要补到blocksize的整数倍
37+ int64_t isTailAxis; // 是否为尾轴场景
38+ int64_t preAxisSize; // 合轴后axis前面轴的大小
39+ int64_t postAxisSize; // 合轴后axis后面轴的大小
40+ int64_t mxScaleSize; // scale数据大小
41+ int64_t tilingKey;
42+ float dstTypeMax;
43+ float invDstTypeMax;
44+};
45+ 
46+struct DynamicMxQuant4OptimizeTilingData {
47+ int64_t totalCoreNum; // 总核数
48+ int64_t usedCoreNum; // 实际使用的核数
49+ int64_t roundMode; // 数据类型转换的模式
50+ int64_t dstType; // 输出y的数据类型
51+ int64_t blockSize; // 进行微缩的数据块大小
52+ int64_t isPad; // 量化轴最后一个block无法被blockSize整除时,为True
53+ int64_t tailBlockSize; // 指定量化轴最后一个blockSize大小
54+ int64_t scaleAlg;
55+ int64_t tilingKey;
56+ int64_t quantAxisSize; // 优化非尾轴模板量化轴大小
57+ int64_t preAxisSize; // 合轴后axis前面轴的大小
58+ int64_t postAxisSize; // 合轴后axis后面轴的大小
59+ int64_t mAlignSize; // 量化轴对齐blockSize之后元素个数
60+ int64_t nAlignSize; // 融合尾轴对齐32,64,128之后元素个数
61+ int64_t mAlignBlockCount; // 量化轴对齐blockSize之后block的个数,实际上等于blockSizeNumInAxis
62+ int64_t nAlignBlockCount; // 融合尾轴对齐32,64,128之后block的个数
63+ int64_t mAlignGroupCount; // 量化轴对齐blockSize*2(一个Group)之后Group的个数
64+ int64_t quantAxisIsOdd; // 量化轴是否是奇数,如果是奇数,则有些group会有一个全0的dummy block
65+ int64_t totalGroupNum; // 当前shape总共需要多少个Group才能计算完
66+ int64_t groupPerCore; // 每个核计算多少个Group
67+ int64_t groupPerTail; // 尾核计算多少个Group
68+ int64_t groupPerUb; // 每个UB可以放下多少个Group
69+ int64_t totalBlockNum; // 总共处理的block数量,此处为对齐成group之后的block数量
70+ int64_t blockNumPerTask; // 每个任务处理多少个blcok
71+ int64_t totalTaskNum; // 总任务数量,用总共处理的block数量除以每个任务处理多少个block
72+ int64_t rowPerHeadCore;
73+ int64_t rowPerTailCore;
74+ int64_t needPadPostAxis; // 融合尾轴是否需要对齐
75+ float dstTypeMax;
76+ float invDstTypeMax;
77+};
78+ 
79+struct DynamicMxQuantTailAxisTilingData {
80+ int64_t tilingKey;
81+ int64_t ubSize;
82+ int64_t roundMode;
83+ int64_t blockSize;
84+ int64_t totalCoreNum;
85+ int64_t usedCoreNum;
86+ int64_t rowTileNum; // row 方向上的切核数
87+ int64_t colTileNum; // col 方向上的切核数
88+ int64_t rowNum; // 合轴之后 -2 轴大小
89+ int64_t colNum; // 合轴之后 -1 轴大小
90+ int64_t colNormalBlockNum; // 列方向头核处理的块数 (1 x 256)
91+ int64_t colTailLen; // 列方向尾块长度
92+ int64_t rowNormalBlockNum; // 行方向头核处理的块数 (1 行)
93+ int64_t rowTailLen; // 行方向尾块长度
94+ int64_t maxUbBlockNum; // UB最大能放下的处理块数 (1 x 32) (8 的倍数)
95+ float dstTypeMax;
96+ float invDstTypeMax;
97+};
98+ 
99+#endif // OPS_NN_DYNAMIC_MX_QUANT_WITH_DUAL_AXIS_H