已合并
[kernel] embedding_hash_table and sparse_tensor_dense_mat_mul modify simt api #4154
majiayuan创建于 4月23日
[kernel] embedding_hash_table and sparse_tensor_dense_mat_mul modify simt api #4154
已合并
共 10 个文件变更+105-105
| @@ -13,6 +13,8 @@ | |||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | + | ||
| 17 | + | ||
| 16 | 18 | ||
| 17 | static constexpr uint8_t VALID_FLAG_MASK = 0b00000001; | 19 | static constexpr uint8_t VALID_FLAG_MASK = 0b00000001; |
| 18 | 20 | ||
| @@ -25,10 +27,10 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(EMBEDDING_THREAD_NUM) inline void ComputeAda | |||
| 25 | __gm__ T* gmMaxGradNorm, __gm__ T* gmMOut, __gm__ T* gmVOut, __gm__ T* gmBeta1PowerOut, __gm__ T* gmBeta2PowerOut, | 27 | __gm__ T* gmMaxGradNorm, __gm__ T* gmMOut, __gm__ T* gmVOut, __gm__ T* gmBeta1PowerOut, __gm__ T* gmBeta2PowerOut, |
| 26 | __gm__ T* gmMaxGradNormOut) | 28 | __gm__ T* gmMaxGradNormOut) |
| 27 | { | 29 | { |
| 28 | - int32_t threadXIdx = Simt::GetThreadIdx<0>(); | 30 | + int32_t threadXIdx = threadIdx.x; |
| 29 | - int32_t threadYIdx = Simt::GetThreadIdx<1>(); | 31 | + int32_t threadYIdx = threadIdx.y; |
| 30 | - int32_t threadXNum = Simt::GetThreadNum<0>(); | 32 | + int32_t threadXNum = blockDim.x; |
| 31 | - int32_t threadYNum = Simt::GetThreadNum<1>(); | 33 | + int32_t threadYNum = blockDim.y; |
| 32 | 34 | ||
| 33 | int64_t tableAddr = *(reinterpret_cast<__gm__ int64_t*>(gmTableIn[0])); | 35 | int64_t tableAddr = *(reinterpret_cast<__gm__ int64_t*>(gmTableIn[0])); |
| 34 | __gm__ uint8_t *table = reinterpret_cast<__gm__ uint8_t*>(tableAddr); | 36 | __gm__ uint8_t *table = reinterpret_cast<__gm__ uint8_t*>(tableAddr); |
| @@ -101,10 +103,10 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(EMBEDDING_THREAD_NUM) inline void ComputeAda | |||
| 101 | 103 | ||
| 102 | float denom = 1.0; | 104 | float denom = 1.0; |
| 103 | if (amsgrad != 0) { | 105 | if (amsgrad != 0) { |
| 104 | - maxGradNormLocal = Simt::Max(maxGradNormLocal, vOutLocal); | 106 | + maxGradNormLocal = fmaxf(maxGradNormLocal, vOutLocal); |
| 105 | - denom = Simt::Sqrt(-maxGradNormLocal / (beta2PowerLocal + (-1))) + epsilonLocal; | 107 | + denom = sqrtf(-maxGradNormLocal / (beta2PowerLocal + (-1))) + epsilonLocal; |
| 106 | } else { | 108 | } else { |
| 107 | - denom = Simt::Sqrt(-vOutLocal / (beta2PowerLocal + (-1))) + epsilonLocal; | 109 | + denom = sqrtf(-vOutLocal / (beta2PowerLocal + (-1))) + epsilonLocal; |
| 108 | } | 110 | } |
| 109 | 111 | ||
| 110 | value = value + (lrLocal * mOutLocal / (beta1PowerLocal + (-1))) / denom; | 112 | value = value + (lrLocal * mOutLocal / (beta1PowerLocal + (-1))) / denom; |
| @@ -183,8 +185,8 @@ public: | |||
| 183 | 185 | ||
| 184 | __aicore__ inline void Process() | 186 | __aicore__ inline void Process() |
| 185 | { | 187 | { |
| 186 | - Simt::VF_CALL<ComputeAdamW<T>>( | 188 | + asc_vf_call<ComputeAdamW<T>>( |
| 187 | - Simt::Dim3{static_cast<uint32_t>(blockX_), static_cast<uint32_t>(blockY_)}, tableSize_, keyNum_, unusedKey, | 189 | + dim3{static_cast<uint32_t>(blockX_), static_cast<uint32_t>(blockY_)}, tableSize_, keyNum_, unusedKey, |
| 188 | bucketSizeByte, xLoopSize_, embeddingDim_, maximize_, amsgrad_, gmTableIn_.GetPhyAddr(0), | 190 | bucketSizeByte, xLoopSize_, embeddingDim_, maximize_, amsgrad_, gmTableIn_.GetPhyAddr(0), |
| 189 | gmKeys_.GetPhyAddr(0), gmM_.GetPhyAddr(0), gmV_.GetPhyAddr(0), gmBeta1Power_.GetPhyAddr(0), | 191 | gmKeys_.GetPhyAddr(0), gmM_.GetPhyAddr(0), gmV_.GetPhyAddr(0), gmBeta1Power_.GetPhyAddr(0), |
| 190 | gmBeta2Power_.GetPhyAddr(0), gmLr_.GetPhyAddr(0), gmWeightDecay_.GetPhyAddr(0), gmBeta1_.GetPhyAddr(0), | 192 | gmBeta2Power_.GetPhyAddr(0), gmLr_.GetPhyAddr(0), gmWeightDecay_.GetPhyAddr(0), gmBeta1_.GetPhyAddr(0), |
| @@ -18,6 +18,8 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | + | ||
| 22 | + | ||
| 21 | 23 | ||
| 22 | namespace EmbeddingHashTableExportAicore { | 24 | namespace EmbeddingHashTableExportAicore { |
| 23 | using namespace AscendC; | 25 | using namespace AscendC; |
| @@ -176,7 +178,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(1) inline void SaveToCoreSyncWorkspace( | |||
| 176 | return; | 178 | return; |
| 177 | } | 179 | } |
| 178 | 180 | ||
| 179 | - if (Simt::GetThreadIdx() == 0) { | 181 | + if (threadIdx.x == 0) { |
| 180 | coreSyncWorkspaceGm[tableIndx * maxCoreNum + blockIdx] = threadCountKeysToExportUB[maxThreadNum]; | 182 | coreSyncWorkspaceGm[tableIndx * maxCoreNum + blockIdx] = threadCountKeysToExportUB[maxThreadNum]; |
| 181 | } | 183 | } |
| 182 | } | 184 | } |
| @@ -190,8 +192,8 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(1) inline void AtomicSubToGm( | |||
| 190 | return; | 192 | return; |
| 191 | } | 193 | } |
| 192 | 194 | ||
| 193 | - if (Simt::GetThreadIdx() == 0) { | 195 | + if (threadIdx.x == 0) { |
| 194 | - Simt::AtomicSub(tableHandleStructGm + SIZE_ALL_NO_EXPORT_IDX, threadCountReFreshExportFlagUB[maxThreadNum]); | 196 | + asc_atomic_sub(tableHandleStructGm + SIZE_ALL_NO_EXPORT_IDX, threadCountReFreshExportFlagUB[maxThreadNum]); |
| 195 | } | 197 | } |
| 196 | } | 198 | } |
| 197 | 199 | ||
| @@ -205,20 +207,18 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void CountP | |||
| 205 | return; | 207 | return; |
| 206 | } | 208 | } |
| 207 | 209 | ||
| 208 | - if (Simt::GetThreadIdx() >= maxThreadNum) { | 210 | + if (threadIdx.x >= maxThreadNum) { |
| 209 | return; | 211 | return; |
| 210 | } | 212 | } |
| 211 | 213 | ||
| 212 | - int64_t curThreadProcessKeys = | 214 | + int64_t curThreadProcessKeys = threadIdx.x < (usedThreadNum - 1) ? normalThreadProcessKeys : tailThreadProcessKeys; |
| 213 | - Simt::GetThreadIdx() < (usedThreadNum - 1) ? normalThreadProcessKeys : tailThreadProcessKeys; | ||
| 214 | int64_t keysNumToExport = 0; | 215 | int64_t keysNumToExport = 0; |
| 215 | - if (Simt::GetThreadIdx() < usedThreadNum) { | 216 | + if (threadIdx.x < usedThreadNum) { |
| 216 | __gm__ uint8_t* tableAddrU8 = reinterpret_cast<__gm__ uint8_t*>(tableAddr); | 217 | __gm__ uint8_t* tableAddrU8 = reinterpret_cast<__gm__ uint8_t*>(tableAddr); |
| 217 | 218 | ||
| 218 | for (int64_t i = 0; i < curThreadProcessKeys; i++) { | 219 | for (int64_t i = 0; i < curThreadProcessKeys; i++) { |
| 219 | uint8_t flag = tableAddrU8 | 220 | uint8_t flag = tableAddrU8 |
| 220 | - [keyWidthByte * | 221 | + [keyWidthByte * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i) + |
| 221 | - (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i) + | ||
| 222 | KEY_FLAG_OFFSET_OF_BYTE]; | 222 | KEY_FLAG_OFFSET_OF_BYTE]; |
| 223 | 223 | ||
| 224 | if ((flag & VALID_FLAG_MASK) && !(flag & EVICTED_FLAG_MASK) && | 224 | if ((flag & VALID_FLAG_MASK) && !(flag & EVICTED_FLAG_MASK) && |
| @@ -227,7 +227,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void CountP | |||
| 227 | } | 227 | } |
| 228 | } | 228 | } |
| 229 | } | 229 | } |
| 230 | - threadCountKeysToExportUB[Simt::GetThreadIdx()] = keysNumToExport; | 230 | + threadCountKeysToExportUB[threadIdx.x] = keysNumToExport; |
| 231 | } | 231 | } |
| 232 | 232 | ||
| 233 | template <typename T> | 233 | template <typename T> |
| @@ -239,10 +239,10 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void CalcOf | |||
| 239 | for (int32_t i = 0; i < blockIdx; i++) { | 239 | for (int32_t i = 0; i < blockIdx; i++) { |
| 240 | offset += coreSyncWorkspaceGm[tableIndx * maxCoreNum + i]; | 240 | offset += coreSyncWorkspaceGm[tableIndx * maxCoreNum + i]; |
| 241 | } | 241 | } |
| 242 | - for (int32_t i = 0; i < Simt::GetThreadIdx(); i++) { | 242 | + for (int32_t i = 0; i < threadIdx.x; i++) { |
| 243 | offset += threadCountKeysToExportUB[i]; | 243 | offset += threadCountKeysToExportUB[i]; |
| 244 | } | 244 | } |
| 245 | - threadCountKeysToExportSumUB[Simt::GetThreadIdx()] = offset; | 245 | + threadCountKeysToExportSumUB[threadIdx.x] = offset; |
| 246 | } | 246 | } |
| 247 | 247 | ||
| 248 | template <typename T> | 248 | template <typename T> |
| @@ -258,11 +258,11 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void Export | |||
| 258 | return; | 258 | return; |
| 259 | } | 259 | } |
| 260 | 260 | ||
| 261 | - if (Simt::GetThreadIdx() >= usedThreadNum) { | 261 | + if (threadIdx.x >= usedThreadNum) { |
| 262 | return; | 262 | return; |
| 263 | } | 263 | } |
| 264 | 264 | ||
| 265 | - int64_t offset = threadCountKeysToExportSumUB[Simt::GetThreadIdx()]; | 265 | + int64_t offset = threadCountKeysToExportSumUB[threadIdx.x]; |
| 266 | 266 | ||
| 267 | __gm__ int64_t* tableAddrI64 = reinterpret_cast<__gm__ int64_t*>(tableAddr); | 267 | __gm__ int64_t* tableAddrI64 = reinterpret_cast<__gm__ int64_t*>(tableAddr); |
| 268 | __gm__ uint64_t* tableAddrU64 = reinterpret_cast<__gm__ uint64_t*>(tableAddr); | 268 | __gm__ uint64_t* tableAddrU64 = reinterpret_cast<__gm__ uint64_t*>(tableAddr); |
| @@ -272,22 +272,18 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void Export | |||
| 272 | int64_t curThreadRefreshExportFlagNum = 0; | 272 | int64_t curThreadRefreshExportFlagNum = 0; |
| 273 | int64_t positionIndex = 0; | 273 | int64_t positionIndex = 0; |
| 274 | 274 | ||
| 275 | - int64_t curThreadProcessKeys = | 275 | + int64_t curThreadProcessKeys = threadIdx.x < (usedThreadNum - 1) ? normalThreadProcessKeys : tailThreadProcessKeys; |
| 276 | - Simt::GetThreadIdx() < (usedThreadNum - 1) ? normalThreadProcessKeys : tailThreadProcessKeys; | ||
| 277 | for (int64_t i = 0; i < curThreadProcessKeys; i++) { | 276 | for (int64_t i = 0; i < curThreadProcessKeys; i++) { |
| 278 | uint8_t flag = tableAddrU8 | 277 | uint8_t flag = tableAddrU8 |
| 279 | - [keyWidthByte * (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i) + | 278 | + [keyWidthByte * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i) + |
| 280 | KEY_FLAG_OFFSET_OF_BYTE]; | 279 | KEY_FLAG_OFFSET_OF_BYTE]; |
| 281 | if ((flag & VALID_FLAG_MASK) && !(flag & EVICTED_FLAG_MASK) && | 280 | if ((flag & VALID_FLAG_MASK) && !(flag & EVICTED_FLAG_MASK) && |
| 282 | (exportMode != 1 || !(flag & EXPORT_FLAG_MASK))) { | 281 | (exportMode != 1 || !(flag & EXPORT_FLAG_MASK))) { |
| 283 | int64_t key = tableAddrI64 | 282 | int64_t key = tableAddrI64 |
| 284 | - [keyWidthByteD8 * | 283 | + [keyWidthByteD8 * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i)]; |
| 285 | - (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i)]; | ||
| 286 | outKeyGm[offset + positionIndex] = key; | 284 | outKeyGm[offset + positionIndex] = key; |
| 287 | uint64_t counter = tableAddrU64 | 285 | uint64_t counter = tableAddrU64 |
| 288 | - [keyWidthByteD8 * | 286 | + [keyWidthByteD8 * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i) + 1]; |
| 289 | - (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i) + | ||
| 290 | - 1]; | ||
| 291 | outCounterGm[offset + positionIndex] = counter; | 287 | outCounterGm[offset + positionIndex] = counter; |
| 292 | 288 | ||
| 293 | if (FILTER_FLAG_MASK & flag) { | 289 | if (FILTER_FLAG_MASK & flag) { |
| @@ -297,22 +293,20 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_THREAD_LAUNCH_BOUND) inline void Export | |||
| 297 | } | 293 | } |
| 298 | for (int64_t j = 0; j < embeddingDims; j++) { | 294 | for (int64_t j = 0; j < embeddingDims; j++) { |
| 299 | outValueGm[(offset + positionIndex) * embeddingDims + j] = tableAddrT | 295 | outValueGm[(offset + positionIndex) * embeddingDims + j] = tableAddrT |
| 300 | - [keyWidthByteDT * | 296 | + [keyWidthByteDT * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i) + |
| 301 | - (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i) + | ||
| 302 | KEY_VALUE_OFFSET_OF_BYTE / sizeof(T) + j]; | 297 | KEY_VALUE_OFFSET_OF_BYTE / sizeof(T) + j]; |
| 303 | } | 298 | } |
| 304 | // 刷新导出flag, 只在第一次导出时刷新 | 299 | // 刷新导出flag, 只在第一次导出时刷新 |
| 305 | if (!(flag & EXPORT_FLAG_MASK)) { | 300 | if (!(flag & EXPORT_FLAG_MASK)) { |
| 306 | tableAddrU8 | 301 | tableAddrU8 |
| 307 | - [keyWidthByte * | 302 | + [keyWidthByte * (blockIdx * normalCoreProcessKeys + threadIdx.x * normalThreadProcessKeys + i) + |
| 308 | - (blockIdx * normalCoreProcessKeys + Simt::GetThreadIdx() * normalThreadProcessKeys + i) + | ||
| 309 | KEY_FLAG_OFFSET_OF_BYTE] |= EXPORT_FLAG_MASK; | 303 | KEY_FLAG_OFFSET_OF_BYTE] |= EXPORT_FLAG_MASK; |
| 310 | curThreadRefreshExportFlagNum++; | 304 | curThreadRefreshExportFlagNum++; |
| 311 | } | 305 | } |
| 312 | positionIndex++; | 306 | positionIndex++; |
| 313 | } | 307 | } |
| 314 | } | 308 | } |
| 315 | - threadCountReFreshExportFlagUB[Simt::GetThreadIdx()] = curThreadRefreshExportFlagNum; | 309 | + threadCountReFreshExportFlagUB[threadIdx.x] = curThreadRefreshExportFlagNum; |
| 316 | } | 310 | } |
| 317 | 311 | ||
| 318 | template <typename T> | 312 | template <typename T> |
| @@ -322,24 +316,23 @@ __aicore__ inline void EmbeddingHashTableExport<T>::Process() | |||
| 322 | Duplicate(threadCountKeysToExportUB_, int64_t(0), maxThreadNum_ * BUFFER_LENGTH); | 316 | Duplicate(threadCountKeysToExportUB_, int64_t(0), maxThreadNum_ * BUFFER_LENGTH); |
| 323 | Duplicate(threadCountReFreshExportFlagUB_, int64_t(0), maxThreadNum_ * BUFFER_LENGTH); | 317 | Duplicate(threadCountReFreshExportFlagUB_, int64_t(0), maxThreadNum_ * BUFFER_LENGTH); |
| 324 | SingleTableCompute(tableIndx); | 318 | SingleTableCompute(tableIndx); |
| 325 | - Simt::VF_CALL<CountPerThread<T>>( | 319 | + asc_vf_call<CountPerThread<T>>( |
| 326 | - Simt::Dim3{static_cast<uint32_t>(maxThreadNum_)}, maxCoreNum_, maxThreadNum_, blockIdx_, | 320 | + dim3{static_cast<uint32_t>(maxThreadNum_)}, maxCoreNum_, maxThreadNum_, blockIdx_, usedCoreNum_, |
| 327 | - usedCoreNum_, usedThreadNum_, normalThreadProcessKeys_, tailThreadProcessKeys_, tableAddr_, keyWidthByte_, | 321 | + usedThreadNum_, normalThreadProcessKeys_, tailThreadProcessKeys_, tableAddr_, keyWidthByte_, |
| 328 | normalCoreProcessKeys_, exportMode_, (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr()); | 322 | normalCoreProcessKeys_, exportMode_, (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr()); |
| 329 | ReduceSum<int64_t>( | 323 | ReduceSum<int64_t>( |
| 330 | threadCountKeysToExportUB_[maxThreadNum_], threadCountKeysToExportUB_, threadCountReFreshExportFlagUB_, | 324 | threadCountKeysToExportUB_[maxThreadNum_], threadCountKeysToExportUB_, threadCountReFreshExportFlagUB_, |
| 331 | usedThreadNum_); | 325 | usedThreadNum_); |
| 332 | - Simt::VF_CALL<SaveToCoreSyncWorkspace<T>>( | 326 | + asc_vf_call<SaveToCoreSyncWorkspace<T>>( |
| 333 | - Simt::Dim3{static_cast<uint32_t>(1)}, maxCoreNum_, maxThreadNum_, tableIndx, blockIdx_, | 327 | + dim3{static_cast<uint32_t>(1)}, maxCoreNum_, maxThreadNum_, tableIndx, blockIdx_, usedCoreNum_, |
| 334 | - usedCoreNum_, coreSyncWorkspaceGm_.GetPhyAddr(0), | 328 | + coreSyncWorkspaceGm_.GetPhyAddr(0), (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr()); |
| 335 | - (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr()); | ||
| 336 | SyncAll(); | 329 | SyncAll(); |
| 337 | - Simt::VF_CALL<CalcOffset<T>>( | 330 | + asc_vf_call<CalcOffset<T>>( |
| 338 | - Simt::Dim3{static_cast<uint32_t>(maxThreadNum_)}, maxCoreNum_, maxThreadNum_, tableIndx, | 331 | + dim3{static_cast<uint32_t>(maxThreadNum_)}, maxCoreNum_, maxThreadNum_, tableIndx, blockIdx_, |
| 339 | - blockIdx_, coreSyncWorkspaceGm_.GetPhyAddr(0), (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr(), | 332 | + coreSyncWorkspaceGm_.GetPhyAddr(0), (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr(), |
| 340 | (__ubuf__ int64_t*)threadCountKeysToExportSumUB_.GetPhyAddr()); | 333 | (__ubuf__ int64_t*)threadCountKeysToExportSumUB_.GetPhyAddr()); |
| 341 | - Simt::VF_CALL<ExportPerThread<T>>( | 334 | + asc_vf_call<ExportPerThread<T>>( |
| 342 | - Simt::Dim3{static_cast<uint32_t>(maxThreadNum_)}, blockIdx_, usedCoreNum_, usedThreadNum_, | 335 | + dim3{static_cast<uint32_t>(maxThreadNum_)}, blockIdx_, usedCoreNum_, usedThreadNum_, |
| 343 | normalThreadProcessKeys_, tailThreadProcessKeys_, tableAddr_, keyWidthByte_, normalCoreProcessKeys_, | 336 | normalThreadProcessKeys_, tailThreadProcessKeys_, tableAddr_, keyWidthByte_, normalCoreProcessKeys_, |
| 344 | exportMode_, keyWidthByteD8_, keyWidthByteDT_, embeddingDims_, coreSyncWorkspaceGm_.GetPhyAddr(0), | 337 | exportMode_, keyWidthByteD8_, keyWidthByteDT_, embeddingDims_, coreSyncWorkspaceGm_.GetPhyAddr(0), |
| 345 | (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr(), outKeyGm_.GetPhyAddr(0), | 338 | (__ubuf__ int64_t*)threadCountKeysToExportUB_.GetPhyAddr(), outKeyGm_.GetPhyAddr(0), |
| @@ -349,8 +342,8 @@ __aicore__ inline void EmbeddingHashTableExport<T>::Process() | |||
| 349 | ReduceSum<int64_t>( | 342 | ReduceSum<int64_t>( |
| 350 | threadCountReFreshExportFlagUB_[maxThreadNum_], threadCountReFreshExportFlagUB_, threadCountKeysToExportUB_, | 343 | threadCountReFreshExportFlagUB_[maxThreadNum_], threadCountReFreshExportFlagUB_, threadCountKeysToExportUB_, |
| 351 | usedThreadNum_); | 344 | usedThreadNum_); |
| 352 | - Simt::VF_CALL<AtomicSubToGm<T>>( | 345 | + asc_vf_call<AtomicSubToGm<T>>( |
| 353 | - Simt::Dim3{static_cast<uint32_t>(1)}, maxCoreNum_, maxThreadNum_, blockIdx_, usedCoreNum_, | 346 | + dim3{static_cast<uint32_t>(1)}, maxCoreNum_, maxThreadNum_, blockIdx_, usedCoreNum_, |
| 354 | tableHandleStructGm_.GetPhyAddr(0), (__ubuf__ int64_t*)threadCountReFreshExportFlagUB_.GetPhyAddr()); | 347 | tableHandleStructGm_.GetPhyAddr(0), (__ubuf__ int64_t*)threadCountReFreshExportFlagUB_.GetPhyAddr()); |
| 355 | SyncAll(); | 348 | SyncAll(); |
| 356 | } | 349 | } |
| @@ -18,6 +18,8 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | + | ||
| 22 | + | ||
| 21 | 23 | ||
| 22 | namespace EmbeddingHashTable { | 24 | namespace EmbeddingHashTable { |
| 23 | using namespace AscendC; | 25 | using namespace AscendC; |
| @@ -155,8 +157,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM_LAUNCH_BOUND) inline void SingleT | |||
| 155 | { | 157 | { |
| 156 | __gm__ int64_t* tableHandle = | 158 | __gm__ int64_t* tableHandle = |
| 157 | reinterpret_cast<__gm__ int64_t*>(reinterpret_cast<__gm__ uint8_t*>(tableHandlesGm[tableIdx])); | 159 | reinterpret_cast<__gm__ int64_t*>(reinterpret_cast<__gm__ uint8_t*>(tableHandlesGm[tableIdx])); |
| 158 | - for (int64_t i = blockIdx * Simt::GetThreadNum() + Simt::GetThreadIdx(); i < keyNum; | 160 | + for (int64_t i = blockIdx * blockDim.x + threadIdx.x; i < keyNum; i = i + blockNum * blockDim.x) { |
| 159 | - i = i + blockNum * Simt::GetThreadNum()) { | ||
| 160 | int64_t insertKey = keyGm[i]; | 161 | int64_t insertKey = keyGm[i]; |
| 161 | uint32_t hashValue = Hashtbl::MurmurHash3(keyGm + i, INT64_TYPE_BYTES, 0); | 162 | uint32_t hashValue = Hashtbl::MurmurHash3(keyGm + i, INT64_TYPE_BYTES, 0); |
| 162 | int64_t hashTabIdx = hashValue % bucketSize; | 163 | int64_t hashTabIdx = hashValue % bucketSize; |
| @@ -170,14 +171,14 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM_LAUNCH_BOUND) inline void SingleT | |||
| 170 | break; | 171 | break; |
| 171 | } | 172 | } |
| 172 | // 插入key值序列 | 173 | // 插入key值序列 |
| 173 | - const int32_t originalFlag = Simt::AtomicCas( | 174 | + const int32_t originalFlag = asc_atomic_cas( |
| 174 | reinterpret_cast<__gm__ int32_t*>(tableGm + blockOffset + TABLE_FLAG_OFFSET), static_cast<int32_t>(0), | 175 | reinterpret_cast<__gm__ int32_t*>(tableGm + blockOffset + TABLE_FLAG_OFFSET), static_cast<int32_t>(0), |
| 175 | BIG_ENDIAN_ONE); | 176 | BIG_ENDIAN_ONE); |
| 176 | 177 | ||
| 177 | int64_t keyOffset = blockOffset + (KEY_OFFSET * INT64_TYPE_BYTES); | 178 | int64_t keyOffset = blockOffset + (KEY_OFFSET * INT64_TYPE_BYTES); |
| 178 | if (0 == originalFlag) { | 179 | if (0 == originalFlag) { |
| 179 | *reinterpret_cast<__gm__ int64_t*>(tableGm + keyOffset) = insertKey; | 180 | *reinterpret_cast<__gm__ int64_t*>(tableGm + keyOffset) = insertKey; |
| 180 | - Simt::ThreadFence(); | 181 | + __threadfence(); |
| 181 | *reinterpret_cast<__gm__ int32_t*>(tableGm + keyOffset + TABLE_STATE_OFFSET) = 1; | 182 | *reinterpret_cast<__gm__ int32_t*>(tableGm + keyOffset + TABLE_STATE_OFFSET) = 1; |
| 182 | isInsertSucc = true; | 183 | isInsertSucc = true; |
| 183 | isNewKey = true; | 184 | isNewKey = true; |
| @@ -203,13 +204,13 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM_LAUNCH_BOUND) inline void SingleT | |||
| 203 | // 更新tableHandle struct | 204 | // 更新tableHandle struct |
| 204 | if (isNewKey) { | 205 | if (isNewKey) { |
| 205 | // 刷新total_hash_addr地址 --> (key不存在, tablesize++, noexportsize不变) | 206 | // 刷新total_hash_addr地址 --> (key不存在, tablesize++, noexportsize不变) |
| 206 | - AscendC::Simt::AtomicAdd(tableHandle + HANDLE_SIZE_ALL_OFFSET, INT64_ONE); | 207 | + asc_atomic_add(tableHandle + HANDLE_SIZE_ALL_OFFSET, INT64_ONE); |
| 207 | } else { | 208 | } else { |
| 208 | // 刷新no_export_hash_addr地址 --> (key存在, tablesize不变, noexportsize--) | 209 | // 刷新no_export_hash_addr地址 --> (key存在, tablesize不变, noexportsize--) |
| 209 | int64_t flagOffset = blockOffset + ((FLAG_OFFSET + 1) * INT64_TYPE_BYTES - 1); | 210 | int64_t flagOffset = blockOffset + ((FLAG_OFFSET + 1) * INT64_TYPE_BYTES - 1); |
| 210 | __gm__ uint8_t* filterFlagValue = reinterpret_cast<__gm__ uint8_t*>(tableGm + flagOffset); | 211 | __gm__ uint8_t* filterFlagValue = reinterpret_cast<__gm__ uint8_t*>(tableGm + flagOffset); |
| 211 | if (!(*filterFlagValue & EXPORT_FLAG_MASK)) { // means change flag from 0 to 1(1 means cannot be exported) | 212 | if (!(*filterFlagValue & EXPORT_FLAG_MASK)) { // means change flag from 0 to 1(1 means cannot be exported) |
| 212 | - AscendC::Simt::AtomicSub(tableHandle + HANDLE_SIZE_ALL_NOEXPORT_OFFSET, INT64_ONE); | 213 | + asc_atomic_sub(tableHandle + HANDLE_SIZE_ALL_NOEXPORT_OFFSET, INT64_ONE); |
| 213 | } | 214 | } |
| 214 | } | 215 | } |
| 215 | // 插入counter值 | 216 | // 插入counter值 |
| @@ -262,9 +263,9 @@ __aicore__ inline void EmbeddingHashTableImport<T>::Process() | |||
| 262 | filterFlagGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(filterFlagsListGm_.GetDataPtr<uint8_t>(idx))); | 263 | filterFlagGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(filterFlagsListGm_.GetDataPtr<uint8_t>(idx))); |
| 263 | valueGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(valuesListGm_.GetDataPtr<T>(idx))); | 264 | valueGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(valuesListGm_.GetDataPtr<T>(idx))); |
| 264 | 265 | ||
| 265 | - Simt::VF_CALL<SingleTableImportCompute<T>>( | 266 | + asc_vf_call<SingleTableImportCompute<T>>( |
| 266 | - Simt::Dim3{static_cast<uint32_t>(THREAD_NUM)}, idx, keyNum, embeddingDim_, blockSize_, bucketSize_, | 267 | + dim3{static_cast<uint32_t>(THREAD_NUM)}, idx, keyNum, embeddingDim_, blockSize_, bucketSize_, bitWidth_, |
| 267 | - bitWidth_, unusedKey_, blockIdx_, blockNum_, tableHandlesGm_.GetPhyAddr(0), keyGm_.GetPhyAddr(0), | 268 | + unusedKey_, blockIdx_, blockNum_, tableHandlesGm_.GetPhyAddr(0), keyGm_.GetPhyAddr(0), |
| 268 | counterGm_.GetPhyAddr(0), filterFlagGm_.GetPhyAddr(0), valueGm_.GetPhyAddr(0), tableGm_.GetPhyAddr(0)); | 269 | counterGm_.GetPhyAddr(0), filterFlagGm_.GetPhyAddr(0), valueGm_.GetPhyAddr(0), tableGm_.GetPhyAddr(0)); |
| 269 | } | 270 | } |
| 270 | } | 271 | } |
Mhash/embedding_hash_table_lookup_or_insert/op_kernel/arch35/kernel_lookup_or_insert_general.h+16-16
| @@ -27,10 +27,10 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsert( | |||
| 27 | __gm__ uint8_t* pTable, __gm__ int64_t* pKeys, __gm__ float* pValues, __ubuf__ int64_t* pThreadInsertCounts) | 27 | __gm__ uint8_t* pTable, __gm__ int64_t* pKeys, __gm__ float* pValues, __ubuf__ int64_t* pThreadInsertCounts) |
| 28 | { | 28 | { |
| 29 | // 每core线程划分为(x,y),每threadXNum个x对应1个y,共启动threadXNum*threadYNum个线程 | 29 | // 每core线程划分为(x,y),每threadXNum个x对应1个y,共启动threadXNum*threadYNum个线程 |
| 30 | - uint32_t threadXIdx = static_cast<uint32_t>(Simt::GetThreadIdx<0>()); | 30 | + uint32_t threadXIdx = static_cast<uint32_t>(threadIdx.x); |
| 31 | - uint32_t threadYIdx = static_cast<uint32_t>(Simt::GetThreadIdx<1>()); | 31 | + uint32_t threadYIdx = static_cast<uint32_t>(threadIdx.y); |
| 32 | - uint32_t threadXNum = static_cast<uint32_t>(Simt::GetThreadNum<0>()); | 32 | + uint32_t threadXNum = static_cast<uint32_t>(blockDim.x); |
| 33 | - uint32_t threadYNum = static_cast<uint32_t>(Simt::GetThreadNum<1>()); | 33 | + uint32_t threadYNum = static_cast<uint32_t>(blockDim.y); |
| 34 | 34 | ||
| 35 | int64_t insertCounts = 0; // 各线程自有变量,记录insert的次数 | 35 | int64_t insertCounts = 0; // 各线程自有变量,记录insert的次数 |
| 36 | for (uint32_t i = threadYIdx + blockIdx * threadYNum; i < keyNum; i += blockNum * threadYNum) { | 36 | for (uint32_t i = threadYIdx + blockIdx * threadYNum; i < keyNum; i += blockNum * threadYNum) { |
| @@ -63,14 +63,14 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsert( | |||
| 63 | detectCounts++; | 63 | detectCounts++; |
| 64 | 64 | ||
| 65 | // 由于AtmoicCas限制,用int32来cas第20~23字节的BIG_ENDIAN_ONE那个位置 | 65 | // 由于AtmoicCas限制,用int32来cas第20~23字节的BIG_ENDIAN_ONE那个位置 |
| 66 | - const int32_t casOrigFlag = AscendC::Simt::AtomicCas( | 66 | + const int32_t casOrigFlag = asc_atomic_cas( |
| 67 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), static_cast<int32_t>(0), | 67 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), static_cast<int32_t>(0), |
| 68 | BIG_ENDIAN_ONE); | 68 | BIG_ENDIAN_ONE); |
| 69 | 69 | ||
| 70 | if (casOrigFlag == 0) { | 70 | if (casOrigFlag == 0) { |
| 71 | // 可以插入 | 71 | // 可以插入 |
| 72 | *reinterpret_cast<__gm__ int64_t*>(pCurrBucket) = insertKey; | 72 | *reinterpret_cast<__gm__ int64_t*>(pCurrBucket) = insertKey; |
| 73 | - Simt::ThreadFence(); | 73 | + __threadfence(); |
| 74 | *reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_STATE_OFFSET) = 1; | 74 | *reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_STATE_OFFSET) = 1; |
| 75 | succ = true; | 75 | succ = true; |
| 76 | insertCounts++; | 76 | insertCounts++; |
| @@ -89,7 +89,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsert( | |||
| 89 | *reinterpret_cast<__gm__ volatile int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32); | 89 | *reinterpret_cast<__gm__ volatile int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32); |
| 90 | if ((currFlag & EVICTED_FLAG_MASK) != 0) { | 90 | if ((currFlag & EVICTED_FLAG_MASK) != 0) { |
| 91 | auto newFlag = currFlag ^ EVICTED_FLAG_MASK; | 91 | auto newFlag = currFlag ^ EVICTED_FLAG_MASK; |
| 92 | - auto oldFlag = Simt::AtomicCas( | 92 | + auto oldFlag = asc_atomic_cas( |
| 93 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), | 93 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), |
| 94 | static_cast<int32_t>(currFlag), newFlag); | 94 | static_cast<int32_t>(currFlag), newFlag); |
| 95 | if ((oldFlag & EVICTED_FLAG_MASK) != 0) { | 95 | if ((oldFlag & EVICTED_FLAG_MASK) != 0) { |
| @@ -112,7 +112,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsert( | |||
| 112 | pCurrBucket = pTable + currIdx * bucketSize; | 112 | pCurrBucket = pTable + currIdx * bucketSize; |
| 113 | if (threadXIdx == 0) { | 113 | if (threadXIdx == 0) { |
| 114 | // 由控制线程来执行bucket的counter++操作 | 114 | // 由控制线程来执行bucket的counter++操作 |
| 115 | - Simt::AtomicAdd( | 115 | + asc_atomic_add( |
| 116 | reinterpret_cast<__gm__ int64_t*>(pCurrBucket + COUNTER_OFFSET), static_cast<int64_t>(1)); | 116 | reinterpret_cast<__gm__ int64_t*>(pCurrBucket + COUNTER_OFFSET), static_cast<int64_t>(1)); |
| 117 | } | 117 | } |
| 118 | for (size_t j = threadXIdx; j < embeddingDim; j += threadXNum) { | 118 | for (size_t j = threadXIdx; j < embeddingDim; j += threadXNum) { |
| @@ -141,15 +141,15 @@ public: | |||
| 141 | reinterpret_cast<__ubuf__ int64_t*>(threadInsertCountsLocal.GetPhyAddr()); | 141 | reinterpret_cast<__ubuf__ int64_t*>(threadInsertCountsLocal.GetPhyAddr()); |
| 142 | 142 | ||
| 143 | if (filterKeyFlag_) { | 143 | if (filterKeyFlag_) { |
| 144 | - Simt::VF_CALL<ComputeLookupOrInsert<true>>( | 144 | + asc_vf_call<ComputeLookupOrInsert<true>>( |
| 145 | - Simt::Dim3{threadXNum_, threadYNum_}, blockIdx_, blockNum_, bucketSize_, tableSize_, embeddingDim_, | 145 | + dim3{threadXNum_, threadYNum_}, blockIdx_, blockNum_, bucketSize_, tableSize_, embeddingDim_, keyNum_, |
| 146 | - keyNum_, defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, | 146 | + defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, |
| 147 | - pValues_, pThreadInsertCounts); | 147 | + pThreadInsertCounts); |
| 148 | } else { | 148 | } else { |
| 149 | - Simt::VF_CALL<ComputeLookupOrInsert<false>>( | 149 | + asc_vf_call<ComputeLookupOrInsert<false>>( |
| 150 | - Simt::Dim3{threadXNum_, threadYNum_}, blockIdx_, blockNum_, bucketSize_, tableSize_, embeddingDim_, | 150 | + dim3{threadXNum_, threadYNum_}, blockIdx_, blockNum_, bucketSize_, tableSize_, embeddingDim_, keyNum_, |
| 151 | - keyNum_, defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, | 151 | + defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, |
| 152 | - pValues_, pThreadInsertCounts); | 152 | + pThreadInsertCounts); |
| 153 | } | 153 | } |
| 154 | 154 | ||
| 155 | // SIMD汇总写回tableHandle的那几个统计字段的值 | 155 | // SIMD汇总写回tableHandle的那几个统计字段的值 |
Mhash/embedding_hash_table_lookup_or_insert/op_kernel/arch35/kernel_lookup_or_insert_opt_dim.h+11-11
| @@ -34,9 +34,9 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsertOptDim | |||
| 34 | __gm__ uint8_t* pTable, __gm__ int64_t* pKeys, __gm__ float* pValues, __ubuf__ int64_t* pThreadInsertCounts) | 34 | __gm__ uint8_t* pTable, __gm__ int64_t* pKeys, __gm__ float* pValues, __ubuf__ int64_t* pThreadInsertCounts) |
| 35 | { | 35 | { |
| 36 | // 每core线程划分为(x,y),每threadXNum个x对应1个y,共启动threadXNum*threadYNum个线程 | 36 | // 每core线程划分为(x,y),每threadXNum个x对应1个y,共启动threadXNum*threadYNum个线程 |
| 37 | - uint32_t threadXIdx = static_cast<uint32_t>(Simt::GetThreadIdx<0>()); | 37 | + uint32_t threadXIdx = static_cast<uint32_t>(threadIdx.x); |
| 38 | - uint32_t threadYIdx = static_cast<uint32_t>(Simt::GetThreadIdx<1>()); | 38 | + uint32_t threadYIdx = static_cast<uint32_t>(threadIdx.y); |
| 39 | - uint32_t threadYNum = static_cast<uint32_t>(Simt::GetThreadNum<1>()); | 39 | + uint32_t threadYNum = static_cast<uint32_t>(blockDim.y); |
| 40 | 40 | ||
| 41 | int64_t insertCounts = 0; // 各线程自有变量,记录insert的次数 | 41 | int64_t insertCounts = 0; // 各线程自有变量,记录insert的次数 |
| 42 | for (uint32_t i = threadYIdx + blockIdx * threadYNum; i < keyNum; i += blockNum * threadYNum) { | 42 | for (uint32_t i = threadYIdx + blockIdx * threadYNum; i < keyNum; i += blockNum * threadYNum) { |
| @@ -67,14 +67,14 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsertOptDim | |||
| 67 | detectCounts++; | 67 | detectCounts++; |
| 68 | 68 | ||
| 69 | // 由于AtmoicCas限制,用int32来cas第20~23字节的BIG_ENDIAN_ONE那个位置 | 69 | // 由于AtmoicCas限制,用int32来cas第20~23字节的BIG_ENDIAN_ONE那个位置 |
| 70 | - const int32_t casOrigFlag = AscendC::Simt::AtomicCas( | 70 | + const int32_t casOrigFlag = asc_atomic_cas( |
| 71 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), static_cast<int32_t>(0), | 71 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), static_cast<int32_t>(0), |
| 72 | BIG_ENDIAN_ONE); | 72 | BIG_ENDIAN_ONE); |
| 73 | 73 | ||
| 74 | if (casOrigFlag == 0) { | 74 | if (casOrigFlag == 0) { |
| 75 | // 可以插入 | 75 | // 可以插入 |
| 76 | *reinterpret_cast<__gm__ int64_t*>(pCurrBucket) = insertKey; | 76 | *reinterpret_cast<__gm__ int64_t*>(pCurrBucket) = insertKey; |
| 77 | - Simt::ThreadFence(); | 77 | + __threadfence(); |
| 78 | *reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_STATE_OFFSET) = 1; | 78 | *reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_STATE_OFFSET) = 1; |
| 79 | succ = true; | 79 | succ = true; |
| 80 | insertCounts++; | 80 | insertCounts++; |
| @@ -93,7 +93,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsertOptDim | |||
| 93 | *reinterpret_cast<__gm__ volatile int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32); | 93 | *reinterpret_cast<__gm__ volatile int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32); |
| 94 | if ((currFlag & EVICTED_FLAG_MASK) != 0) { | 94 | if ((currFlag & EVICTED_FLAG_MASK) != 0) { |
| 95 | auto newFlag = currFlag ^ EVICTED_FLAG_MASK; | 95 | auto newFlag = currFlag ^ EVICTED_FLAG_MASK; |
| 96 | - auto oldFlag = Simt::AtomicCas( | 96 | + auto oldFlag = asc_atomic_cas( |
| 97 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), | 97 | reinterpret_cast<__gm__ int32_t*>(pCurrBucket + TABLE_FLAG_OFFSET_FOR_B32), |
| 98 | static_cast<int32_t>(currFlag), newFlag); | 98 | static_cast<int32_t>(currFlag), newFlag); |
| 99 | if ((oldFlag & EVICTED_FLAG_MASK) != 0) { | 99 | if ((oldFlag & EVICTED_FLAG_MASK) != 0) { |
| @@ -115,7 +115,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsertOptDim | |||
| 115 | pCurrBucket = pTable + currIdx * bucketSize; | 115 | pCurrBucket = pTable + currIdx * bucketSize; |
| 116 | if (threadXIdx == 0) { | 116 | if (threadXIdx == 0) { |
| 117 | // 由控制线程来执行bucket的counter++操作 | 117 | // 由控制线程来执行bucket的counter++操作 |
| 118 | - Simt::AtomicAdd( | 118 | + asc_atomic_add( |
| 119 | reinterpret_cast<__gm__ int64_t*>(pCurrBucket + COUNTER_OFFSET), static_cast<int64_t>(1)); | 119 | reinterpret_cast<__gm__ int64_t*>(pCurrBucket + COUNTER_OFFSET), static_cast<int64_t>(1)); |
| 120 | } | 120 | } |
| 121 | __gm__ float* pCurrValue = | 121 | __gm__ float* pCurrValue = |
| @@ -132,13 +132,13 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) void ComputeLookupOrInsertOptDim | |||
| 132 | 132 | ||
| 133 | 133 | ||
| 134 | if (macro_f == 0) { \ | 134 | if (macro_f == 0) { \ |
| 135 | - Simt::VF_CALL<ComputeLookupOrInsertOptDim<macro_d, false>>( \ | 135 | + asc_vf_call<ComputeLookupOrInsertOptDim<macro_d, false>>( \ |
| 136 | - Simt::Dim3{macro_d, THREAD_NUM / macro_d}, blockIdx_, blockNum_, bucketSize_, tableSize_, keyNum_, \ | 136 | + dim3{macro_d, THREAD_NUM / macro_d}, blockIdx_, blockNum_, bucketSize_, tableSize_, keyNum_, \ |
| 137 | defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, \ | 137 | defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, \ |
| 138 | macro_pcounts); \ | 138 | macro_pcounts); \ |
| 139 | } else { \ | 139 | } else { \ |
| 140 | - Simt::VF_CALL<ComputeLookupOrInsertOptDim<macro_d, true>>( \ | 140 | + asc_vf_call<ComputeLookupOrInsertOptDim<macro_d, true>>( \ |
| 141 | - Simt::Dim3{macro_d, THREAD_NUM / macro_d}, blockIdx_, blockNum_, bucketSize_, tableSize_, keyNum_, \ | 141 | + dim3{macro_d, THREAD_NUM / macro_d}, blockIdx_, blockNum_, bucketSize_, tableSize_, keyNum_, \ |
| 142 | defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, \ | 142 | defaultKeyOrValue_, defaultKey_, defaultValue_, filterKey_, pTableHandle_, pTable_, pKeys_, pValues_, \ |
| 143 | macro_pcounts); \ | 143 | macro_pcounts); \ |
| 144 | } | 144 | } |
| @@ -19,6 +19,8 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | + | ||
| 23 | + | ||
| 22 | 24 | ||
| 23 | namespace Hashtbl { | 25 | namespace Hashtbl { |
| 24 | using namespace AscendC; | 26 | using namespace AscendC; |
| @@ -17,6 +17,7 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 20 | 21 | ||
| 21 | namespace InitEmbeddingHashTable { | 22 | namespace InitEmbeddingHashTable { |
| 22 | using namespace AscendC; | 23 | using namespace AscendC; |
| @@ -43,8 +44,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(THREAD_NUM) inline void InitCompute( | |||
| 43 | uint32_t blockIdx, uint32_t blockNum, __gm__ int64_t* tableHanldeGm, __gm__ Tvalue* sampledValuesGm, | 44 | uint32_t blockIdx, uint32_t blockNum, __gm__ int64_t* tableHanldeGm, __gm__ Tvalue* sampledValuesGm, |
| 44 | __gm__ uint8_t* outputGm) | 45 | __gm__ uint8_t* outputGm) |
| 45 | { | 46 | { |
| 46 | - for (int64_t i = blockIdx * Simt::GetThreadNum() + Simt::GetThreadIdx(); i < bucketSize; | 47 | + for (int64_t i = blockIdx * blockDim.x + threadIdx.x; i < bucketSize; i += blockNum * blockDim.x) { |
| 47 | - i += blockNum * Simt::GetThreadNum()) { | ||
| 48 | // SetKey(-1) | 48 | // SetKey(-1) |
| 49 | int64_t keyOffset = (bucketLength * i + KEY_OFFSET) * INT64_PER_BYTE; | 49 | int64_t keyOffset = (bucketLength * i + KEY_OFFSET) * INT64_PER_BYTE; |
| 50 | __gm__ int64_t* key = reinterpret_cast<__gm__ int64_t*>(outputGm + keyOffset); | 50 | __gm__ int64_t* key = reinterpret_cast<__gm__ int64_t*>(outputGm + keyOffset); |
| @@ -99,8 +99,8 @@ public: | |||
| 99 | } | 99 | } |
| 100 | __aicore__ inline void Process() | 100 | __aicore__ inline void Process() |
| 101 | { | 101 | { |
| 102 | - Simt::VF_CALL<InitCompute<Tkey, Tvalue>>( | 102 | + asc_vf_call<InitCompute<Tkey, Tvalue>>( |
| 103 | - Simt::Dim3{static_cast<uint32_t>(useThreadNum)}, embeddingDim, bucketSize, bucketLength, initializerMode, | 103 | + dim3{static_cast<uint32_t>(useThreadNum)}, embeddingDim, bucketSize, bucketLength, initializerMode, |
| 104 | constantValue, blockIdx, blockNum, tableHanldeGm.GetPhyAddr(0), sampledValuesGm.GetPhyAddr(0), | 104 | constantValue, blockIdx, blockNum, tableHanldeGm.GetPhyAddr(0), sampledValuesGm.GetPhyAddr(0), |
| 105 | outputGm.GetPhyAddr(0)); | 105 | outputGm.GetPhyAddr(0)); |
| 106 | } | 106 | } |
| @@ -110,8 +110,8 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_MAX_THREAD_NUM) inline void ComputeB16( | |||
| 110 | __gm__ T_SUM* workspaceGmAddr) | 110 | __gm__ T_SUM* workspaceGmAddr) |
| 111 | { | 111 | { |
| 112 | // 总共有usedCoreNum*ThreadNum个线程,每个线程所在的位置为currCoreIdx*ThreadNum+LocalThreadIdx | 112 | // 总共有usedCoreNum*ThreadNum个线程,每个线程所在的位置为currCoreIdx*ThreadNum+LocalThreadIdx |
| 113 | - for (int32_t elemIdx = currCoreIdx * Simt::GetThreadNum() + Simt::GetThreadIdx(); elemIdx < elemNum; | 113 | + for (int32_t elemIdx = currCoreIdx * blockDim.x + threadIdx.x; elemIdx < elemNum; |
| 114 | - elemIdx += usedCoreNum * Simt::GetThreadNum()) { | 114 | + elemIdx += usedCoreNum * blockDim.x) { |
| 115 | // 计算索引i、j、k,用于找到v1=x1(i,k)和v2=x2GmAddr(k,j) | 115 | // 计算索引i、j、k,用于找到v1=x1(i,k)和v2=x2GmAddr(k,j) |
| 116 | // i、j、k 统一转成int32类型 | 116 | // i、j、k 统一转成int32类型 |
| 117 | int32_t x1VecIdx = elemIdx / p; | 117 | int32_t x1VecIdx = elemIdx / p; |
| @@ -132,7 +132,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_MAX_THREAD_NUM) inline void ComputeB16( | |||
| 132 | } | 132 | } |
| 133 | // 累加到对应位置 | 133 | // 累加到对应位置 |
| 134 | __gm__ T_SUM* outAddr = workspaceGmAddr + i * p + j; | 134 | __gm__ T_SUM* outAddr = workspaceGmAddr + i * p + j; |
| 135 | - Simt::AtomicAdd(outAddr, v1 * v2); | 135 | + asc_atomic_add(outAddr, v1 * v2); |
| 136 | } | 136 | } |
| 137 | } | 137 | } |
| 138 | 138 | ||
| @@ -144,10 +144,10 @@ __aicore__ inline void SparseTensorDenseMatMulB16<T_IDX, T_VAL, T_SUM, ADJ_A, AD | |||
| 144 | __gm__ T_VAL* x1ValuesGmAddr = (__gm__ T_VAL*)x1ValuesGm_.GetPhyAddr(); | 144 | __gm__ T_VAL* x1ValuesGmAddr = (__gm__ T_VAL*)x1ValuesGm_.GetPhyAddr(); |
| 145 | __gm__ T_VAL* x2GmAddr = (__gm__ T_VAL*)x2Gm_.GetPhyAddr(); | 145 | __gm__ T_VAL* x2GmAddr = (__gm__ T_VAL*)x2Gm_.GetPhyAddr(); |
| 146 | __gm__ T_SUM* workspaceGmAddr = (__gm__ T_SUM*)workspaceGm_.GetPhyAddr(); | 146 | __gm__ T_SUM* workspaceGmAddr = (__gm__ T_SUM*)workspaceGm_.GetPhyAddr(); |
| 147 | - Simt::VF_CALL<ComputeB16<T_IDX, T_VAL, T_SUM, ADJ_A, ADJ_B>>( | 147 | + asc_vf_call<ComputeB16<T_IDX, T_VAL, T_SUM, ADJ_A, ADJ_B>>( |
| 148 | - Simt::Dim3{SIMT_MAX_THREAD_NUM, 1, 1}, tilingData_->computeUsedCoreNum, blockIdx_, | 148 | + dim3{SIMT_MAX_THREAD_NUM, 1, 1}, tilingData_->computeUsedCoreNum, blockIdx_, |
| 149 | - tilingData_->computeTotalElemNum, tilingData_->computeM, tilingData_->computeN, | 149 | + tilingData_->computeTotalElemNum, tilingData_->computeM, tilingData_->computeN, tilingData_->computeP, |
| 150 | - tilingData_->computeP, x1IndicesGmAddr, x1ValuesGmAddr, x2GmAddr, workspaceGmAddr); | 150 | + x1IndicesGmAddr, x1ValuesGmAddr, x2GmAddr, workspaceGmAddr); |
| 151 | } | 151 | } |
| 152 | SyncAll(); | 152 | SyncAll(); |
| 153 | if (blockIdx_ < tilingData_->initAndOutUsedCoreNum) { | 153 | if (blockIdx_ < tilingData_->initAndOutUsedCoreNum) { |
| @@ -25,8 +25,8 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_MAX_THREAD_NUM) inline void ComputeB32( | |||
| 25 | const int32_t p, __gm__ T_IDX* x1Indices, __gm__ T_VAL* x1Values, __gm__ T_VAL* x2, __gm__ T_VAL* y) | 25 | const int32_t p, __gm__ T_IDX* x1Indices, __gm__ T_VAL* x1Values, __gm__ T_VAL* x2, __gm__ T_VAL* y) |
| 26 | { | 26 | { |
| 27 | // 总共有usedCoreNum*ThreadNum个线程,每个线程所在的位置为currCoreIdx*ThreadNum+LocalThreadIdx | 27 | // 总共有usedCoreNum*ThreadNum个线程,每个线程所在的位置为currCoreIdx*ThreadNum+LocalThreadIdx |
| 28 | - for (int32_t elemIdx = currCoreIdx * Simt::GetThreadNum() + Simt::GetThreadIdx(); elemIdx < elemNum; | 28 | + for (int32_t elemIdx = currCoreIdx * blockDim.x + threadIdx.x; elemIdx < elemNum; |
| 29 | - elemIdx += usedCoreNum * Simt::GetThreadNum()) { | 29 | + elemIdx += usedCoreNum * blockDim.x) { |
| 30 | // 计算索引i、j、k,用于找到v1=x1(i,k)和v2=x2(k,j) | 30 | // 计算索引i、j、k,用于找到v1=x1(i,k)和v2=x2(k,j) |
| 31 | // i、j、k 统一转成int32类型 | 31 | // i、j、k 统一转成int32类型 |
| 32 | int32_t x1VecIdx = elemIdx / p; | 32 | int32_t x1VecIdx = elemIdx / p; |
| @@ -49,7 +49,7 @@ __simt_vf__ __aicore__ LAUNCH_BOUND(SIMT_MAX_THREAD_NUM) inline void ComputeB32( | |||
| 49 | } | 49 | } |
| 50 | // 累加到对应位置 | 50 | // 累加到对应位置 |
| 51 | __gm__ T_VAL* outAddr = y + i * p + j; | 51 | __gm__ T_VAL* outAddr = y + i * p + j; |
| 52 | - Simt::AtomicAdd(outAddr, v1 * v2); | 52 | + asc_atomic_add(outAddr, v1 * v2); |
| 53 | } | 53 | } |
| 54 | } | 54 | } |
| 55 | 55 | ||
| @@ -118,10 +118,10 @@ __aicore__ inline void SparseTensorDenseMatMulB32<T_IDX, T_VAL, ADJ_A, ADJ_B>::P | |||
| 118 | __gm__ T_VAL* x1ValuesGmAddr = (__gm__ T_VAL*)x1ValuesGm_.GetPhyAddr(); | 118 | __gm__ T_VAL* x1ValuesGmAddr = (__gm__ T_VAL*)x1ValuesGm_.GetPhyAddr(); |
| 119 | __gm__ T_VAL* x2GmAddr = (__gm__ T_VAL*)x2Gm_.GetPhyAddr(); | 119 | __gm__ T_VAL* x2GmAddr = (__gm__ T_VAL*)x2Gm_.GetPhyAddr(); |
| 120 | __gm__ T_VAL* yGmAddr = (__gm__ T_VAL*)yGm_.GetPhyAddr(); | 120 | __gm__ T_VAL* yGmAddr = (__gm__ T_VAL*)yGm_.GetPhyAddr(); |
| 121 | - Simt::VF_CALL<ComputeB32<T_IDX, T_VAL, ADJ_A, ADJ_B>>( | 121 | + asc_vf_call<ComputeB32<T_IDX, T_VAL, ADJ_A, ADJ_B>>( |
唐 | |||
| 122 | - Simt::Dim3{SIMT_MAX_THREAD_NUM, 1, 1}, tilingData_->computeUsedCoreNum, currCoreIdx_, | 122 | + dim3{SIMT_MAX_THREAD_NUM, 1, 1}, tilingData_->computeUsedCoreNum, currCoreIdx_, |
| 123 | - tilingData_->computeTotalElemNum, tilingData_->computeM, tilingData_->computeN, | 123 | + tilingData_->computeTotalElemNum, tilingData_->computeM, tilingData_->computeN, tilingData_->computeP, |
| 124 | - tilingData_->computeP, x1IndicesGmAddr, x1ValuesGmAddr, x2GmAddr, yGmAddr); | 124 | + x1IndicesGmAddr, x1ValuesGmAddr, x2GmAddr, yGmAddr); |
| 125 | } | 125 | } |
| 126 | } | 126 | } |
| 127 | 127 | ||
| @@ -17,6 +17,9 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 20 | 23 | ||
| 21 | namespace SparseTensorDenseMatMul { | 24 | namespace SparseTensorDenseMatMul { |
| 22 | 25 | ||
| @@ -25,4 +28,3 @@ constexpr int32_t BUFFER_NUM = 2; | |||
| 25 | constexpr int32_t INDICES_DIM_1 = 2; | 28 | constexpr int32_t INDICES_DIM_1 = 2; |
| 26 | 29 | ||
| 27 | } // namespace SparseTensorDenseMatMul | 30 | } // namespace SparseTensorDenseMatMul |
| 28 | - | ||
为什么新接口不是驼峰的?