已合并
ScatterNd修复UB空间分配和Mask创建错误 #4019
zhang-wenbo-beat创建于 4月20日
ScatterNd修复UB空间分配和Mask创建错误 #4019
已合并
zhang-wenbo-beat创建于 4月20日
2 个文件变更+29-19
@@ -180,9 +180,15 @@ ge::graphStatus ScatterNdTiling::getRestAvailableSize(uint64_t sampleNum, uint64
180 OP_CHECK_IF(indicesDtypeSize <= 0, OP_LOGE(opName, "get indicesType size fail."),180 OP_CHECK_IF(indicesDtypeSize <= 0, OP_LOGE(opName, "get indicesType size fail."),
181 return ge::GRAPH_FAILED);181 return ge::GRAPH_FAILED);
182 auto ubBlock = static_cast<uint64_t>(Ops::Base::GetUbBlockSize(context_));182 auto ubBlock = static_cast<uint64_t>(Ops::Base::GetUbBlockSize(context_));
183- uint64_t occupy = sampleNum * Ops::Base::CeilAlign(FLOAT_BYTES * postAxisSize_, ubBlock) + sampleNum * Ops::Base::CeilAlign(FLOAT_BYTES * postAxisSize_, ubBlock) + 183+ uint64_t occupy = sampleNum * Ops::Base::CeilAlign(FLOAT_BYTES * postAxisSize_, ubBlock) +
184- Ops::Base::CeilAlign(sampleNum * indicesDtypeSize, ubBlock) * THREE + Ops::Base::CeilAlign(sampleNum * UINT32_BYTES, ubBlock) + TWO * TWO * UB_AGLIN_VALUE +184+ sampleNum * Ops::Base::CeilAlign(FLOAT_BYTES * postAxisSize_, ubBlock) +
185- TWO * UB_AGLIN_VALUE + GetSortTmpSize(idType, sampleNum, false) + Ops::Base::CeilAlign(postAxisSize_ * FLOAT_BYTES, ubBlock);185+ Ops::Base::CeilAlign(sampleNum * indicesDtypeSize * rankSize_, ubBlock) +
186+ Ops::Base::CeilAlign(sampleNum * indicesDtypeSize, ubBlock) * THREE +
187+ Ops::Base::CeilAlign(sampleNum * UINT32_BYTES, ubBlock) +
188+ TWO * TWO * UB_AGLIN_VALUE +
189+ TWO * UB_AGLIN_VALUE +
190+ GetSortTmpSize(idType, sampleNum, false) +
191+ Ops::Base::CeilAlign(postAxisSize_ * FLOAT_BYTES, ubBlock);
186 return originalSize - occupy;192 return originalSize - occupy;
187}193}
188 194 
@@ -211,7 +217,7 @@ void ScatterNdTiling::BlockTiling() {
211 217 
212ge::graphStatus ScatterNdTiling::UbTiling() {218ge::graphStatus ScatterNdTiling::UbTiling() {
213 // halfUbSize for double buffer219 // halfUbSize for double buffer
214- auto halfUbSize = ubSize_ / BUFFER_NUM;220+ auto halfUbSize = (ubSize_ - TWO * TWO * UB_AGLIN_VALUE) / BUFFER_NUM;
215 auto indiceNum = indiceShapeSize_ / rankSize_;221 auto indiceNum = indiceShapeSize_ / rankSize_;
216 sliceSize_ = updateShapeSize_ / indiceNum;222 sliceSize_ = updateShapeSize_ / indiceNum;
217 OP_CHECK_IF(sliceSize_ == static_cast<uint64_t>(0),223 OP_CHECK_IF(sliceSize_ == static_cast<uint64_t>(0),
@@ -246,7 +252,7 @@ ge::graphStatus ScatterNdTiling::ScatterNdDeterministicTiling()
246 }252 }
247 if (DeterministicFlag) {253 if (DeterministicFlag) {
248 OP_LOGD(opName, "ScatterNd Deterministic Non-Quant branch start");254 OP_LOGD(opName, "ScatterNd Deterministic Non-Quant branch start");
249- ubSize_ = static_cast<uint32_t>(ubSize_ / BUFFER_NUM);255+ ubSize_ = static_cast<uint32_t>((ubSize_ - TWO * TWO * UB_AGLIN_VALUE) / BUFFER_NUM);
250 perCoreHandleCol_ = Ops::Base::CeilDiv(afterAxis_, coreNum_);256 perCoreHandleCol_ = Ops::Base::CeilDiv(afterAxis_, coreNum_);
251 logicCoreNum_ = Ops::Base::CeilDiv(afterAxis_, perCoreHandleCol_);257 logicCoreNum_ = Ops::Base::CeilDiv(afterAxis_, perCoreHandleCol_);
252 tailCoreHandleCol_ = static_cast<uint64_t>(afterAxis_) - (logicCoreNum_ - static_cast<uint64_t>(1)) * perCoreHandleCol_;258 tailCoreHandleCol_ = static_cast<uint64_t>(afterAxis_) - (logicCoreNum_ - static_cast<uint64_t>(1)) * perCoreHandleCol_;
@@ -31,6 +31,7 @@ constexpr uint64_t THREE = 3;
31constexpr uint32_t SORT_STAT_PADDING = 64;31constexpr uint32_t SORT_STAT_PADDING = 64;
32constexpr uint64_t UB_AGLIN_VALUE = 32;32constexpr uint64_t UB_AGLIN_VALUE = 32;
33constexpr uint16_t MAX_RANK_COUNT = 7;33constexpr uint16_t MAX_RANK_COUNT = 7;
34+constexpr uint16_t MAX_SHAPE_RANK = 8;
34 35 
35static constexpr MicroAPI::CastTrait castTraitFP322INT32 = {MicroAPI::RegLayout::UNKNOWN, MicroAPI::SatMode::SAT,36static constexpr MicroAPI::CastTrait castTraitFP322INT32 = {MicroAPI::RegLayout::UNKNOWN, MicroAPI::SatMode::SAT,
36 MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT};37 MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT};
@@ -126,7 +127,7 @@ private:
126 uint64_t postVarAlignSize_{0};127 uint64_t postVarAlignSize_{0};
127 uint64_t postVarAlignSizeFp32_{0};128 uint64_t postVarAlignSizeFp32_{0};
128 uint64_t strideList[MAX_RANK_COUNT];129 uint64_t strideList[MAX_RANK_COUNT];
129- uint64_t outputShape[MAX_RANK_COUNT];130+ uint64_t outputShape[MAX_SHAPE_RANK];
130 uint64_t initPerCore{0};131 uint64_t initPerCore{0};
131 uint64_t initCoreReal{0};132 uint64_t initCoreReal{0};
132 133 
@@ -167,7 +168,7 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Init(GM_ADDR indices, G
167 pipe_.InitBuffer(indicesBuf_, tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U));168 pipe_.InitBuffer(indicesBuf_, tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U));
168 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));169 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));
169 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U));170 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U));
170- pipe_.InitBuffer(outputShapeBuf_, MAX_RANK_COUNT * sizeof(U));171+ pipe_.InitBuffer(outputShapeBuf_, MAX_SHAPE_RANK * sizeof(U));
171 172 
172 SyncAll();173 SyncAll();
173 return;174 return;
@@ -206,7 +207,7 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Init(GM_ADDR indices, G
206 pipe_.InitBuffer(indicesQue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U), UB_AGLIN_VALUE));207 pipe_.InitBuffer(indicesQue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * tilingData_.rankSize * sizeof(U), UB_AGLIN_VALUE));
207 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));208 pipe_.InitBuffer(outOfstBuf_, tilingData_.indicesUbFactor * sizeof(U));
208 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U));209 pipe_.InitBuffer(calcBuf_, MAX_RANK_COUNT * sizeof(U));
209- pipe_.InitBuffer(outputShapeBuf_, MAX_RANK_COUNT * sizeof(U));210+ pipe_.InitBuffer(outputShapeBuf_, MAX_SHAPE_RANK * sizeof(U));
210 pipe_.InitBuffer(sortedIndicesQue_, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE) + SORT_STAT_PADDING);211 pipe_.InitBuffer(sortedIndicesQue_, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE) + SORT_STAT_PADDING);
211 pipe_.InitBuffer(updatesOriginIdexQue_, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(uint32_t), UB_AGLIN_VALUE));212 pipe_.InitBuffer(updatesOriginIdexQue_, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(uint32_t), UB_AGLIN_VALUE));
212 pipe_.InitBuffer(updateSumIdxQueue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE));213 pipe_.InitBuffer(updateSumIdxQueue_, DOUBLE_BUF, ops::CeilAlign(tilingData_.indicesUbFactor * sizeof(U), UB_AGLIN_VALUE));
@@ -217,12 +218,15 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Init(GM_ADDR indices, G
217template<typename T, typename U>218template<typename T, typename U>
218__aicore__ inline void ScatterNdDeterministicImpl<T, U>::InitWspZero()219__aicore__ inline void ScatterNdDeterministicImpl<T, U>::InitWspZero()
219{220{
220- InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));221+ if (GetBlockIdx() < tilingData_.logicCoreNum)
221- auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));222+ {
222- SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);223+ InitGlobalMemory(workspaceCoreLogicCoreSumId_, tilingData_.perCoreHandleIndices, (U)(-1));
223- WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);224+ auto vWaitMte3EventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
224- InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));225+ SetFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
225- InitGlobalMemory(workspaceMaxValueCount_, tilingData_.rankFusedAxis, (uint32_t)(0));226+ WaitFlag<HardEvent::MTE3_V>(vWaitMte3EventID);
227+ InitGlobalMemory(workspaceLogicCoreSumValue_, tilingData_.perCoreHandleIndices * postVarAlignSizeFp32_, (float)(0));
228+ InitGlobalMemory(workspaceMaxValueCount_, tilingData_.rankFusedAxis, (uint32_t)(0));
229+ }
226 InitGlobalMemory(workspaceMaxValueInit_, initCoreReal, (float)(0));230 InitGlobalMemory(workspaceMaxValueInit_, initCoreReal, (float)(0));
227 InitGlobalMemory(workspaceInt32ResInit_, initCoreReal, (int)(0));231 InitGlobalMemory(workspaceInt32ResInit_, initCoreReal, (int)(0));
228}232}
@@ -277,15 +281,15 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::ComputOutOfset(
277 IndexRegType indexReg;281 IndexRegType indexReg;
278 AscendC::MicroAPI::MaskReg pregLoop;282 AscendC::MicroAPI::MaskReg pregLoop;
279 AscendC::MicroAPI::MaskReg cmpMask;283 AscendC::MicroAPI::MaskReg cmpMask;
280- AscendC::MicroAPI::MaskReg invalidMask; 284+ AscendC::MicroAPI::MaskReg invalidMask;
281 285 
282 for (uint16_t i = 0; i < loopCnt; i++) {286 for (uint16_t i = 0; i < loopCnt; i++) {
283 if constexpr (IsSameType<U, int64_t>::value) {287 if constexpr (IsSameType<U, int64_t>::value) {
284 pregLoop = AscendC::MicroAPI::UpdateMask<U, AscendC::MicroAPI::RegTraitNumTwo>(dataLen);288 pregLoop = AscendC::MicroAPI::UpdateMask<U, AscendC::MicroAPI::RegTraitNumTwo>(dataLen);
285- invalidMask = AscendC::MicroAPI::UpdateMask<U, AscendC::MicroAPI::RegTraitNumTwo>(dataLen);289+ invalidMask = AscendC::MicroAPI::CreateMask<U, MicroAPI::MaskPattern::ALLF, AscendC::MicroAPI::RegTraitNumTwo>();
286 } else {290 } else {
287 pregLoop = AscendC::MicroAPI::UpdateMask<U>(dataLen);291 pregLoop = AscendC::MicroAPI::UpdateMask<U>(dataLen);
288- invalidMask = AscendC::MicroAPI::UpdateMask<U>(dataLen);292+ invalidMask = AscendC::MicroAPI::CreateMask<U, MicroAPI::MaskPattern::ALLF>();
289 }293 }
290 AscendC::MicroAPI::Duplicate(outReg, 0, pregLoop);294 AscendC::MicroAPI::Duplicate(outReg, 0, pregLoop);
291 AscendC::MicroAPI::Arange(orderReg, i * vfLen);295 AscendC::MicroAPI::Arange(orderReg, i * vfLen);
@@ -373,7 +377,7 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::ProcessAtomicAdd()
373 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {377 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {
374 calcLocal(i) = tilingData_.strideList[i];378 calcLocal(i) = tilingData_.strideList[i];
375 }379 }
376- for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {380+ for (int32_t i = 0; i < MAX_SHAPE_RANK; i++) {
377 outputShapeLocal(i) = tilingData_.outPutShape[i];381 outputShapeLocal(i) = tilingData_.outPutShape[i];
378 }382 }
379 for (int64_t indicesLoopIdx = 0; indicesLoopIdx < indicesLoopSize - 1; ++indicesLoopIdx) {383 for (int64_t indicesLoopIdx = 0; indicesLoopIdx < indicesLoopSize - 1; ++indicesLoopIdx) {
@@ -959,7 +963,7 @@ __aicore__ inline void ScatterNdDeterministicImpl<T, U>::Process()
959 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {963 for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {
960 calcLocal(i) = tilingData_.strideList[i];964 calcLocal(i) = tilingData_.strideList[i];
961 }965 }
962- for (int32_t i = 0; i < MAX_RANK_COUNT; i++) {966+ for (int32_t i = 0; i < MAX_SHAPE_RANK; i++) {
963 outputShapeLocal(i) = tilingData_.outPutShape[i];967 outputShapeLocal(i) = tilingData_.outPutShape[i];
964 }968 }
965 969