已合并
docs(tbmm,tqbmm): 资料修改(T-C量化约束、batchSplitFactor说明、标点规范等) #10800
jgx创建于 19 天前
docs(tbmm,tqbmm): 资料修改(T-C量化约束、batchSplitFactor说明、标点规范等) #10800
已合并
共 7 个文件变更+50-41
| @@ -124,11 +124,11 @@ | |||
| 124 | - 支持非连续tensor。 | 124 | - 支持非连续tensor。 |
| 125 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 | 125 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 |
| 126 | - 当x1的输入shape为(B, M, K)时,K <= 65535;当x1的输入shape为(M, B, K)时,B * K <= 65535。 | 126 | - 当x1的输入shape为(B, M, K)时,K <= 65535;当x1的输入shape为(M, B, K)时,B * K <= 65535。 |
| 127 | - - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536,且仅支持输入为FLOAT16和输出为INT8的类型推导。 | 127 | + - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536,且仅支持输入为FLOAT16和输出为INT8的类型推导。 |
| 128 | - <term>Ascend 950PR/Ascend 950DT</term>: | 128 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| 129 | - 当scale不为空时,batchSplitFactor只能等于1,且仅支持输入为FLOAT16和输出为INT8的类型推导。 | 129 | - 当scale不为空时,batchSplitFactor只能等于1,且仅支持输入为FLOAT16和输出为INT8的类型推导。 |
| 130 | - bias为预留参数,当前暂不支持。 | 130 | - bias为预留参数,当前暂不支持。 |
| 131 | - | 131 | + |
| 132 | ## 调用说明 | 132 | ## 调用说明 |
| 133 | 133 | ||
| 134 | | 调用方式 | 样例代码 | 说明 | | 134 | | 调用方式 | 样例代码 | 说明 | |
| @@ -232,8 +232,8 @@ aclnnStatus aclnnTransposeBatchMatMul( | |||
| 232 | <ul> | 232 | <ul> |
| 233 | <li>当batchSplitFactor大于1时,out的输出shape为(batchSplitFactor, M, B * N / batchSplitFactor)。</li> | 233 | <li>当batchSplitFactor大于1时,out的输出shape为(batchSplitFactor, M, B * N / batchSplitFactor)。</li> |
| 234 | <ul> | 234 | <ul> |
| 235 | - <li> 示例一: M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 2时,out的输出shape大小为(2, 32, 1024)。</li> | 235 | + <li> 示例一:M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 2时,out的输出shape大小为(2, 32, 1024)。</li> |
| 236 | - <li> 示例二: M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 4时,out的输出shape大小为(4, 32, 512)。</li> | 236 | + <li> 示例二:M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 4时,out的输出shape大小为(4, 32, 512)。</li> |
| 237 | </ul> | 237 | </ul> |
| 238 | </ul> | 238 | </ul> |
| 239 | </td> | 239 | </td> |
| @@ -270,8 +270,8 @@ aclnnStatus aclnnTransposeBatchMatMul( | |||
| 270 | <td>传入的x1、x2或out是空指针。</td> | 270 | <td>传入的x1、x2或out是空指针。</td> |
| 271 | </tr> | 271 | </tr> |
| 272 | <tr> | 272 | <tr> |
| 273 | - <td rowspan="6">ACLNN_ERR_PARAM_INVALID</td> | 273 | + <td rowspan="5">ACLNN_ERR_PARAM_INVALID</td> |
| 274 | - <td rowspan="6">161002</td> | 274 | + <td rowspan="5">161002</td> |
| 275 | <td>x1、x2或out的数据类型不在支持的范围内。</td> | 275 | <td>x1、x2或out的数据类型不在支持的范围内。</td> |
| 276 | </tr> | 276 | </tr> |
| 277 | <tr> | 277 | <tr> |
| @@ -339,9 +339,9 @@ aclnnStatus aclnnTransposeBatchMatMul( | |||
| 339 | <!-- npu="A3,910b" id7 --> | 339 | <!-- npu="A3,910b" id7 --> |
| 340 | - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>: | 340 | - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>: |
| 341 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 | 341 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 |
| 342 | - - 当x1的输入shape为(B, M, K)时,需要K <= 65535;当x1的输入shape为(M, B, K)且B * K > 65535时,不支持传入scale,并且batchSplitFactor只能等于1, permX1必须为[1, 0, 2]。 | 342 | + - 当x1的输入shape为(B, M, K)时,需要K <= 65535;当x1的输入shape为(M, B, K)且B * K > 65535时,不支持传入scale,并且batchSplitFactor只能等于1,permX1必须为[1, 0, 2]。 |
| 343 | - - 当permX2输入为[0, 2, 1]时,不支持传入scale,并且batchSplitFactor只能等于1, permX1必须为[1, 0, 2]。 | 343 | + - 当permX2输入为[0, 2, 1]时,不支持传入scale,并且batchSplitFactor只能等于1,permX1必须为[1, 0, 2]。 |
| 344 | - - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536,且仅支持输入为FLOAT16和输出为INT8的类型推导。 | 344 | + - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536,且仅支持输入为FLOAT16和输出为INT8的类型推导。 |
| 345 | <!-- end id7 --> | 345 | <!-- end id7 --> |
| 346 | <!-- npu="950" id8 --> | 346 | <!-- npu="950" id8 --> |
| 347 | - <term>Ascend 950PR/Ascend 950DT</term>: | 347 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| @@ -112,7 +112,7 @@ aclnnStatus aclnnTransposeBatchMatMulWeightNz( | |||
| 112 | <li>数据类型需要与x1满足数据类型推导规则(参见<a href="../../../docs/zh/context/deduction_relationship.md">互推导关系</a>和<a href="#约束说明">约束说明</a>)。</li> | 112 | <li>数据类型需要与x1满足数据类型推导规则(参见<a href="../../../docs/zh/context/deduction_relationship.md">互推导关系</a>和<a href="#约束说明">约束说明</a>)。</li> |
| 113 | <li>x2的Reduce维度需要与x1的Reduce维度大小相等。</li> | 113 | <li>x2的Reduce维度需要与x1的Reduce维度大小相等。</li> |
| 114 | <li>不支持输入x1,x2分别为BFLOAT16和FLOAT16的数据类型推导。</li> | 114 | <li>不支持输入x1,x2分别为BFLOAT16和FLOAT16的数据类型推导。</li> |
| 115 | - <li>NZ格式各个维度表示:(b, n1,k1,k0,n0),其中k0 = 16, n0为16。x1 shape中的k和x2 shape中的k1需要满足以下关系:ceil(k,k0) = k1,x2 shape中的n1与out的n满足以下关系: ceil(n, n0) = n1。</li> | 115 | + <li>NZ格式各个维度表示:(b,n1,k1,k0,n0),其中k0 = 16,n0为16。x1 shape中的k和x2 shape中的k1需要满足以下关系:ceil(k, k0) = k1,x2 shape中的n1与out的n满足以下关系:ceil(n, n0) = n1。</li> |
| 116 | </ul> | 116 | </ul> |
| 117 | </td> | 117 | </td> |
| 118 | <td>BFLOAT16、FLOAT16</td> | 118 | <td>BFLOAT16、FLOAT16</td> |
| @@ -232,8 +232,8 @@ aclnnStatus aclnnTransposeBatchMatMulWeightNz( | |||
| 232 | <ul> | 232 | <ul> |
| 233 | <li>当batchSplitFactor大于1时,out的输出shape为(batchSplitFactor, M, B * N / batchSplitFactor)。</li> | 233 | <li>当batchSplitFactor大于1时,out的输出shape为(batchSplitFactor, M, B * N / batchSplitFactor)。</li> |
| 234 | <ul> | 234 | <ul> |
| 235 | - <li> 示例一: M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 2时,out的输出shape大小为(2, 32, 1024)。</li> | 235 | + <li> 示例一:M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 2时,out的输出shape大小为(2, 32, 1024)。</li> |
| 236 | - <li> 示例二: M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 4时,out的输出shape大小为(4, 32, 512)。</li> | 236 | + <li> 示例二:M, K, N, B = 32, 512, 128, 16;batchSplitFactor = 4时,out的输出shape大小为(4, 32, 512)。</li> |
| 237 | </ul> | 237 | </ul> |
| 238 | </ul> | 238 | </ul> |
| 239 | </td> | 239 | </td> |
| @@ -290,8 +290,8 @@ aclnnStatus aclnnTransposeBatchMatMulWeightNz( | |||
| 290 | <td>传入的x1、x2或out是空指针。</td> | 290 | <td>传入的x1、x2或out是空指针。</td> |
| 291 | </tr> | 291 | </tr> |
| 292 | <tr> | 292 | <tr> |
| 293 | - <td rowspan="7">ACLNN_ERR_PARAM_INVALID</td> | 293 | + <td rowspan="6">ACLNN_ERR_PARAM_INVALID</td> |
| 294 | - <td rowspan="7">161002</td> | 294 | + <td rowspan="6">161002</td> |
| 295 | <td>x1、x2或out的数据类型不在支持的范围内。</td> | 295 | <td>x1、x2或out的数据类型不在支持的范围内。</td> |
| 296 | </tr> | 296 | </tr> |
| 297 | <tr> | 297 | <tr> |
| @@ -376,8 +376,8 @@ aclnnStatus aclnnTransposeBatchMatMulWeightNz( | |||
| 376 | - 当scale不为空时,batchSplitFactor只能等于1,且仅支持输入为FLOAT16和输出为INT8的类型推导。 | 376 | - 当scale不为空时,batchSplitFactor只能等于1,且仅支持输入为FLOAT16和输出为INT8的类型推导。 |
| 377 | 377 | ||
| 378 | <!-- end id8 --> | 378 | <!-- end id8 --> |
| 379 | -- self只支持3维, mat2只支持昇腾私有格式,调用此接口之前,必须完成mat2从ND到昇腾私有格式的转换。 | 379 | +- x1只支持3维,x2只支持昇腾私有格式,调用此接口之前,必须完成x2从ND到昇腾私有格式的转换。 |
| 380 | -- 不支持mat2最后两根轴其中一根轴为1,即k=1或者n=1。 | 380 | +- 不支持x2最后两根轴其中一根轴为1,即k=1或者n=1。 |
| 381 | 381 | ||
| 382 | ## 调用示例 | 382 | ## 调用示例 |
| 383 | 383 | ||
| @@ -13,7 +13,7 @@ | |||
| 13 | 13 | ||
| 14 | ## 功能说明 | 14 | ## 功能说明 |
| 15 | 15 | ||
| 16 | -- 算子功能:完成张量x1与张量x2量化的矩阵乘计算,支持K-C、MX[量化模式](../../docs/zh/context/quant_mode_introduction.md)。仅支持三维的Tensor传入。Tensor支持转置,转置序列根据传入的序列进行变更。permX1代表张量x1的转置序列,支持[1,0,2],permX2代表张量x2的转置序列,K-C模式支持[0,1,2],MX模式支持[0,1,2]和[0,2,1],permY表示矩阵乘输出矩阵的转置序列,当前仅支持[1,0,2],序列值为0的是batch维度,其余两个维度做矩阵乘法。x1Scale和x2Scale表示输出矩阵的量化系数;bias为预留参数,当前暂不支持,详细约束条件可见约束说明或者[aclnnTransposeQuantBatchMatMul](docs/aclnnTransposeQuantBatchMatMul.md)调用说明文档。 | 16 | +- 算子功能:完成张量x1与张量x2量化的矩阵乘计算,支持K-C、MX、T-C[量化模式](../../docs/zh/context/quant_mode_introduction.md)。仅支持三维的Tensor传入。Tensor支持转置,转置序列根据传入的序列进行变更。permX1代表张量x1的转置序列,支持[1,0,2],permX2代表张量x2的转置序列,K-C模式支持[0,1,2],MX和T-C模式支持[0,1,2]和[0,2,1],permY表示矩阵乘输出矩阵的转置序列,当前仅支持[1,0,2],序列值为0的是batch维度,其余两个维度做矩阵乘法。x1Scale和x2Scale表示输出矩阵的量化系数;bias为预留参数,当前暂不支持,详细约束条件可见约束说明或者[aclnnTransposeQuantBatchMatMul](docs/aclnnTransposeQuantBatchMatMul.md)调用说明文档。 |
| 17 | 17 | ||
| 18 | - 示例: | 18 | - 示例: |
| 19 | 假设x1的shape是[M, B, K],x2的shape是[B, K, N],x1Scale和x2Scale不为None,batchSplitFactor等于1时,计算输出out的shape是[M, B, N]。 | 19 | 假设x1的shape是[M, B, K],x2的shape是[B, K, N],x1Scale和x2Scale不为None,batchSplitFactor等于1时,计算输出out的shape是[M, B, N]。 |
| @@ -40,14 +40,14 @@ | |||
| 40 | <td>x1</td> | 40 | <td>x1</td> |
| 41 | <td>输入</td> | 41 | <td>输入</td> |
| 42 | <td>矩阵乘运算中的左矩阵。</td> | 42 | <td>矩阵乘运算中的左矩阵。</td> |
| 43 | - <td>FLOAT8_E5M2, FLOAT8_E4M3FN</td> | 43 | + <td>FLOAT8_E5M2, FLOAT8_E4M3FN, FLOAT4_E2M1, HIFLOAT8</td> |
| 44 | <td>ND</td> | 44 | <td>ND</td> |
| 45 | </tr> | 45 | </tr> |
| 46 | <tr> | 46 | <tr> |
| 47 | <td>x2</td> | 47 | <td>x2</td> |
| 48 | <td>输入</td> | 48 | <td>输入</td> |
| 49 | <td>矩阵乘运算中的右矩阵。</td> | 49 | <td>矩阵乘运算中的右矩阵。</td> |
| 50 | - <td>FLOAT8_E5M2, FLOAT8_E4M3FN</td> | 50 | + <td>FLOAT8_E5M2, FLOAT8_E4M3FN, FLOAT4_E2M1, HIFLOAT8</td> |
| 51 | <td>ND</td> | 51 | <td>ND</td> |
| 52 | </tr> | 52 | </tr> |
| 53 | <tr> | 53 | <tr> |
| @@ -61,14 +61,14 @@ | |||
| 61 | <td>x1Scale</td> | 61 | <td>x1Scale</td> |
| 62 | <td>输入</td> | 62 | <td>输入</td> |
| 63 | <td>量化参数的缩放因子。</td> | 63 | <td>量化参数的缩放因子。</td> |
| 64 | - <td>FLOAT32、FLOAT8_E8M0</td> | 64 | + <td>FLOAT32、FLOAT8_E8M0、UINT64、INT64</td> |
| 65 | <td>ND</td> | 65 | <td>ND</td> |
| 66 | </tr> | 66 | </tr> |
| 67 | <tr> | 67 | <tr> |
| 68 | <td>x2Scale</td> | 68 | <td>x2Scale</td> |
| 69 | <td>输入</td> | 69 | <td>输入</td> |
| 70 | <td>量化参数的缩放因子。</td> | 70 | <td>量化参数的缩放因子。</td> |
| 71 | - <td>FLOAT32、FLOAT8_E8M0</td> | 71 | + <td>FLOAT32、FLOAT8_E8M0、UINT64、INT64</td> |
| 72 | <td>ND</td> | 72 | <td>ND</td> |
| 73 | </tr> | 73 | </tr> |
| 74 | <tr> | 74 | <tr> |
| @@ -109,7 +109,7 @@ | |||
| 109 | <tr> | 109 | <tr> |
| 110 | <td>batchSplitFactor</td> | 110 | <td>batchSplitFactor</td> |
| 111 | <td>输入</td> | 111 | <td>输入</td> |
| 112 | - <td>用于指定矩阵乘输出矩阵中N维的切分大小。</td> | 112 | + <td>用于指定矩阵乘输出矩阵中B维的切分大小,当前仅支持取值为1。</td> |
| 113 | <td>INT32</td> | 113 | <td>INT32</td> |
| 114 | <td>-</td> | 114 | <td>-</td> |
| 115 | </tr> | 115 | </tr> |
| @@ -117,7 +117,7 @@ | |||
| 117 | <td>y</td> | 117 | <td>y</td> |
| 118 | <td>输出</td> | 118 | <td>输出</td> |
| 119 | <td>矩阵乘运算的计算结果。</td> | 119 | <td>矩阵乘运算的计算结果。</td> |
| 120 | - <td>FLOAT16, BFLOAT16</td> | 120 | + <td>FLOAT16, BFLOAT16, HIFLOAT8</td> |
| 121 | <td>ND</td> | 121 | <td>ND</td> |
| 122 | </tr> | 122 | </tr> |
| 123 | </tbody></table> | 123 | </tbody></table> |
| @@ -127,10 +127,19 @@ | |||
| 127 | ## 约束说明 | 127 | ## 约束说明 |
| 128 | 128 | ||
| 129 | - <term>Ascend 950PR/Ascend 950DT</term>: | 129 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| 130 | - - permX1和permY支持[1, 0, 2] | 130 | + - permX1和permY支持[1, 0, 2]。 |
| 131 | - - K-C量化场景,permX2支持输入[0, 1, 2];MX量化场景,permX2支持输入[0, 1, 2]或[0, 2, 1]。 | 131 | + - K-C量化场景,permX2支持输入[0, 1, 2];MX和T-C量化场景,permX2支持输入[0, 1, 2]或[0, 2, 1]。 |
| 132 | - - K-C量化场景,K仅支持512,N仅支持128。x1Scale和x2Scale为1维,并且x1Scale为(M,), x2Scale为(N,),group_size仅支持配置为0,其他取值不生效。 | 132 | + - K-C[量化模式](../../docs/zh/context/quant_mode_introduction.md),K仅支持512,N仅支持128。x1Scale和x2Scale为1维,并且x1Scale为(M,),x2Scale为(N,),group_size仅支持配置为0,其他取值不生效。x1/x2输入支持FLOAT8_E5M2、FLOAT8_E4M3FN两种类型,x1Scale/x2Scale仅支持FLOAT32类型。 |
| 133 | - - MX量化场景,K仅支持64的倍数。x1Scale和x2Scale为4维,并且x1Scale为(M, B, K/64, 2), x2Scale为(B, K/64, N, 2)或(B, N, K/64, 2),group_size的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。 | 133 | + - MX[量化模式](../../docs/zh/context/quant_mode_introduction.md),支持MXFP8和MXFP4两种数据类型。K仅支持64的倍数。x1Scale和x2Scale为4维,并且x1Scale为(M, B, K/64, 2),x2Scale为(B, K/64, N, 2)或(B, N, K/64, 2),group_size的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。x1/x2输入支持FLOAT8_E4M3FN、FLOAT4_E2M1数据类型,x1Scale/x2Scale仅支持FLOAT8_E8M0类型。 |
| 134 | + - T-C[量化模式](../../docs/zh/context/quant_mode_introduction.md),仅支持静态量化。x1Scale支持配置为空或(1,),非空时仅支持UINT64/INT64类型;x2Scale为1维且为(N,),需经过[trans_quant_param](../../quant/trans_quant_param_v2/docs/aclnnTransQuantParamV2.md)预处理转换为UINT64/INT64类型;x1/x2仅支持HIFLOAT8类型,group_size配置不生效。 | ||
| 135 | + - groupSize相关约束: | ||
| 136 | + - 仅在MX量化场景中生效。 | ||
| 137 | + - 传入的groupSize内部会按如下公式分解得到groupSizeM、groupSizeN、groupSizeK,groupSizeM和groupSizeN取值为0时由接口根据scale的shape推断,groupSizeK仅支持32。原理:假设groupSizeM=0,表示M方向量化分组值由接口推断,推断公式为groupSizeM = M / scaleM(需保证M能被scaleM整除),其中M与x1 shape中的M一致,scaleM与x1Scale shape中的M一致。 | ||
| 138 | + | ||
| 139 | + $$ | ||
| 140 | + groupSize = groupSizeK | groupSizeN << 16 | groupSizeM << 32 | ||
| 141 | + $$ | ||
| 142 | + - 不支持空Tensor。 | ||
| 134 | 143 | ||
| 135 | ## 调用说明 | 144 | ## 调用说明 |
| 136 | 145 | ||
| @@ -289,11 +289,11 @@ aclnnStatus aclnnTransposeQuantBatchMatMul( | |||
| 289 | <tr> | 289 | <tr> |
| 290 | <td>ACLNN_ERR_PARAM_NULLPTR</td> | 290 | <td>ACLNN_ERR_PARAM_NULLPTR</td> |
| 291 | <td>161001</td> | 291 | <td>161001</td> |
| 292 | - <td>传入的x1、x2、out、x1Scale、x2Scale、permX1、permX2、permY是空指针。</td> | 292 | + <td>传入的x1、x2、out、x2Scale、permX1、permX2、permY是空指针;除T-C量化场景外,x1Scale是空指针。</td> |
| 293 | </tr> | 293 | </tr> |
| 294 | <tr> | 294 | <tr> |
| 295 | - <td rowspan="8">ACLNN_ERR_PARAM_INVALID</td> | 295 | + <td rowspan="5">ACLNN_ERR_PARAM_INVALID</td> |
| 296 | - <td rowspan="8">161002</td> | 296 | + <td rowspan="5">161002</td> |
| 297 | <td>x1、x2、x1Scale、x2Scale或out的数据类型不在支持的范围内。</td> | 297 | <td>x1、x2、x1Scale、x2Scale或out的数据类型不在支持的范围内。</td> |
| 298 | </tr> | 298 | </tr> |
| 299 | <tr> | 299 | <tr> |
| @@ -303,7 +303,7 @@ aclnnStatus aclnnTransposeQuantBatchMatMul( | |||
| 303 | <td>x1、x2、permX1、permX2、permY的维度大小不等于3。</td> | 303 | <td>x1、x2、permX1、permX2、permY的维度大小不等于3。</td> |
| 304 | </tr> | 304 | </tr> |
| 305 | <tr> | 305 | <tr> |
| 306 | - <td>batchSplitFactor不在支持的范围内</td> | 306 | + <td>batchSplitFactor不在支持的范围内。</td> |
| 307 | </tr> | 307 | </tr> |
| 308 | <tr> | 308 | <tr> |
| 309 | <td>permX1、permX2、permY的取值不在支持的范围内。</td> | 309 | <td>permX1、permX2、permY的取值不在支持的范围内。</td> |
| @@ -363,8 +363,8 @@ aclnnStatus aclnnTransposeQuantBatchMatMul( | |||
| 363 | <!-- npu="950" id7 --> | 363 | <!-- npu="950" id7 --> |
| 364 | - <term>Ascend 950PR/Ascend 950DT</term>: | 364 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| 365 | - K-C[量化模式](../../../docs/zh/context/quant_mode_introduction.md),K仅支持512,N仅支持128。x1Scale和x2Scale仅支持1维,并且x1Scale要求shape为(M,), x2Scale要求shape为(N,),groupSize仅支持配置为0,其他取值不生效。x1/x2输入支持FLOAT8_E5M2、FLOAT8_E4M3FN两种类型,x1Scale/x2Scale仅支持FLOAT32类型。 | 365 | - K-C[量化模式](../../../docs/zh/context/quant_mode_introduction.md),K仅支持512,N仅支持128。x1Scale和x2Scale仅支持1维,并且x1Scale要求shape为(M,), x2Scale要求shape为(N,),groupSize仅支持配置为0,其他取值不生效。x1/x2输入支持FLOAT8_E5M2、FLOAT8_E4M3FN两种类型,x1Scale/x2Scale仅支持FLOAT32类型。 |
| 366 | - - MX[量化模式](../../../docs/zh/context/quant_mode_introduction.md),支持MXFP8和MXFP4两种数据类型。K仅支持64的倍数。 x1Scale和x2Scale仅支持4维,并且x1Scale要求shape为(M, B, K/64, 2),当permX2为[0, 1, 2]时,x2Scale要求shape为(B, K/64, N, 2);当permX2为[0, 2, 1]时,x2Scale要求shape为(B, N, K/64, 2)。groupSize的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。x1/x2输入支持FLOAT8_E4M3FN、FLOAT4_E2M1数据类型,x1Scale/x2Scale仅支持FLOAT8_E8M0类型。 | 366 | + - MX[量化模式](../../../docs/zh/context/quant_mode_introduction.md),支持MXFP8和MXFP4两种数据类型。K仅支持64的倍数。x1Scale和x2Scale仅支持4维,并且x1Scale要求shape为(M, B, K/64, 2),当permX2为[0, 1, 2]时,x2Scale要求shape为(B, K/64, N, 2);当permX2为[0, 2, 1]时,x2Scale要求shape为(B, N, K/64, 2)。groupSize的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。x1/x2输入支持FLOAT8_E4M3FN、FLOAT4_E2M1数据类型,x1Scale/x2Scale仅支持FLOAT8_E8M0类型。 |
| 367 | - - T-C[量化模式](../../../docs/zh/context/quant_mode_introduction.md),仅支持静态量化,x1Scale支持配置为空或(1,),x2Scale要求shape为(N,),groupSize配置不生效。x1Scale非空时仅支持INT64/INT64类型,x2Scale需经过[trans_quant_param](../../../quant/trans_quant_param_v2/docs/aclnnTransQuantParamV2.md)预处理转换为UINT64/INT64类型,x1/x2仅支持HIFLOAT8类型。 | 367 | + - T-C[量化模式](../../../docs/zh/context/quant_mode_introduction.md),仅支持静态量化,x1Scale支持配置为空或(1,),x2Scale要求shape为(N,),groupSize配置不生效。x1Scale非空时仅支持UINT64/INT64类型,x2Scale需经过[trans_quant_param](../../../quant/trans_quant_param_v2/docs/aclnnTransQuantParamV2.md)预处理转换为UINT64/INT64类型,x1/x2仅支持HIFLOAT8类型。 |
| 368 | - groupSize相关约束: | 368 | - groupSize相关约束: |
| 369 | - 仅在MX量化场景中生效。 | 369 | - 仅在MX量化场景中生效。 |
| 370 | - 传入的groupSize内部会按如下公式分解得到groupSizeM、groupSizeN、groupSizeK,当其中有1个或多个为0,会根据x1/x2/x1Scale/x2Scale输入shape重新设置groupSizeM、groupSizeN、groupSizeK用于计算。原理:假设groupSizeM=0,表示M方向量化分组值由接口推断,推断公式为groupSizeM = M / scaleM(需保证M能被scaleM整除),其中M与x1 shape中的M一致,scaleM与x1Scale shape中的M一致。 | 370 | - 传入的groupSize内部会按如下公式分解得到groupSizeM、groupSizeN、groupSizeK,当其中有1个或多个为0,会根据x1/x2/x1Scale/x2Scale输入shape重新设置groupSizeM、groupSizeN、groupSizeK用于计算。原理:假设groupSizeM=0,表示M方向量化分组值由接口推断,推断公式为groupSizeM = M / scaleM(需保证M能被scaleM整除),其中M与x1 shape中的M一致,scaleM与x1Scale shape中的M一致。 |
| @@ -277,8 +277,8 @@ aclnnStatus aclnnTransposeQuantBatchMatMulWeightNz( | |||
| 277 | <td>传入的x1、x2、out、x1Scale、x2Scale、permX1、permX2、permY是空指针。</td> | 277 | <td>传入的x1、x2、out、x1Scale、x2Scale、permX1、permX2、permY是空指针。</td> |
| 278 | </tr> | 278 | </tr> |
| 279 | <tr> | 279 | <tr> |
| 280 | - <td rowspan="8">ACLNN_ERR_PARAM_INVALID</td> | 280 | + <td rowspan="6">ACLNN_ERR_PARAM_INVALID</td> |
| 281 | - <td rowspan="8">161002</td> | 281 | + <td rowspan="6">161002</td> |
| 282 | <td>x1、x2、x1Scale、x2Scale或out的数据类型不在支持的范围内。</td> | 282 | <td>x1、x2、x1Scale、x2Scale或out的数据类型不在支持的范围内。</td> |
| 283 | </tr> | 283 | </tr> |
| 284 | <tr> | 284 | <tr> |
| @@ -288,7 +288,7 @@ aclnnStatus aclnnTransposeQuantBatchMatMulWeightNz( | |||
| 288 | <td>x1或x2的ViewShape的维度大小不等于3。</td> | 288 | <td>x1或x2的ViewShape的维度大小不等于3。</td> |
| 289 | </tr> | 289 | </tr> |
| 290 | <tr> | 290 | <tr> |
| 291 | - <td>batchSplitFactor不在支持的范围内</td> | 291 | + <td>batchSplitFactor不在支持的范围内。</td> |
| 292 | </tr> | 292 | </tr> |
| 293 | <tr> | 293 | <tr> |
| 294 | <td>permX1、permX2、permY的取值不在支持的范围内。</td> | 294 | <td>permX1、permX2、permY的取值不在支持的范围内。</td> |
| @@ -350,7 +350,7 @@ aclnnStatus aclnnTransposeQuantBatchMatMulWeightNz( | |||
| 350 | 350 | ||
| 351 | <!-- npu="950" id7 --> | 351 | <!-- npu="950" id7 --> |
| 352 | - <term>Ascend 950PR/Ascend 950DT</term>: | 352 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| 353 | - - x1只支持3维, x2只支持昇腾私有格式,调用此接口之前,必须完成x2从ND到昇腾私有格式的转换。 | 353 | + - x1只支持3维,x2只支持昇腾私有格式,调用此接口之前,必须完成x2从ND到昇腾私有格式的转换。 |
| 354 | - K仅支持64的倍数。groupSize的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。 | 354 | - K仅支持64的倍数。groupSize的groupSizeM和groupSizeN仅支持0或1,groupSizeK仅支持32。 |
| 355 | - groupSize相关约束: | 355 | - groupSize相关约束: |
| 356 | - 仅在MX量化场景中生效。 | 356 | - 仅在MX量化场景中生效。 |
| @@ -30,7 +30,7 @@ | |||
| 30 | y[m, n] = \sum_{j=0}^{K/32-1} \left(\left(\sum_{k=0}^{31} x1[m, j \times 32 + k] \cdot x2[j \times 32 + k, n]\right) \cdot x1Scale[m, j] \cdot x2Scale[j, n]\right) | 30 | y[m, n] = \sum_{j=0}^{K/32-1} \left(\left(\sum_{k=0}^{31} x1[m, j \times 32 + k] \cdot x2[j \times 32 + k, n]\right) \cdot x1Scale[m, j] \cdot x2Scale[j, n]\right) |
| 31 | $$ | 31 | $$ |
| 32 | 32 | ||
| 33 | - 其中K为矩阵乘的K轴长度,x1Scale、x2Scale为FLOAT8_E8M0编码的MX量化缩放因子,矩阵乘中间结果按K轴每32个元素一组进行缩放累加。 | 33 | + 其中K为矩阵乘的K轴长度,x1Scale、x2Scale为torch.float8_e8m0fnu编码的MX量化缩放因子,矩阵乘中间结果按K轴每32个元素一组进行缩放累加。 |
| 34 | 34 | ||
| 35 | - 示例:假设x1的shape是(M, B, K),x2的shape是(B, K, N),输出y的shape是(M, B, N)。 | 35 | - 示例:假设x1的shape是(M, B, K),x2的shape是(B, K, N),输出y的shape是(M, B, N)。 |
| 36 | 36 | ||
| @@ -61,7 +61,7 @@ cann_ops_nn.transpose_quant_batch_mat_mul( | |||
| 61 | 61 | ||
| 62 | | 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) | | 62 | | 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) | |
| 63 | | --- | --- | --- | --- | --- | --- | | 63 | | --- | --- | --- | --- | --- | --- | |
| 64 | -| x1 | Tensor | 必选 | 矩阵乘运算中的左矩阵,shape为(M, B, K)。MXFP4场景Tensor最后一维为FP4拼包后的物理长度K/2。 | torch.float8_e4m3fn;MXFP4场景为实际存储类型(如torch.uint8),并通过x1_dtype指定为FLOAT4_E2M1 | 3维,(M, B, K) | | 64 | +| x1 | Tensor | 必选 | 矩阵乘运算中的左矩阵,shape为(M, B, K)。MXFP4场景Tensor最后一维为FP4拼包后的物理长度K/2。 | torch.float8_e4m3fn;MXFP4场景为实际存储类型(如torch.uint8),并通过x1_dtype指定为torch_npu.float4_e2m1fn_x2 | 3维,(M, B, K) | |
| 65 | | x2 | Tensor | 必选 | 矩阵乘运算中的右矩阵,数据类型与x1一致,K轴长度与x1一致。perm_x2为[0, 1, 2]时shape为(B, K, N);perm_x2为[0, 2, 1]时shape为(B, N, K)。MXFP4场景Tensor最后一维为FP4拼包后的物理长度(N/2或K/2)。 | 同x1 | 3维 | | 65 | | x2 | Tensor | 必选 | 矩阵乘运算中的右矩阵,数据类型与x1一致,K轴长度与x1一致。perm_x2为[0, 1, 2]时shape为(B, K, N);perm_x2为[0, 2, 1]时shape为(B, N, K)。MXFP4场景Tensor最后一维为FP4拼包后的物理长度(N/2或K/2)。 | 同x1 | 3维 | |
| 66 | | dtype | int | 必选 | 输出y的数据类型枚举值:1表示torch.float16,27表示torch.bfloat16。 | int64 | - | | 66 | | dtype | int | 必选 | 输出y的数据类型枚举值:1表示torch.float16,27表示torch.bfloat16。 | int64 | - | |
| 67 | | bias | Tensor | 可选 | 矩阵乘运算后累加的偏置。预留参数,当前暂不支持,必须传入None。 | - | - | | 67 | | bias | Tensor | 可选 | 矩阵乘运算后累加的偏置。预留参数,当前暂不支持,必须传入None。 | - | - | |
| @@ -97,7 +97,7 @@ cann_ops_nn.transpose_quant_batch_mat_mul( | |||
| 97 | - batch_split_factor当前仅支持取值1。 | 97 | - batch_split_factor当前仅支持取值1。 |
| 98 | - bias为预留参数,当前暂不支持。 | 98 | - bias为预留参数,当前暂不支持。 |
| 99 | - 不支持空Tensor。 | 99 | - 不支持空Tensor。 |
| 100 | -- MXFP4场景数据按两个FLOAT4_E2M1拼包存储(Tensor最后一维为物理长度,即逻辑长度的一半):x1的最后一维为K/2;x2在perm_x2为[0, 1, 2]时最后一维为N/2,在perm_x2为[0, 2, 1]时最后一维为K/2。此时Tensor实际存储类型(如torch.uint8)无法自动推导出FP4类型,必须通过x1_dtype、x2_dtype指定为torch_npu.float4_e2m1fn_x2。 | 100 | +- MXFP4场景数据以torch_npu.float4_e2m1fn_x2格式存储(两个FP4元素拼包,Tensor最后一维为物理长度,即逻辑长度的一半):x1的最后一维为K/2;x2在perm_x2为[0, 1, 2]时最后一维为N/2,在perm_x2为[0, 2, 1]时最后一维为K/2。此时Tensor实际存储类型(如torch.uint8)无法自动推导出FP4类型,必须通过x1_dtype、x2_dtype指定为torch_npu.float4_e2m1fn_x2。 |
| 101 | - 仅x2支持FRACTAL_NZ格式(仅MX量化模式)。 | 101 | - 仅x2支持FRACTAL_NZ格式(仅MX量化模式)。 |
| 102 | 102 | ||
| 103 | ## 确定性计算 | 103 | ## 确定性计算 |
| @@ -140,7 +140,7 @@ cann_ops_nn.transpose_quant_batch_mat_mul( | |||
| 140 | import cann_ops_nn | 140 | import cann_ops_nn |
| 141 | 141 | ||
| 142 | M, B, K, N = 64, 16, 128, 256 | 142 | M, B, K, N = 64, 16, 128, 256 |
| 143 | - # FP4按两个FLOAT4_E2M1拼包存储,Tensor最后一维为物理长度(逻辑长度的一半) | 143 | + # FP4以torch_npu.float4_e2m1fn_x2格式存储(两个FP4元素拼包),Tensor最后一维为物理长度(逻辑长度的一半) |
| 144 | x1 = torch.randint(0, 256, (M, B, K // 2), dtype=torch.uint8).npu() | 144 | x1 = torch.randint(0, 256, (M, B, K // 2), dtype=torch.uint8).npu() |
| 145 | x2 = torch.randint(0, 256, (B, K, N // 2), dtype=torch.uint8).npu() | 145 | x2 = torch.randint(0, 256, (B, K, N // 2), dtype=torch.uint8).npu() |
| 146 | x1_scale = torch.ones(M, B, K // 64, 2, dtype=torch.float8_e8m0fnu).npu() | 146 | x1_scale = torch.ones(M, B, K // 64, 2, dtype=torch.float8_e8m0fnu).npu() |