已合并
修复连续使用initGlobalMemory初始化不同值之间的同步问题 #2939
zhang-wenbo-beat创建于 3月19日
修复连续使用initGlobalMemory初始化不同值之间的同步问题 #2939
已合并
共 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 | ||
| 202 | template<typename T, typename U> | 212 | template<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 | } |