已合并
fix: 修复QBMM伪量化场景空tensor提前返回路径误触发(合入9.1.0分支) #6208
zhoushaolong创建于 6月17日
fix: 修复QBMM伪量化场景空tensor提前返回路径误触发(合入9.1.0分支) #6208
已合并
zhoushaolong创建于 6月17日
5 个文件变更+8-6
@@ -229,7 +229,7 @@ aclnnStatus aclnnQuantMatmulV3(
229 - scale数据类型支持UINT64、INT64、FLOAT32、BFLOAT16229 - scale数据类型支持UINT64、INT64、FLOAT32、BFLOAT16
230 - scale支持INT32、BFLOAT16、FLOAT32230 - scale支持INT32、BFLOAT16、FLOAT32
231 - out支持FLOAT16、INT8、BFLOAT16、INT32231 - out支持FLOAT16、INT8、BFLOAT16、INT32
232- - x2仅支持ND格式,当输入x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor232+ - x2仅支持ND格式,全量化场景下,当输入x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor
233 233 
234- **返回值:**234- **返回值:**
235 235 
@@ -317,7 +317,7 @@ aclnnStatus aclnnQuantMatmulV4(
317 - x2数据类型支持INT8、INT4。317 - x2数据类型支持INT8、INT4。
318 - bias数据类型支持INT32,BFLOAT16,FLOAT16,FLOAT32。318 - bias数据类型支持INT32,BFLOAT16,FLOAT16,FLOAT32。
319 - out数据类型支持FLOAT16、INT8、BFLOAT16、INT32。319 - out数据类型支持FLOAT16、INT8、BFLOAT16、INT32。
320- - x2仅支持ND格式,当输入x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor。320+ - x2仅支持ND格式,全量化场景下,当输入x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor。
321 321 
322- **返回值:**322- **返回值:**
323 323 
@@ -1820,7 +1820,8 @@ static aclnnStatus aclnnQuantMatmulGetWorkspaceSizeCommonProcess(TupleTensor man
1820 bool &transposeX2 = std::get<INDEX_X2_IN_MANDTORY_TUPLE>(boolsTrans);1820 bool &transposeX2 = std::get<INDEX_X2_IN_MANDTORY_TUPLE>(boolsTrans);
1821 bool isA8W4F = isA8W4Float(x1, x2);1821 bool isA8W4F = isA8W4Float(x1, x2);
1822 bool isA8W4I = isA8W4Int(x1, x2);1822 bool isA8W4I = isA8W4Int(x1, x2);
1823- if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) {1823+ bool isPseudoQuant = isA8W4F || isA8W4I;
1824+ if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && !isPseudoQuant) {
1824 auto x1DimNum = x1->GetViewShape().GetDimNum();1825 auto x1DimNum = x1->GetViewShape().GetDimNum();
1825 auto inputSizeM = transposeX1 ? x1->GetViewShape().GetDim(x1DimNum - 1) :1826 auto inputSizeM = transposeX1 ? x1->GetViewShape().GetDim(x1DimNum - 1) :
1826 x1->GetViewShape().GetDim(x1DimNum - PENULTIMATE_DIM);1827 x1->GetViewShape().GetDim(x1DimNum - PENULTIMATE_DIM);
@@ -479,7 +479,7 @@ aclnnStatus aclnnQuantMatmulV5(
479 479 
480 - 上表数据类型列中的角标“1”代表该系列不支持的数据类型。480 - 上表数据类型列中的角标“1”代表该系列不支持的数据类型。
481 - 输入参数x1、x2均不支持INT32类型。481 - 输入参数x1、x2均不支持INT32类型。
482- - x2仅支持ND格式,当输入参数x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor。482+ - x2仅支持ND格式,全量化场景下,当输入参数x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor。
483 483 
484 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term><term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>484 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term><term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>
485 485 
@@ -1028,11 +1028,12 @@ static aclnnStatus aclnnQuantMatmulGetWorkspaceSizeCommonProcess(TupleInput &inp
1028 bool &transposeX2 = std::get<INDEX_X2_IN_INPUT_TUPLE>(boolsTrans);1028 bool &transposeX2 = std::get<INDEX_X2_IN_INPUT_TUPLE>(boolsTrans);
1029 bool isA8W4 = false;1029 bool isA8W4 = false;
1030 bool isA8W4F = isA8W4Float(x1, x2);1030 bool isA8W4F = isA8W4Float(x1, x2);
1031+ bool isPseudoQuant = isA8W4F || isA8W4Int(x1, x2);
1031 if (isA8W4F) {1032 if (isA8W4F) {
1032 CHECK_RET(A8W4InferGroupSize(groupSize), ACLNN_ERR_PARAM_INVALID);1033 CHECK_RET(A8W4InferGroupSize(groupSize), ACLNN_ERR_PARAM_INVALID);
1033 OP_LOGD("Infer groupSize success. groupSize: %ld.", groupSize);1034 OP_LOGD("Infer groupSize success. groupSize: %ld.", groupSize);
1034 }1035 }
1035- if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) {1036+ if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510 && !isPseudoQuant) {
1036 auto x1DimNum = x1->GetViewShape().GetDimNum();1037 auto x1DimNum = x1->GetViewShape().GetDimNum();
1037 auto inputSizeM = transposeX1 ? x1->GetViewShape().GetDim(x1DimNum - 1) :1038 auto inputSizeM = transposeX1 ? x1->GetViewShape().GetDim(x1DimNum - 1) :
1038 x1->GetViewShape().GetDim(x1DimNum - PENULTIMATE_DIM);1039 x1->GetViewShape().GetDim(x1DimNum - PENULTIMATE_DIM);
@@ -1165,4 +1166,4 @@ aclnnStatus aclnnQuantMatmulV5(void *workspace, uint64_t workspaceSize, aclOpExe
1165 1166 
1166#ifdef __cplusplus1167#ifdef __cplusplus
1167}1168}
1168-#endif1169+#endif