已合并
scatterupdate算子输入类型为float8时,动态编译报错修复 #3447
xiaodong666创建于 4月1日
scatterupdate算子输入类型为float8时,动态编译报错修复 #3447
已合并
xiaodong666创建于 4月1日
1 个文件变更+23-23
@@ -229,63 +229,63 @@ __aicore__ inline void SortSimdTensor(
229 }229 }
230}230}
231 231 
232-template <typename IDX_SIZE_T>232+template <typename IDX_SIZE_T, typename VAR_T>
233__aicore__ inline void DeterministicSimdSort(233__aicore__ inline void DeterministicSimdSort(
234 GM_ADDR var, GM_ADDR indices, GM_ADDR updates, GM_ADDR userWs, const ScatterUpdateTilingData& tilingData, TPipe &pipe)234 GM_ADDR var, GM_ADDR indices, GM_ADDR updates, GM_ADDR userWs, const ScatterUpdateTilingData& tilingData, TPipe &pipe)
235{235{
236 if (tilingData.indicesCastMode == CAST_NOT_CAST) {236 if (tilingData.indicesCastMode == CAST_NOT_CAST) {
237- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);237+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);
238 op.Init(var, indices, updates, userWs);238 op.Init(var, indices, updates, userWs);
239 op.Process();239 op.Process();
240 } else if (tilingData.indicesCastMode == CAST_INT32_TO_INT16) {240 } else if (tilingData.indicesCastMode == CAST_INT32_TO_INT16) {
241- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT32_TO_INT16> op(tilingData, pipe);241+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT32_TO_INT16> op(tilingData, pipe);
242 op.Init(var, indices, updates, userWs);242 op.Init(var, indices, updates, userWs);
243 op.Process();243 op.Process();
244 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT32) {244 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT32) {
245- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int32_t, CAST_INT64_TO_INT32> op(tilingData, pipe);245+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int32_t, CAST_INT64_TO_INT32> op(tilingData, pipe);
246 op.Init(var, indices, updates, userWs);246 op.Init(var, indices, updates, userWs);
247 op.Process();247 op.Process();
248 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT16) {248 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT16) {
249- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT64_TO_INT16> op(tilingData, pipe);249+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT64_TO_INT16> op(tilingData, pipe);
250 op.Init(var, indices, updates, userWs);250 op.Init(var, indices, updates, userWs);
251 op.Process();251 op.Process();
252 } else if (tilingData.indicesCastMode == CAST_INT32_TO_UINT8) {252 } else if (tilingData.indicesCastMode == CAST_INT32_TO_UINT8) {
253- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT32_TO_UINT8> op(tilingData, pipe);253+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT32_TO_UINT8> op(tilingData, pipe);
254 op.Init(var, indices, updates, userWs);254 op.Init(var, indices, updates, userWs);
255 op.Process();255 op.Process();
256 } else if (tilingData.indicesCastMode == CAST_INT64_TO_UINT8) {256 } else if (tilingData.indicesCastMode == CAST_INT64_TO_UINT8) {
257- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT64_TO_UINT8> op(tilingData, pipe);257+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT64_TO_UINT8> op(tilingData, pipe);
258 op.Init(var, indices, updates, userWs);258 op.Init(var, indices, updates, userWs);
259 op.Process();259 op.Process();
260 }260 }
261}261}
262 262 
263-template <typename IDX_SIZE_T>263+template <typename IDX_SIZE_T, typename VAR_T>
264__aicore__ inline void DeterministicSimtSort(264__aicore__ inline void DeterministicSimtSort(
265 GM_ADDR var, GM_ADDR indices, GM_ADDR updates, GM_ADDR userWs, const ScatterUpdateTilingData& tilingData, TPipe &pipe)265 GM_ADDR var, GM_ADDR indices, GM_ADDR updates, GM_ADDR userWs, const ScatterUpdateTilingData& tilingData, TPipe &pipe)
266{266{
267 if (tilingData.indicesCastMode == CAST_NOT_CAST) {267 if (tilingData.indicesCastMode == CAST_NOT_CAST) {
268- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);268+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);
269 op.Init(var, indices, updates, userWs);269 op.Init(var, indices, updates, userWs);
270 op.Process();270 op.Process();
271 } else if (tilingData.indicesCastMode == CAST_INT32_TO_INT16) {271 } else if (tilingData.indicesCastMode == CAST_INT32_TO_INT16) {
272- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT32_TO_INT16> op(tilingData, pipe);272+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT32_TO_INT16> op(tilingData, pipe);
273 op.Init(var, indices, updates, userWs);273 op.Init(var, indices, updates, userWs);
274 op.Process();274 op.Process();
275 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT32) {275 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT32) {
276- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int32_t, CAST_INT64_TO_INT32> op(tilingData, pipe);276+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int32_t, CAST_INT64_TO_INT32> op(tilingData, pipe);
277 op.Init(var, indices, updates, userWs);277 op.Init(var, indices, updates, userWs);
278 op.Process();278 op.Process();
279 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT16) {279 } else if (tilingData.indicesCastMode == CAST_INT64_TO_INT16) {
280- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT64_TO_INT16> op(tilingData, pipe);280+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, int16_t, CAST_INT64_TO_INT16> op(tilingData, pipe);
281 op.Init(var, indices, updates, userWs);281 op.Init(var, indices, updates, userWs);
282 op.Process();282 op.Process();
283 } else if (tilingData.indicesCastMode == CAST_INT32_TO_UINT8) {283 } else if (tilingData.indicesCastMode == CAST_INT32_TO_UINT8) {
284- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT32_TO_UINT8> op(tilingData, pipe);284+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT32_TO_UINT8> op(tilingData, pipe);
285 op.Init(var, indices, updates, userWs);285 op.Init(var, indices, updates, userWs);
286 op.Process();286 op.Process();
287 } else if (tilingData.indicesCastMode == CAST_INT64_TO_UINT8) {287 } else if (tilingData.indicesCastMode == CAST_INT64_TO_UINT8) {
288- ScatterUpdateDeterministicSimt<DTYPE_VAR, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT64_TO_UINT8> op(tilingData, pipe);288+ ScatterUpdateDeterministicSimt<VAR_T, DTYPE_INDICES, IDX_SIZE_T, false, uint8_t, CAST_INT64_TO_UINT8> op(tilingData, pipe);
289 op.Init(var, indices, updates, userWs);289 op.Init(var, indices, updates, userWs);
290 op.Process();290 op.Process();
291 }291 }
@@ -311,19 +311,19 @@ extern "C" __global__ __aicore__ void scatter_update(GM_ADDR var, GM_ADDR indice
311 311 
312 ///////////////////// SIMT //////////////////////////312 ///////////////////// SIMT //////////////////////////
313 if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR32_SCALAR)) {313 if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR32_SCALAR)) {
314- ScatterUpdateSimt<DTYPE_INDICES, DTYPE_VAR, uint32_t, true> op(tilingData);314+ ScatterUpdateSimt<DTYPE_INDICES, VAR_T, uint32_t, true> op(tilingData);
315 op.Init(var, indices, updates, userWs);315 op.Init(var, indices, updates, userWs);
316 op.Process();316 op.Process();
317 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR32_TENSOR)) {317 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR32_TENSOR)) {
318- ScatterUpdateSimt<DTYPE_INDICES, DTYPE_VAR, uint32_t, false> op(tilingData);318+ ScatterUpdateSimt<DTYPE_INDICES, VAR_T, uint32_t, false> op(tilingData);
319 op.Init(var, indices, updates, userWs);319 op.Init(var, indices, updates, userWs);
320 op.Process();320 op.Process();
321 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR64_SCALAR)) {321 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR64_SCALAR)) {
322- ScatterUpdateSimt<DTYPE_INDICES, DTYPE_VAR, uint64_t, true> op(tilingData);322+ ScatterUpdateSimt<DTYPE_INDICES, VAR_T, uint64_t, true> op(tilingData);
323 op.Init(var, indices, updates, userWs);323 op.Init(var, indices, updates, userWs);
324 op.Process();324 op.Process();
325 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR64_TENSOR)) {325 } else if (TILING_KEY_IS(TILING_KEY_SIMT_ADDR64_TENSOR)) {
326- ScatterUpdateSimt<DTYPE_INDICES, DTYPE_VAR, uint64_t, false> op(tilingData);326+ ScatterUpdateSimt<DTYPE_INDICES, VAR_T, uint64_t, false> op(tilingData);
327 op.Init(var, indices, updates, userWs);327 op.Init(var, indices, updates, userWs);
328 op.Process();328 op.Process();
329 329 
@@ -359,20 +359,20 @@ extern "C" __global__ __aicore__ void scatter_update(GM_ADDR var, GM_ADDR indice
359 359 
360 ///////////////////// Deterministic //////////////////////////360 ///////////////////// Deterministic //////////////////////////
361 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMD_SPLITCOL)) {361 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMD_SPLITCOL)) {
362- ScatterUpdateDeterministicSimd<DTYPE_VAR, DTYPE_INDICES, int64_t, true, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);362+ ScatterUpdateDeterministicSimd<VAR_T, DTYPE_INDICES, int64_t, true, DTYPE_INDICES, CAST_NOT_CAST> op(tilingData, pipe);
363 op.Init(var, indices, updates, userWs);363 op.Init(var, indices, updates, userWs);
364 op.Process();364 op.Process();
365 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMD_SPLITROW)) {365 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMD_SPLITROW)) {
366 if (tilingData.isIndicesSizeInt64) {366 if (tilingData.isIndicesSizeInt64) {
367- DeterministicSimdSort<int64_t>(var, indices, updates, userWs, tilingData, pipe);367+ DeterministicSimdSort<int64_t, VAR_T>(var, indices, updates, userWs, tilingData, pipe);
368 } else {368 } else {
369- DeterministicSimdSort<int32_t>(var, indices, updates, userWs, tilingData, pipe);369+ DeterministicSimdSort<int32_t, VAR_T>(var, indices, updates, userWs, tilingData, pipe);
370 }370 }
371 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMT)) {371 } else if (TILING_KEY_IS(TILING_KEY_DETERMINISTIC_SIMT)) {
372 if (tilingData.isIndicesSizeInt64) {372 if (tilingData.isIndicesSizeInt64) {
373- DeterministicSimtSort<int64_t>(var, indices, updates, userWs, tilingData, pipe);373+ DeterministicSimtSort<int64_t, VAR_T>(var, indices, updates, userWs, tilingData, pipe);
374 } else {374 } else {
375- DeterministicSimtSort<int32_t>(var, indices, updates, userWs, tilingData, pipe);375+ DeterministicSimtSort<int32_t, VAR_T>(var, indices, updates, userWs, tilingData, pipe);
376 }376 }
377 } else if (TILING_KEY_IS(TILING_KEY_0)) {377 } else if (TILING_KEY_IS(TILING_KEY_0)) {
378 return;378 return;