已合并
add dynamic_mx_quant dst_type_value && rewrite tail_axis #3145
pyongq创建于 3月25日
add dynamic_mx_quant dst_type_value && rewrite tail_axis #3145
已合并
共 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 aclTensor | 472 | // 创建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 compatibility | 65 | * @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 | */ |
| 61 | REG_OP(DynamicMxQuant) | 68 | REG_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 ge | 79 | } // namespace ge |
| 75 | 80 | ||
| 76 | 81 | ||
Mquant/dynamic_mx_quant/op_host/arch35/dynamic_mx_quant_optimize_tiling_arch35.cpp+99-456文件内容审核中,请稍后刷新重试
Aquant/dynamic_mx_quant/op_host/arch35/dynamic_mx_quant_tail_axis_tiling_arch35.cpp+230-0文件内容审核中,请稍后刷新重试
| @@ -20,6 +20,7 @@ namespace ops { | |||
| 20 | static constexpr int32_t DEFAULT_BLOCK_SIZE = 32; | 20 | static constexpr int32_t DEFAULT_BLOCK_SIZE = 32; |
| 21 | static constexpr int32_t DEFAULT_DST_TYPE = 40; | 21 | static constexpr int32_t DEFAULT_DST_TYPE = 40; |
| 22 | static constexpr int32_t DEFAULT_SCALE_ALG = 0; | 22 | static constexpr int32_t DEFAULT_SCALE_ALG = 0; |
| 23 | +static constexpr float DEFAULT_DST_TYPE_MAX = 0.0; | ||
| 23 | class DynamicMxQuant : public OpDef { | 24 | class DynamicMxQuant : public OpDef { |
| 24 | public: | 25 | public: |
| 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 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 21 | + | ||
| 20 | namespace DynamicMxQuant { | 22 | namespace DynamicMxQuant { |
| 21 | 23 | ||
| 22 | template <typename Tp, Tp v> | 24 | template <typename Tp, Tp v> |
| @@ -32,12 +34,30 @@ struct IsSame<Tp, Tp> : public trueType {}; | |||
| 32 | 34 | ||
| 33 | constexpr int64_t DB_BUFFER = 2; | 35 | constexpr int64_t DB_BUFFER = 2; |
| 34 | constexpr int64_t DIM2 = 2; | 36 | constexpr int64_t DIM2 = 2; |
| 37 | +constexpr int64_t DIGIT_ZERO = 0; | ||
| 38 | +constexpr int64_t DIGIT_ONE = 1; | ||
| 35 | constexpr int64_t DIGIT_TWO = 2; | 39 | constexpr int64_t DIGIT_TWO = 2; |
| 36 | constexpr int64_t DIGIT_FOUR = 4; | 40 | constexpr int64_t DIGIT_FOUR = 4; |
| 41 | +constexpr int64_t DIGIT_EIGHT = 8; | ||
| 37 | constexpr int64_t DIGIT_SIXTY_THREE = 63; | 42 | constexpr 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; | ||
| 38 | constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64; | 59 | constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64; |
| 39 | constexpr int64_t OUT_ELE_NUM_ONE_BLK_FP8 = 32; | 60 | constexpr int64_t OUT_ELE_NUM_ONE_BLK_FP8 = 32; |
| 40 | -constexpr int64_t OUT_ALL = 256; | ||
| 41 | constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81; | 61 | constexpr uint16_t NAN_CUSTOMIZATION = 0x7f81; |
| 42 | constexpr uint32_t NAN_CUSTOMIZATION_FP32 = 0x7f810000; | 62 | constexpr uint32_t NAN_CUSTOMIZATION_FP32 = 0x7f810000; |
| 43 | constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80; | 63 | constexpr uint16_t MAX_EXP_FOR_BF16 = 0x7f80; |
| @@ -73,6 +93,9 @@ constexpr int32_t SCALE_BUFFER_SIZE = 16 * 1024; | |||
| 73 | constexpr int32_t MAX_MTE_BLOCK_COUNT = 4095; | 93 | constexpr int32_t MAX_MTE_BLOCK_COUNT = 4095; |
| 74 | constexpr uint16_t NAN_CUSTOMIZATION_PACK = 0x00007f81; | 94 | constexpr uint16_t NAN_CUSTOMIZATION_PACK = 0x00007f81; |
| 75 | constexpr uint16_t ABS_MASK_FOR_16BIT = 0x7fff; | 95 | constexpr 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; | ||
| 76 | constexpr uint32_t MAN_MASK_FLOAT = 0x007fffff; | 99 | constexpr uint32_t MAN_MASK_FLOAT = 0x007fffff; |
| 77 | constexpr uint32_t FP32_EXP_BIAS_CUBLAS = 0x00007f00; | 100 | constexpr uint32_t FP32_EXP_BIAS_CUBLAS = 0x00007f00; |
| 78 | constexpr uint32_t FP8_E5M2_MAX = 0x37924925; // 1/57344的float32表示 57334是E5M2所能表示的最大值 | 101 | constexpr 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 | + | ||
| 117 | template <typename T> | 160 | template <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 | ||
| 167 | template <AscendC::RoundMode roundMode, typename outType, typename inType> | 210 | template <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 | ||
| 211 | template <AscendC::RoundMode roundMode, typename outType, typename inType, typename calcTypeInt> | 253 | template <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 DynamicMxQuant | 280 | } // namespace DynamicMxQuant |
| @@ -20,6 +20,7 @@ | |||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 23 | 24 | ||
| 24 | namespace DynamicMxQuant { | 25 | namespace DynamicMxQuant { |
| 25 | using namespace AscendC; | 26 | using 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 | ||
| 121 | template <typename T, typename U, const bool ISTAIL> | 129 | template <typename T, typename U, const bool ISTAIL> |
Mquant/dynamic_mx_quant/op_kernel/arch35/dynamic_mx_quant_not_tail_axis_optimize.h+358-78文件内容审核中,请稍后刷新重试
| @@ -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 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 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 | + | ||
Mquant/dynamic_mx_quant/tests/ut/op_host/arch35/test_dynamic_mx_quant_tiling.cpp+29-19文件内容审核中,请稍后刷新重试