已合并
修复连续使用initGlobalMemory初始化不同值之间的同步问题 #2939
zhang-wenbo-beat创建于 3月19日
修复连续使用initGlobalMemory初始化不同值之间的同步问题 #2939
已合并
zhang-wenbo-beat创建于 3月19日
共 4 个文件变更+40-24
@@ -159,6 +159,9 @@ __aicore__ inline void InplaceIndexAddDeterminstic<VAR_T, IDX_T>::Init(
159 GetBlockIdx() * tilingData_.eachCoreIndexCount, tilingData_.eachCoreIndexCount);159 GetBlockIdx() * tilingData_.eachCoreIndexCount, tilingData_.eachCoreIndexCount);
160 160 
161 InitGlobalMemory(updateSumWsGm_, sumWsSize_, (float)(0));161 InitGlobalMemory(updateSumWsGm_, sumWsSize_, (float)(0));
162+ auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
163+ SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
164+ WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
162 InitGlobalMemory(updateSumIdxWsGm_, tilingData_.eachCoreIndexCount, (IDX_T)(-1));165 InitGlobalMemory(updateSumIdxWsGm_, tilingData_.eachCoreIndexCount, (IDX_T)(-1));
163 AscendC::SyncAll();166 AscendC::SyncAll();
164}167}
@@ -168,6 +168,9 @@ template<typename T, typename U, uint32_t scatterOp>
168__aicore__ inline void ScatterAddDeterministicImpl<T, U, scatterOp>::InitWspZero()168__aicore__ inline void ScatterAddDeterministicImpl<T, U, scatterOp>::InitWspZero()
169{169{
170 InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));170 InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));
171+ auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
172+ SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
173+ WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
171 InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));174 InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));
172 InitGlobalMemory(workspaceMaxValueCount_, tilingData_.varShape[0], (uint32_t)(0));175 InitGlobalMemory(workspaceMaxValueCount_, tilingData_.varShape[0], (uint32_t)(0));
173 InitGlobalMemory(workspaceMaxValue_, tilingData_.varShape[1] * tilingData_.varShape[0], (float)(0));176 InitGlobalMemory(workspaceMaxValue_, tilingData_.varShape[1] * tilingData_.varShape[0], (float)(0));
@@ -171,6 +171,34 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Init(GM_ADDR indices, G
171 // 量化分支171 // 量化分支
172 shiftOffset_ = UB_AGLIN_VALUE / sizeof(U);172 shiftOffset_ = UB_AGLIN_VALUE / sizeof(U);
173 indicesUbLoop_ = GetBlockIdx() == tilingData_.logicCoreNum - 1 ? tilingData_.tailCoreIndicesLoopSize : tilingData_.indicesLoopSize;173 indicesUbLoop_ = GetBlockIdx() == tilingData_.logicCoreNum - 1 ? tilingData_.tailCoreIndicesLoopSize : tilingData_.indicesLoopSize;
174+
175+ indicesGm_.SetGlobalBuffer((__gm__ U *)(indices) + GetBlockIdx() * tilingData_.perCoreHandleIndices * tilingData_.rankSize);
176+ updatesGm_.SetGlobalBuffer((__gm__ T *)(updates) + GetBlockIdx() * tilingData_.perCoreHandleIndices * tilingData_.afterAxis);
177+ 
178+ workspaceMaxValueCount_.SetGlobalBuffer((__gm__ uint32_t *)(workspace));
179+
180+ workspaceMaxValue_.SetGlobalBuffer((__gm__ float *)workspace + tilingData_.rankFusedAxis);
181+ workspaceMaxValueInit_.SetGlobalBuffer((__gm__ float *)workspace + tilingData_.rankFusedAxis + GetBlockIdx() * initPerCore);
182+
183+ workspaceInt32Res_.SetGlobalBuffer((__gm__ int32_t *)workspace + tilingData_.rankFusedAxis +
184+ tilingData_.rankFusedAxis * tilingData_.afterAxis);
185+ workspaceInt32ResInit_.SetGlobalBuffer((__gm__ int32_t *)workspace + tilingData_.rankFusedAxis +
186+ tilingData_.rankFusedAxis * tilingData_.afterAxis + GetBlockIdx() * initPerCore);
187+
188+ workspaceLogicCoreSumValue_.SetGlobalBuffer((__gm__ float*)workspace + tilingData_.rankFusedAxis +
189+ tilingData_.rankFusedAxis * tilingData_.afterAxis * DOUBLE + GetBlockIdx() * tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_);
190+
191+ __gm__ float* workspaceCoreLogicCoreSumIdStart = (__gm__ float*)workspace + tilingData_.rankFusedAxis +
192+ tilingData_.rankFusedAxis * tilingData_.afterAxis * DOUBLE +
193+ tilingData_.logicCoreNum * tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_;
194+ workspaceCoreLogicCoreSumId_.SetGlobalBuffer((__gm__ U*)workspaceCoreLogicCoreSumIdStart + GetBlockIdx() * tilingData_.perCoreHandleIndices);
195+ 
196+ InitWspZero();
197+ if (tilingData_.isIdxSplit == 1) {
198+ InitGlobalMemory(yGmInit_, initCoreReal, (T)0);
199+ }
200+ SyncAll();
201+ 
174 pipe_.InitBuffer(indicesQue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U), UB_AGLIN_VALUE));202 pipe_.InitBuffer(indicesQue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U), UB_AGLIN_VALUE));
175 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));203 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));
176 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U)); 204 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U));
@@ -179,30 +207,15 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Init(GM_ADDR indices, G
179 pipe_.InitBuffer(updateSumIdxQueue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE));207 pipe_.InitBuffer(updateSumIdxQueue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE));
180 pipe_.InitBuffer(updatesQueue_, DOUBLE_BUF, tilingData_.indicesUbFactor * postVarAlignSize_ * sizeof(T));208 pipe_.InitBuffer(updatesQueue_, DOUBLE_BUF, tilingData_.indicesUbFactor * postVarAlignSize_ * sizeof(T));
181 pipe_.InitBuffer(updateSumQue_, DOUBLE_BUF, tilingData_.indicesUbFactor * postVarAlignSizeFp32_ * sizeof(float));209 pipe_.InitBuffer(updateSumQue_, DOUBLE_BUF, tilingData_.indicesUbFactor * postVarAlignSizeFp32_ * sizeof(float));
182-
183- indicesGm_.SetGlobalBuffer((__gm__ U *)(indices) + GetBlockIdx() * tilingData_.perCoreHandleIndices * tilingData_.rankSize);
184- updatesGm_.SetGlobalBuffer((__gm__ T *)(updates) + GetBlockIdx() * tilingData_.perCoreHandleIndices * tilingData_.afterAxis);
185- 
186- workspaceMaxValueCount_.SetGlobalBuffer((__gm__ uint32_t *)(workspace));
187- workspaceMaxValue_.SetGlobalBuffer((__gm__ float *)workspace + tilingData_.rankFusedAxis);
188- workspaceMaxValueInit_.SetGlobalBuffer((__gm__ float *)workspace + tilingData_.rankFusedAxis + GetBlockIdx() * initPerCore);
189- workspaceInt32Res_.SetGlobalBuffer((__gm__ int32_t *)workspace + tilingData_.rankFusedAxis +
190- tilingData_.rankFusedAxis * tilingData_.afterAxis);
191- workspaceInt32ResInit_.SetGlobalBuffer((__gm__ int32_t *)workspace + tilingData_.rankFusedAxis +
192- tilingData_.rankFusedAxis * tilingData_.afterAxis + GetBlockIdx() * initPerCore);
193- workspaceLogicCoreSumValue_.SetGlobalBuffer((__gm__ float*)workspace + tilingData_.rankFusedAxis +
194- tilingData_.rankFusedAxis * tilingData_.afterAxis * DOUBLE + GetBlockIdx() * tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_);
195-
196- __gm__ float* workspaceCoreLogicCoreSumIdStart = (__gm__ float*)workspace + tilingData_.rankFusedAxis +
197- tilingData_.rankFusedAxis * tilingData_.afterAxis * DOUBLE +
198- tilingData_.logicCoreNum * tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_;
199- workspaceCoreLogicCoreSumId_.SetGlobalBuffer((__gm__ U*)workspaceCoreLogicCoreSumIdStart + GetBlockIdx() * tilingData_.perCoreHandleIndices);
200}210}
201 211 
202template<typename T, typename U>212template<typename T, typename U>
203__aicore__ inline void ScatterNdDeterministicImpl<T, U>::InitWspZero()213__aicore__ inline void ScatterNdDeterministicImpl<T, U>::InitWspZero()
204{214{
205 InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));215 InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));
216+ auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
217+ SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
218+ WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
206 InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));219 InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));
207 InitGlobalMemory(workspaceMaxValueCount_, tilingData_.rankFusedAxis, (uint32_t)(0));220 InitGlobalMemory(workspaceMaxValueCount_, tilingData_.rankFusedAxis, (uint32_t)(0));
208 InitGlobalMemory(workspaceMaxValueInit_, initCoreReal, (float)(0));221 InitGlobalMemory(workspaceMaxValueInit_, initCoreReal, (float)(0));
@@ -913,12 +926,6 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Process()
913 return;926 return;
914 }927 }
915 928 
916- InitWspZero();
917- if (tilingData_.isIdxSplit == 1) {
918- InitGlobalMemory(yGmInit_, initCoreReal, (T)0);
919- }
920- SyncAll();
921- 
922 LocalTensor<U> calcLocal = calcBuf_.Get<U>();929 LocalTensor<U> calcLocal = calcBuf_.Get<U>();
923 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {930 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {
924 calcLocal(i) = tilingData_.strideList[i];931 calcLocal(i) = tilingData_.strideList[i];
@@ -160,6 +160,9 @@ __aicore__ inline void ScatterNdAddDeterministic<T, U, CAST_T, castType>::Init(G
160 GetBlockIdx() * tilingData_.eachCoreIndexCount, tilingData_.eachCoreIndexCount);160 GetBlockIdx() * tilingData_.eachCoreIndexCount, tilingData_.eachCoreIndexCount);
161 161
162 InitGlobalMemory(updateSumWsGm_, sumWsSize_, (float)(0));162 InitGlobalMemory(updateSumWsGm_, sumWsSize_, (float)(0));
163+ auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
164+ SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
165+ WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
163 InitGlobalMemory(updateSumIdxWsGm_, tilingData_.eachCoreIndexCount, (U)(-1));166 InitGlobalMemory(updateSumIdxWsGm_, tilingData_.eachCoreIndexCount, (U)(-1));
164 AscendC::SyncAll();167 AscendC::SyncAll();
165}168}