已合并
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
HKFLYE创建于 8月17日
【FEATURE】A2/A3 ACLNN Addmm/Baddbmm/AddmmWeightNz 支持 broadcast bias 路由 GemmV3 #8784
已合并
共 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场景 直接走gemmV3 | 281 | + (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 NN | 335 | } // namespace NN |
| 346 | -} // namespace Ops | 336 | +} // 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的值是否符合预期 |
| 27 | bool CheckCubeMathTypeForMm(const op::DataType cubeTensorDtype, int8_t cubeMathType); | 27 | bool CheckCubeMathTypeForMm(const op::DataType cubeTensorDtype, int8_t cubeMathType); |
| 28 | 28 | ||
| 29 | -// 校验针对addmm算子的输入shape检验其是否需要广播,是否可以广播需要预先校验 | 29 | +// 校验Addmm/Baddbmm的bias是否需要相对矩阵乘输出进行广播 |
| 30 | bool CheckAddmmTensorShapeNeedBroadcast(const aclTensor* mat1, const aclTensor* mat2, const aclTensor* self); | 30 | bool 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 | // 返回芯片对应支持的数据类型 |
| 37 | const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion(); | 36 | const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion(); |
| @@ -751,6 +751,156 @@ | |||
| 751 | "value": false | 751 | "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": false | 751 | "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 extensions | 54 | + 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 | 59 | ||
| 60 | static_assert(sizeof(GemmV3TilingData) % sizeof(uint64_t) == 0, "GemmV3TilingData must be 8-byte aligned"); | 60 | static_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场景 直接走gemmV3 | 530 | + 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) { | ||
| 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 * bias | 539 | // 执行 Muls: out1 = beta * bias |
| 536 | // 非inplace接口不能改变输入tensor,isMulsInplace=false | 540 | // 非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 | ||
| 1542 | TEST_F(l2_addmm_test, addmm_910b_bf16_bf16_use_fp32_add_with_self_need_broadcast) | 1542 | TEST_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的情况 | ||
| 661 | TEST_F(l2_addmmWeightNz_test, addmm_NZ_910b_FP32_FP16_USE_FP32_ADD) | 660 | TEST_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); |
InplaceAddmm如果路由到broadcast bias的GemmV3实现,存在越界写风险:inplace接口的输出缓冲区与输入矩阵共享内存,若输入shape小于广播后的输出shape,内核按完整输出写回会越界。