已合并
ScatterNd修复UB空间分配和Mask创建错误 #4019
zhang-wenbo-beat创建于 4月20日
ScatterNd修复UB空间分配和Mask创建错误 #4019
已合并
共 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 | ||
| 212 | ge::graphStatus ScatterNdTiling::UbTiling() { | 218 | ge::graphStatus ScatterNdTiling::UbTiling() { |
| 213 | // halfUbSize for double buffer | 219 | // 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; | |||
| 31 | constexpr uint32_t SORT_STAT_PADDING = 64; | 31 | constexpr uint32_t SORT_STAT_PADDING = 64; |
| 32 | constexpr uint64_t UB_AGLIN_VALUE = 32; | 32 | constexpr uint64_t UB_AGLIN_VALUE = 32; |
| 33 | constexpr uint16_t MAX_RANK_COUNT = 7; | 33 | constexpr uint16_t MAX_RANK_COUNT = 7; |
| 34 | +constexpr uint16_t MAX_SHAPE_RANK = 8; | ||
| 34 | 35 | ||
| 35 | static constexpr MicroAPI::CastTrait castTraitFP322INT32 = {MicroAPI::RegLayout::UNKNOWN, MicroAPI::SatMode::SAT, | 36 | static 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 | |||
| 217 | template<typename T, typename U> | 218 | template<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 | ||