草稿
[WIP]scatter_nd_max新增simt sort模板 #3
klein8793创建于 5月21日
[WIP]scatter_nd_max新增simt sort模板 #3
草稿
共 12 个文件变更+665-17
| @@ -38,6 +38,7 @@ struct ScatterNdCommonSimtTilingData{ | |||
| 38 | 38 | ||
| 39 | struct ScatterNdCommonSimtSortTilingData{ | 39 | struct ScatterNdCommonSimtSortTilingData{ |
| 40 | uint64_t strideList[MAX_RANK_COUNT_NUM]; | 40 | uint64_t strideList[MAX_RANK_COUNT_NUM]; |
| 41 | + uint64_t outPutShape[MAX_SHAPE_RANK_NUM]; | ||
| 41 | int64_t indicesFactor; | 42 | int64_t indicesFactor; |
| 42 | int64_t afterAxis; | 43 | int64_t afterAxis; |
| 43 | int64_t varInAxis; | 44 | int64_t varInAxis; |
| @@ -28,14 +28,14 @@ | |||
| 28 | 28 | ||
| 29 | 29 | ||
| 30 | 30 | ||
| 31 | - | 31 | +#define TPL_MODE_TEMPLATE_SIMT_SORT 5 |
| 32 | 32 | ||
| 33 | 33 | ||
| 34 | namespace ScatterNdCommon { | 34 | namespace ScatterNdCommon { |
| 35 | 35 | ||
| 36 | ASCENDC_TPL_ARGS_DECL( | 36 | ASCENDC_TPL_ARGS_DECL( |
| 37 | ScatterNdMax, | 37 | ScatterNdMax, |
| 38 | - ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, 2, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMD_SORT, TPL_MODE_TEMPLATE_SIMT), | 38 | + ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, 2, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMD_SORT, TPL_MODE_TEMPLATE_SIMT_SORT, TPL_MODE_TEMPLATE_SIMT), |
| 39 | ASCENDC_TPL_UINT_DECL(CAST_MODE, 3, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), | 39 | ASCENDC_TPL_UINT_DECL(CAST_MODE, 3, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), |
| 40 | ASCENDC_TPL_UINT_DECL(ADDR_MODE, 1, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64) | 40 | ASCENDC_TPL_UINT_DECL(ADDR_MODE, 1, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64) |
| 41 | ); | 41 | ); |
| @@ -49,6 +49,13 @@ ASCENDC_TPL_SEL( | |||
| 49 | ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), | 49 | ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), |
| 50 | ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimdSortTilingData) | 50 | ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimdSortTilingData) |
| 51 | ), | 51 | ), |
| 52 | + ASCENDC_TPL_ARGS_SEL( | ||
| 53 | + ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), | ||
| 54 | + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT_SORT), | ||
| 55 | + ASCENDC_TPL_UINT_SEL(CAST_MODE, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), | ||
| 56 | + ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), | ||
| 57 | + ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimtSortTilingData) | ||
| 58 | + ), | ||
| 52 | ASCENDC_TPL_ARGS_SEL( | 59 | ASCENDC_TPL_ARGS_SEL( |
| 53 | ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), | 60 | ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), |
| 54 | ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT), | 61 | ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT), |
| @@ -28,14 +28,14 @@ | |||
| 28 | 28 | ||
| 29 | 29 | ||
| 30 | 30 | ||
| 31 | - | 31 | +#define TPL_MODE_TEMPLATE_SIMT_SORT 5 |
| 32 | 32 | ||
| 33 | 33 | ||
| 34 | namespace ScatterNdCommon { | 34 | namespace ScatterNdCommon { |
| 35 | 35 | ||
| 36 | ASCENDC_TPL_ARGS_DECL( | 36 | ASCENDC_TPL_ARGS_DECL( |
| 37 | ScatterNdMin, | 37 | ScatterNdMin, |
| 38 | - ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, 2, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMD_SORT, TPL_MODE_TEMPLATE_SIMT), | 38 | + ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, 2, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMD_SORT, TPL_MODE_TEMPLATE_SIMT_SORT, TPL_MODE_TEMPLATE_SIMT), |
| 39 | ASCENDC_TPL_UINT_DECL(CAST_MODE, 3, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), | 39 | ASCENDC_TPL_UINT_DECL(CAST_MODE, 3, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), |
| 40 | ASCENDC_TPL_UINT_DECL(ADDR_MODE, 1, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64) | 40 | ASCENDC_TPL_UINT_DECL(ADDR_MODE, 1, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64) |
| 41 | ); | 41 | ); |
| @@ -49,6 +49,13 @@ ASCENDC_TPL_SEL( | |||
| 49 | ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), | 49 | ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), |
| 50 | ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimdSortTilingData) | 50 | ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimdSortTilingData) |
| 51 | ), | 51 | ), |
| 52 | + ASCENDC_TPL_ARGS_SEL( | ||
| 53 | + ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), | ||
| 54 | + ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT_SORT), | ||
| 55 | + ASCENDC_TPL_UINT_SEL(CAST_MODE, ASCENDC_TPL_UI_LIST, CAST_0, CAST_1, CAST_2, CAST_3, CAST_4, CAST_5), | ||
| 56 | + ASCENDC_TPL_UINT_SEL(ADDR_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_ADDR_INT32, TPL_MODE_ADDR_INT64), | ||
| 57 | + ASCENDC_TPL_TILING_STRUCT_SEL(ScatterNdCommonSimtSortTilingData) | ||
| 58 | + ), | ||
| 52 | ASCENDC_TPL_ARGS_SEL( | 59 | ASCENDC_TPL_ARGS_SEL( |
| 53 | ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), | 60 | ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_AIV_ONLY), |
| 54 | ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT), | 61 | ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, TPL_MODE_TEMPLATE_SIMT), |
| @@ -57,7 +57,7 @@ IMPL_OP_OPTILING(ScatterNdMax) | |||
| 57 | .TilingParse<ScatterNdCommonCompileInfo>(TilingPrepare4ScatterNdMax); | 57 | .TilingParse<ScatterNdCommonCompileInfo>(TilingPrepare4ScatterNdMax); |
| 58 | 58 | ||
| 59 | REGISTER_TILING_TEMPLATE("ScatterNdMax", ScatterNdMaxSimdSortTiling, 2); | 59 | REGISTER_TILING_TEMPLATE("ScatterNdMax", ScatterNdMaxSimdSortTiling, 2); |
| 60 | - | 60 | +REGISTER_TILING_TEMPLATE("ScatterNdMax", ScatterNdMaxSimtSortTiling, 5); |
| 61 | REGISTER_TILING_TEMPLATE("ScatterNdMax", ScatterNdMaxSimtTiling, 8); | 61 | REGISTER_TILING_TEMPLATE("ScatterNdMax", ScatterNdMaxSimtTiling, 8); |
| 62 | 62 | ||
| 63 | 63 | ||
| @@ -18,7 +18,7 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | - | 21 | +#include "index/scatter_nd_common/op_host/arch35/scatter_nd_common_simt_sort_tiling.h" |
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | namespace optiling { | 24 | namespace optiling { |
| @@ -33,6 +33,16 @@ public: | |||
| 33 | {} | 33 | {} |
| 34 | }; | 34 | }; |
| 35 | 35 | ||
| 36 | +// ---------------------------ScatterNdMax Simt Sort Tiling--------------------------- | ||
| 37 | +class ScatterNdMaxSimtSortTiling : public ScatterNdCommonSimtSortTiling | ||
| 38 | +{ | ||
| 39 | +public: | ||
| 40 | + explicit ScatterNdMaxSimtSortTiling(gert::TilingContext* context) : ScatterNdCommonSimtSortTiling(context) | ||
| 41 | + {} | ||
| 42 | + ~ScatterNdMaxSimtSortTiling() override | ||
| 43 | + {} | ||
| 44 | +}; | ||
| 45 | + | ||
| 36 | // ---------------------------ScatterNdMax Simd Sort Tiling--------------------------- | 46 | // ---------------------------ScatterNdMax Simd Sort Tiling--------------------------- |
| 37 | class ScatterNdMaxSimdSortTiling : public ScatterNdCommonSimdSortTiling | 47 | class ScatterNdMaxSimdSortTiling : public ScatterNdCommonSimdSortTiling |
| 38 | { | 48 | { |
| @@ -16,6 +16,7 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | 21 | ||
| 21 | 22 | ||
| @@ -125,6 +126,70 @@ __global__ __aicore__ void scatter_nd_max( | |||
| 125 | op.Init(x, indices, updates, y); | 126 | op.Init(x, indices, updates, y); |
| 126 | op.Process(); | 127 | op.Process(); |
| 127 | } | 128 | } |
| 129 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_0 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 130 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 131 | + return; | ||
| 132 | + } else { | ||
| 133 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 134 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, DTYPE_INDICES, CAST_0, MODE_MAX> op(tilingData, pipe); | ||
| 135 | + op.Init(x, indices, updates, y, workspace); | ||
| 136 | + op.Process(); | ||
| 137 | + } | ||
| 138 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_1 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 139 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 140 | + return; | ||
| 141 | + } else { | ||
| 142 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 143 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int16_t, CAST_1, MODE_MAX> op(tilingData, pipe); | ||
| 144 | + op.Init(x, indices, updates, y, workspace); | ||
| 145 | + op.Process(); | ||
| 146 | + } | ||
| 147 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_2 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 148 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 149 | + return; | ||
| 150 | + } else { | ||
| 151 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 152 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int32_t, CAST_2, MODE_MAX> op(tilingData, pipe); | ||
| 153 | + op.Init(x, indices, updates, y, workspace); | ||
| 154 | + op.Process(); | ||
| 155 | + } | ||
| 156 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_3 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 157 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 158 | + return; | ||
| 159 | + } else { | ||
| 160 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 161 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int16_t, CAST_3, MODE_MAX> op(tilingData, pipe); | ||
| 162 | + op.Init(x, indices, updates, y, workspace); | ||
| 163 | + op.Process(); | ||
| 164 | + } | ||
| 165 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_4 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 166 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 167 | + return; | ||
| 168 | + } else { | ||
| 169 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 170 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, uint8_t, CAST_4, MODE_MAX> op(tilingData, pipe); | ||
| 171 | + op.Init(x, indices, updates, y, workspace); | ||
| 172 | + op.Process(); | ||
| 173 | + } | ||
| 174 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_5 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 175 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 176 | + return; | ||
| 177 | + } else { | ||
| 178 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 179 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, uint8_t, CAST_5, MODE_MAX> op(tilingData, pipe); | ||
| 180 | + op.Init(x, indices, updates, y, workspace); | ||
| 181 | + op.Process(); | ||
| 182 | + } | ||
| 183 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_0 && ADDR_MODE == TPL_MODE_ADDR_INT64) { | ||
| 184 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 185 | + return; | ||
| 186 | + } else { | ||
| 187 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 188 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint64_t, DTYPE_INDICES, CAST_0, MODE_MAX> op(tilingData, pipe); | ||
| 189 | + op.Init(x, indices, updates, y, workspace); | ||
| 190 | + op.Process(); | ||
| 191 | + } | ||
| 128 | } | 192 | } |
| 129 | 193 | ||
| 194 | + | ||
| 130 | } | 195 | } |
| @@ -57,7 +57,7 @@ IMPL_OP_OPTILING(ScatterNdMin) | |||
| 57 | .TilingParse<ScatterNdCommonCompileInfo>(TilingPrepare4ScatterNdMin); | 57 | .TilingParse<ScatterNdCommonCompileInfo>(TilingPrepare4ScatterNdMin); |
| 58 | 58 | ||
| 59 | REGISTER_TILING_TEMPLATE("ScatterNdMin", ScatterNdMinSimdSortTiling, 2); | 59 | REGISTER_TILING_TEMPLATE("ScatterNdMin", ScatterNdMinSimdSortTiling, 2); |
| 60 | - | 60 | +REGISTER_TILING_TEMPLATE("ScatterNdMin", ScatterNdMinSimtSortTiling, 5); |
| 61 | REGISTER_TILING_TEMPLATE("ScatterNdMin", ScatterNdMinSimtTiling, 8); | 61 | REGISTER_TILING_TEMPLATE("ScatterNdMin", ScatterNdMinSimtTiling, 8); |
| 62 | 62 | ||
| 63 | 63 | ||
| @@ -18,7 +18,7 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | - | 21 | +#include "index/scatter_nd_common/op_host/arch35/scatter_nd_common_simt_sort_tiling.h" |
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | namespace optiling { | 24 | namespace optiling { |
| @@ -33,15 +33,15 @@ public: | |||
| 33 | {} | 33 | {} |
| 34 | }; | 34 | }; |
| 35 | 35 | ||
| 36 | -// // ---------------------------ScatterNdMin Simt Sort Tiling--------------------------- | 36 | +// ---------------------------ScatterNdMin Simt Sort Tiling--------------------------- |
| 37 | -// class ScatterNdMinSimtSortTiling : public ScatterNdCommonSimtSortTiling | 37 | +class ScatterNdMinSimtSortTiling : public ScatterNdCommonSimtSortTiling |
| 38 | -// { | 38 | +{ |
| 39 | -// public: | 39 | +public: |
| 40 | -// explicit ScatterNdMinSimtSortTiling(gert::TilingContext* context) : ScatterNdCommonSimtSortTiling(context) | 40 | + explicit ScatterNdMinSimtSortTiling(gert::TilingContext* context) : ScatterNdCommonSimtSortTiling(context) |
| 41 | -// {} | 41 | + {} |
| 42 | -// ~ScatterNdMinSimtSortTiling() override | 42 | + ~ScatterNdMinSimtSortTiling() override |
| 43 | -// {} | 43 | + {} |
| 44 | -// }; | 44 | +}; |
| 45 | 45 | ||
| 46 | // ---------------------------ScatterNdMin Simd Sort Tiling--------------------------- | 46 | // ---------------------------ScatterNdMin Simd Sort Tiling--------------------------- |
| 47 | class ScatterNdMinSimdSortTiling : public ScatterNdCommonSimdSortTiling | 47 | class ScatterNdMinSimdSortTiling : public ScatterNdCommonSimdSortTiling |
| @@ -16,6 +16,7 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | 21 | ||
| 21 | 22 | ||
| @@ -125,6 +126,69 @@ __global__ __aicore__ void scatter_nd_min( | |||
| 125 | op.Init(x, indices, updates, y); | 126 | op.Init(x, indices, updates, y); |
| 126 | op.Process(); | 127 | op.Process(); |
| 127 | } | 128 | } |
| 129 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_0 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 130 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 131 | + return; | ||
| 132 | + } else { | ||
| 133 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 134 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, DTYPE_INDICES, CAST_0, MODE_MIN> op(tilingData, pipe); | ||
| 135 | + op.Init(x, indices, updates, y, workspace); | ||
| 136 | + op.Process(); | ||
| 137 | + } | ||
| 138 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_1 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 139 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 140 | + return; | ||
| 141 | + } else { | ||
| 142 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 143 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int16_t, CAST_1, MODE_MIN> op(tilingData, pipe); | ||
| 144 | + op.Init(x, indices, updates, y, workspace); | ||
| 145 | + op.Process(); | ||
| 146 | + } | ||
| 147 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_2 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 148 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 149 | + return; | ||
| 150 | + } else { | ||
| 151 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 152 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int32_t, CAST_2, MODE_MIN> op(tilingData, pipe); | ||
| 153 | + op.Init(x, indices, updates, y, workspace); | ||
| 154 | + op.Process(); | ||
| 155 | + } | ||
| 156 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_3 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 157 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 158 | + return; | ||
| 159 | + } else { | ||
| 160 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 161 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, int16_t, CAST_3, MODE_MIN> op(tilingData, pipe); | ||
| 162 | + op.Init(x, indices, updates, y, workspace); | ||
| 163 | + op.Process(); | ||
| 164 | + } | ||
| 165 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_4 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 166 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 167 | + return; | ||
| 168 | + } else { | ||
| 169 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 170 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, uint8_t, CAST_4, MODE_MIN> op(tilingData, pipe); | ||
| 171 | + op.Init(x, indices, updates, y, workspace); | ||
| 172 | + op.Process(); | ||
| 173 | + } | ||
| 174 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_5 && ADDR_MODE == TPL_MODE_ADDR_INT32) { | ||
| 175 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 176 | + return; | ||
| 177 | + } else { | ||
| 178 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 179 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint32_t, uint8_t, CAST_5, MODE_MIN> op(tilingData, pipe); | ||
| 180 | + op.Init(x, indices, updates, y, workspace); | ||
| 181 | + op.Process(); | ||
| 182 | + } | ||
| 183 | + } else if constexpr (TEMPLATE_MODE == TPL_MODE_TEMPLATE_SIMT_SORT && CAST_MODE == CAST_0 && ADDR_MODE == TPL_MODE_ADDR_INT64) { | ||
| 184 | + if constexpr (std::is_same<int8_t, DTYPE_VAR>::value || std::is_same<int16_t, DTYPE_VAR>::value) { | ||
| 185 | + return; | ||
| 186 | + } else { | ||
| 187 | + GET_TILING_DATA_WITH_STRUCT(ScatterNdCommonSimtSortTilingData, tilingData, tiling); | ||
| 188 | + ScatterNdCommon::ScatterNdCommonSimtSort<DTYPE_VAR, DTYPE_INDICES, uint64_t, DTYPE_INDICES, CAST_0, MODE_MIN> op(tilingData, pipe); | ||
| 189 | + op.Init(x, indices, updates, y, workspace); | ||
| 190 | + op.Process(); | ||
| 191 | + } | ||
| 128 | } | 192 | } |
| 129 | 193 | ||
| 130 | } | 194 | } |