已合并
scatterupdate算子输入类型为float8时,动态编译报错修复 #3447
xiaodong666创建于 4月1日
scatterupdate算子输入类型为float8时,动态编译报错修复 #3447
已合并
共 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; |