已合并
fix: 修复QBMM伪量化场景空tensor提前返回路径误触发(合入9.1.0分支) #6208
zhoushaolong创建于 6月17日
fix: 修复QBMM伪量化场景空tensor提前返回路径误触发(合入9.1.0分支) #6208
已合并
共 5 个文件变更+8-6
| @@ -229,7 +229,7 @@ aclnnStatus aclnnQuantMatmulV3( | |||
| 229 | - scale数据类型支持UINT64、INT64、FLOAT32、BFLOAT16 | 229 | - scale数据类型支持UINT64、INT64、FLOAT32、BFLOAT16 |
| 230 | - scale支持INT32、BFLOAT16、FLOAT32 | 230 | - scale支持INT32、BFLOAT16、FLOAT32 |
| 231 | - out支持FLOAT16、INT8、BFLOAT16、INT32 | 231 | - out支持FLOAT16、INT8、BFLOAT16、INT32 |
| 232 | - - x2仅支持ND格式,当输入x1为m=0的空tensor或x2为n=0的空tensor时,输出为空tensor | 232 | + - 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 | 1167 | ||
| 1167 | } | 1168 | } |
| 1168 | -#endif | 1169 | +#endif |