| @@ -61,7 +61,6 @@ | |||
| 61 | </tr> | 61 | </tr> |
| 62 | </tbody></table> | 62 | </tbody></table> |
| 63 | 63 | ||
| 64 | -- Ascend 950PR/Ascend 950DT:不支持INT16、INT8、UINT8。 | ||
| 65 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16、INT8、UINT8。 | 64 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16、INT8、UINT8。 |
| 66 | 65 | ||
| 67 | ## 约束说明 | 66 | ## 约束说明 |
| @@ -105,6 +105,12 @@ static ge::graphStatus ForeachExpTilingFunc(gert::TilingContext* context) | |||
| 105 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_FLOAT16); | 105 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_FLOAT16); |
| 106 | } else if (dataType == ge::DT_BF16) { | 106 | } else if (dataType == ge::DT_BF16) { |
| 107 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_BF16); | 107 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_BF16); |
| 108 | + } else if (dataType == ge::DT_INT16) { | ||
| 109 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_INT16); | ||
| 110 | + } else if (dataType == ge::DT_INT8) { | ||
| 111 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_INT8); | ||
| 112 | + } else if (dataType == ge::DT_UINT8) { | ||
| 113 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXP_TPL_SCH_MODE_UINT8); | ||
| 108 | } else { | 114 | } else { |
| 109 | OP_LOGE(context, "unsupported dtype for foreach_exp"); | 115 | OP_LOGE(context, "unsupported dtype for foreach_exp"); |
| 110 | return ge::GRAPH_FAILED; | 116 | return ge::GRAPH_FAILED; |
| @@ -93,6 +93,99 @@ | |||
| 93 | } | 93 | } |
| 94 | ] | 94 | ] |
| 95 | ] | 95 | ] |
| 96 | + }, | ||
| 97 | + { | ||
| 98 | + "bin_filename": "ForeachExp_362812f0686499ffb98bc8a4178b5f86", | ||
| 99 | + "inputs": [ | ||
| 100 | + [ | ||
| 101 | + { | ||
| 102 | + "name": "x", | ||
| 103 | + "index": 0, | ||
| 104 | + "dtype": "int16", | ||
| 105 | + "format": "ND", | ||
| 106 | + "paramType": "dynamic", | ||
| 107 | + "shape": [ | ||
| 108 | + -2 | ||
| 109 | + ] | ||
| 110 | + } | ||
| 111 | + ] | ||
| 112 | + ], | ||
| 113 | + "outputs": [ | ||
| 114 | + [ | ||
| 115 | + { | ||
| 116 | + "name": "y", | ||
| 117 | + "index": 0, | ||
| 118 | + "dtype": "float32", | ||
| 119 | + "format": "ND", | ||
| 120 | + "paramType": "dynamic", | ||
| 121 | + "shape": [ | ||
| 122 | + -2 | ||
| 123 | + ] | ||
| 124 | + } | ||
| 125 | + ] | ||
| 126 | + ] | ||
| 127 | + }, | ||
| 128 | + { | ||
| 129 | + "bin_filename": "ForeachExp_62d040a3da22820d0c7ce25a5fa81716", | ||
| 130 | + "inputs": [ | ||
| 131 | + [ | ||
| 132 | + { | ||
| 133 | + "name": "x", | ||
| 134 | + "index": 0, | ||
| 135 | + "dtype": "int8", | ||
| 136 | + "format": "ND", | ||
| 137 | + "paramType": "dynamic", | ||
| 138 | + "shape": [ | ||
| 139 | + -2 | ||
| 140 | + ] | ||
| 141 | + } | ||
| 142 | + ] | ||
| 143 | + ], | ||
| 144 | + "outputs": [ | ||
| 145 | + [ | ||
| 146 | + { | ||
| 147 | + "name": "y", | ||
| 148 | + "index": 0, | ||
| 149 | + "dtype": "float32", | ||
| 150 | + "format": "ND", | ||
| 151 | + "paramType": "dynamic", | ||
| 152 | + "shape": [ | ||
| 153 | + -2 | ||
| 154 | + ] | ||
| 155 | + } | ||
| 156 | + ] | ||
| 157 | + ] | ||
| 158 | + }, | ||
| 159 | + { | ||
| 160 | + "bin_filename": "ForeachExp_d65408408eba4f1d75a15523e39691a2", | ||
| 161 | + "inputs": [ | ||
| 162 | + [ | ||
| 163 | + { | ||
| 164 | + "name": "x", | ||
| 165 | + "index": 0, | ||
| 166 | + "dtype": "uint8", | ||
| 167 | + "format": "ND", | ||
| 168 | + "paramType": "dynamic", | ||
| 169 | + "shape": [ | ||
| 170 | + -2 | ||
| 171 | + ] | ||
| 172 | + } | ||
| 173 | + ] | ||
| 174 | + ], | ||
| 175 | + "outputs": [ | ||
| 176 | + [ | ||
| 177 | + { | ||
| 178 | + "name": "y", | ||
| 179 | + "index": 0, | ||
| 180 | + "dtype": "float32", | ||
| 181 | + "format": "ND", | ||
| 182 | + "paramType": "dynamic", | ||
| 183 | + "shape": [ | ||
| 184 | + -2 | ||
| 185 | + ] | ||
| 186 | + } | ||
| 187 | + ] | ||
| 188 | + ] | ||
| 96 | } | 189 | } |
| 97 | ] | 190 | ] |
| 98 | -} | 191 | +} |
| @@ -42,7 +42,10 @@ public: | |||
| 42 | this->AICore().AddConfig("ascend910b"); | 42 | this->AICore().AddConfig("ascend910b"); |
| 43 | 43 | ||
| 44 | OpAICoreConfig regbaseCfg; | 44 | OpAICoreConfig regbaseCfg; |
| 45 | - std::vector<ge::DataType> tensor_dtype_list_ascend950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}; | 45 | + std::vector<ge::DataType> tensor_dtype_list_ascend950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, |
| 46 | + ge::DT_INT16, ge::DT_INT8, ge::DT_UINT8}; | ||
| 47 | + std::vector<ge::DataType> output_dtype_list_ascend950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, | ||
| 48 | + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT}; | ||
| 46 | std::vector<ge::Format> format_list_ascend950(tensor_dtype_list_ascend950.size(), ge::FORMAT_ND); | 49 | std::vector<ge::Format> format_list_ascend950(tensor_dtype_list_ascend950.size(), ge::FORMAT_ND); |
| 47 | regbaseCfg.DynamicCompileStaticFlag(true) | 50 | regbaseCfg.DynamicCompileStaticFlag(true) |
| 48 | .DynamicFormatFlag(false) | 51 | .DynamicFormatFlag(false) |
| @@ -58,7 +61,7 @@ public: | |||
| 58 | .AutoContiguous(); | 61 | .AutoContiguous(); |
| 59 | regbaseCfg.Output("y") | 62 | regbaseCfg.Output("y") |
| 60 | .ParamType(DYNAMIC) | 63 | .ParamType(DYNAMIC) |
| 61 | - .DataType(tensor_dtype_list_ascend950) | 64 | + .DataType(output_dtype_list_ascend950) |
| 62 | .Format(format_list_ascend950) | 65 | .Format(format_list_ascend950) |
| 63 | .UnknownShapeFormat(format_list_ascend950) | 66 | .UnknownShapeFormat(format_list_ascend950) |
| 64 | .AutoContiguous(); | 67 | .AutoContiguous(); |
| @@ -1,9 +1,9 @@ | |||
| 1 | # ---------------------------------------------------------------------------- | 1 | # ---------------------------------------------------------------------------- |
| 2 | # Copyright (c) 2026 Huawei Technologies Co., Ltd. | 2 | # Copyright (c) 2026 Huawei Technologies Co., Ltd. |
| 3 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of |
| 4 | # CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | # CANN Open Software License Agreement Version 2.0 (the "License"). |
| 5 | # Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | # Please refer to the License for details. You may not use this file except in compliance with the License. |
| 6 | -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, |
| 7 | # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. |
| 8 | # See LICENSE in the root of the software repository for the full text of the License. | 8 | # See LICENSE in the root of the software repository for the full text of the License. |
| 9 | # ---------------------------------------------------------------------------- | 9 | # ---------------------------------------------------------------------------- |
| @@ -16,4 +16,5 @@ add_kernel_sources( | |||
| 16 | KERNEL_SRC arch35/foreach_exp.cpp | 16 | KERNEL_SRC arch35/foreach_exp.cpp |
| 17 | COMPUTE_UNITS ascend950 | 17 | COMPUTE_UNITS ascend950 |
| 18 | AUTO_SYNC false | 18 | AUTO_SYNC false |
| 19 | + OPTIONS "--cce-use-fast-math=false" | ||
| 19 | ) | 20 | ) |
| @@ -19,6 +19,9 @@ enum class ForeachExpTilingKey : uint32_t { | |||
| 19 | TILING_KEY_FLOAT = 0, | 19 | TILING_KEY_FLOAT = 0, |
| 20 | TILING_KEY_FLOAT16 = 1, | 20 | TILING_KEY_FLOAT16 = 1, |
| 21 | TILING_KEY_BF16 = 2, | 21 | TILING_KEY_BF16 = 2, |
| 22 | + TILING_KEY_INT16 = 3, | ||
| 23 | + TILING_KEY_INT8 = 4, | ||
| 24 | + TILING_KEY_UINT8 = 5, | ||
| 22 | }; | 25 | }; |
| 23 | 26 | ||
| 24 | template <uint32_t schMode> | 27 | template <uint32_t schMode> |
| @@ -30,10 +33,16 @@ __global__ __aicore__ void foreach_exp(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, | |||
| 30 | const ForeachExpTilingData* tilingGm = &tilingData; | 33 | const ForeachExpTilingData* tilingGm = &tilingData; |
| 31 | 34 | ||
| 32 | if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_FLOAT)) { | 35 | if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_FLOAT)) { |
| 33 | - NsForeachExp::Process<float>(x, y, tilingGm); | 36 | + NsForeachExp::Process<float, float>(x, y, tilingGm); |
| 34 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_FLOAT16)) { | 37 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_FLOAT16)) { |
| 35 | - NsForeachExp::Process<half>(x, y, tilingGm); | 38 | + NsForeachExp::Process<half, half>(x, y, tilingGm); |
| 36 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_BF16)) { | 39 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_BF16)) { |
| 37 | - NsForeachExp::Process<bfloat16_t>(x, y, tilingGm); | 40 | + NsForeachExp::Process<bfloat16_t, bfloat16_t>(x, y, tilingGm); |
| 41 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_INT16)) { | ||
| 42 | + NsForeachExp::Process<int16_t, float>(x, y, tilingGm); | ||
| 43 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_INT8)) { | ||
| 44 | + NsForeachExp::Process<int8_t, float>(x, y, tilingGm); | ||
| 45 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpTilingKey::TILING_KEY_UINT8)) { | ||
| 46 | + NsForeachExp::Process<uint8_t, float>(x, y, tilingGm); | ||
| 38 | } | 47 | } |
| 39 | } | 48 | } |
| @@ -50,7 +50,10 @@ __simt_callee__ inline __gm__ T* SimtGetTensorAddr(GM_ADDR tensorListPtr, int64_ | |||
| 50 | * \brief Promote half/bf16 to float32 for precision | 50 | * \brief Promote half/bf16 to float32 for precision |
| 51 | */ | 51 | */ |
| 52 | template <typename T> | 52 | template <typename T> |
| 53 | -__simt_callee__ inline float PromoteToFloat(T val); | 53 | +__simt_callee__ inline float PromoteToFloat(T val) |
| 54 | +{ | ||
| 55 | + return static_cast<float>(val); | ||
| 56 | +} | ||
| 54 | 57 | ||
| 55 | template <> | 58 | template <> |
| 56 | __simt_callee__ inline float PromoteToFloat<half>(half val) | 59 | __simt_callee__ inline float PromoteToFloat<half>(half val) |
| @@ -70,54 +73,30 @@ __simt_callee__ inline float PromoteToFloat<float>(float val) | |||
| 70 | return val; | 73 | return val; |
| 71 | } | 74 | } |
| 72 | 75 | ||
| 73 | -/** | ||
| 74 | - * \brief Cast float32 back to half/bf16 | ||
| 75 | - */ | ||
| 76 | -template <typename T> | ||
| 77 | -__simt_callee__ inline T CastFromFloat(float val); | ||
| 78 | - | ||
| 79 | -template <> | ||
| 80 | -__simt_callee__ inline half CastFromFloat<half>(float val) | ||
| 81 | -{ | ||
| 82 | - return static_cast<half>(val); | ||
| 83 | -} | ||
| 84 | - | ||
| 85 | -template <> | ||
| 86 | -__simt_callee__ inline bfloat16_t CastFromFloat<bfloat16_t>(float val) | ||
| 87 | -{ | ||
| 88 | - return static_cast<bfloat16_t>(val); | ||
| 89 | -} | ||
| 90 | - | ||
| 91 | -template <> | ||
| 92 | -__simt_callee__ inline float CastFromFloat<float>(float val) | ||
| 93 | -{ | ||
| 94 | - return val; | ||
| 95 | -} | ||
| 96 | - | ||
| 97 | /** | 76 | /** |
| 98 | * \brief SIMT VF kernel: compute exp for all elements across all tensors | 77 | * \brief SIMT VF kernel: compute exp for all elements across all tensors |
| 99 | */ | 78 | */ |
| 100 | -template <typename T> | 79 | +template <typename X_T, typename Y_T> |
| 101 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachExpSimt(int32_t tensorId, int64_t count, | 80 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachExpSimt(int32_t tensorId, int64_t count, |
| 102 | GM_ADDR xList, GM_ADDR yList) | 81 | GM_ADDR xList, GM_ADDR yList) |
| 103 | { | 82 | { |
| 104 | - __gm__ T* xData = SimtGetTensorAddr<T>(xList, tensorId); | 83 | + __gm__ X_T* xData = SimtGetTensorAddr<X_T>(xList, tensorId); |
| 105 | - __gm__ T* yData = SimtGetTensorAddr<T>(yList, tensorId); | 84 | + __gm__ Y_T* yData = SimtGetTensorAddr<Y_T>(yList, tensorId); |
| 106 | uint64_t tid = static_cast<uint64_t>(AscendC::Simt::GetBlockIdx() * AscendC::Simt::GetThreadNum() + | 85 | uint64_t tid = static_cast<uint64_t>(AscendC::Simt::GetBlockIdx() * AscendC::Simt::GetThreadNum() + |
| 107 | AscendC::Simt::GetThreadIdx()); | 86 | AscendC::Simt::GetThreadIdx()); |
| 108 | uint64_t stride = static_cast<uint64_t>(AscendC::Simt::GetThreadNum() * AscendC::Simt::GetBlockNum()); | 87 | uint64_t stride = static_cast<uint64_t>(AscendC::Simt::GetThreadNum() * AscendC::Simt::GetBlockNum()); |
| 109 | for (uint64_t idx = tid; idx < static_cast<uint64_t>(count); idx += stride) { | 88 | for (uint64_t idx = tid; idx < static_cast<uint64_t>(count); idx += stride) { |
| 110 | - T xVal = xData[idx]; | 89 | + X_T xVal = xData[idx]; |
| 111 | - float xFloat = PromoteToFloat<T>(xVal); | 90 | + float xFloat = PromoteToFloat<X_T>(xVal); |
| 112 | float yFloat = expf(xFloat); | 91 | float yFloat = expf(xFloat); |
| 113 | - yData[idx] = CastFromFloat<T>(yFloat); | 92 | + yData[idx] = static_cast<Y_T>(yFloat); |
| 114 | } | 93 | } |
| 115 | } | 94 | } |
| 116 | 95 | ||
| 117 | /** | 96 | /** |
| 118 | * \brief Process entry: launch SIMT VF for foreach_exp | 97 | * \brief Process entry: launch SIMT VF for foreach_exp |
| 119 | */ | 98 | */ |
| 120 | -template <typename T> | 99 | +template <typename X_T, typename Y_T> |
| 121 | __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpTilingData* tilingGm) | 100 | __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpTilingData* tilingGm) |
| 122 | { | 101 | { |
| 123 | for (int32_t tensorId = 0; tensorId < tilingGm->tensorCount; tensorId++) { | 102 | for (int32_t tensorId = 0; tensorId < tilingGm->tensorCount; tensorId++) { |
| @@ -125,7 +104,7 @@ __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpTilingData* | |||
| 125 | if (count <= 0) { | 104 | if (count <= 0) { |
| 126 | continue; | 105 | continue; |
| 127 | } | 106 | } |
| 128 | - AscendC::Simt::VF_CALL<OpForeachExpSimt<T>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x, y); | 107 | + AscendC::Simt::VF_CALL<OpForeachExpSimt<X_T, Y_T>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x, y); |
| 129 | } | 108 | } |
| 130 | } | 109 | } |
| 131 | 110 | ||
| @@ -21,13 +21,19 @@ | |||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 24 | 27 | ||
| 25 | -ASCENDC_TPL_ARGS_DECL(ForeachExp, | 28 | +ASCENDC_TPL_ARGS_DECL(ForeachExp, ASCENDC_TPL_UINT_DECL(schMode, 6, ASCENDC_TPL_UI_LIST, FOREACH_EXP_TPL_SCH_MODE_FLOAT, |
这里的6表示的是支持类型数量的位宽,2^6可表示64种,可以减少到3,2^3就可以表示8种了,会减少Key的长度,便于阅读 ![]() ![]() | |||
| 26 | - ASCENDC_TPL_UINT_DECL(schMode, 3, ASCENDC_TPL_UI_LIST, FOREACH_EXP_TPL_SCH_MODE_FLOAT, | 29 | + FOREACH_EXP_TPL_SCH_MODE_FLOAT16, FOREACH_EXP_TPL_SCH_MODE_BF16, |
| 27 | - FOREACH_EXP_TPL_SCH_MODE_FLOAT16, FOREACH_EXP_TPL_SCH_MODE_BF16)); | 30 | + FOREACH_EXP_TPL_SCH_MODE_INT16, FOREACH_EXP_TPL_SCH_MODE_INT8, |
| 31 | + FOREACH_EXP_TPL_SCH_MODE_UINT8)); | ||
| 28 | 32 | ||
| 29 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, FOREACH_EXP_TPL_SCH_MODE_FLOAT, | 33 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, FOREACH_EXP_TPL_SCH_MODE_FLOAT, |
| 30 | FOREACH_EXP_TPL_SCH_MODE_FLOAT16, | 34 | FOREACH_EXP_TPL_SCH_MODE_FLOAT16, |
| 31 | - FOREACH_EXP_TPL_SCH_MODE_BF16))); | 35 | + FOREACH_EXP_TPL_SCH_MODE_BF16, FOREACH_EXP_TPL_SCH_MODE_INT16, |
| 36 | + FOREACH_EXP_TPL_SCH_MODE_INT8, | ||
| 37 | + FOREACH_EXP_TPL_SCH_MODE_UINT8))); | ||
| 32 | 38 | ||
| 33 | -#endif // FOREACH_EXP_TILING_KEY_H | 39 | +#endif // FOREACH_EXP_TILING_KEY_H |
| @@ -103,6 +103,12 @@ static ge::graphStatus ForeachExpm1TilingFunc(gert::TilingContext* context) | |||
| 103 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16); | 103 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16); |
| 104 | } else if (dataType == ge::DT_BF16) { | 104 | } else if (dataType == ge::DT_BF16) { |
| 105 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_BF16); | 105 | tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_BF16); |
| 106 | + } else if (dataType == ge::DT_INT16) { | ||
| 107 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_INT16); | ||
| 108 | + } else if (dataType == ge::DT_INT8) { | ||
| 109 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_INT8); | ||
| 110 | + } else if (dataType == ge::DT_UINT8) { | ||
| 111 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_EXPM1_TPL_SCH_MODE_UINT8); | ||
| 106 | } else { | 112 | } else { |
| 107 | OP_LOGE(context, "unsupported dtype for foreach_expm1"); | 113 | OP_LOGE(context, "unsupported dtype for foreach_expm1"); |
| 108 | return ge::GRAPH_FAILED; | 114 | return ge::GRAPH_FAILED; |
| @@ -121,4 +127,4 @@ IMPL_OP_OPTILING(ForeachExpm1) | |||
| 121 | .Tiling(ForeachExpm1TilingFunc) | 127 | .Tiling(ForeachExpm1TilingFunc) |
| 122 | .TilingParse<ForeachExpm1CompileInfo>(TilingParseForForeachExpm1); | 128 | .TilingParse<ForeachExpm1CompileInfo>(TilingParseForForeachExpm1); |
| 123 | 129 | ||
| 124 | -} // namespace optiling | 130 | +} // namespace optiling |
| @@ -93,6 +93,99 @@ | |||
| 93 | } | 93 | } |
| 94 | ] | 94 | ] |
| 95 | ] | 95 | ] |
| 96 | + }, | ||
| 97 | + { | ||
| 98 | + "bin_filename": "ForeachExpm1_26cd2233fb00bd84b497158224c89a2f", | ||
| 99 | + "inputs": [ | ||
| 100 | + [ | ||
| 101 | + { | ||
| 102 | + "name": "x", | ||
| 103 | + "index": 0, | ||
| 104 | + "dtype": "int16", | ||
| 105 | + "format": "ND", | ||
| 106 | + "paramType": "dynamic", | ||
| 107 | + "shape": [ | ||
| 108 | + -2 | ||
| 109 | + ] | ||
| 110 | + } | ||
| 111 | + ] | ||
| 112 | + ], | ||
| 113 | + "outputs": [ | ||
| 114 | + [ | ||
| 115 | + { | ||
| 116 | + "name": "y", | ||
| 117 | + "index": 0, | ||
| 118 | + "dtype": "float32", | ||
| 119 | + "format": "ND", | ||
| 120 | + "paramType": "dynamic", | ||
| 121 | + "shape": [ | ||
| 122 | + -2 | ||
| 123 | + ] | ||
| 124 | + } | ||
| 125 | + ] | ||
| 126 | + ] | ||
| 127 | + }, | ||
| 128 | + { | ||
| 129 | + "bin_filename": "ForeachExpm1_91a8140448c01c25be8a2753016fe86b", | ||
| 130 | + "inputs": [ | ||
| 131 | + [ | ||
| 132 | + { | ||
| 133 | + "name": "x", | ||
| 134 | + "index": 0, | ||
| 135 | + "dtype": "int8", | ||
| 136 | + "format": "ND", | ||
| 137 | + "paramType": "dynamic", | ||
| 138 | + "shape": [ | ||
| 139 | + -2 | ||
| 140 | + ] | ||
| 141 | + } | ||
| 142 | + ] | ||
| 143 | + ], | ||
| 144 | + "outputs": [ | ||
| 145 | + [ | ||
| 146 | + { | ||
| 147 | + "name": "y", | ||
| 148 | + "index": 0, | ||
| 149 | + "dtype": "float32", | ||
| 150 | + "format": "ND", | ||
| 151 | + "paramType": "dynamic", | ||
| 152 | + "shape": [ | ||
| 153 | + -2 | ||
| 154 | + ] | ||
| 155 | + } | ||
| 156 | + ] | ||
| 157 | + ] | ||
| 158 | + }, | ||
| 159 | + { | ||
| 160 | + "bin_filename": "ForeachExpm1_cadd0d4087956fb8cb0aaf94e2fa0dcd", | ||
| 161 | + "inputs": [ | ||
| 162 | + [ | ||
| 163 | + { | ||
| 164 | + "name": "x", | ||
| 165 | + "index": 0, | ||
| 166 | + "dtype": "uint8", | ||
| 167 | + "format": "ND", | ||
| 168 | + "paramType": "dynamic", | ||
| 169 | + "shape": [ | ||
| 170 | + -2 | ||
| 171 | + ] | ||
| 172 | + } | ||
| 173 | + ] | ||
| 174 | + ], | ||
| 175 | + "outputs": [ | ||
| 176 | + [ | ||
| 177 | + { | ||
| 178 | + "name": "y", | ||
| 179 | + "index": 0, | ||
| 180 | + "dtype": "float32", | ||
| 181 | + "format": "ND", | ||
| 182 | + "paramType": "dynamic", | ||
| 183 | + "shape": [ | ||
| 184 | + -2 | ||
| 185 | + ] | ||
| 186 | + } | ||
| 187 | + ] | ||
| 188 | + ] | ||
| 96 | } | 189 | } |
| 97 | ] | 190 | ] |
| 98 | -} | 191 | +} |
| @@ -44,7 +44,10 @@ public: | |||
| 44 | this->AICore().AddConfig("ascend910b"); | 44 | this->AICore().AddConfig("ascend910b"); |
| 45 | 45 | ||
| 46 | OpAICoreConfig regbaseCfg; | 46 | OpAICoreConfig regbaseCfg; |
| 47 | - std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}; | 47 | + std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, |
同上:这里的list如果和上面的tensor_dtype_list一致是不是可以用同一份 ![]() ![]() | |||
| 48 | + ge::DT_INT16, ge::DT_INT8, ge::DT_UINT8}; | ||
| 49 | + std::vector<ge::DataType> output_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, | ||
| 50 | + ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT}; | ||
| 48 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); | 51 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); |
| 49 | regbaseCfg.DynamicCompileStaticFlag(true) | 52 | regbaseCfg.DynamicCompileStaticFlag(true) |
| 50 | .DynamicFormatFlag(false) | 53 | .DynamicFormatFlag(false) |
| @@ -61,7 +64,7 @@ public: | |||
| 61 | .AutoContiguous(); | 64 | .AutoContiguous(); |
| 62 | regbaseCfg.Output("y") | 65 | regbaseCfg.Output("y") |
| 63 | .ParamType(DYNAMIC) | 66 | .ParamType(DYNAMIC) |
| 64 | - .DataType(tensor_dtype_list_950) | 67 | + .DataType(output_dtype_list_950) |
| 65 | .Format(format_list_950) | 68 | .Format(format_list_950) |
| 66 | .UnknownShapeFormat(format_list_950) | 69 | .UnknownShapeFormat(format_list_950) |
| 67 | .AutoContiguous(); | 70 | .AutoContiguous(); |
| @@ -5,7 +5,7 @@ | |||
| 5 | # Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | # Please refer to the License for details. You may not use this file except in compliance with the License. |
| 6 | # THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | # THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, |
| 7 | # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. |
| 8 | -# See LICENSE in the software repository for the full text of the License. | 8 | +# See LICENSE in the root of the software repository for the full text of the License. |
| 9 | # ---------------------------------------------------------------------------------------------------------- | 9 | # ---------------------------------------------------------------------------------------------------------- |
| 10 | # Generated By CANNBot | 10 | # Generated By CANNBot |
| 11 | 11 | ||
| @@ -18,4 +18,5 @@ add_kernel_sources( | |||
| 18 | KERNEL_SRC arch35/foreach_expm1.cpp | 18 | KERNEL_SRC arch35/foreach_expm1.cpp |
| 19 | COMPUTE_UNITS ascend950 | 19 | COMPUTE_UNITS ascend950 |
| 20 | AUTO_SYNC false | 20 | AUTO_SYNC false |
| 21 | -) | 21 | + OPTIONS "--cce-use-fast-math=false" |
| 22 | +) | ||
| @@ -21,6 +21,9 @@ enum class ForeachExpm1TilingKey : uint32_t { | |||
| 21 | TILING_KEY_FLOAT = 0, | 21 | TILING_KEY_FLOAT = 0, |
| 22 | TILING_KEY_FLOAT16 = 1, | 22 | TILING_KEY_FLOAT16 = 1, |
| 23 | TILING_KEY_BF16 = 2, | 23 | TILING_KEY_BF16 = 2, |
| 24 | + TILING_KEY_INT16 = 3, | ||
| 25 | + TILING_KEY_INT8 = 4, | ||
| 26 | + TILING_KEY_UINT8 = 5, | ||
| 24 | }; | 27 | }; |
| 25 | 28 | ||
| 26 | template <uint32_t schMode> | 29 | template <uint32_t schMode> |
| @@ -32,10 +35,16 @@ __global__ __aicore__ void foreach_expm1(GM_ADDR x, GM_ADDR y, GM_ADDR workspace | |||
| 32 | const ForeachExpm1TilingData* tilingGm = &tilingData; | 35 | const ForeachExpm1TilingData* tilingGm = &tilingData; |
| 33 | 36 | ||
| 34 | if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_FLOAT)) { | 37 | if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_FLOAT)) { |
| 35 | - NsForeachExpm1::Process<float>(x, y, tilingGm); | 38 | + NsForeachExpm1::Process<float, float>(x, y, tilingGm); |
| 36 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_FLOAT16)) { | 39 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_FLOAT16)) { |
| 37 | - NsForeachExpm1::Process<half>(x, y, tilingGm); | 40 | + NsForeachExpm1::Process<half, half>(x, y, tilingGm); |
| 38 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_BF16)) { | 41 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_BF16)) { |
| 39 | - NsForeachExpm1::Process<bfloat16_t>(x, y, tilingGm); | 42 | + NsForeachExpm1::Process<bfloat16_t, bfloat16_t>(x, y, tilingGm); |
| 43 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_INT16)) { | ||
| 44 | + NsForeachExpm1::Process<int16_t, float>(x, y, tilingGm); | ||
| 45 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_INT8)) { | ||
| 46 | + NsForeachExpm1::Process<int8_t, float>(x, y, tilingGm); | ||
| 47 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachExpm1TilingKey::TILING_KEY_UINT8)) { | ||
| 48 | + NsForeachExpm1::Process<uint8_t, float>(x, y, tilingGm); | ||
| 40 | } | 49 | } |
| 41 | } | 50 | } |
| @@ -48,48 +48,37 @@ __simt_callee__ inline __gm__ T* SimtGetTensorAddr(GM_ADDR tensorListPtr, int64_ | |||
| 48 | } | 48 | } |
| 49 | 49 | ||
| 50 | /** | 50 | /** |
| 51 | - * \brief Compute expm1 with type promotion for half/bfloat16 | 51 | + * \brief Compute expm1 in float32 and cast to the registered output type |
| 52 | - * \note half/bfloat16 are promoted to float32 for computation, then cast back | ||
| 53 | */ | 52 | */ |
| 54 | -template <typename T> | 53 | +template <typename X_T, typename Y_T> |
| 55 | -__simt_callee__ inline T ComputeExpm1(T x) | 54 | +__simt_callee__ inline Y_T ComputeExpm1(X_T x) |
| 56 | { | 55 | { |
| 57 | - // Generic path: cast to float, compute expm1f, cast back | ||
| 58 | float xFloat = static_cast<float>(x); | 56 | float xFloat = static_cast<float>(x); |
| 59 | float resultFloat = expm1f(xFloat); | 57 | float resultFloat = expm1f(xFloat); |
| 60 | - return static_cast<T>(resultFloat); | 58 | + return static_cast<Y_T>(resultFloat); |
| 61 | -} | ||
| 62 | - | ||
| 63 | -/** | ||
| 64 | - * \brief Specialization for float: direct expm1f call | ||
| 65 | - */ | ||
| 66 | -template <> | ||
| 67 | -__simt_callee__ inline float ComputeExpm1<float>(float x) | ||
| 68 | -{ | ||
| 69 | - return expm1f(x); | ||
| 70 | } | 59 | } |
| 71 | 60 | ||
| 72 | /** | 61 | /** |
| 73 | * \brief SIMT VF kernel: compute expm1 for all elements across all tensors | 62 | * \brief SIMT VF kernel: compute expm1 for all elements across all tensors |
| 74 | */ | 63 | */ |
| 75 | -template <typename T> | 64 | +template <typename X_T, typename Y_T> |
| 76 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachExpm1Simt(int32_t tensorId, int64_t count, | 65 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachExpm1Simt(int32_t tensorId, int64_t count, |
| 77 | GM_ADDR xList, GM_ADDR yList) | 66 | GM_ADDR xList, GM_ADDR yList) |
| 78 | { | 67 | { |
| 79 | - __gm__ T* xData = SimtGetTensorAddr<T>(xList, tensorId); | 68 | + __gm__ X_T* xData = SimtGetTensorAddr<X_T>(xList, tensorId); |
| 80 | - __gm__ T* yData = SimtGetTensorAddr<T>(yList, tensorId); | 69 | + __gm__ Y_T* yData = SimtGetTensorAddr<Y_T>(yList, tensorId); |
| 81 | uint64_t tid = static_cast<uint64_t>(AscendC::Simt::GetBlockIdx() * AscendC::Simt::GetThreadNum() + | 70 | uint64_t tid = static_cast<uint64_t>(AscendC::Simt::GetBlockIdx() * AscendC::Simt::GetThreadNum() + |
| 82 | AscendC::Simt::GetThreadIdx()); | 71 | AscendC::Simt::GetThreadIdx()); |
| 83 | uint64_t stride = static_cast<uint64_t>(AscendC::Simt::GetThreadNum() * AscendC::Simt::GetBlockNum()); | 72 | uint64_t stride = static_cast<uint64_t>(AscendC::Simt::GetThreadNum() * AscendC::Simt::GetBlockNum()); |
| 84 | for (uint64_t idx = tid; idx < static_cast<uint64_t>(count); idx += stride) { | 73 | for (uint64_t idx = tid; idx < static_cast<uint64_t>(count); idx += stride) { |
| 85 | - yData[idx] = ComputeExpm1<T>(xData[idx]); | 74 | + yData[idx] = ComputeExpm1<X_T, Y_T>(xData[idx]); |
| 86 | } | 75 | } |
| 87 | } | 76 | } |
| 88 | 77 | ||
| 89 | /** | 78 | /** |
| 90 | * \brief Process entry: launch SIMT VF for foreach_expm1 | 79 | * \brief Process entry: launch SIMT VF for foreach_expm1 |
| 91 | */ | 80 | */ |
| 92 | -template <typename T> | 81 | +template <typename X_T, typename Y_T> |
| 93 | __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpm1TilingData* tilingGm) | 82 | __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpm1TilingData* tilingGm) |
| 94 | { | 83 | { |
| 95 | for (int32_t tensorId = 0; tensorId < tilingGm->tensorCount; tensorId++) { | 84 | for (int32_t tensorId = 0; tensorId < tilingGm->tensorCount; tensorId++) { |
| @@ -97,7 +86,7 @@ __aicore__ inline void Process(GM_ADDR x, GM_ADDR y, const ForeachExpm1TilingDat | |||
| 97 | if (count <= 0) { | 86 | if (count <= 0) { |
| 98 | continue; | 87 | continue; |
| 99 | } | 88 | } |
| 100 | - AscendC::Simt::VF_CALL<OpForeachExpm1Simt<T>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x, y); | 89 | + AscendC::Simt::VF_CALL<OpForeachExpm1Simt<X_T, Y_T>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x, y); |
| 101 | } | 90 | } |
| 102 | } | 91 | } |
| 103 | 92 | ||
| @@ -23,13 +23,19 @@ | |||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 26 | 29 | ||
| 27 | ASCENDC_TPL_ARGS_DECL(ForeachExpm1, | 30 | ASCENDC_TPL_ARGS_DECL(ForeachExpm1, |
| 28 | - ASCENDC_TPL_UINT_DECL(schMode, 3, ASCENDC_TPL_UI_LIST, FOREACH_EXPM1_TPL_SCH_MODE_FLOAT, | 31 | + ASCENDC_TPL_UINT_DECL(schMode, 6, ASCENDC_TPL_UI_LIST, FOREACH_EXPM1_TPL_SCH_MODE_FLOAT, |
| 29 | - FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16, FOREACH_EXPM1_TPL_SCH_MODE_BF16)); | 32 | + FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16, FOREACH_EXPM1_TPL_SCH_MODE_BF16, |
| 33 | + FOREACH_EXPM1_TPL_SCH_MODE_INT16, FOREACH_EXPM1_TPL_SCH_MODE_INT8, | ||
| 34 | + FOREACH_EXPM1_TPL_SCH_MODE_UINT8)); | ||
| 30 | 35 | ||
| 31 | -ASCENDC_TPL_SEL( | 36 | +ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL( |
| 32 | - ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, FOREACH_EXPM1_TPL_SCH_MODE_FLOAT, | 37 | + schMode, ASCENDC_TPL_UI_LIST, FOREACH_EXPM1_TPL_SCH_MODE_FLOAT, FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16, |
| 33 | - FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16, FOREACH_EXPM1_TPL_SCH_MODE_BF16))); | 38 | + FOREACH_EXPM1_TPL_SCH_MODE_BF16, FOREACH_EXPM1_TPL_SCH_MODE_INT16, FOREACH_EXPM1_TPL_SCH_MODE_INT8, |
| 39 | + FOREACH_EXPM1_TPL_SCH_MODE_UINT8))); | ||
| 34 | 40 | ||
| 35 | -#endif // FOREACH_EXPM1_TILING_KEY_H | 41 | +#endif // FOREACH_EXPM1_TILING_KEY_H |
| @@ -68,7 +68,6 @@ | |||
| 68 | </tr> | 68 | </tr> |
| 69 | </tbody></table> | 69 | </tbody></table> |
| 70 | 70 | ||
| 71 | -- Ascend 950PR/Ascend 950DT:不支持INT16。 | ||
| 72 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16。 | 71 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16。 |
| 73 | 72 | ||
| 74 | ## 约束说明 | 73 | ## 约束说明 |
| @@ -43,6 +43,9 @@ static ge::graphStatus GetTilingKeyByDtype(gert::TilingContext* context, ge::Dat | |||
| 43 | case ge::DT_BF16: | 43 | case ge::DT_BF16: |
| 44 | tilingKey = 2; | 44 | tilingKey = 2; |
| 45 | return ge::GRAPH_SUCCESS; | 45 | return ge::GRAPH_SUCCESS; |
| 46 | + case ge::DT_INT16: | ||
| 47 | + tilingKey = 3; | ||
| 48 | + return ge::GRAPH_SUCCESS; | ||
| 46 | default: | 49 | default: |
| 47 | OP_LOGE(context, "unsupported dtype: %d", static_cast<int32_t>(dtype)); | 50 | OP_LOGE(context, "unsupported dtype: %d", static_cast<int32_t>(dtype)); |
| 48 | return ge::GRAPH_FAILED; | 51 | return ge::GRAPH_FAILED; |
| @@ -123,6 +123,47 @@ | |||
| 123 | } | 123 | } |
| 124 | ] | 124 | ] |
| 125 | ] | 125 | ] |
| 126 | + }, | ||
| 127 | + { | ||
| 128 | + "bin_filename": "ForeachRoundOffNumber_Int16", | ||
| 129 | + "inputs": [ | ||
| 130 | + [ | ||
| 131 | + { | ||
| 132 | + "name": "x", | ||
| 133 | + "index": 0, | ||
| 134 | + "dtype": "int16", | ||
| 135 | + "format": "ND", | ||
| 136 | + "paramType": "dynamic", | ||
| 137 | + "shape": [ | ||
| 138 | + -2 | ||
| 139 | + ] | ||
| 140 | + } | ||
| 141 | + ], | ||
| 142 | + { | ||
| 143 | + "name": "roundMode", | ||
| 144 | + "index": 1, | ||
| 145 | + "dtype": "int8", | ||
| 146 | + "format": "ND", | ||
| 147 | + "paramType": "required", | ||
| 148 | + "shape": [ | ||
| 149 | + -2 | ||
| 150 | + ] | ||
| 151 | + } | ||
| 152 | + ], | ||
| 153 | + "outputs": [ | ||
| 154 | + [ | ||
| 155 | + { | ||
| 156 | + "name": "y", | ||
| 157 | + "index": 0, | ||
| 158 | + "dtype": "int16", | ||
| 159 | + "format": "ND", | ||
| 160 | + "paramType": "dynamic", | ||
| 161 | + "shape": [ | ||
| 162 | + -2 | ||
| 163 | + ] | ||
| 164 | + } | ||
| 165 | + ] | ||
| 166 | + ] | ||
| 126 | } | 167 | } |
| 127 | ] | 168 | ] |
| 128 | -} | 169 | +} |
| @@ -46,7 +46,7 @@ public: | |||
| 46 | this->AICore().AddConfig("ascend910b"); | 46 | this->AICore().AddConfig("ascend910b"); |
| 47 | 47 | ||
| 48 | OpAICoreConfig aicoreConfig950; | 48 | OpAICoreConfig aicoreConfig950; |
| 49 | - std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}; | 49 | + std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT16}; |
| 50 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); | 50 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); |
| 51 | std::vector<ge::DataType> scalarDtypeList950{tensor_dtype_list_950.size(), ge::DT_INT8}; | 51 | std::vector<ge::DataType> scalarDtypeList950{tensor_dtype_list_950.size(), ge::DT_INT8}; |
| 52 | aicoreConfig950.DynamicCompileStaticFlag(true) | 52 | aicoreConfig950.DynamicCompileStaticFlag(true) |
| @@ -25,6 +25,7 @@ enum class ForeachRoundOffNumberTilingKey : uint32_t { | |||
| 25 | TILING_KEY_FLOAT = 0, | 25 | TILING_KEY_FLOAT = 0, |
| 26 | TILING_KEY_FLOAT16 = 1, | 26 | TILING_KEY_FLOAT16 = 1, |
| 27 | TILING_KEY_BF16 = 2, | 27 | TILING_KEY_BF16 = 2, |
| 28 | + TILING_KEY_INT16 = 3, | ||
| 28 | }; | 29 | }; |
| 29 | 30 | ||
| 30 | template <uint32_t schMode> | 31 | template <uint32_t schMode> |
| @@ -40,5 +41,7 @@ __global__ __aicore__ void foreach_round_off_number(GM_ADDR x, GM_ADDR roundMode | |||
| 40 | NsForeachRoundOffNumber::Process<half>(x, roundMode, y, &tilingData); | 41 | NsForeachRoundOffNumber::Process<half>(x, roundMode, y, &tilingData); |
| 41 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachRoundOffNumberTilingKey::TILING_KEY_BF16)) { | 42 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachRoundOffNumberTilingKey::TILING_KEY_BF16)) { |
| 42 | NsForeachRoundOffNumber::Process<bfloat16_t>(x, roundMode, y, &tilingData); | 43 | NsForeachRoundOffNumber::Process<bfloat16_t>(x, roundMode, y, &tilingData); |
| 44 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachRoundOffNumberTilingKey::TILING_KEY_INT16)) { | ||
| 45 | + NsForeachRoundOffNumber::Process<int16_t>(x, roundMode, y, &tilingData); | ||
| 43 | } | 46 | } |
🟠 High Priority 变更行:foreach_round_off_number.cpp 第 28 行新增 受影响行为/契约:本 PR 同时修改了 proto( 失败模式:当 schMode==3 时 ![]() ![]() 不准确? | |||
| 44 | } | 47 | } |
| @@ -62,8 +62,11 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachRoundOffNum | |||
| 62 | yData[idx] = rintf(xVal); | 62 | yData[idx] = rintf(xVal); |
| 63 | } else if constexpr (std::is_same_v<T, half>) { | 63 | } else if constexpr (std::is_same_v<T, half>) { |
| 64 | yData[idx] = hrint(xVal); | 64 | yData[idx] = hrint(xVal); |
| 65 | - } else { | 65 | + } else if constexpr (std::is_same_v<T, bfloat16_t>) { |
| 66 | yData[idx] = static_cast<bfloat16_t>(rintf(static_cast<float>(xVal))); | 66 | yData[idx] = static_cast<bfloat16_t>(rintf(static_cast<float>(xVal))); |
| 67 | + } else { | ||
| 68 | + // Rounding an integer is an identity operation. | ||
| 69 | + yData[idx] = xVal; | ||
| 67 | } | 70 | } |
| 68 | } | 71 | } |
| 69 | } | 72 | } |
| @@ -21,15 +21,18 @@ | |||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | + | ||
| 24 | 25 | ||
| 25 | -ASCENDC_TPL_ARGS_DECL(ForeachRoundOffNumber, ASCENDC_TPL_UINT_DECL(schMode, 3, ASCENDC_TPL_UI_LIST, | 26 | +ASCENDC_TPL_ARGS_DECL(ForeachRoundOffNumber, ASCENDC_TPL_UINT_DECL(schMode, 4, ASCENDC_TPL_UI_LIST, |
| 26 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT, | 27 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT, |
| 27 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16, | 28 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16, |
| 28 | - FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16)); | 29 | + FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16, |
| 30 | + FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_INT16)); | ||
| 29 | 31 | ||
| 30 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, | 32 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, |
| 31 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT, | 33 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT, |
| 32 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16, | 34 | FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16, |
| 33 | - FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16))); | 35 | + FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16, |
| 36 | + FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_INT16))); | ||
| 34 | 37 | ||
| 35 | 38 | ||
| @@ -81,7 +81,6 @@ | |||
| 81 | 81 | ||
| 82 | ## 约束说明 | 82 | ## 约束说明 |
| 83 | 83 | ||
| 84 | -- Ascend 950PR/Ascend 950DT:不支持INT16、INT8、UINT8。 | ||
| 85 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16、INT8、UINT8。 | 84 | - Kirin X90/Kirin 9030处理器系列产品:不支持BFLOAT16、INT16、INT8、UINT8。 |
| 86 | - 输出不支持非连续Tensor。 | 85 | - 输出不支持非连续Tensor。 |
| 87 | 86 | ||
| @@ -35,6 +35,74 @@ constexpr int32_t MAX_TENSOR_NUM = 256; | |||
| 35 | 35 | ||
| 36 | struct ForeachSubListCompileInfo {}; | 36 | struct ForeachSubListCompileInfo {}; |
| 37 | 37 | ||
| 38 | +static ge::DataType GetExpectedAlphaDtype(ge::DataType inputDtype) | ||
| 39 | +{ | ||
| 40 | + switch (inputDtype) { | ||
| 41 | + case ge::DT_FLOAT16: | ||
| 42 | + return ge::DT_FLOAT16; | ||
| 43 | + case ge::DT_FLOAT: | ||
| 44 | + case ge::DT_BF16: | ||
| 45 | + return ge::DT_FLOAT; | ||
| 46 | + case ge::DT_INT32: | ||
| 47 | + case ge::DT_INT16: | ||
| 48 | + case ge::DT_INT8: | ||
| 49 | + case ge::DT_UINT8: | ||
| 50 | + return ge::DT_INT32; | ||
| 51 | + default: | ||
| 52 | + return ge::DT_UNDEFINED; | ||
| 53 | + } | ||
| 54 | +} | ||
| 55 | + | ||
| 56 | +static ge::graphStatus CheckInputCompatibility(gert::TilingContext* context, uint64_t tensorNum, | ||
| 57 | + ge::DataType inputDtype) | ||
| 58 | +{ | ||
| 59 | + auto computeNodeInfoPtr = context->GetComputeNodeInfo(); | ||
| 60 | + OP_CHECK_NULL_WITH_CONTEXT(context, computeNodeInfoPtr); | ||
| 61 | + auto x2InstanceInfoPtr = computeNodeInfoPtr->GetInputInstanceInfo(INPUT_IDX_X2); | ||
| 62 | + OP_CHECK_NULL_WITH_CONTEXT(context, x2InstanceInfoPtr); | ||
| 63 | + OP_CHECK_IF(x2InstanceInfoPtr->GetInstanceNum() != tensorNum, | ||
| 64 | + OP_LOGE(context, "x1 and x2 must contain the same number of tensors, but got %lu and %lu", tensorNum, | ||
| 65 | + x2InstanceInfoPtr->GetInstanceNum()), | ||
| 66 | + return ge::GRAPH_FAILED); | ||
| 67 | + | ||
| 68 | + for (uint64_t i = 0; i < tensorNum; i++) { | ||
| 69 | + auto x1Desc = context->GetDynamicInputDesc(INPUT_IDX_X1, i); | ||
| 70 | + auto x2Desc = context->GetDynamicInputDesc(INPUT_IDX_X2, i); | ||
| 71 | + OP_CHECK_NULL_WITH_CONTEXT(context, x1Desc); | ||
| 72 | + OP_CHECK_NULL_WITH_CONTEXT(context, x2Desc); | ||
| 73 | + OP_CHECK_IF(x1Desc->GetDataType() != inputDtype || x2Desc->GetDataType() != inputDtype, | ||
| 74 | + OP_LOGE(context, "x1[%lu] and x2[%lu] must have dtype %d, but got %d and %d", i, i, | ||
| 75 | + static_cast<int32_t>(inputDtype), static_cast<int32_t>(x1Desc->GetDataType()), | ||
| 76 | + static_cast<int32_t>(x2Desc->GetDataType())), | ||
| 77 | + return ge::GRAPH_FAILED); | ||
| 78 | + | ||
| 79 | + auto x1Shape = context->GetDynamicInputShape(INPUT_IDX_X1, i); | ||
| 80 | + auto x2Shape = context->GetDynamicInputShape(INPUT_IDX_X2, i); | ||
| 81 | + OP_CHECK_NULL_WITH_CONTEXT(context, x1Shape); | ||
| 82 | + OP_CHECK_NULL_WITH_CONTEXT(context, x2Shape); | ||
| 83 | + OP_CHECK_IF(x1Shape->GetStorageShape() != x2Shape->GetStorageShape(), | ||
| 84 | + OP_LOGE(context, "x1[%lu] and x2[%lu] must have the same storage shape", i, i), | ||
| 85 | + return ge::GRAPH_FAILED); | ||
| 86 | + } | ||
| 87 | + | ||
| 88 | + auto alphaDesc = context->GetRequiredInputDesc(INPUT_IDX_ALPHA); | ||
| 89 | + OP_CHECK_NULL_WITH_CONTEXT(context, alphaDesc); | ||
| 90 | + ge::DataType expectedAlphaDtype = GetExpectedAlphaDtype(inputDtype); | ||
| 91 | + OP_CHECK_IF(alphaDesc->GetDataType() != expectedAlphaDtype, | ||
| 92 | + OP_LOGE(context, "alpha dtype must be %d when x1/x2 dtype is %d, but got %d", | ||
| 93 | + static_cast<int32_t>(expectedAlphaDtype), static_cast<int32_t>(inputDtype), | ||
| 94 | + static_cast<int32_t>(alphaDesc->GetDataType())), | ||
| 95 | + return ge::GRAPH_FAILED); | ||
| 96 | + | ||
| 97 | + auto alphaShape = context->GetRequiredInputShape(INPUT_IDX_ALPHA); | ||
| 98 | + OP_CHECK_NULL_WITH_CONTEXT(context, alphaShape); | ||
| 99 | + OP_CHECK_IF(alphaShape->GetStorageShape().GetShapeSize() != 1, | ||
| 100 | + OP_LOGE(context, "alpha must contain exactly one element, but got %ld", | ||
| 101 | + alphaShape->GetStorageShape().GetShapeSize()), | ||
| 102 | + return ge::GRAPH_FAILED); | ||
| 103 | + return ge::GRAPH_SUCCESS; | ||
| 104 | +} | ||
| 105 | + | ||
| 38 | static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum) | 106 | static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum) |
| 39 | { | 107 | { |
| 40 | fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo(); | 108 | fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo(); |
| @@ -63,6 +131,8 @@ static ge::graphStatus ForeachSubListTilingFunc(gert::TilingContext* context) | |||
| 63 | auto inputDesc = context->GetDynamicInputDesc(INPUT_IDX_X1, 0); | 131 | auto inputDesc = context->GetDynamicInputDesc(INPUT_IDX_X1, 0); |
| 64 | OP_CHECK_NULL_WITH_CONTEXT(context, inputDesc); | 132 | OP_CHECK_NULL_WITH_CONTEXT(context, inputDesc); |
| 65 | ge::DataType dataType = inputDesc->GetDataType(); | 133 | ge::DataType dataType = inputDesc->GetDataType(); |
| 134 | + OP_CHECK_IF(CheckInputCompatibility(context, tensorNum, dataType) != ge::GRAPH_SUCCESS, | ||
| 135 | + OP_LOGE(context, "input compatibility check failed"), return ge::GRAPH_FAILED); | ||
| 66 | 136 | ||
| 67 | ForeachSubListTilingData* tiling = context->GetTilingData<ForeachSubListTilingData>(); | 137 | ForeachSubListTilingData* tiling = context->GetTilingData<ForeachSubListTilingData>(); |
| 68 | OP_CHECK_NULL_WITH_CONTEXT(context, tiling); | 138 | OP_CHECK_NULL_WITH_CONTEXT(context, tiling); |
| @@ -109,6 +179,12 @@ static ge::graphStatus ForeachSubListTilingFunc(gert::TilingContext* context) | |||
| 109 | tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_INT32); | 179 | tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_INT32); |
| 110 | } else if (dataType == ge::DT_BF16) { | 180 | } else if (dataType == ge::DT_BF16) { |
| 111 | tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_BF16); | 181 | tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_BF16); |
| 182 | + } else if (dataType == ge::DT_INT16) { | ||
| 183 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_INT16); | ||
| 184 | + } else if (dataType == ge::DT_INT8) { | ||
| 185 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_INT8); | ||
| 186 | + } else if (dataType == ge::DT_UINT8) { | ||
| 187 | + tilingKey = GET_TPL_TILING_KEY(FOREACH_SUB_LIST_TPL_SCH_MODE_UINT8); | ||
| 112 | } else { | 188 | } else { |
| 113 | OP_LOGE(context, "unsupported dtype for foreach_sub_list"); | 189 | OP_LOGE(context, "unsupported dtype for foreach_sub_list"); |
| 114 | return ge::GRAPH_FAILED; | 190 | return ge::GRAPH_FAILED; |
| @@ -212,6 +212,165 @@ | |||
| 212 | } | 212 | } |
| 213 | ] | 213 | ] |
| 214 | ] | 214 | ] |
| 215 | + }, | ||
| 216 | + { | ||
| 217 | + "bin_filename": "ForeachSubList_7d016a74d29151527b4aa7688ff29b96", | ||
| 218 | + "inputs": [ | ||
| 219 | + [ | ||
| 220 | + { | ||
| 221 | + "name": "x1", | ||
| 222 | + "index": 0, | ||
| 223 | + "dtype": "int16", | ||
| 224 | + "format": "ND", | ||
| 225 | + "paramType": "dynamic", | ||
| 226 | + "shape": [ | ||
| 227 | + -2 | ||
| 228 | + ] | ||
| 229 | + } | ||
| 230 | + ], | ||
| 231 | + [ | ||
| 232 | + { | ||
| 233 | + "name": "x2", | ||
| 234 | + "index": 1, | ||
| 235 | + "dtype": "int16", | ||
| 236 | + "format": "ND", | ||
| 237 | + "paramType": "dynamic", | ||
| 238 | + "shape": [ | ||
| 239 | + -2 | ||
| 240 | + ] | ||
| 241 | + } | ||
| 242 | + ], | ||
| 243 | + { | ||
| 244 | + "name": "alpha", | ||
| 245 | + "index": 2, | ||
| 246 | + "dtype": "int32", | ||
| 247 | + "format": "ND", | ||
| 248 | + "paramType": "required", | ||
| 249 | + "shape": [ | ||
| 250 | + -2 | ||
| 251 | + ] | ||
| 252 | + } | ||
| 253 | + ], | ||
| 254 | + "outputs": [ | ||
| 255 | + [ | ||
| 256 | + { | ||
| 257 | + "name": "y", | ||
| 258 | + "index": 0, | ||
| 259 | + "dtype": "int16", | ||
| 260 | + "format": "ND", | ||
| 261 | + "paramType": "dynamic", | ||
| 262 | + "shape": [ | ||
| 263 | + -2 | ||
| 264 | + ] | ||
| 265 | + } | ||
| 266 | + ] | ||
| 267 | + ] | ||
| 268 | + }, | ||
| 269 | + { | ||
| 270 | + "bin_filename": "ForeachSubList_96e246186a2a9dfa8992f03882a8f144", | ||
| 271 | + "inputs": [ | ||
| 272 | + [ | ||
| 273 | + { | ||
| 274 | + "name": "x1", | ||
| 275 | + "index": 0, | ||
| 276 | + "dtype": "int8", | ||
| 277 | + "format": "ND", | ||
| 278 | + "paramType": "dynamic", | ||
| 279 | + "shape": [ | ||
| 280 | + -2 | ||
| 281 | + ] | ||
| 282 | + } | ||
| 283 | + ], | ||
| 284 | + [ | ||
| 285 | + { | ||
| 286 | + "name": "x2", | ||
| 287 | + "index": 1, | ||
| 288 | + "dtype": "int8", | ||
| 289 | + "format": "ND", | ||
| 290 | + "paramType": "dynamic", | ||
| 291 | + "shape": [ | ||
| 292 | + -2 | ||
| 293 | + ] | ||
| 294 | + } | ||
| 295 | + ], | ||
| 296 | + { | ||
| 297 | + "name": "alpha", | ||
| 298 | + "index": 2, | ||
| 299 | + "dtype": "int32", | ||
| 300 | + "format": "ND", | ||
| 301 | + "paramType": "required", | ||
| 302 | + "shape": [ | ||
| 303 | + -2 | ||
| 304 | + ] | ||
| 305 | + } | ||
| 306 | + ], | ||
| 307 | + "outputs": [ | ||
| 308 | + [ | ||
| 309 | + { | ||
| 310 | + "name": "y", | ||
| 311 | + "index": 0, | ||
| 312 | + "dtype": "int8", | ||
| 313 | + "format": "ND", | ||
| 314 | + "paramType": "dynamic", | ||
| 315 | + "shape": [ | ||
| 316 | + -2 | ||
| 317 | + ] | ||
| 318 | + } | ||
| 319 | + ] | ||
| 320 | + ] | ||
| 321 | + }, | ||
| 322 | + { | ||
| 323 | + "bin_filename": "ForeachSubList_6605e2a61c010ba4cf39ca420adecc84", | ||
| 324 | + "inputs": [ | ||
| 325 | + [ | ||
| 326 | + { | ||
| 327 | + "name": "x1", | ||
| 328 | + "index": 0, | ||
| 329 | + "dtype": "uint8", | ||
| 330 | + "format": "ND", | ||
| 331 | + "paramType": "dynamic", | ||
| 332 | + "shape": [ | ||
| 333 | + -2 | ||
| 334 | + ] | ||
| 335 | + } | ||
| 336 | + ], | ||
| 337 | + [ | ||
| 338 | + { | ||
| 339 | + "name": "x2", | ||
| 340 | + "index": 1, | ||
| 341 | + "dtype": "uint8", | ||
| 342 | + "format": "ND", | ||
| 343 | + "paramType": "dynamic", | ||
| 344 | + "shape": [ | ||
| 345 | + -2 | ||
| 346 | + ] | ||
| 347 | + } | ||
| 348 | + ], | ||
| 349 | + { | ||
| 350 | + "name": "alpha", | ||
| 351 | + "index": 2, | ||
| 352 | + "dtype": "int32", | ||
| 353 | + "format": "ND", | ||
| 354 | + "paramType": "required", | ||
| 355 | + "shape": [ | ||
| 356 | + -2 | ||
| 357 | + ] | ||
| 358 | + } | ||
| 359 | + ], | ||
| 360 | + "outputs": [ | ||
| 361 | + [ | ||
| 362 | + { | ||
| 363 | + "name": "y", | ||
| 364 | + "index": 0, | ||
| 365 | + "dtype": "uint8", | ||
| 366 | + "format": "ND", | ||
| 367 | + "paramType": "dynamic", | ||
| 368 | + "shape": [ | ||
| 369 | + -2 | ||
| 370 | + ] | ||
| 371 | + } | ||
| 372 | + ] | ||
| 373 | + ] | ||
| 215 | } | 374 | } |
| 216 | ] | 375 | ] |
| 217 | } | 376 | } |
| @@ -65,7 +65,8 @@ private: | |||
| 65 | OpAICoreConfig GetA5CoreConfig() const | 65 | OpAICoreConfig GetA5CoreConfig() const |
| 66 | { | 66 | { |
| 67 | OpAICoreConfig config_950; | 67 | OpAICoreConfig config_950; |
| 68 | - std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT32, ge::DT_BF16}; | 68 | + std::vector<ge::DataType> tensor_dtype_list_950 = {ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT32, ge::DT_BF16, |
| 69 | + ge::DT_INT16, ge::DT_INT8, ge::DT_UINT8}; | ||
| 69 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); | 70 | std::vector<ge::Format> format_list_950(tensor_dtype_list_950.size(), ge::FORMAT_ND); |
| 70 | std::vector<ge::DataType> scalar_tensor_dtype_list_950; | 71 | std::vector<ge::DataType> scalar_tensor_dtype_list_950; |
| 71 | std::for_each(tensor_dtype_list_950.cbegin(), tensor_dtype_list_950.cend(), | 72 | std::for_each(tensor_dtype_list_950.cbegin(), tensor_dtype_list_950.cend(), |
| @@ -22,6 +22,9 @@ enum class ForeachSubListTilingKey : uint32_t { | |||
| 22 | TILING_KEY_FLOAT16 = 1, | 22 | TILING_KEY_FLOAT16 = 1, |
| 23 | TILING_KEY_INT32 = 2, | 23 | TILING_KEY_INT32 = 2, |
| 24 | TILING_KEY_BF16 = 3, | 24 | TILING_KEY_BF16 = 3, |
| 25 | + TILING_KEY_INT16 = 4, | ||
| 26 | + TILING_KEY_INT8 = 5, | ||
| 27 | + TILING_KEY_UINT8 = 6, | ||
| 25 | }; | 28 | }; |
| 26 | 29 | ||
| 27 | template <uint32_t schMode> | 30 | template <uint32_t schMode> |
| @@ -34,12 +37,18 @@ __global__ __aicore__ void foreach_sub_list(GM_ADDR x1, GM_ADDR x2, GM_ADDR alph | |||
| 34 | const ForeachSubListTilingData* tilingGm = &tilingData; | 37 | const ForeachSubListTilingData* tilingGm = &tilingData; |
| 35 | 38 | ||
| 36 | if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_FLOAT)) { | 39 | if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_FLOAT)) { |
| 37 | - NsForeachSubList::Process<float>(x1, x2, alpha, y, tilingGm); | 40 | + NsForeachSubList::Process<float, float>(x1, x2, alpha, y, tilingGm); |
| 38 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_FLOAT16)) { | 41 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_FLOAT16)) { |
| 39 | - NsForeachSubList::Process<half>(x1, x2, alpha, y, tilingGm); | 42 | + NsForeachSubList::Process<half, half>(x1, x2, alpha, y, tilingGm); |
| 40 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_INT32)) { | 43 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_INT32)) { |
| 41 | - NsForeachSubList::Process<int32_t>(x1, x2, alpha, y, tilingGm); | 44 | + NsForeachSubList::Process<int32_t, int32_t>(x1, x2, alpha, y, tilingGm); |
| 42 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_BF16)) { | 45 | } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_BF16)) { |
| 43 | - NsForeachSubList::Process<bfloat16_t>(x1, x2, alpha, y, tilingGm); | 46 | + NsForeachSubList::Process<bfloat16_t, float>(x1, x2, alpha, y, tilingGm); |
| 47 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_INT16)) { | ||
| 48 | + NsForeachSubList::Process<int16_t, int32_t>(x1, x2, alpha, y, tilingGm); | ||
| 49 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_INT8)) { | ||
| 50 | + NsForeachSubList::Process<int8_t, int32_t>(x1, x2, alpha, y, tilingGm); | ||
| 51 | + } else if constexpr (schMode == static_cast<uint32_t>(ForeachSubListTilingKey::TILING_KEY_UINT8)) { | ||
| 52 | + NsForeachSubList::Process<uint8_t, int32_t>(x1, x2, alpha, y, tilingGm); | ||
| 44 | } | 53 | } |
| 45 | } | 54 | } |
| @@ -40,13 +40,13 @@ __simt_callee__ inline __gm__ T* SimtGetTensorAddr(GM_ADDR tensorListPtr, int64_ | |||
| 40 | return reinterpret_cast<__gm__ T*>(*(tensorPtr + idx)); | 40 | return reinterpret_cast<__gm__ T*>(*(tensorPtr + idx)); |
| 41 | } | 41 | } |
| 42 | 42 | ||
| 43 | -template <typename T> | 43 | +template <typename T, typename AlphaT> |
| 44 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachSubListSimt(int32_t tensorId, int64_t count, | 44 | __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachSubListSimt(int32_t tensorId, int64_t count, |
| 45 | GM_ADDR x1List, GM_ADDR x2List, | 45 | GM_ADDR x1List, GM_ADDR x2List, |
| 46 | GM_ADDR alpha, GM_ADDR yList) | 46 | GM_ADDR alpha, GM_ADDR yList) |
| 47 | { | 47 | { |
| 48 | - __gm__ T* alphaPtr = reinterpret_cast<__gm__ T*>(alpha); | 48 | + __gm__ AlphaT* alphaPtr = reinterpret_cast<__gm__ AlphaT*>(alpha); |
| 49 | - T alphaVal = alphaPtr[0]; | 49 | + AlphaT alphaVal = alphaPtr[0]; |
| 50 | __gm__ T* x1Data = SimtGetTensorAddr<T>(x1List, tensorId); | 50 | __gm__ T* x1Data = SimtGetTensorAddr<T>(x1List, tensorId); |
| 51 | __gm__ T* x2Data = SimtGetTensorAddr<T>(x2List, tensorId); | 51 | __gm__ T* x2Data = SimtGetTensorAddr<T>(x2List, tensorId); |
| 52 | __gm__ T* yData = SimtGetTensorAddr<T>(yList, tensorId); | 52 | __gm__ T* yData = SimtGetTensorAddr<T>(yList, tensorId); |
| @@ -60,13 +60,16 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void OpForeachSubListSimt | |||
| 60 | float alphaF = static_cast<float>(alphaVal); | 60 | float alphaF = static_cast<float>(alphaVal); |
| 61 | float result = x1Val - alphaF * x2Val; | 61 | float result = x1Val - alphaF * x2Val; |
| 62 | yData[idx] = static_cast<T>(result); | 62 | yData[idx] = static_cast<T>(result); |
| 63 | + } else if constexpr (std::is_same_v<T, int16_t> || std::is_same_v<T, int8_t> || std::is_same_v<T, uint8_t>) { | ||
| 64 | + int32_t result = static_cast<int32_t>(x1Data[idx]) - alphaVal * static_cast<int32_t>(x2Data[idx]); | ||
| 65 | + yData[idx] = static_cast<T>(result); | ||
| 63 | } else { | 66 | } else { |
| 64 | yData[idx] = x1Data[idx] - alphaVal * x2Data[idx]; | 67 | yData[idx] = x1Data[idx] - alphaVal * x2Data[idx]; |
| 65 | } | 68 | } |
| 66 | } | 69 | } |
| 67 | } | 70 | } |
| 68 | 71 | ||
| 69 | -template <typename T> | 72 | +template <typename T, typename AlphaT> |
| 70 | __aicore__ inline void Process(GM_ADDR x1, GM_ADDR x2, GM_ADDR alpha, GM_ADDR y, | 73 | __aicore__ inline void Process(GM_ADDR x1, GM_ADDR x2, GM_ADDR alpha, GM_ADDR y, |
| 71 | const ForeachSubListTilingData* tilingGm) | 74 | const ForeachSubListTilingData* tilingGm) |
| 72 | { | 75 | { |
| @@ -75,8 +78,8 @@ __aicore__ inline void Process(GM_ADDR x1, GM_ADDR x2, GM_ADDR alpha, GM_ADDR y, | |||
| 75 | if (count <= 0) { | 78 | if (count <= 0) { |
| 76 | continue; | 79 | continue; |
| 77 | } | 80 | } |
| 78 | - AscendC::Simt::VF_CALL<OpForeachSubListSimt<T>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x1, x2, alpha, | 81 | + AscendC::Simt::VF_CALL<OpForeachSubListSimt<T, AlphaT>>(AscendC::Simt::Dim3(THREAD_NUM), tensorId, count, x1, |
| 79 | - y); | 82 | + x2, alpha, y); |
| 80 | } | 83 | } |
🟠 High Priority 变更行:foreach_sub_list_simt.h 第 69 行将 受影响行为/契约:但真正执行计算的 VF 内核 失败模式:(1) INT16/INT8/UINT8 输入时 alpha 为 int32 标量,内核只读低 2/1 字节,任何超出 int16/int8/uint8 范围的值(如 alpha=70000、alpha=300)都会被错误解释,且 建议:把 VF 内核改为 ![]() ![]() 不准确? | |||
| 81 | } | 84 | } |
| 82 | 85 | ||
| @@ -19,14 +19,19 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 22 | 25 | ||
| 23 | ASCENDC_TPL_ARGS_DECL(ForeachSubList, | 26 | ASCENDC_TPL_ARGS_DECL(ForeachSubList, |
| 24 | - ASCENDC_TPL_UINT_DECL(schMode, 4, ASCENDC_TPL_UI_LIST, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT, | 27 | + ASCENDC_TPL_UINT_DECL(schMode, 7, ASCENDC_TPL_UI_LIST, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT, |
| 25 | FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16, FOREACH_SUB_LIST_TPL_SCH_MODE_INT32, | 28 | FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16, FOREACH_SUB_LIST_TPL_SCH_MODE_INT32, |
| 26 | - FOREACH_SUB_LIST_TPL_SCH_MODE_BF16)); | 29 | + FOREACH_SUB_LIST_TPL_SCH_MODE_BF16, FOREACH_SUB_LIST_TPL_SCH_MODE_INT16, |
| 30 | + FOREACH_SUB_LIST_TPL_SCH_MODE_INT8, FOREACH_SUB_LIST_TPL_SCH_MODE_UINT8)); | ||
| 27 | 31 | ||
| 28 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL( | 32 | ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL( |
| 29 | schMode, ASCENDC_TPL_UI_LIST, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16, | 33 | schMode, ASCENDC_TPL_UI_LIST, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT, FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16, |
| 30 | - FOREACH_SUB_LIST_TPL_SCH_MODE_INT32, FOREACH_SUB_LIST_TPL_SCH_MODE_BF16))); | 34 | + FOREACH_SUB_LIST_TPL_SCH_MODE_INT32, FOREACH_SUB_LIST_TPL_SCH_MODE_BF16, FOREACH_SUB_LIST_TPL_SCH_MODE_INT16, |
| 35 | + FOREACH_SUB_LIST_TPL_SCH_MODE_INT8, FOREACH_SUB_LIST_TPL_SCH_MODE_UINT8))); | ||
| 31 | 36 | ||
| 32 | 37 | ||
| @@ -372,7 +372,7 @@ TEST_F(ForeachSubListTilingTest, test_tiling_bf16_001) | |||
| 372 | .PlatformInfo(reinterpret_cast<char*>(&platform_info)) | 372 | .PlatformInfo(reinterpret_cast<char*>(&platform_info)) |
| 373 | .NodeInputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) | 373 | .NodeInputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) |
| 374 | .NodeInputTd(1, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) | 374 | .NodeInputTd(1, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) |
| 375 | - .NodeInputTd(2, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) | 375 | + .NodeInputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND) |
| 376 | .NodeOutputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) | 376 | .NodeOutputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_ND) |
| 377 | .TilingData(param.get()) | 377 | .TilingData(param.get()) |
| 378 | .Workspace(ws_size) | 378 | .Workspace(ws_size) |


这里的list如果和上面的tensor_dtype_list一致是不是可以用同一份