已合并
[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
已合并
majiayuan创建于 4月23日
共 10 个文件变更+105-105
@@ -13,6 +13,8 @@
13 13 
14#include "kernel_operator.h"14#include "kernel_operator.h"
15#include "embedding_common.h"15#include "embedding_common.h"
16+#include "simt_api/asc_simt.h"
17+#include "simt_api/math_functions.h"
16 18 
17static constexpr uint8_t VALID_FLAG_MASK = 0b00000001;19static 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#include "kernel_operator.h"19#include "kernel_operator.h"
20#include "kernel_operator_list_tensor_intf.h"20#include "kernel_operator_list_tensor_intf.h"
21+#include "simt_api/asc_simt.h"
22+#include "simt_api/device_atomic_functions.h"
21 23 
22namespace EmbeddingHashTableExportAicore {24namespace EmbeddingHashTableExportAicore {
23using namespace AscendC;25using 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 = tableAddrU8220 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 
233template <typename T>233template <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 
248template <typename T>248template <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 = tableAddrU8277 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 = tableAddrI64282 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 = tableAddrU64285 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] = tableAddrT295 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 tableAddrU8301 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 
318template <typename T>312template <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#include "kernel_operator.h"18#include "kernel_operator.h"
19#include "kernel_operator_list_tensor_intf.h"19#include "kernel_operator_list_tensor_intf.h"
20#include "../../inc/hashtable_common.h"20#include "../../inc/hashtable_common.h"
21+#include "simt_api/asc_simt.h"
22+#include "simt_api/device_atomic_functions.h"
21 23 
22namespace EmbeddingHashTable {24namespace EmbeddingHashTable {
23using namespace AscendC;25using 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 struct204 // 更新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}
@@ -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的那几个统计字段的值
@@ -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#define CALL_COMPUTE_VF(macro_d, macro_f, macro_pcounts) \133#define CALL_COMPUTE_VF(macro_d, macro_f, macro_pcounts) \
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#include "op_kernel/math_util.h"19#include "op_kernel/math_util.h"
20#include "op_kernel/platform_util.h"20#include "op_kernel/platform_util.h"
21#include "../../inc/hashtable_common.h"21#include "../../inc/hashtable_common.h"
22+#include "simt_api/asc_simt.h"
23+#include "simt_api/device_atomic_functions.h"
22 24 
23namespace Hashtbl {25namespace Hashtbl {
24using namespace AscendC;26using namespace AscendC;
@@ -17,6 +17,7 @@
17#define OPS_BUILT_IN_TBE_IMPL_ASCENDC_INIT_EMBEDDING_HASHTABLE_INIT_EMBEDDING_HASH_TABLE_H17#define OPS_BUILT_IN_TBE_IMPL_ASCENDC_INIT_EMBEDDING_HASHTABLE_INIT_EMBEDDING_HASH_TABLE_H
18 18 
19#include "kernel_operator.h"19#include "kernel_operator.h"
20+#include "simt_api/asc_simt.h"
20 21 
21namespace InitEmbeddingHashTable {22namespace InitEmbeddingHashTable {
22using namespace AscendC;23using 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+LocalThreadIdx112 // 总共有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+LocalThreadIdx27 // 总共有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>>(
唐
唐唐超5月11日

为什么新接口不是驼峰的?

likedislike
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#include "kernel_operator.h"17#include "kernel_operator.h"
18#include "kernel_tiling/kernel_tiling.h"18#include "kernel_tiling/kernel_tiling.h"
19#include "sparse_tensor_dense_mat_mul_tiling_def.h"19#include "sparse_tensor_dense_mat_mul_tiling_def.h"
20+#include "simt_api/asc_fp16.h"
21+#include "simt_api/asc_simt.h"
22+#include "simt_api/device_atomic_functions.h"
20 23 
21namespace SparseTensorDenseMatMul {24namespace SparseTensorDenseMatMul {
22 25 
@@ -25,4 +28,3 @@ constexpr int32_t BUFFER_NUM = 2;
25constexpr int32_t INDICES_DIM_1 = 2;28constexpr int32_t INDICES_DIM_1 = 2;
26 29 
27} // namespace SparseTensorDenseMatMul30} // namespace SparseTensorDenseMatMul
28-