已合并
解决foreach_exp/foreach_expm1/foreach_round_off_number/foreach_sub_list算子已知问题 #10178
surezz创建于 25 天前
解决foreach_exp/foreach_expm1/foreach_round_off_number/foreach_sub_list算子已知问题 #10178
已合并
surezz创建于 25 天前
共 30 个文件变更+613-106
@@ -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,
sgqt12
sgqt12sgqt1220 天前

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

likedislike
surezz
surezz
20 天前 评论:
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.cpp16 KERNEL_SRC arch35/foreach_exp.cpp
17 COMPUTE_UNITS ascend95017 COMPUTE_UNITS ascend950
18 AUTO_SYNC false18 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 
24template <uint32_t schMode>27template <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 precision50 * \brief Promote half/bf16 to float32 for precision
51 */51 */
52template <typename T>52template <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 
55template <>58template <>
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 tensors77 * \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_exp97 * \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#define FOREACH_EXP_TPL_SCH_MODE_FLOAT 021#define FOREACH_EXP_TPL_SCH_MODE_FLOAT 0
22#define FOREACH_EXP_TPL_SCH_MODE_FLOAT16 122#define FOREACH_EXP_TPL_SCH_MODE_FLOAT16 1
23#define FOREACH_EXP_TPL_SCH_MODE_BF16 223#define FOREACH_EXP_TPL_SCH_MODE_BF16 2
24+#define FOREACH_EXP_TPL_SCH_MODE_INT16 3
25+#define FOREACH_EXP_TPL_SCH_MODE_INT8 4
26+#define FOREACH_EXP_TPL_SCH_MODE_UINT8 5
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,
sgqt12
sgqt12sgqt1220 天前

这里的6表示的是支持类型数量的位宽,2^6可表示64种,可以减少到3,2^3就可以表示8种了,会减少Key的长度,便于阅读

likedislike
surezz
surezz
20 天前 评论:
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 
29ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, FOREACH_EXP_TPL_SCH_MODE_FLOAT,33ASCENDC_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_H39+#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 optiling130+} // 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,
sgqt12
sgqt12sgqt1220 天前

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

likedislike
surezz
surezz
20 天前 评论:
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 CANNBot10# Generated By CANNBot
11 11 
@@ -18,4 +18,5 @@ add_kernel_sources(
18 KERNEL_SRC arch35/foreach_expm1.cpp18 KERNEL_SRC arch35/foreach_expm1.cpp
19 COMPUTE_UNITS ascend95019 COMPUTE_UNITS ascend950
20 AUTO_SYNC false20 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 
26template <uint32_t schMode>29template <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/bfloat1651+ * \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 tensors62 * \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_expm179 * \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#define FOREACH_EXPM1_TPL_SCH_MODE_FLOAT 023#define FOREACH_EXPM1_TPL_SCH_MODE_FLOAT 0
24#define FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16 124#define FOREACH_EXPM1_TPL_SCH_MODE_FLOAT16 1
25#define FOREACH_EXPM1_TPL_SCH_MODE_BF16 225#define FOREACH_EXPM1_TPL_SCH_MODE_BF16 2
26+#define FOREACH_EXPM1_TPL_SCH_MODE_INT16 3
27+#define FOREACH_EXPM1_TPL_SCH_MODE_INT8 4
28+#define FOREACH_EXPM1_TPL_SCH_MODE_UINT8 5
26 29 
27ASCENDC_TPL_ARGS_DECL(ForeachExpm1,30ASCENDC_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_H41+#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 
30template <uint32_t schMode>31template <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 }
atomgit-bot
atomgit-botatomgit-bot25 天前

🟠 High Priority

变更行:foreach_round_off_number.cpp 第 28 行新增 TILING_KEY_INT16 = 3,但同一文件第 38–44 行的 if constexpr 分派链只覆盖 FLOAT/FLOAT16/BF16,没有为 schMode==3 增加 Process<int16_t> 分支。

受影响行为/契约:本 PR 同时修改了 proto(DYNAMIC_INPUT(x, ...DT_INT16))、def(ascend950 配置加入 DT_INT16)、tiling(GetTilingKeyByDtype 将 DT_INT16 映射为 tilingKey=3)和 binary json(新增 ForeachRoundOffNumber_Int16 条目),int16 路径已完全对外放开,运行时必然加载 schMode=3 的内核实例。

失败模式:当 schMode==3 时 if constexpr 全部分支为假,内核函数体为空——不读取 x、不写 y,算子返回成功但输出 GM 缓冲区从未被写入,得到未初始化/陈旧数据。这是新增 int16 功能被静默破坏的高概率运行时错误(silent data corruption)。对照 foreach_exp.cpp / foreach_expm1.cpp 在本次 PR 中均为每个新增 key 补齐了对应 Process<T> 分支,唯独此处遗漏,且新增的单测只覆盖 host 侧 tiling key 选择(3),无法发现内核分派缺失。

likedislike
不准确?
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#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT 021#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT 0
22#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16 122#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_FLOAT16 1
23#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16 223#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_BF16 2
24+#define FOREACH_ROUND_OFF_NUMBER_TPL_SCH_MODE_INT16 3
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 
30ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST,32ASCENDC_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#endif // FOREACH_ROUND_OFF_NUMBER_TILING_KEY_H38#endif // FOREACH_ROUND_OFF_NUMBER_TILING_KEY_H
@@ -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 
36struct ForeachSubListCompileInfo {};36struct 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+ 
38static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum)106static 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() const65 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 
27template <uint32_t schMode>30template <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 }
atomgit-bot
atomgit-botatomgit-bot25 天前

🟠 High Priority

变更行:foreach_sub_list_simt.h 第 69 行将 Process 改为 template <typename T, typename AlphaT>,foreach_sub_list.cpp 第 40–52 行也按 <T, AlphaT> 分派(FLOAT16→<half,half>、BF16→<bfloat16_t,float>、INT16→<int16_t,int32_t>、INT8→<int8_t,int32_t>、UINT8→<uint8_t,int32_t>),tiling 侧新增 CheckInputCompatibility 强制 alpha 类型与 GetExpectedAlphaDtype 一致(INT16/INT8/UINT8 要求 DT_INT32、BF16 要求 DT_FLOAT)。

受影响行为/契约:但真正执行计算的 VF 内核 OpForeachSubListSimt<T>(第 43 行)仍是单模板参数,第 48–49 行 reinterpret_cast<__gm__ T*>(alpha); T alphaVal = alphaPtr[0]; 仍按 T 读取 alpha;第 78 行 VF_CALL<OpForeachSubListSimt<T>> 也没有把 AlphaT 传下去。AlphaT 是完全无效的“死”参数。

失败模式:(1) INT16/INT8/UINT8 输入时 alpha 为 int32 标量,内核只读低 2/1 字节,任何超出 int16/int8/uint8 范围的值(如 alpha=70000、alpha=300)都会被错误解释,且 x1Data[idx] - alphaVal * x2Data[idx] 的乘法在错误位宽上进行;(2) BF16 输入时 tiling/单测已把 alpha 契约改为 float32,但内核仍按 bf16 读取 float32 缓冲区的前 2 字节(小端下是尾数低位,如 1.0f=0x3F800000 读成 0x0000→0.0),alpha 完全失效。这些是本次新增/强化契约引入的错误,host 侧新增单测只校验 tiling key 与 alpha 类型校验逻辑,无法覆盖内核读取行为。

建议:把 VF 内核改为 template <typename T, typename AlphaT> __simt_vf__ ... OpForeachSubListSimt(...),用 __gm__ AlphaT* alphaPtr = reinterpret_cast<__gm__ AlphaT*>(alpha); AlphaT alphaVal = alphaPtr[0]; 读取 alpha,并把 Process 里的 VF_CALL<OpForeachSubListSimt<T, AlphaT>> 一并修改;同时为整数路径补充内核级数值测试。

likedislike
不准确?
81}84}
82 85 
@@ -19,14 +19,19 @@
19#define FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16 119#define FOREACH_SUB_LIST_TPL_SCH_MODE_FLOAT16 1
20#define FOREACH_SUB_LIST_TPL_SCH_MODE_INT32 220#define FOREACH_SUB_LIST_TPL_SCH_MODE_INT32 2
21#define FOREACH_SUB_LIST_TPL_SCH_MODE_BF16 321#define FOREACH_SUB_LIST_TPL_SCH_MODE_BF16 3
22+#define FOREACH_SUB_LIST_TPL_SCH_MODE_INT16 4
23+#define FOREACH_SUB_LIST_TPL_SCH_MODE_INT8 5
24+#define FOREACH_SUB_LIST_TPL_SCH_MODE_UINT8 6
22 25 
23ASCENDC_TPL_ARGS_DECL(ForeachSubList,26ASCENDC_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 
28ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(32ASCENDC_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#endif // FOREACH_SUB_LIST_TILING_KEY_H37#endif // FOREACH_SUB_LIST_TILING_KEY_H
@@ -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)