已合并
engram反向超时DFX #11577
luozhonglin创建于 23 天前
engram反向超时DFX #11577
已合并
共 5 个文件变更+139-91
| @@ -489,6 +489,50 @@ static ge::graphStatus CheckAttrs(const gert::TilingContext *context) | |||
| 489 | return ge::GRAPH_SUCCESS; | 489 | return ge::GRAPH_SUCCESS; |
| 490 | } | 490 | } |
| 491 | 491 | ||
| 492 | +constexpr int64_t SORT_UB_HISTOGRAM_BINS = 256; // 镜像 HISTOGRAM_BINS | ||
| 493 | +constexpr int64_t SORT_UB_TILE_ELEMENTS = 4096; // 镜像 DEFAULT_TILE_SIZE | ||
| 494 | +constexpr int64_t SORT_UB_SINGLE_CORE_ELEMENTS = 8192; // 镜像 MAX_SINGLE_CORE_ELEMENTS | ||
| 495 | +constexpr int64_t SORT_UB_SIMT_SLOT_ALIGN = 32; // 镜像 SIMT_SLOT_ALIGN | ||
| 496 | +constexpr int64_t SORT_UB_SIMT_STAGING_PAD_BYTES = | ||
| 497 | + SORT_UB_HISTOGRAM_BINS * (SORT_UB_SIMT_SLOT_ALIGN - sizeof(int32_t)) + | ||
| 498 | + SORT_UB_SIMT_SLOT_ALIGN; // 镜像 SIMT_STAGING_PAD_BYTES | ||
| 499 | +constexpr int64_t SORT_UB_TMP_BUCKET_BYTES = 512; // 镜像 Mc2Kernel::SORT_TMP_BUCKET_BYTES | ||
| 500 | +constexpr int64_t SORT_UB_TMP_BYTES_PER_ELEM = 7; // 镜像 Mc2Kernel::SORT_TMP_BYTES_PER_ELEM | ||
| 501 | +constexpr int64_t SORT_UB_COUNT_ALIGN = 32; // 镜像 Mc2Kernel::SORT_COUNT_ALIGN | ||
| 502 | + | ||
| 503 | +static uint32_t CalcSortUbBytes(uint32_t numCores) | ||
| 504 | +{ | ||
| 505 | + auto alignUb = [](uint64_t bytes) -> uint32_t { | ||
| 506 | + return static_cast<uint32_t>((bytes + Mc2Kernel::UB_ALIGN - 1) / Mc2Kernel::UB_ALIGN * Mc2Kernel::UB_ALIGN); | ||
| 507 | + }; | ||
| 508 | + const uint32_t tileAlignBytes = alignUb(SORT_UB_TILE_ELEMENTS * sizeof(int32_t)); | ||
| 509 | + const uint32_t phaseHistBytes = 2U * tileAlignBytes; | ||
| 510 | + const uint32_t phasePrefixBytes = numCores * SORT_UB_HISTOGRAM_BINS * sizeof(int32_t) + tileAlignBytes; | ||
| 511 | + const uint32_t simtPadBytes = alignUb(tileAlignBytes + SORT_UB_SIMT_STAGING_PAD_BYTES); | ||
| 512 | + const uint32_t phaseScatterBytes = | ||
| 513 | + 3U * tileAlignBytes + 2U * SORT_UB_HISTOGRAM_BINS * sizeof(uint32_t) + 2U * simtPadBytes; | ||
| 514 | + const uint32_t tmpAlignedCount = static_cast<uint32_t>((SORT_UB_SINGLE_CORE_ELEMENTS + SORT_UB_COUNT_ALIGN - 1) / | ||
| 515 | + SORT_UB_COUNT_ALIGN * SORT_UB_COUNT_ALIGN); | ||
| 516 | + const uint32_t sortTmpBytes = static_cast<uint32_t>(SORT_UB_TMP_BUCKET_BYTES) + | ||
| 517 | + static_cast<uint32_t>(SORT_UB_TMP_BYTES_PER_ELEM) * tmpAlignedCount; | ||
| 518 | + const uint32_t scWorkspaceBytes = 3U * alignUb(SORT_UB_SINGLE_CORE_ELEMENTS * sizeof(int32_t)) + sortTmpBytes; | ||
| 519 | + uint32_t workspaceBytes = phaseHistBytes; | ||
| 520 | + if (phasePrefixBytes > workspaceBytes) { | ||
| 521 | + workspaceBytes = phasePrefixBytes; | ||
| 522 | + } | ||
| 523 | + if (phaseScatterBytes > workspaceBytes) { | ||
| 524 | + workspaceBytes = phaseScatterBytes; | ||
| 525 | + } | ||
| 526 | + if (scWorkspaceBytes > workspaceBytes) { | ||
| 527 | + workspaceBytes = scWorkspaceBytes; | ||
| 528 | + } | ||
| 529 | + // keyBuffer + sortedKeyBuffer + histogram/cumulative/prefix/histWide/prefixQueue 五小 buf | ||
| 530 | + const uint32_t smallBufsBytes = static_cast<uint32_t>(2U * SORT_UB_SINGLE_CORE_ELEMENTS) + | ||
| 531 | + static_cast<uint32_t>(SORT_UB_HISTOGRAM_BINS) * | ||
| 532 | + (2U * sizeof(int32_t) + sizeof(uint16_t) + 2U * sizeof(int32_t)); | ||
| 533 | + return workspaceBytes + smallBufsBytes; | ||
| 534 | +} | ||
| 535 | + | ||
| 492 | static ge::graphStatus SetPlatformInfo(gert::TilingContext *context, EngramFetchGradTilingData &tilingData) | 536 | static ge::graphStatus SetPlatformInfo(gert::TilingContext *context, EngramFetchGradTilingData &tilingData) |
| 493 | { | 537 | { |
| 494 | const char *nodeName = context->GetNodeName(); | 538 | const char *nodeName = context->GetNodeName(); |
| @@ -564,7 +608,7 @@ static ge::graphStatus SetTilingData(const gert::TilingContext *context, EngramF | |||
| 564 | auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE); | 608 | auto gradUniqueDesc = context->GetOutputDesc(OUT_GRAD_UNIQUE); |
| 565 | tilingData.outputDtype = static_cast<int32_t>(gradUniqueDesc->GetDataType()); | 609 | tilingData.outputDtype = static_cast<int32_t>(gradUniqueDesc->GetDataType()); |
| 566 | // 半精度输出时 FlushAccum 借用 entryBuf_ 尾部(20KB 偏移后)作 flush cast 双缓冲, | 610 | // 半精度输出时 FlushAccum 借用 entryBuf_ 尾部(20KB 偏移后)作 flush cast 双缓冲, |
| 567 | - // 需 20KB + 2*Align32(hiddenDim*2) 不越界(SEC-4.2③/TIL-3①;Kernel 侧另有 RUNTIME_ABORT 兜底) | 611 | + // 需 20KB + 2*Align32(hiddenDim*2) 不越界(SEC-4.2③/TIL-3①;Kernel 侧另有 ascendc_assert 兜底) |
| 568 | if (tilingData.outputDtype != static_cast<int32_t>(ge::DT_FLOAT)) { | 612 | if (tilingData.outputDtype != static_cast<int32_t>(ge::DT_FLOAT)) { |
| 569 | int64_t flushCastNeed = | 613 | int64_t flushCastNeed = |
| 570 | Mc2Kernel::FLUSH_CAST_HEAD_BYTES + 2 * AlignTo(static_cast<int64_t>(hiddenDim) * 2, Mc2Kernel::UB_ALIGN); | 614 | Mc2Kernel::FLUSH_CAST_HEAD_BYTES + 2 * AlignTo(static_cast<int64_t>(hiddenDim) * 2, Mc2Kernel::UB_ALIGN); |
| @@ -619,12 +663,32 @@ static ge::graphStatus SetTilingData(const gert::TilingContext *context, EngramF | |||
| 619 | static_cast<int64_t>(Mc2Kernel::UB_ALIGN))); | 663 | static_cast<int64_t>(Mc2Kernel::UB_ALIGN))); |
| 620 | bool needCast = (tilingData.inputDtype != static_cast<int32_t>(ge::DT_FLOAT)); | 664 | bool needCast = (tilingData.inputDtype != static_cast<int32_t>(ge::DT_FLOAT)); |
| 621 | uint32_t availableForPool = static_cast<uint32_t>(tilingData.ubSize) - static_cast<uint32_t>(permanentUb); | 665 | uint32_t availableForPool = static_cast<uint32_t>(tilingData.ubSize) - static_cast<uint32_t>(permanentUb); |
| 622 | - uint32_t availableForCast = (availableForPool > Mc2Kernel::COMM_BUF_BYTES + accumNeed) ? | 666 | + // sort 阶段池峰值 fail-fast:Kernel 侧 poolSize = max(sortUb, unique 峰值),常驻区随 rankSize |
| 623 | - (availableForPool - Mc2Kernel::COMM_BUF_BYTES - accumNeed) : | 667 | + // 增长会先挤破 sort 池(sortUb 仅依赖核数,本平台 177152),必须与 unique 一并建模 |
| 624 | - 0U; | 668 | + uint32_t sortUbBytes = CalcSortUbBytes(tilingData.aivNum); |
| 669 | + OP_TILING_CHECK(sortUbBytes > availableForPool, | ||
| 670 | + OP_LOGE(nodeName, | ||
| 671 | + "sort-phase UB pool overflow: sortUb=%u exceeds availableForPool=%u " | ||
| 672 | + "(aivNum=%u, rankSize=%u, permanentUb=%llu)", | ||
| 673 | + sortUbBytes, availableForPool, tilingData.aivNum, rankSize, permanentUb), | ||
| 674 | + return ge::GRAPH_FAILED); | ||
| 675 | + uint32_t uniqueEntryBytes = Mc2Kernel::FLUSH_CAST_HEAD_BYTES; | ||
| 676 | + if (tilingData.outputDtype != static_cast<int32_t>(ge::DT_FLOAT)) { | ||
| 677 | + uniqueEntryBytes += 2U * static_cast<uint32_t>(AlignTo(hiddenDim * 2, Mc2Kernel::UB_ALIGN)); | ||
| 678 | + } | ||
| 625 | // cast 缓冲行 stride 同样 32B 对齐(fp32 行) | 679 | // cast 缓冲行 stride 同样 32B 对齐(fp32 行) |
| 626 | uint32_t fp32RowStride = static_cast<uint32_t>( | 680 | uint32_t fp32RowStride = static_cast<uint32_t>( |
| 627 | AlignTo(static_cast<int64_t>(tilingData.hiddenDim) * sizeof(float), static_cast<int64_t>(Mc2Kernel::UB_ALIGN))); | 681 | AlignTo(static_cast<int64_t>(tilingData.hiddenDim) * sizeof(float), static_cast<int64_t>(Mc2Kernel::UB_ALIGN))); |
| 682 | + uint32_t minUniqueNeed = | ||
| 683 | + Mc2Kernel::GRAD_BUF_BYTES + uniqueEntryBytes + accumNeed + (needCast ? 2U * fp32RowStride : 0U); | ||
| 684 | + OP_TILING_CHECK(minUniqueNeed > availableForPool, | ||
| 685 | + OP_LOGE(nodeName, | ||
| 686 | + "unique-phase UB pool overflow: grad=%u + entry=%u + accum=%u + minCast=%u " | ||
| 687 | + "exceeds availableForPool=%u (hiddenDim=%lld, inputDtype=%d, outputDtype=%d)", | ||
| 688 | + Mc2Kernel::GRAD_BUF_BYTES, uniqueEntryBytes, accumNeed, needCast ? 2U * fp32RowStride : 0U, | ||
| 689 | + availableForPool, hiddenDim, tilingData.inputDtype, tilingData.outputDtype), | ||
| 690 | + return ge::GRAPH_FAILED); | ||
| 691 | + uint32_t availableForCast = availableForPool - Mc2Kernel::GRAD_BUF_BYTES - uniqueEntryBytes - accumNeed; | ||
| 628 | uint32_t maxByCast = needCast ? (availableForCast / (fp32RowStride * Mc2Kernel::ACCUM_BUF_COPIES)) : maxByPong; | 692 | uint32_t maxByCast = needCast ? (availableForCast / (fp32RowStride * Mc2Kernel::ACCUM_BUF_COPIES)) : maxByPong; |
| 629 | uint32_t gradSubBatch = maxByPong; | 693 | uint32_t gradSubBatch = maxByPong; |
| 630 | if (maxByCast < gradSubBatch) { | 694 | if (maxByCast < gradSubBatch) { |
| @@ -55,9 +55,16 @@ __aicore__ inline void EngramFetchGradSyncFunc() | |||
| 55 | 55 | ||
| 56 | constexpr uint32_t ENGRAM_GRAD_TIMEOUT_US = 60U * 1000U * 1000U; | 56 | constexpr uint32_t ENGRAM_GRAD_TIMEOUT_US = 60U * 1000U * 1000U; |
| 57 | constexpr uint32_t ENGRAM_GRAD_CYCLES_PER_US = 1000U; | 57 | constexpr uint32_t ENGRAM_GRAD_CYCLES_PER_US = 1000U; |
| 58 | -constexpr uint32_t COMM_RETRY_COUNT = 3U; | ||
| 59 | constexpr int32_t GRAD_CREDIT_READ_SENTINEL = -1; | 58 | constexpr int32_t GRAD_CREDIT_READ_SENTINEL = -1; |
| 60 | 59 | ||
| 60 | +// 超时部位标记:用于 TimeoutCheck 定位卡死等待点 | ||
| 61 | +enum TimeoutSite { | ||
| 62 | + TIMEOUT_CREDIT_READ_WAIT = 1, // CompleteCreditCounter 等待异步读完成 | ||
| 63 | + TIMEOUT_STATUS_FLAG_WAIT = 2, // WaitAllStatusFlags 等待跨 rank barrier | ||
| 64 | + TIMEOUT_SEND_CREDIT_WAIT = 3, // SendGradRemote 发送端等待对端 credit | ||
| 65 | + TIMEOUT_RECV_COUNTER_WAIT = 4, // RecvGradFromPeers 接收端等待对端写计数 | ||
| 66 | +}; | ||
| 67 | + | ||
| 61 | class EngramFetchGradArch35 { | 68 | class EngramFetchGradArch35 { |
| 62 | public: | 69 | public: |
| 63 | __aicore__ inline EngramFetchGradArch35() = default; | 70 | __aicore__ inline EngramFetchGradArch35() = default; |
| @@ -72,7 +79,7 @@ public: | |||
| 72 | private: | 79 | private: |
| 73 | __aicore__ inline void WriteNbiChecked(uint64_t handle, GM_ADDR dst, GM_ADDR src, uint64_t len); | 80 | __aicore__ inline void WriteNbiChecked(uint64_t handle, GM_ADDR dst, GM_ADDR src, uint64_t len); |
| 74 | __aicore__ inline void DrainChecked(uint64_t handle); | 81 | __aicore__ inline void DrainChecked(uint64_t handle); |
| 75 | - __aicore__ inline void TimeoutCheck(uint64_t startTime); | 82 | + __aicore__ inline void TimeoutCheck(uint64_t startTime, TimeoutSite site); |
| 76 | __aicore__ inline void UnsortGrad(); | 83 | __aicore__ inline void UnsortGrad(); |
| 77 | __aicore__ inline uint32_t LoadGradChunk(int64_t pos, int64_t end, LocalTensor<uint8_t> &buf, int32_t bufIdx, | 84 | __aicore__ inline uint32_t LoadGradChunk(int64_t pos, int64_t end, LocalTensor<uint8_t> &buf, int32_t bufIdx, |
| 78 | LocalTensor<int32_t> &idxUb, uint32_t tokensPerBuf); | 85 | LocalTensor<int32_t> &idxUb, uint32_t tokensPerBuf); |
| @@ -133,6 +140,7 @@ private: | |||
| 133 | uint64_t ubSize_{0}; | 140 | uint64_t ubSize_{0}; |
| 134 | uint32_t tileBytes_{0}; | 141 | uint32_t tileBytes_{0}; |
| 135 | uint32_t gradSubBatch_{Mc2Kernel::GRAD_SUB_BATCH}; | 142 | uint32_t gradSubBatch_{Mc2Kernel::GRAD_SUB_BATCH}; |
| 143 | + uint32_t uniqueEntryBytes_{0}; | ||
| 136 | 144 | ||
| 137 | uint64_t barrierFlagOffset_{0}; | 145 | uint64_t barrierFlagOffset_{0}; |
| 138 | uint64_t tokenWriteOffset_{0}; | 146 | uint64_t tokenWriteOffset_{0}; |
| @@ -186,38 +194,21 @@ private: | |||
| 186 | __aicore__ inline void EngramFetchGradArch35::WriteNbiChecked(uint64_t handle, GM_ADDR dst, GM_ADDR src, uint64_t len) | 194 | __aicore__ inline void EngramFetchGradArch35::WriteNbiChecked(uint64_t handle, GM_ADDR dst, GM_ADDR src, uint64_t len) |
| 187 | { | 195 | { |
| 188 | int32_t ret = hcomm_.WriteNbi(handle, dst, src, len); | 196 | int32_t ret = hcomm_.WriteNbi(handle, dst, src, len); |
| 189 | - if (ret != 0) { | 197 | + ascendc_assert(ret == 0, "WriteNbi failed, ret=%d, rankId=%u, aivId=%u", ret, rankId_, aivId_); |
| 190 | - for (uint32_t i = 0; i < COMM_RETRY_COUNT; i++) { | ||
| 191 | - ret = hcomm_.WriteNbi(handle, dst, src, len); | ||
| 192 | - if (ret == 0) { | ||
| 193 | - return; | ||
| 194 | - } | ||
| 195 | - } | ||
| 196 | - RUNTIME_ABORT("WriteNbi failed after %u retries, ret=%d, rankId=%u, aivId=%u", COMM_RETRY_COUNT, ret, rankId_, | ||
| 197 | - aivId_); | ||
| 198 | - } | ||
| 199 | } | 198 | } |
| 200 | 199 | ||
| 201 | __aicore__ inline void EngramFetchGradArch35::DrainChecked(uint64_t handle) | 200 | __aicore__ inline void EngramFetchGradArch35::DrainChecked(uint64_t handle) |
| 202 | { | 201 | { |
| 203 | int32_t ret = hcomm_.Drain(handle); | 202 | int32_t ret = hcomm_.Drain(handle); |
| 204 | - if (ret != 0) { | 203 | + ascendc_assert(ret == 0, "Drain failed, ret=%d, rankId=%u, aivId=%u", ret, rankId_, aivId_); |
| 205 | - for (uint32_t i = 0; i < COMM_RETRY_COUNT; i++) { | ||
| 206 | - ret = hcomm_.Drain(handle); | ||
| 207 | - if (ret == 0) { | ||
| 208 | - return; | ||
| 209 | - } | ||
| 210 | - } | ||
| 211 | - RUNTIME_ABORT("DrainChecked failed after %u retries, ret=%d, handle=%llu", COMM_RETRY_COUNT, ret, handle); | ||
| 212 | - } | ||
| 213 | } | 204 | } |
| 214 | 205 | ||
| 215 | -__aicore__ inline void EngramFetchGradArch35::TimeoutCheck(uint64_t startTime) | 206 | +__aicore__ inline void EngramFetchGradArch35::TimeoutCheck(uint64_t startTime, TimeoutSite site) |
| 216 | { | 207 | { |
| 217 | uint64_t nowUs = static_cast<uint64_t>(AscendC::GetSystemCycle()) / ENGRAM_GRAD_CYCLES_PER_US; | 208 | uint64_t nowUs = static_cast<uint64_t>(AscendC::GetSystemCycle()) / ENGRAM_GRAD_CYCLES_PER_US; |
| 218 | - if ((nowUs - startTime) >= ENGRAM_GRAD_TIMEOUT_US) { | 209 | + ascendc_assert((nowUs - startTime) < ENGRAM_GRAD_TIMEOUT_US, |
| 219 | - RUNTIME_ABORT("timeout, rankId=%u, aivId=%u, elapsed=%llu us", rankId_, aivId_, nowUs - startTime); | 210 | + "timeout, tag=%d, rankId=%u, aivId=%u, elapsed=%llu us\n", static_cast<int>(site), rankId_, aivId_, |
| 220 | - } | 211 | + nowUs - startTime); |
| 221 | } | 212 | } |
| 222 | 213 | ||
| 223 | __aicore__ inline GM_ADDR EngramFetchGradArch35::GetRemoteWinAddr(uint32_t dstRank, uint64_t offset) | 214 | __aicore__ inline GM_ADDR EngramFetchGradArch35::GetRemoteWinAddr(uint32_t dstRank, uint64_t offset) |
| @@ -281,9 +272,7 @@ __aicore__ inline void EngramFetchGradArch35::PrefetchCreditCounter(uint32_t dst | |||
| 281 | (static_cast<uint64_t>(rankId_) * sendersPerRank_ + senderIdx) * STATE_OFFSET; | 272 | (static_cast<uint64_t>(rankId_) * sendersPerRank_ + senderIdx) * STATE_OFFSET; |
| 282 | uint64_t handle = GetCommHandle(dstRank, senderIdx); | 273 | uint64_t handle = GetCommHandle(dstRank, senderIdx); |
| 283 | int32_t ret = hcomm_.ReadNbi(handle, scratchAddr, remoteCounterAddr, sizeof(int32_t)); | 274 | int32_t ret = hcomm_.ReadNbi(handle, scratchAddr, remoteCounterAddr, sizeof(int32_t)); |
| 284 | - if (ret != 0) { | 275 | + ascendc_assert(ret == 0, "CreditRead launch failed, ret=%d, rankId=%u, dstRank=%u", ret, rankId_, dstRank); |
| 285 | - RUNTIME_ABORT("CreditRead launch failed, ret=%d, rankId=%u, dstRank=%u", ret, rankId_, dstRank); | ||
| 286 | - } | ||
| 287 | creditReadInFlight_ = true; | 276 | creditReadInFlight_ = true; |
| 288 | } | 277 | } |
| 289 | 278 | ||
| @@ -300,7 +289,7 @@ __aicore__ inline int32_t EngramFetchGradArch35::CompleteCreditCounter(uint64_t | |||
| 300 | DataCopyPad(creditLocal, scratchGM, cpParams, cpPad); | 289 | DataCopyPad(creditLocal, scratchGM, cpParams, cpPad); |
| 301 | EngramFetchGradSyncFunc<HardEvent::MTE2_S>(); | 290 | EngramFetchGradSyncFunc<HardEvent::MTE2_S>(); |
| 302 | value = creditLocal.GetValue(0); | 291 | value = creditLocal.GetValue(0); |
| 303 | - TimeoutCheck(startTime); | 292 | + TimeoutCheck(startTime, TIMEOUT_CREDIT_READ_WAIT); |
| 304 | } | 293 | } |
| 305 | creditReadInFlight_ = false; | 294 | creditReadInFlight_ = false; |
| 306 | return value; | 295 | return value; |
| @@ -329,7 +318,7 @@ __aicore__ inline void EngramFetchGradArch35::WaitAllStatusFlags(GM_ADDR statusW | |||
| 329 | int32_t flagVal = slotGM.GetValue(0); | 318 | int32_t flagVal = slotGM.GetValue(0); |
| 330 | sumOfFlag += flagVal; | 319 | sumOfFlag += flagVal; |
| 331 | } | 320 | } |
| 332 | - TimeoutCheck(startTime); | 321 | + TimeoutCheck(startTime, TIMEOUT_STATUS_FLAG_WAIT); |
| 333 | } | 322 | } |
| 334 | } | 323 | } |
| 335 | 324 | ||
| @@ -442,6 +431,17 @@ __aicore__ inline static uint32_t AccumBufBytes(uint32_t hiddenDim) | |||
| 442 | return Ceil(hiddenDim * sizeof(float), UB_ALIGN) * UB_ALIGN * Mc2Kernel::ACCUM_BUF_COPIES; | 431 | return Ceil(hiddenDim * sizeof(float), UB_ALIGN) * UB_ALIGN * Mc2Kernel::ACCUM_BUF_COPIES; |
| 443 | } | 432 | } |
| 444 | 433 | ||
| 434 | +__aicore__ inline static uint32_t UniqueEntryBytes(uint32_t hiddenDim, int32_t outputDtype) | ||
| 435 | +{ | ||
| 436 | + uint32_t bytes = Mc2Kernel::FLUSH_CAST_HEAD_BYTES; | ||
| 437 | + if (outputDtype != Mc2Kernel::ENGRAM_DT_FLOAT) { | ||
| 438 | + uint32_t castHalfBytes = | ||
| 439 | + (hiddenDim * 2U + Mc2Kernel::UB_ALIGN - 1U) / Mc2Kernel::UB_ALIGN * Mc2Kernel::UB_ALIGN; | ||
| 440 | + bytes += 2U * castHalfBytes; | ||
| 441 | + } | ||
| 442 | + return bytes; | ||
| 443 | +} | ||
| 444 | + | ||
| 445 | __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR gradFetched, GM_ADDR permOut, | 445 | __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR gradFetched, GM_ADDR permOut, |
| 446 | GM_ADDR sendCountsOut, GM_ADDR recvCountsOut, | 446 | GM_ADDR sendCountsOut, GM_ADDR recvCountsOut, |
| 447 | GM_ADDR recvLocalEntryOut, GM_ADDR numRecvOut, GM_ADDR gradUniqueOut, | 447 | GM_ADDR recvLocalEntryOut, GM_ADDR numRecvOut, GM_ADDR gradUniqueOut, |
| @@ -469,9 +469,8 @@ __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR | |||
| 469 | numRanks_ = ctxPtr_->rankSize; | 469 | numRanks_ = ctxPtr_->rankSize; |
| 470 | // commContext 为设备侧外部数据,与 Host 侧 tiling 的 rankSize(sendCounts.dim0/8)互为独立来源, | 470 | // commContext 为设备侧外部数据,与 Host 侧 tiling 的 rankSize(sendCounts.dim0/8)互为独立来源, |
| 471 | // 必须一致性校验,否则 workspace 的 displs 区按 Host rankSize 规划而 Kernel 按 numRanks_ 写入会越界 | 471 | // 必须一致性校验,否则 workspace 的 displs 区按 Host rankSize 规划而 Kernel 按 numRanks_ 写入会越界 |
| 472 | - if (numRanks_ == 0U || numRanks_ > Mc2Kernel::MAX_QP_SIZE || numRanks_ != tilingData->rankSize) { | 472 | + ascendc_assert(numRanks_ != 0U && numRanks_ <= Mc2Kernel::MAX_QP_SIZE && numRanks_ == tilingData->rankSize, |
| 473 | - RUNTIME_ABORT("invalid rankSize: commContext=%u, tiling=%u", numRanks_, tilingData->rankSize); | 473 | + "invalid rankSize: commContext=%u, tiling=%u", numRanks_, tilingData->rankSize); |
| 474 | - } | ||
| 475 | channelsPerRank_ = ctxPtr_->channelsPerRank; | 474 | channelsPerRank_ = ctxPtr_->channelsPerRank; |
| 476 | if (channelsPerRank_ == 0) { | 475 | if (channelsPerRank_ == 0) { |
| 477 | channelsPerRank_ = 1; | 476 | channelsPerRank_ = 1; |
| @@ -516,10 +515,11 @@ __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR | |||
| 516 | if (numSendCores_ > halfBlocks) { | 515 | if (numSendCores_ > halfBlocks) { |
| 517 | numSendCores_ = halfBlocks; | 516 | numSendCores_ = halfBlocks; |
| 518 | } | 517 | } |
| 519 | - numRecvCores_ = totalBlocks_ - numSendCores_; | 518 | + |
| 520 | - if (numRecvCores_ < 1U) { | 519 | + numRecvCores_ = (totalBlocks_ > numSendCores_ + 1U) ? (totalBlocks_ - numSendCores_ - 1U) : 0U; |
| 520 | + if (numRecvCores_ == 0U) { | ||
| 521 | numRecvCores_ = 1U; | 521 | numRecvCores_ = 1U; |
| 522 | - numSendCores_ = totalBlocks_ - 1U; | 522 | + numSendCores_ = (totalBlocks_ > 1U) ? (totalBlocks_ - 1U) : 1U; |
| 523 | } | 523 | } |
| 524 | isSender_ = (aivId_ < numSendCores_) || (totalBlocks_ <= 1U); | 524 | isSender_ = (aivId_ < numSendCores_) || (totalBlocks_ <= 1U); |
| 525 | isReceiver_ = (aivId_ >= numSendCores_ && aivId_ < totalBlocks_ - 1U) || (totalBlocks_ <= 1U); | 525 | isReceiver_ = (aivId_ >= numSendCores_ && aivId_ < totalBlocks_ - 1U) || (totalBlocks_ <= 1U); |
| @@ -619,7 +619,9 @@ __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR | |||
| 619 | uint32_t maxByPong = MaxGradRowsPerPing(static_cast<uint32_t>(hiddenBytes_), gradSubBatch_); | 619 | uint32_t maxByPong = MaxGradRowsPerPing(static_cast<uint32_t>(hiddenBytes_), gradSubBatch_); |
| 620 | uint32_t castBufSize = (inputDtype_ != Mc2Kernel::ENGRAM_DT_FLOAT) ? CastBufBytes(hiddenDim_, maxByPong) : 0U; | 620 | uint32_t castBufSize = (inputDtype_ != Mc2Kernel::ENGRAM_DT_FLOAT) ? CastBufBytes(hiddenDim_, maxByPong) : 0U; |
| 621 | uint32_t accumBufSize = AccumBufBytes(hiddenDim_); | 621 | uint32_t accumBufSize = AccumBufBytes(hiddenDim_); |
| 622 | - uint32_t uniqueBufSize = Mc2Kernel::COMM_BUF_BYTES + castBufSize + accumBufSize; | 622 | + uniqueEntryBytes_ = UniqueEntryBytes(hiddenDim_, outputDtype_); |
| 623 | + // unique 阶段池峰值:grad 整缓冲 + entryBuf 收缩区 + cast + accum(与 Process 二次 InitBuffer 布局一致) | ||
| 624 | + uint32_t uniqueBufSize = Mc2Kernel::GRAD_BUF_BYTES + uniqueEntryBytes_ + castBufSize + accumBufSize; | ||
| 623 | uint32_t poolSize = sortUbSize; | 625 | uint32_t poolSize = sortUbSize; |
| 624 | if (uniqueBufSize > poolSize) { | 626 | if (uniqueBufSize > poolSize) { |
| 625 | poolSize = uniqueBufSize; | 627 | poolSize = uniqueBufSize; |
| @@ -627,9 +629,8 @@ __aicore__ inline void EngramFetchGradArch35::Init(GM_ADDR commContext, GM_ADDR | |||
| 627 | // UB 池预算自检:池 + 常驻四缓冲必须落在 SetLocalMemorySize 授权范围内,超限确定性失败 | 629 | // UB 池预算自检:池 + 常驻四缓冲必须落在 SetLocalMemorySize 授权范围内,超限确定性失败 |
| 628 | uint64_t permanentUsed = Mc2Kernel::HCOMM_INIT_SIZE + statusBufSize + tempBufSize + indicesBufSize; | 630 | uint64_t permanentUsed = Mc2Kernel::HCOMM_INIT_SIZE + statusBufSize + tempBufSize + indicesBufSize; |
| 629 | uint64_t budgetLeft = (ubSize_ > permanentUsed) ? (ubSize_ - permanentUsed) : 0U; | 631 | uint64_t budgetLeft = (ubSize_ > permanentUsed) ? (ubSize_ - permanentUsed) : 0U; |
| 630 | - if (static_cast<uint64_t>(poolSize) > budgetLeft) { | 632 | + ascendc_assert(static_cast<uint64_t>(poolSize) <= budgetLeft, |
| 631 | - RUNTIME_ABORT("UB pool overflow: pool=%u, permanent=%llu, ubSize=%llu", poolSize, permanentUsed, ubSize_); | 633 | + "UB pool overflow: pool=%u, permanent=%llu, ubSize=%llu", poolSize, permanentUsed, ubSize_); |
| 632 | - } | ||
| 633 | tpipe_->InitBufPool(sortPool_, poolSize); | 634 | tpipe_->InitBufPool(sortPool_, poolSize); |
| 634 | sortPool_.InitBuffer(entryBuf_, Mc2Kernel::ENTRY_BUF_BYTES); | 635 | sortPool_.InitBuffer(entryBuf_, Mc2Kernel::ENTRY_BUF_BYTES); |
| 635 | sortPool_.InitBuffer(gradBuf_, Mc2Kernel::GRAD_BUF_BYTES); | 636 | sortPool_.InitBuffer(gradBuf_, Mc2Kernel::GRAD_BUF_BYTES); |
| @@ -920,7 +921,7 @@ __aicore__ inline void EngramFetchGradArch35::SendGradRemote(uint32_t dstRank, u | |||
| 920 | localWriteCnt - static_cast<uint32_t>(remoteReadCnt) >= slotsPerSender) { | 921 | localWriteCnt - static_cast<uint32_t>(remoteReadCnt) >= slotsPerSender) { |
| 921 | PrefetchCreditCounter(dstRank, senderIdx); | 922 | PrefetchCreditCounter(dstRank, senderIdx); |
| 922 | remoteReadCnt = CompleteCreditCounter(startTime); | 923 | remoteReadCnt = CompleteCreditCounter(startTime); |
| 923 | - TimeoutCheck(startTime); | 924 | + TimeoutCheck(startTime, TIMEOUT_SEND_CREDIT_WAIT); |
| 924 | } | 925 | } |
| 925 | } | 926 | } |
| 926 | 927 | ||
| @@ -940,20 +941,8 @@ __aicore__ inline void EngramFetchGradArch35::SendGradRemote(uint32_t dstRank, u | |||
| 940 | (static_cast<uint64_t>(rankId_) * sendersPerRank_ + senderIdx) * STATE_OFFSET; | 941 | (static_cast<uint64_t>(rankId_) * sendersPerRank_ + senderIdx) * STATE_OFFSET; |
| 941 | int32_t ret = hcomm_.WriteWithNotifyNbi(handle, remoteSlotAddr, srcAddr, dataBytes, remoteCounterAddr, | 942 | int32_t ret = hcomm_.WriteWithNotifyNbi(handle, remoteSlotAddr, srcAddr, dataBytes, remoteCounterAddr, |
| 942 | static_cast<uint64_t>(localWriteCnt + 1)); | 943 | static_cast<uint64_t>(localWriteCnt + 1)); |
| 943 | - if (ret != 0) { | 944 | + ascendc_assert(ret == 0, "WriteWithNotifyNbi failed, ret=%d, tag=ExTok_data, rankId=%u, dstRank=%u", ret, |
| 944 | - for (uint32_t i = 0; i < COMM_RETRY_COUNT; i++) { | 945 | + rankId_, dstRank); |
| 945 | - ret = hcomm_.WriteWithNotifyNbi(handle, remoteSlotAddr, srcAddr, dataBytes, remoteCounterAddr, | ||
| 946 | - static_cast<uint64_t>(localWriteCnt + 1)); | ||
| 947 | - if (ret == 0) { | ||
| 948 | - break; | ||
| 949 | - } | ||
| 950 | - } | ||
| 951 | - if (ret != 0) { | ||
| 952 | - RUNTIME_ABORT( | ||
| 953 | - "WriteWithNotifyNbi failed after %u retries, ret=%d, tag=ExTok_data, rankId=%u, dstRank=%u", | ||
| 954 | - COMM_RETRY_COUNT, ret, rankId_, dstRank); | ||
| 955 | - } | ||
| 956 | - } | ||
| 957 | 946 | ||
| 958 | localWriteCnt++; | 947 | localWriteCnt++; |
| 959 | totalSent += chunkLen; | 948 | totalSent += chunkLen; |
| @@ -961,10 +950,9 @@ __aicore__ inline void EngramFetchGradArch35::SendGradRemote(uint32_t dstRank, u | |||
| 961 | RetireCreditCounter(); | 950 | RetireCreditCounter(); |
| 962 | 951 | ||
| 963 | // 单核 remote handle 数随 numRanks_/numSendCores_ 配置增长,必须守卫固定数组边界 | 952 | // 单核 remote handle 数随 numRanks_/numSendCores_ 配置增长,必须守卫固定数组边界 |
| 964 | - if (pendingHandleCount_ >= Mc2Kernel::MAX_PENDING_HANDLES) { | 953 | + ascendc_assert(pendingHandleCount_ < Mc2Kernel::MAX_PENDING_HANDLES, |
| 965 | - RUNTIME_ABORT("pendingHandles overflow: count=%u, max=%u, rankId=%u, dstRank=%u", pendingHandleCount_, | 954 | + "pendingHandles overflow: count=%u, max=%u, rankId=%u, dstRank=%u", pendingHandleCount_, |
| 966 | - Mc2Kernel::MAX_PENDING_HANDLES, rankId_, dstRank); | 955 | + Mc2Kernel::MAX_PENDING_HANDLES, rankId_, dstRank); |
| 967 | - } | ||
| 968 | pendingHandles_[pendingHandleCount_] = handle; | 956 | pendingHandles_[pendingHandleCount_] = handle; |
| 969 | pendingHandleCount_++; | 957 | pendingHandleCount_++; |
| 970 | } | 958 | } |
| @@ -984,7 +972,7 @@ __aicore__ inline void EngramFetchGradArch35::RecvGradFromPeers() | |||
| 984 | } | 972 | } |
| 985 | 973 | ||
| 986 | GM_ADDR localWinBase = (GM_ADDR)ctxPtr_->commBuffer[rankId_]; | 974 | GM_ADDR localWinBase = (GM_ADDR)ctxPtr_->commBuffer[rankId_]; |
| 987 | - uint32_t recvIdx = aivId_ - numSendCores_; | 975 | + uint32_t recvIdx = (aivId_ > numSendCores_) ? (aivId_ - numSendCores_) : 0U; |
| 988 | uint32_t totalWorkUnits = (numRanks_ - 1U) * sendersPerRank_; | 976 | uint32_t totalWorkUnits = (numRanks_ - 1U) * sendersPerRank_; |
| 989 | if (totalWorkUnits == 0U) { | 977 | if (totalWorkUnits == 0U) { |
| 990 | return; | 978 | return; |
| @@ -1034,7 +1022,7 @@ __aicore__ inline void EngramFetchGradArch35::RecvGradFromPeers() | |||
| 1034 | int32_t remoteWriteCnt = ReadLocalCounter(localWinBase, tokenWriteOffset_, srcRank, si); | 1022 | int32_t remoteWriteCnt = ReadLocalCounter(localWinBase, tokenWriteOffset_, srcRank, si); |
| 1035 | while (remoteWriteCnt <= 0 || static_cast<uint32_t>(remoteWriteCnt) <= localReadCnt) { | 1023 | while (remoteWriteCnt <= 0 || static_cast<uint32_t>(remoteWriteCnt) <= localReadCnt) { |
| 1036 | remoteWriteCnt = ReadLocalCounter(localWinBase, tokenWriteOffset_, srcRank, si); | 1024 | remoteWriteCnt = ReadLocalCounter(localWinBase, tokenWriteOffset_, srcRank, si); |
| 1037 | - TimeoutCheck(startTime); | 1025 | + TimeoutCheck(startTime, TIMEOUT_RECV_COUNTER_WAIT); |
| 1038 | } | 1026 | } |
| 1039 | 1027 | ||
| 1040 | uint32_t availSlots = static_cast<uint32_t>(remoteWriteCnt) - localReadCnt; | 1028 | uint32_t availSlots = static_cast<uint32_t>(remoteWriteCnt) - localReadCnt; |
| @@ -1303,7 +1291,7 @@ __aicore__ inline void EngramFetchGradArch35::Process() | |||
| 1303 | RunSort(numRecv); | 1291 | RunSort(numRecv); |
| 1304 | 1292 | ||
| 1305 | sortPool_.Reset(); | 1293 | sortPool_.Reset(); |
| 1306 | - sortPool_.InitBuffer(entryBuf_, Mc2Kernel::ENTRY_BUF_BYTES); | 1294 | + sortPool_.InitBuffer(entryBuf_, uniqueEntryBytes_); |
| 1307 | sortPool_.InitBuffer(gradBuf_, Mc2Kernel::GRAD_BUF_BYTES); | 1295 | sortPool_.InitBuffer(gradBuf_, Mc2Kernel::GRAD_BUF_BYTES); |
| 1308 | uint32_t maxByPong = MaxGradRowsPerPing(static_cast<uint32_t>(hiddenBytes_), gradSubBatch_); | 1296 | uint32_t maxByPong = MaxGradRowsPerPing(static_cast<uint32_t>(hiddenBytes_), gradSubBatch_); |
| 1309 | if (inputDtype_ != Mc2Kernel::ENGRAM_DT_FLOAT) { | 1297 | if (inputDtype_ != Mc2Kernel::ENGRAM_DT_FLOAT) { |
| @@ -1313,6 +1301,7 @@ __aicore__ inline void EngramFetchGradArch35::Process() | |||
| 1313 | uniqueScatter_.SetCastBuf(castFp32Buf_); | 1301 | uniqueScatter_.SetCastBuf(castFp32Buf_); |
| 1314 | uniqueScatter_.SetAccumBuf(accumBuf_); | 1302 | uniqueScatter_.SetAccumBuf(accumBuf_); |
| 1315 | uniqueScatter_.SetGradSubBatch(gradSubBatch_); | 1303 | uniqueScatter_.SetGradSubBatch(gradSubBatch_); |
| 1304 | + uniqueScatter_.SetEntryBufBytes(uniqueEntryBytes_); | ||
| 1316 | 1305 | ||
| 1317 | uniqueScatter_.CountUniquesParallel(numRecv, recvLocalEntryOutGM_, coreStartGM_, segCountGM_); | 1306 | uniqueScatter_.CountUniquesParallel(numRecv, recvLocalEntryOutGM_, coreStartGM_, segCountGM_); |
| 1318 | uniqueScatter_.ZeroGradUnique(numRecv, gradUniqueOutGM_); | 1307 | uniqueScatter_.ZeroGradUnique(numRecv, gradUniqueOutGM_); |
| @@ -669,11 +669,11 @@ __aicore__ inline void EngramFetchGradSort::ProcessHist(uint32_t byteRound, uint | |||
| 669 | AscendC::Duplicate(coreSumAccum, (int32_t)0, HISTOGRAM_BINS); | 669 | AscendC::Duplicate(coreSumAccum, (int32_t)0, HISTOGRAM_BINS); |
| 670 | 670 | ||
| 671 | for (uint32_t batch = 0; batch < batchCount; batch++) { | 671 | for (uint32_t batch = 0; batch < batchCount; batch++) { |
| 672 | - uint32_t firstTile = batch * coreCount_; | ||
| 673 | - uint32_t batchCores = MinU32(tileCount_ - firstTile, coreCount_); | ||
| 674 | uint32_t coreId = AscendC::GetBlockIdx(); | 672 | uint32_t coreId = AscendC::GetBlockIdx(); |
| 675 | - if (coreId < batchCores) { | 673 | + uint32_t coreStart = coreId * batchCount; |
| 676 | - uint32_t tileId = firstTile + coreId; | 674 | + uint32_t myTiles = (coreStart < tileCount_) ? (MinU32(coreStart + batchCount, tileCount_) - coreStart) : 0U; |
| 675 | + if (batch < myTiles) { | ||
| 676 | + uint32_t tileId = coreStart + batch; | ||
| 677 | uint32_t offset = tileId * tileElements_; | 677 | uint32_t offset = tileId * tileElements_; |
| 678 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); | 678 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); |
| 679 | 679 | ||
| @@ -738,10 +738,12 @@ __aicore__ inline void EngramFetchGradSort::ProcessScatter( | |||
| 738 | event_t evtMte2Mte3 = static_cast<event_t>(pipe.FetchEventID(AscendC::HardEvent::MTE2_MTE3)); | 738 | event_t evtMte2Mte3 = static_cast<event_t>(pipe.FetchEventID(AscendC::HardEvent::MTE2_MTE3)); |
| 739 | 739 | ||
| 740 | for (uint32_t batch = 0; batch < batchCount; batch++) { | 740 | for (uint32_t batch = 0; batch < batchCount; batch++) { |
| 741 | - uint32_t firstTile = batch * coreCount_; | ||
| 742 | - uint32_t batchCores = MinU32(tileCount_ - firstTile, coreCount_); | ||
| 743 | uint32_t coreId = AscendC::GetBlockIdx(); | 741 | uint32_t coreId = AscendC::GetBlockIdx(); |
| 744 | 742 | ||
| 743 | + uint32_t coreStart = coreId * batchCount; | ||
| 744 | + uint32_t myTiles = (coreStart < tileCount_) ? (MinU32(coreStart + batchCount, tileCount_) - coreStart) : 0U; | ||
| 745 | + bool tileActive = (batch < myTiles); | ||
| 746 | + | ||
| 745 | AscendC::LocalTensor<int32_t> valueLocal = ValsUb(); | 747 | AscendC::LocalTensor<int32_t> valueLocal = ValsUb(); |
| 746 | AscendC::LocalTensor<int32_t> indexLocal = IdxsUb(); | 748 | AscendC::LocalTensor<int32_t> indexLocal = IdxsUb(); |
| 747 | AscendC::LocalTensor<uint8_t> keyLocal = keyBuffer_.Get<uint8_t>(); | 749 | AscendC::LocalTensor<uint8_t> keyLocal = keyBuffer_.Get<uint8_t>(); |
| @@ -750,8 +752,8 @@ __aicore__ inline void EngramFetchGradSort::ProcessScatter( | |||
| 750 | AscendC::LocalTensor<int32_t> prefixLocal = prefixBuffer_.Get<int32_t>(); | 752 | AscendC::LocalTensor<int32_t> prefixLocal = prefixBuffer_.Get<int32_t>(); |
| 751 | AscendC::LocalTensor<int32_t> offsetLocal = histogramBuffer_.Get<int32_t>(); | 753 | AscendC::LocalTensor<int32_t> offsetLocal = histogramBuffer_.Get<int32_t>(); |
| 752 | 754 | ||
| 753 | - if (coreId < batchCores) { | 755 | + if (tileActive) { |
| 754 | - uint32_t tileId = firstTile + coreId; | 756 | + uint32_t tileId = coreStart + batch; |
| 755 | uint32_t offset = tileId * tileElements_; | 757 | uint32_t offset = tileId * tileElements_; |
| 756 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); | 758 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); |
| 757 | 759 | ||
| @@ -771,8 +773,8 @@ __aicore__ inline void EngramFetchGradSort::ProcessScatter( | |||
| 771 | // entry (same-core MTE3->MTE3_S->MTE2 ordering from the prefix phase), and the | 773 | // entry (same-core MTE3->MTE3_S->MTE2 ordering from the prefix phase), and the |
| 772 | // round input data was ordered by the previous round's trailing SyncAll. | 774 | // round input data was ordered by the previous round's trailing SyncAll. |
| 773 | 775 | ||
| 774 | - if (coreId < batchCores) { | 776 | + if (tileActive) { |
| 775 | - uint32_t tileId = firstTile + coreId; | 777 | + uint32_t tileId = coreStart + batch; |
| 776 | uint32_t offset = tileId * tileElements_; | 778 | uint32_t offset = tileId * tileElements_; |
| 777 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); | 779 | uint32_t tileLen = MinU32(elementCount_ - offset, tileElements_); |
| 778 | 780 | ||
| @@ -71,6 +71,10 @@ public: | |||
| 71 | { | 71 | { |
| 72 | gradSubBatch_ = batch; | 72 | gradSubBatch_ = batch; |
| 73 | } | 73 | } |
| 74 | + __aicore__ inline void SetEntryBufBytes(uint32_t bytes) | ||
| 75 | + { | ||
| 76 | + entryBufBytes_ = bytes; | ||
| 77 | + } | ||
| 74 | 78 | ||
| 75 | // entryBuf_ int32 slot layout: one ENTRY_BATCH_CAP-sized slot per array. | 79 | // entryBuf_ int32 slot layout: one ENTRY_BATCH_CAP-sized slot per array. |
| 76 | __aicore__ inline AscendC::LocalTensor<int32_t> CompUb() | 80 | __aicore__ inline AscendC::LocalTensor<int32_t> CompUb() |
| @@ -142,6 +146,7 @@ private: | |||
| 142 | AscendC::TBuf<> *castBuf_{nullptr}; | 146 | AscendC::TBuf<> *castBuf_{nullptr}; |
| 143 | AscendC::TBuf<> *accumBuf_{nullptr}; | 147 | AscendC::TBuf<> *accumBuf_{nullptr}; |
| 144 | uint32_t gradSubBatch_{Mc2Kernel::GRAD_SUB_BATCH}; | 148 | uint32_t gradSubBatch_{Mc2Kernel::GRAD_SUB_BATCH}; |
| 149 | + uint32_t entryBufBytes_{Mc2Kernel::ENTRY_BUF_BYTES}; | ||
| 145 | 150 | ||
| 146 | int32_t myPreCoreOffset_{0}; | 151 | int32_t myPreCoreOffset_{0}; |
| 147 | int32_t mySegCount_{0}; | 152 | int32_t mySegCount_{0}; |
| @@ -468,10 +473,9 @@ __aicore__ inline void EngramFetchGradUnique::FlushAccum(GM_ADDR gradUniqueOutGM | |||
| 468 | Mc2Kernel::UB_ALIGN * Mc2Kernel::UB_ALIGN; | 473 | Mc2Kernel::UB_ALIGN * Mc2Kernel::UB_ALIGN; |
| 469 | // 双缓冲借用区必须完整落在 entryBuf_ 尾部内(Host 侧已按 hiddenDim 上界拒绝超限 shape,此处兜底) | 474 | // 双缓冲借用区必须完整落在 entryBuf_ 尾部内(Host 侧已按 hiddenDim 上界拒绝超限 shape,此处兜底) |
| 470 | uint32_t castTailBytes = flushCastOffset + 2U * castHalfBytes; | 475 | uint32_t castTailBytes = flushCastOffset + 2U * castHalfBytes; |
| 471 | - if (castTailBytes > Mc2Kernel::ENTRY_BUF_BYTES) { | 476 | + ascendc_assert(castTailBytes <= entryBufBytes_, |
| 472 | - RUNTIME_ABORT("FlushAccum cast staging overflow: need %u bytes, entryBuf=%u bytes, hiddenDim=%u", | 477 | + "FlushAccum cast staging overflow: need %u bytes, entryBuf=%u bytes, hiddenDim=%u", |
| 473 | - castTailBytes, Mc2Kernel::ENTRY_BUF_BYTES, static_cast<uint32_t>(hiddenDim_)); | 478 | + castTailBytes, entryBufBytes_, static_cast<uint32_t>(hiddenDim_)); |
| 474 | - } | ||
| 475 | AscendC::LocalTensor<uint8_t> flushCastBuf = entryRaw[flushCastOffset + bufIdx * castHalfBytes]; | 479 | AscendC::LocalTensor<uint8_t> flushCastBuf = entryRaw[flushCastOffset + bufIdx * castHalfBytes]; |
| 476 | uint32_t castCount = static_cast<uint32_t>(hiddenDim_); | 480 | uint32_t castCount = static_cast<uint32_t>(hiddenDim_); |
| 477 | 481 | ||
| @@ -19,17 +19,6 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | -// 确定性失败兜底:用于不可恢复的容量/一致性校验(Kernel 侧无返回值路径的最后手段) | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - do { \ | ||
| 26 | - ascendc_assert(false, fmt, ##__VA_ARGS__); \ | ||
| 27 | - while (true) { \ | ||
| 28 | - (void)AscendC::GetSystemCycle(); \ | ||
| 29 | - } \ | ||
| 30 | - } while (0) | ||
| 31 | - | ||
| 32 | - | ||
| 33 | namespace Mc2Kernel { | 22 | namespace Mc2Kernel { |
| 34 | // 共享布局常量已收敛至 engram_fetch_grad_tiling_data.h(单一权威定义),此处仅保留 Kernel 侧私有常量 | 23 | // 共享布局常量已收敛至 engram_fetch_grad_tiling_data.h(单一权威定义),此处仅保留 Kernel 侧私有常量 |
| 35 | constexpr int32_t BITS_PER_BYTE = 8; | 24 | constexpr int32_t BITS_PER_BYTE = 8; |