已合并
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
已合并
HKFLYE创建于 8月17日
共 14 个文件变更+370-68
@@ -26,7 +26,7 @@
26- 接口功能:26- 接口功能:
27计算α与batch1、batch2的矩阵乘结果的乘积,再与β和self的乘积求和。27计算α与batch1、batch2的矩阵乘结果的乘积,再与β和self的乘积求和。
28注意:batch1、batch2必须是三维Tensor,两个shape仅在aclnnBaddbmm支持做broadcast,两个shape在aclnnInplaceBaddbmm做broadcast则会被拦截;28注意:batch1、batch2必须是三维Tensor,两个shape仅在aclnnBaddbmm支持做broadcast,两个shape在aclnnInplaceBaddbmm做broadcast则会被拦截;
29-self必须要支持和batch1@batch2的结果做broadcast。(broadcast,广播机制,是指较小的shape扩展至较大的shape,使两者shape互相兼容,当前仅支持(1,n)的broadcast,即两个Tensor对应的每一维度必须相同或其中一个为1。)29+self必须能broadcast到batch1@batch2的结果shape,支持N方向、M方向、scalar及per-batch scalar等broadcast形式。broadcast是指较小的shape扩展至较大的shape,使两个Tensor的shape互相兼容,即对应的每一维度必须相同或其中一个为1。
30 30 
31- 计算公式:31- 计算公式:
32 32 
@@ -226,7 +226,7 @@ aclnnStatus aclnnInplaceBaddbmm(
226 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;226 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;
227 - cubeMathType=2,当输入数据类型是FLOAT32时,会转换为FLOAT16计算,当输入为其他数据类型时不做处理;227 - cubeMathType=2,当输入数据类型是FLOAT32时,会转换为FLOAT16计算,当输入为其他数据类型时不做处理;
228 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理。228 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理。
229- - cubeMathType=4,输入数类型为FLOAT16/BFLOAT16时addmm过程升精度计算,该情况下当前不支持输入self与matmul计算结果矩阵做broadcast。229+ - cubeMathType=4,当self、batch1和batch2均为FLOAT16或均为BFLOAT16时,批量矩阵乘与bias相加过程使用FLOAT32中间结果计算,支持self相对批量矩阵乘结果进行broadcast,输出数据类型由out指定。
230 230 
231 <!-- end id8 -->231 <!-- end id8 -->
232 <!-- npu="950" id9 -->232 <!-- npu="950" id9 -->
@@ -464,7 +464,7 @@ aclnnStatus aclnnInplaceBaddbmm(
464 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;464 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;
465 - cubeMathType=2,当输入数据类型是FLOAT32,会转换为FLOAT16计算;当输入为其他数据类型时不做处理;465 - cubeMathType=2,当输入数据类型是FLOAT32,会转换为FLOAT16计算;当输入为其他数据类型时不做处理;
466 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理。466 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理。
467- - cubeMathType=4,输入数类型为FLOAT16/BFLOAT16时addmm过程升精度计算,该情况下当前不支持输入self与matmul计算结果矩阵做broadcast。467+ - cubeMathType=4,当selfRef、batch1和batch2均为FLOAT16或均为BFLOAT16时,批量矩阵乘与bias相加过程使用FLOAT32中间结果计算。
468 <!-- end id12 -->468 <!-- end id12 -->
469 <!-- npu="950" id13 -->469 <!-- npu="950" id13 -->
470 - <term>Ascend 950PR/Ascend 950DT</term>:470 - <term>Ascend 950PR/Ascend 950DT</term>:
@@ -186,7 +186,7 @@ static aclnnStatus CheckInputParams(const aclTensor* self, const aclTensor* batc
186 CHECK_RET(CheckFormat(batch1, batch2, out), ACLNN_ERR_PARAM_INVALID);186 CHECK_RET(CheckFormat(batch1, batch2, out), ACLNN_ERR_PARAM_INVALID);
187 187 
188 // 6. 检查cubeMathType是否支持188 // 6. 检查cubeMathType是否支持
189- CHECK_RET(CheckCubeMathTypeForAddMm(batch1, batch2, self, out, cubeMathType), ACLNN_ERR_PARAM_INVALID);189+ CHECK_RET(CheckCubeMathTypeForAddMm(cubeMathType), ACLNN_ERR_PARAM_INVALID);
190 190 
191 return ACLNN_SUCCESS;191 return ACLNN_SUCCESS;
192}192}
@@ -272,11 +272,15 @@ public:
272 }272 }
273 bool enable16In32Out = NeedEnableFp32Output(matA->GetDataType(), matB->GetDataType(), output->GetDataType(),273 bool enable16In32Out = NeedEnableFp32Output(matA->GetDataType(), matB->GetDataType(), output->GetDataType(),
274 cubeMathType);274 cubeMathType);
275+ bool isSupportNpuArch = GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_2201;
276+ op::DataType biasDtype = bias->GetDataType();
277+ bool biasDtypeValid = biasDtype == matA->GetDataType() || biasDtype == op::DataType::DT_FLOAT;
275 bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(matA, matB, bias);278 bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(matA, matB, bias);
276- bool useGemm16In32Out = enable16In32Out && !needBroadcast && bias->GetDataType() == matA->GetDataType() &&279+ bool useGemm16In32Out = enable16In32Out && !needBroadcast && biasDtypeValid && isSupportNpuArch;
277- (GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_2201);280+ bool useGemmFp32Add = CheckGemmV3WithAlphaBeta(bias, matA, matB, cubeMathType) ||
278- // A2/A3上对于 16in32out,且不需要broadcast场景 直接走gemmV3281+ (enable16In32Out && cubeMathType == USE_FP32_ADD && biasDtypeValid && isSupportNpuArch);
279- if (CheckGemmV3WithAlphaBeta(bias, matA, matB, cubeMathType) || useGemm16In32Out) {282+ // USE_FP32_ADD(包括broadcast及16in32out),或无需broadcast的普通16in32out场景走GemmV3
283+ if (useGemmFp32Add || useGemm16In32Out) {
280 const aclTensor* bmmOut = ExecGemmV3WithAlphaBetaOp(bias, matA, matB, alpha, beta, executor,284 const aclTensor* bmmOut = ExecGemmV3WithAlphaBetaOp(bias, matA, matB, alpha, beta, executor,
281 enable16In32Out);285 enable16In32Out);
282 CHECK_RET(bmmOut != nullptr, ACLNN_ERR_INNER_NULLPTR);286 CHECK_RET(bmmOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
@@ -150,10 +150,8 @@ bool CheckAddmmTensorShapeNeedBroadcast(const aclTensor* mat1, const aclTensor*
150 return false;150 return false;
151}151}
152 152 
153-bool CheckCubeMathTypeForAddMm(const aclTensor* mat1, const aclTensor* mat2, const aclTensor* self,153+bool CheckCubeMathTypeForAddMm(int8_t cubeMathType)
154- const aclTensor* out, int8_t cubeMathType)
155{154{
156- (void)out;
157 if (cubeMathType > USE_FP32_ADD) {155 if (cubeMathType > USE_FP32_ADD) {
158 OP_LOGE(ACLNN_ERR_PARAM_INVALID,156 OP_LOGE(ACLNN_ERR_PARAM_INVALID,
159 "The value of cubeMathType only support {0: KEEP_DTYPE, 1: "157 "The value of cubeMathType only support {0: KEEP_DTYPE, 1: "
@@ -171,14 +169,6 @@ bool CheckCubeMathTypeForAddMm(const aclTensor* mat1, const aclTensor* mat2, con
171 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "current platform not support cubeMathType = 4: USE_FP32_ADD.");169 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "current platform not support cubeMathType = 4: USE_FP32_ADD.");
172 return false;170 return false;
173 }171 }
174- // A2平台上,当cubeMathType=USE_FP32_ADD时,当前不支持self与mmout broadcast
175- if (npuArch == NpuArch::DAV_2201) {
176- bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(mat1, mat2, self);
177- OP_CHECK(!needBroadcast,
178- OP_LOGE(ACLNN_ERR_PARAM_INVALID,
179- "when cubeMathType = 4:USE_FP32_ADD, do not support broadcast between self and mmout."),
180- return false;);
181- }
182 return true;172 return true;
183}173}
184 174 
@@ -343,4 +333,4 @@ bool NeedCubeGoHF32(const DataType cubeTensorPromoteType, int8_t cubeMathType)
343}333}
344 334 
345} // namespace NN335} // namespace NN
346-} // namespace Ops336+} // namespace Ops
@@ -26,12 +26,11 @@ bool CheckCubeMathType(const op::DataType cubeTensorDtype, int8_t cubeMathType);
26// 校验针对mm算子 tensor的dtype,cubeMathType的值是否符合预期26// 校验针对mm算子 tensor的dtype,cubeMathType的值是否符合预期
27bool CheckCubeMathTypeForMm(const op::DataType cubeTensorDtype, int8_t cubeMathType);27bool CheckCubeMathTypeForMm(const op::DataType cubeTensorDtype, int8_t cubeMathType);
28 28 
29-// 校验针对addmm算子的输入shape检验其是否需要广播,是否可以广播需要预先校验29+// 校验Addmm/Baddbmm的bias是否需要相对矩阵乘输出进行广播
30bool CheckAddmmTensorShapeNeedBroadcast(const aclTensor* mat1, const aclTensor* mat2, const aclTensor* self);30bool CheckAddmmTensorShapeNeedBroadcast(const aclTensor* mat1, const aclTensor* mat2, const aclTensor* self);
31 31 
32// 校验针对Addmm算子的cubeMathType和平台是否符合预期32// 校验针对Addmm算子的cubeMathType和平台是否符合预期
33-bool CheckCubeMathTypeForAddMm(const aclTensor* mat1, const aclTensor* mat2, const aclTensor* self,33+bool CheckCubeMathTypeForAddMm(int8_t cubeMathType);
34- const aclTensor* out, int8_t cubeMathType);
35 34 
36// 返回芯片对应支持的数据类型35// 返回芯片对应支持的数据类型
37const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion();36const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion();
@@ -751,6 +751,156 @@
751 "value": false751 "value": false
752 }752 }
753 ]753 ]
754+ },
755+ {
756+ "bin_filename": "GemmV3_ND_NZ_ND_ND_FP16_FP16_FP16_FP16",
757+ "simplified_key": "diy,2/29/2/2/1/1/1/1",
758+ "inputs": [
759+ {
760+ "name": "a",
761+ "index": 0,
762+ "dtype": "float16",
763+ "format": "ND",
764+ "paramType": "required",
765+ "shape": [
766+ -2
767+ ]
768+ },
769+ {
770+ "name": "b",
771+ "index": 1,
772+ "dtype": "float16",
773+ "format": "FRACTAL_NZ",
774+ "paramType": "required",
775+ "shape": [
776+ -2
777+ ]
778+ },
779+ {
780+ "name": "c",
781+ "index": 2,
782+ "dtype": "float16",
783+ "format": "ND",
784+ "paramType": "optional",
785+ "shape": [
786+ -2
787+ ]
788+ }
789+ ],
790+ "outputs": [
791+ {
792+ "name": "y",
793+ "index": 0,
794+ "dtype": "float16",
795+ "format": "ND",
796+ "paramType": "required",
797+ "shape": [
798+ -2
799+ ]
800+ }
801+ ],
802+ "attrs": [
803+ {
804+ "name": "alpha",
805+ "dtype": "float",
806+ "value": 1.0
807+ },
808+ {
809+ "name": "beta",
810+ "dtype": "float",
811+ "value": 1.0
812+ },
813+ {
814+ "name": "transpose_a",
815+ "dtype": "bool",
816+ "value": false
817+ },
818+ {
819+ "name": "transpose_b",
820+ "dtype": "bool",
821+ "value": false
822+ },
823+ {
824+ "name": "enable_hf32",
825+ "dtype": "bool",
826+ "value": false
827+ }
828+ ]
829+ },
830+ {
831+ "bin_filename": "GemmV3_ND_NZ_ND_ND_BF16_BF16_BF16_BF16",
832+ "simplified_key": "diy,2/29/2/2/27/27/27/27",
833+ "inputs": [
834+ {
835+ "name": "a",
836+ "index": 0,
837+ "dtype": "bfloat16",
838+ "format": "ND",
839+ "paramType": "required",
840+ "shape": [
841+ -2
842+ ]
843+ },
844+ {
845+ "name": "b",
846+ "index": 1,
847+ "dtype": "bfloat16",
848+ "format": "FRACTAL_NZ",
849+ "paramType": "required",
850+ "shape": [
851+ -2
852+ ]
853+ },
854+ {
855+ "name": "c",
856+ "index": 2,
857+ "dtype": "bfloat16",
858+ "format": "ND",
859+ "paramType": "optional",
860+ "shape": [
861+ -2
862+ ]
863+ }
864+ ],
865+ "outputs": [
866+ {
867+ "name": "y",
868+ "index": 0,
869+ "dtype": "bfloat16",
870+ "format": "ND",
871+ "paramType": "required",
872+ "shape": [
873+ -2
874+ ]
875+ }
876+ ],
877+ "attrs": [
878+ {
879+ "name": "alpha",
880+ "dtype": "float",
881+ "value": 1.0
882+ },
883+ {
884+ "name": "beta",
885+ "dtype": "float",
886+ "value": 1.0
887+ },
888+ {
889+ "name": "transpose_a",
890+ "dtype": "bool",
891+ "value": false
892+ },
893+ {
894+ "name": "transpose_b",
895+ "dtype": "bool",
896+ "value": false
897+ },
898+ {
899+ "name": "enable_hf32",
900+ "dtype": "bool",
901+ "value": false
902+ }
903+ ]
754 }904 }
755 ]905 ]
756-}906+}
@@ -751,6 +751,156 @@
751 "value": false751 "value": false
752 }752 }
753 ]753 ]
754+ },
755+ {
756+ "bin_filename": "GemmV3_ND_NZ_ND_ND_FP16_FP16_FP16_FP16",
757+ "simplified_key": "diy,2/29/2/2/1/1/1/1",
758+ "inputs": [
759+ {
760+ "name": "a",
761+ "index": 0,
762+ "dtype": "float16",
763+ "format": "ND",
764+ "paramType": "required",
765+ "shape": [
766+ -2
767+ ]
768+ },
769+ {
770+ "name": "b",
771+ "index": 1,
772+ "dtype": "float16",
773+ "format": "FRACTAL_NZ",
774+ "paramType": "required",
775+ "shape": [
776+ -2
777+ ]
778+ },
779+ {
780+ "name": "c",
781+ "index": 2,
782+ "dtype": "float16",
783+ "format": "ND",
784+ "paramType": "optional",
785+ "shape": [
786+ -2
787+ ]
788+ }
789+ ],
790+ "outputs": [
791+ {
792+ "name": "y",
793+ "index": 0,
794+ "dtype": "float16",
795+ "format": "ND",
796+ "paramType": "required",
797+ "shape": [
798+ -2
799+ ]
800+ }
801+ ],
802+ "attrs": [
803+ {
804+ "name": "alpha",
805+ "dtype": "float",
806+ "value": 1.0
807+ },
808+ {
809+ "name": "beta",
810+ "dtype": "float",
811+ "value": 1.0
812+ },
813+ {
814+ "name": "transpose_a",
815+ "dtype": "bool",
816+ "value": false
817+ },
818+ {
819+ "name": "transpose_b",
820+ "dtype": "bool",
821+ "value": false
822+ },
823+ {
824+ "name": "enable_hf32",
825+ "dtype": "bool",
826+ "value": false
827+ }
828+ ]
829+ },
830+ {
831+ "bin_filename": "GemmV3_ND_NZ_ND_ND_BF16_BF16_BF16_BF16",
832+ "simplified_key": "diy,2/29/2/2/27/27/27/27",
833+ "inputs": [
834+ {
835+ "name": "a",
836+ "index": 0,
837+ "dtype": "bfloat16",
838+ "format": "ND",
839+ "paramType": "required",
840+ "shape": [
841+ -2
842+ ]
843+ },
844+ {
845+ "name": "b",
846+ "index": 1,
847+ "dtype": "bfloat16",
848+ "format": "FRACTAL_NZ",
849+ "paramType": "required",
850+ "shape": [
851+ -2
852+ ]
853+ },
854+ {
855+ "name": "c",
856+ "index": 2,
857+ "dtype": "bfloat16",
858+ "format": "ND",
859+ "paramType": "optional",
860+ "shape": [
861+ -2
862+ ]
863+ }
864+ ],
865+ "outputs": [
866+ {
867+ "name": "y",
868+ "index": 0,
869+ "dtype": "bfloat16",
870+ "format": "ND",
871+ "paramType": "required",
872+ "shape": [
873+ -2
874+ ]
875+ }
876+ ],
877+ "attrs": [
878+ {
879+ "name": "alpha",
880+ "dtype": "float",
881+ "value": 1.0
882+ },
883+ {
884+ "name": "beta",
885+ "dtype": "float",
886+ "value": 1.0
887+ },
888+ {
889+ "name": "transpose_a",
890+ "dtype": "bool",
891+ "value": false
892+ },
893+ {
894+ "name": "transpose_b",
895+ "dtype": "bool",
896+ "value": false
897+ },
898+ {
899+ "name": "enable_hf32",
900+ "dtype": "bool",
901+ "value": false
902+ }
903+ ]
754 }904 }
755 ]905 ]
756-}906+}
@@ -22,27 +22,28 @@ public:
22 this->Input("a")22 this->Input("a")
23 .ParamType(REQUIRED)23 .ParamType(REQUIRED)
24 .DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,24 .DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,
25- ge::DT_BF16, ge::DT_BF16})25+ ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
26 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,26 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
27- ge::FORMAT_ND, ge::FORMAT_ND});27+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
28 this->Input("b")28 this->Input("b")
29 .ParamType(REQUIRED)29 .ParamType(REQUIRED)
30 .DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,30 .DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,
31- ge::DT_BF16, ge::DT_BF16})31+ ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
32 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ,32 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ,
33- ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ});33+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
34+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_ND});
34 this->Input("c")35 this->Input("c")
35 .ParamType(OPTIONAL)36 .ParamType(OPTIONAL)
36 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT,37 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT,
37- ge::DT_BF16, ge::DT_FLOAT})38+ ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
38 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,39 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
39- ge::FORMAT_ND, ge::FORMAT_ND});40+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
40 this->Output("y")41 this->Output("y")
41 .ParamType(REQUIRED)42 .ParamType(REQUIRED)
42 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT,43 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT,
43- ge::DT_FLOAT, ge::DT_FLOAT})44+ ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT})
44 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,45 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
45- ge::FORMAT_ND, ge::FORMAT_ND});46+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
46 this->Attr("alpha").AttrType(REQUIRED).Float(1.0);47 this->Attr("alpha").AttrType(REQUIRED).Float(1.0);
47 this->Attr("beta").AttrType(REQUIRED).Float(1.0);48 this->Attr("beta").AttrType(REQUIRED).Float(1.0);
48 this->Attr("transpose_a").AttrType(REQUIRED).Bool(false);49 this->Attr("transpose_a").AttrType(REQUIRED).Bool(false);
@@ -258,7 +258,7 @@ ge::graphStatus GemmV3BaseTiling::PostTiling()
258 tilingPtr->alpha = alpha_;258 tilingPtr->alpha = alpha_;
259 tilingPtr->beta = beta_;259 tilingPtr->beta = beta_;
260 tilingPtr->biasBroadcastType = biasBroadcastType_;260 tilingPtr->biasBroadcastType = biasBroadcastType_;
261- tilingPtr->reservedBiasBroadcast = 0;261+ tilingPtr->reserved = 0;
262 tilingPtr->cBatchStride = cBatchStride_;262 tilingPtr->cBatchStride = cBatchStride_;
263 tilingPtr->cMStride = cMStride_;263 tilingPtr->cMStride = cMStride_;
264 tilingPtr->cNStride = cNStride_;264 tilingPtr->cNStride = cNStride_;
@@ -51,10 +51,10 @@ struct alignas(8) GemmV3TilingData {
51 float alpha{0.0f};51 float alpha{0.0f};
52 float beta{0.0f};52 float beta{0.0f};
53 uint32_t biasBroadcastType{BIAS_BCAST_NONE};53 uint32_t biasBroadcastType{BIAS_BCAST_NONE};
54- uint32_t reservedBiasBroadcast{0}; // Reserved for future bias-broadcast extensions54+ uint32_t reserved{0}; // Reserved padding for alignment.
55- uint64_t cBatchStride{0}; // C stride along Batch, in elements; 0 for broadcast.55+ uint64_t cBatchStride{0}; // C stride along Batch, in elements; 0 for broadcast.
56- uint64_t cMStride{0}; // C stride along M, in elements; 0 for broadcast.56+ uint64_t cMStride{0}; // C stride along M, in elements; 0 for broadcast.
57- uint64_t cNStride{0}; // C stride along N, in elements; 0 for broadcast.57+ uint64_t cNStride{0}; // C stride along N, in elements; 0 for broadcast.
58};58};
59#pragma pack(pop)59#pragma pack(pop)
60static_assert(sizeof(GemmV3TilingData) % sizeof(uint64_t) == 0, "GemmV3TilingData must be 8-byte aligned");60static_assert(sizeof(GemmV3TilingData) % sizeof(uint64_t) == 0, "GemmV3TilingData must be 8-byte aligned");
@@ -227,7 +227,7 @@ aclnnStatus aclnnInplaceAddmm(
227 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;227 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;
228 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;228 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;
229 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项。229 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项。
230- - cubeMathType=4,当输入数据类型为FLOAT16/BFLOAT16时addmm过程升精度计算,该情况下当前不支持输入self与matmul计算结果矩阵做broadcast;当输入数据类型为FLOAT32且k轴大于2048时,会使用分组累加进行计算。230+ - cubeMathType=4,当self、mat1和mat2均为FLOAT16或均为BFLOAT16时,矩阵乘与bias相加过程使用FLOAT32中间结果计算,支持self相对矩阵乘结果进行broadcast,输出数据类型由out指定;当输入数据类型为FLOAT32且k轴大于2048时,会使用分组累加进行计算。
231 <!-- end id8 -->231 <!-- end id8 -->
232 <!-- npu="950" id9 -->232 <!-- npu="950" id9 -->
233 - <term>Ascend 950PR/Ascend 950DT</term>:233 - <term>Ascend 950PR/Ascend 950DT</term>:
@@ -443,7 +443,7 @@ aclnnStatus aclnnInplaceAddmm(
443 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;443 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;
444 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;444 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;
445 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项。445 - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项。
446- - cubeMathType=4,当输入数据类型为FLOAT16/BFLOAT16时addmm过程升精度计算,该情况下当前不支持输入self与matmul计算结果矩阵做broadcast;当输入数据类型为FLOAT32且k轴大于2048时,会使用分组累加进行计算。446+ - cubeMathType=4,当selfRef、mat1和mat2均为FLOAT16或均为BFLOAT16时,矩阵乘与bias相加过程使用FLOAT32中间结果计算;当输入数据类型为FLOAT32且k轴大于2048时,会使用分组累加进行计算。
447 <!-- end id11 -->447 <!-- end id11 -->
448 <!-- npu="950" id12 -->448 <!-- npu="950" id12 -->
449 - <term>Ascend 950PR/Ascend 950DT</term>:449 - <term>Ascend 950PR/Ascend 950DT</term>:
@@ -33,6 +33,8 @@
33- 示例:33- 示例:
34 * 对于aclnnAddmmWeightNz接口,self的shape是[n,],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。34 * 对于aclnnAddmmWeightNz接口,self的shape是[n,],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。
35 * 对于aclnnAddmmWeightNz接口,self的shape是[1, n],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。35 * 对于aclnnAddmmWeightNz接口,self的shape是[1, n],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。
36+ * 对于aclnnAddmmWeightNz接口,self的shape是[m, 1],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。
37+ * 对于aclnnAddmmWeightNz接口,self的shape是[1, 1],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n],self的shape能broadcast到[m, n]。
36 * 对于aclnnAddmmWeightNz接口,self的shape是[m, n],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n]。38 * 对于aclnnAddmmWeightNz接口,self的shape是[m, n],mat1的shape是[m, k],mat2的shape是[k, n],mat1和mat2的矩阵乘的结果shape是[m, n]。
37 39 
38## 函数原型40## 函数原型
@@ -90,7 +92,7 @@ aclnnStatus aclnnAddmmWeightNz(
90 <td>输入</td>92 <td>输入</td>
91 <td>表示bias矩阵,公式中的self。</td>93 <td>表示bias矩阵,公式中的self。</td>
92 <td><ul><li>数据类型需要与mat1@mat2满足数据类型推导规则(参见<a href="../../../docs/zh/context/deduction_relationship.md">互推导关系</a>和<a href="#约束说明">约束说明</a>)。</li>94 <td><ul><li>数据类型需要与mat1@mat2满足数据类型推导规则(参见<a href="../../../docs/zh/context/deduction_relationship.md">互推导关系</a>和<a href="#约束说明">约束说明</a>)。</li>
93- <li>需要与mat1@mat2满足<a href="../../../docs/zh/context/broadcast_relationship.md">broadcast关系</a>。</li> <li> self支持shape为(n),(1,n),(m,n)。</li> </ul></td>95+ <li>需要与mat1@mat2满足<a href="../../../docs/zh/context/broadcast_relationship.md">broadcast关系</a>。</li> <li>self支持shape为(n)、(1,n)、(m,1)、(1,1)或(m,n)。</li> </ul></td>
94 <td>BFLOAT16、FLOAT16、FLOAT32</td>96 <td>BFLOAT16、FLOAT16、FLOAT32</td>
95 <td>ND</td>97 <td>ND</td>
96 <td>1-2</td>98 <td>1-2</td>
@@ -158,7 +160,8 @@ aclnnStatus aclnnAddmmWeightNz(
158 <li>0:KEEP_DTYPE,保持输入的数据类型进行计算。</li>160 <li>0:KEEP_DTYPE,保持输入的数据类型进行计算。</li>
159 <li>1:ALLOW_FP32_DOWN_PRECISION,支持将输入数据降精度计算。</li>161 <li>1:ALLOW_FP32_DOWN_PRECISION,支持将输入数据降精度计算。</li>
160 <li>2:USE_FP16,支持将输入降精度至FLOAT16计算。</li>162 <li>2:USE_FP16,支持将输入降精度至FLOAT16计算。</li>
161- <li>3:USE_HF32,支持将输入降精度至数据类型HFLOAT32计算。</li></ul>163+ <li>3:USE_HF32,支持将输入降精度至数据类型HFLOAT32计算。</li>
164+ <li>4:USE_FP32_ADD,支持使用高精度方式进行计算。</li></ul>
162 </td>165 </td>
163 <td>INT8</td>166 <td>INT8</td>
164 <td>-</td>167 <td>-</td>
@@ -191,7 +194,8 @@ aclnnStatus aclnnAddmmWeightNz(
191 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:194 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
192 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;195 - cubeMathType=1,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不做处理;
193 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;196 - cubeMathType=2,当输入数据类型为BFLOAT16时不支持该选项;
194- - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项。197+ - cubeMathType=3,当输入数据类型为FLOAT32时,会转换为HFLOAT32计算,当输入为其他数据类型时不支持该选项;
198+ - cubeMathType=4,当self、mat1和mat2均为FLOAT16或均为BFLOAT16时,矩阵乘与bias相加过程使用FLOAT32中间结果计算,支持self相对矩阵乘结果进行broadcast,输出数据类型由out指定。
195 <!-- end id7 -->199 <!-- end id7 -->
196 <!-- npu="950" id8 -->200 <!-- npu="950" id8 -->
197 - <term>Ascend 950PR/Ascend 950DT</term>:201 - <term>Ascend 950PR/Ascend 950DT</term>:
@@ -186,9 +186,7 @@ static aclnnStatus CheckInputParams(AclnnAddmmTensor& addmmTensor, int8_t cubeMa
186 CHECK_RET(CheckOutShape(addmmTensor.mat1, addmmTensor.mat2, addmmTensor.out), ACLNN_ERR_PARAM_INVALID);186 CHECK_RET(CheckOutShape(addmmTensor.mat1, addmmTensor.mat2, addmmTensor.out), ACLNN_ERR_PARAM_INVALID);
187 187 
188 // 6. 检查是否满足GemmV3BaseKernel条件。188 // 6. 检查是否满足GemmV3BaseKernel条件。
189- CHECK_RET(189+ CHECK_RET(CheckCubeMathTypeForAddMm(cubeMathType), ACLNN_ERR_PARAM_INVALID);
190- CheckCubeMathTypeForAddMm(addmmTensor.mat1, addmmTensor.mat2, addmmTensor.self, addmmTensor.out, cubeMathType),
191- ACLNN_ERR_PARAM_INVALID);
192 return ACLNN_SUCCESS;190 return ACLNN_SUCCESS;
193}191}
194 192 
@@ -521,13 +519,19 @@ public:
521 {519 {
522 auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch();520 auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch();
523 bool isSupportNpuArch = (npuArch == NpuArch::DAV_2201);521 bool isSupportNpuArch = (npuArch == NpuArch::DAV_2201);
524- bool enable16In32Out = NeedEnableFp32Output(matA->GetDataType(), matB->GetDataType(), output->GetDataType(),522+ bool enableFp32Output = NeedEnableFp32Output(matA->GetDataType(), matB->GetDataType(), output->GetDataType(),
525- cubeMathType, nullptr, true);523+ cubeMathType, nullptr, true);
524+ bool enableGemm16In32Out = NeedEnableFp32Output(matA->GetDataType(), matB->GetDataType(), output->GetDataType(),
525+ cubeMathType);
526+ op::DataType biasDtype = bias->GetDataType();
527+ bool biasDtypeValid = biasDtype == matA->GetDataType() || biasDtype == op::DataType::DT_FLOAT;
526 bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(matA, matB, bias);528 bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(matA, matB, bias);
527- bool useGemm16In32Out = enable16In32Out && !needBroadcast && bias->GetDataType() == matA->GetDataType();529+ bool useGemm16In32Out = enableGemm16In32Out && !needBroadcast && biasDtypeValid;
528- // A2/A3上对于 16in32out,且不需要broadcast场景 直接走gemmV3530+ bool useGemmFp32Add = CheckGemmV3WithAlphaBeta(bias, matA, matB, cubeMathType) ||
529- if ((CheckGemmV3WithAlphaBeta(bias, matA, matB, cubeMathType) || useGemm16In32Out) && isSupportNpuArch) {531+ (enableGemm16In32Out && cubeMathType == USE_FP32_ADD && biasDtypeValid);
530- auto outGemmV3 = ExecGemmV3WithAlphaBetaOp(bias, matA, matB, alpha, beta, executor, enable16In32Out);532+ // USE_FP32_ADD(包括broadcast及16in32out),或无需broadcast的普通16in32out场景走GemmV3
533+ if ((useGemmFp32Add || useGemm16In32Out) && isSupportNpuArch) {
范其瑞
范其瑞范其瑞8月17日

InplaceAddmm如果路由到broadcast bias的GemmV3实现,存在越界写风险:inplace接口的输出缓冲区与输入矩阵共享内存,若输入shape小于广播后的输出shape,内核按完整输出写回会越界。

likedislike
HKFLYE
HKFLYE
8月17日 评论:
534+ auto outGemmV3 = ExecGemmV3WithAlphaBetaOp(bias, matA, matB, alpha, beta, executor, enableGemm16In32Out);
531 CHECK_RET(outGemmV3 != nullptr, ACLNN_ERR_INNER_NULLPTR);535 CHECK_RET(outGemmV3 != nullptr, ACLNN_ERR_INNER_NULLPTR);
532 convOut = outGemmV3;536 convOut = outGemmV3;
533 return ACLNN_SUCCESS;537 return ACLNN_SUCCESS;
@@ -535,7 +539,7 @@ public:
535 // 执行 Muls: out1 = beta * bias539 // 执行 Muls: out1 = beta * bias
536 // 非inplace接口不能改变输入tensor,isMulsInplace=false540 // 非inplace接口不能改变输入tensor,isMulsInplace=false
537 const aclTensor* biasCastType = bias;541 const aclTensor* biasCastType = bias;
538- if (enable16In32Out && bias->GetDataType() != DataType::DT_FLOAT) {542+ if (enableFp32Output && bias->GetDataType() != DataType::DT_FLOAT) {
539 biasCastType = l0op::Contiguous(bias, executor);543 biasCastType = l0op::Contiguous(bias, executor);
540 CHECK_RET(biasCastType != nullptr, ACLNN_ERR_INNER_NULLPTR);544 CHECK_RET(biasCastType != nullptr, ACLNN_ERR_INNER_NULLPTR);
541 biasCastType = l0op::Cast(biasCastType, op::DataType::DT_FLOAT, executor);545 biasCastType = l0op::Cast(biasCastType, op::DataType::DT_FLOAT, executor);
@@ -821,9 +825,6 @@ ACLNN_API aclnnStatus aclnnAddmmWeightNzGetWorkspaceSize(const aclTensor* self,
821 825 
822 AclnnAddmmTensor addmmTensor = {self, mat1, mat2, beta, alpha, out};826 AclnnAddmmTensor addmmTensor = {self, mat1, mat2, beta, alpha, out};
823 827 
824- // 路由cubeMathType4到cubeMathType0, 该接口不支持cubeMathType=4的场景
825- cubeMathType = routeCubeMathType4ToCubeMathType0DAV_2201(cubeMathType);
826- 
827 auto ret = AddmmCheckWeightNzParam(addmmTensor, cubeMathType);828 auto ret = AddmmCheckWeightNzParam(addmmTensor, cubeMathType);
828 CHECK_RET(ret == ACLNN_SUCCESS, ret);829 CHECK_RET(ret == ACLNN_SUCCESS, ret);
829 auto uniqueExecutor = CREATE_EXECUTOR();830 auto uniqueExecutor = CREATE_EXECUTOR();
@@ -834,27 +835,31 @@ ACLNN_API aclnnStatus aclnnAddmmWeightNzGetWorkspaceSize(const aclTensor* self,
834 return ACLNN_SUCCESS;835 return ACLNN_SUCCESS;
835 }836 }
836 837 
837- bool enable16In32Out = NeedEnableFp32Output(mat1->GetDataType(), mat2->GetDataType(), out->GetDataType(),838+ bool enableFp32Output = NeedEnableFp32Output(mat1->GetDataType(), mat2->GetDataType(), out->GetDataType(),
838- cubeMathType, nullptr, true);839+ cubeMathType, nullptr, true);
839- bool addmmNeedBroadcast = CheckAddmmTensorShapeNeedBroadcast(mat1, mat2, self);840+ bool enableGemm16In32Out = NeedEnableFp32Output(mat1->GetDataType(), mat2->GetDataType(), out->GetDataType(),
841+ cubeMathType);
840 bool isSupportNpuArch = op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_2201;842 bool isSupportNpuArch = op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_2201;
841- bool useGemm16In32Out = enable16In32Out && !addmmNeedBroadcast && self->GetDataType() == mat1->GetDataType() &&843+ op::DataType biasDtype = self->GetDataType();
842- isSupportNpuArch;844+ bool biasDtypeValid = biasDtype == mat1->GetDataType() || biasDtype == op::DataType::DT_FLOAT;
845+ bool needBroadcast = CheckAddmmTensorShapeNeedBroadcast(mat1, mat2, self);
846+ bool useGemm16In32Out = enableGemm16In32Out && !needBroadcast && biasDtypeValid && isSupportNpuArch;
847+ bool useGemmFp32Add = CheckGemmV3WithAlphaBeta(self, mat1, mat2, cubeMathType) ||
848+ (enableGemm16In32Out && cubeMathType == USE_FP32_ADD && biasDtypeValid && isSupportNpuArch);
843 849 
844 const aclTensor* castOut = nullptr;850 const aclTensor* castOut = nullptr;
845 if (fabs(beta->ToFloat() - 0.0f) <= numeric_limits<float>::epsilon()) {851 if (fabs(beta->ToFloat() - 0.0f) <= numeric_limits<float>::epsilon()) {
846 castOut = MatmulMulProcess(addmmTensor, cubeMathType, uniqueExecutor.get());852 castOut = MatmulMulProcess(addmmTensor, cubeMathType, uniqueExecutor.get());
847- } else if (useGemm16In32Out) {853+ } else if (useGemmFp32Add || useGemm16In32Out) {
848 OP_LOGD("aclnnAddmmWeightNz run in ExecGemmV3WithAlphaBetaOp branch");854 OP_LOGD("aclnnAddmmWeightNz run in ExecGemmV3WithAlphaBetaOp branch");
849- // 16in32out场景优先走gemmV3通路855+ castOut = ExecGemmV3WithAlphaBetaOp(self, mat1, mat2, alpha, beta, uniqueExecutor.get(), enableGemm16In32Out);
850- castOut = ExecGemmV3WithAlphaBetaOp(self, mat1, mat2, alpha, beta, uniqueExecutor.get(), enable16In32Out);
851 } else if (NeedToConvertBias(self, mat1, mat2, beta, alpha) && check16In32Output(mat1, mat2, out)) {856 } else if (NeedToConvertBias(self, mat1, mat2, beta, alpha) && check16In32Output(mat1, mat2, out)) {
852 OP_LOGD("aclnnAddmmWeightNz run in NeedToConvertBias branch");857 OP_LOGD("aclnnAddmmWeightNz run in NeedToConvertBias branch");
853 auto biasMmOut = ExecMmOpWithBias(mat1, mat2, self, out, cubeMathType, uniqueExecutor.get(), false, false);858 auto biasMmOut = ExecMmOpWithBias(mat1, mat2, self, out, cubeMathType, uniqueExecutor.get(), false, false);
854 CHECK_RET(biasMmOut != nullptr, ACLNN_ERR_INNER_NULLPTR);859 CHECK_RET(biasMmOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
855 castOut = l0op::Cast(biasMmOut, out->GetDataType(), uniqueExecutor.get());860 castOut = l0op::Cast(biasMmOut, out->GetDataType(), uniqueExecutor.get());
856 } else {861 } else {
857- castOut = AddMatmulProcess(addmmTensor, cubeMathType, enable16In32Out, uniqueExecutor.get());862+ castOut = AddMatmulProcess(addmmTensor, cubeMathType, enableFp32Output, uniqueExecutor.get());
858 }863 }
859 CHECK_RET(castOut != nullptr, ACLNN_ERR_INNER_NULLPTR);864 CHECK_RET(castOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
860 865 
@@ -1536,7 +1536,7 @@ TEST_F(l2_addmm_test, addmm_910b_fp16_fp16_use_fp32_add_with_self_need_broadcast
1536 1536 
1537 uint64_t workspace_size = 0;1537 uint64_t workspace_size = 0;
1538 aclnnStatus aclRet = ut.TestGetWorkspaceSize(&workspace_size);1538 aclnnStatus aclRet = ut.TestGetWorkspaceSize(&workspace_size);
1539- EXPECT_NE(aclRet, ACL_SUCCESS);1539+ EXPECT_EQ(aclRet, ACL_SUCCESS);
1540}1540}
1541 1541 
1542TEST_F(l2_addmm_test, addmm_910b_bf16_bf16_use_fp32_add_with_self_need_broadcast)1542TEST_F(l2_addmm_test, addmm_910b_bf16_bf16_use_fp32_add_with_self_need_broadcast)
@@ -1554,5 +1554,5 @@ TEST_F(l2_addmm_test, addmm_910b_bf16_bf16_use_fp32_add_with_self_need_broadcast
1554 1554 
1555 uint64_t workspace_size = 0;1555 uint64_t workspace_size = 0;
1556 aclnnStatus aclRet = ut.TestGetWorkspaceSize(&workspace_size);1556 aclnnStatus aclRet = ut.TestGetWorkspaceSize(&workspace_size);
1557- EXPECT_NE(aclRet, ACL_SUCCESS);1557+ EXPECT_EQ(aclRet, ACL_SUCCESS);
1558}1558}
@@ -657,7 +657,6 @@ TEST_F(l2_addmmWeightNz_test, case_empty_tensor_self)
657 EXPECT_EQ(aclRet, ACLNN_ERR_PARAM_NULLPTR);657 EXPECT_EQ(aclRet, ACLNN_ERR_PARAM_NULLPTR);
658}658}
659 659 
660-// 转型cubeMathType为0的情况
661TEST_F(l2_addmmWeightNz_test, addmm_NZ_910b_FP32_FP16_USE_FP32_ADD)660TEST_F(l2_addmmWeightNz_test, addmm_NZ_910b_FP32_FP16_USE_FP32_ADD)
662{661{
663 auto self = TensorDesc({16}, ACL_FLOAT, ACL_FORMAT_ND).ValueRange(0, 2);662 auto self = TensorDesc({16}, ACL_FLOAT, ACL_FORMAT_ND).ValueRange(0, 2);