已合并
engram反向超时DFX #11577
engram反向超时DFX #11577
已合并
luozhonglin创建于 23 天前
共 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+ 
492static ge::graphStatus SetPlatformInfo(gert::TilingContext *context, EngramFetchGradTilingData &tilingData)536static 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 
56constexpr uint32_t ENGRAM_GRAD_TIMEOUT_US = 60U * 1000U * 1000U;56constexpr uint32_t ENGRAM_GRAD_TIMEOUT_US = 60U * 1000U * 1000U;
57constexpr uint32_t ENGRAM_GRAD_CYCLES_PER_US = 1000U;57constexpr uint32_t ENGRAM_GRAD_CYCLES_PER_US = 1000U;
58-constexpr uint32_t COMM_RETRY_COUNT = 3U;
59constexpr int32_t GRAD_CREDIT_READ_SENTINEL = -1;58constexpr 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+ 
61class EngramFetchGradArch35 {68class EngramFetchGradArch35 {
62public:69public:
63 __aicore__ inline EngramFetchGradArch35() = default;70 __aicore__ inline EngramFetchGradArch35() = default;
@@ -72,7 +79,7 @@ public:
72private:79private:
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 the773 // 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#include "engram_fetch_grad_tiling_data.h"20#include "engram_fetch_grad_tiling_data.h"
21 21 
22-// 确定性失败兜底:用于不可恢复的容量/一致性校验(Kernel 侧无返回值路径的最后手段)
23-#ifndef RUNTIME_ABORT
24-#define RUNTIME_ABORT(fmt, ...) \
25- do { \
26- ascendc_assert(false, fmt, ##__VA_ARGS__); \
27- while (true) { \
28- (void)AscendC::GetSystemCycle(); \
29- } \
30- } while (0)
31-#endif
32- 
33namespace Mc2Kernel {22namespace Mc2Kernel {
34// 共享布局常量已收敛至 engram_fetch_grad_tiling_data.h(单一权威定义),此处仅保留 Kernel 侧私有常量23// 共享布局常量已收敛至 engram_fetch_grad_tiling_data.h(单一权威定义),此处仅保留 Kernel 侧私有常量
35constexpr int32_t BITS_PER_BYTE = 8;24constexpr int32_t BITS_PER_BYTE = 8;