已合并
docs(tbmm,tqbmm): 资料修改(T-C量化约束、batchSplitFactor说明、标点规范等) #10800
jgx创建于 19 天前
docs(tbmm,tqbmm): 资料修改(T-C量化约束、batchSplitFactor说明、标点规范等) #10800
已合并
jgx创建于 19 天前
共 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_nn140 import cann_ops_nn
141 141 
142 M, B, K, N = 64, 16, 128, 256142 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()