LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_send_receive - coll_batch_send_recv_group_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 69.0 % 478 330
Test Date: 2026-08-18 17:47:01 Functions: 85.7 % 28 24

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "coll_batch_send_recv_group_executor.h"
      12              : 
      13              : namespace hccl {
      14              : 
      15           64 : CollBatchSendRecvGroupExecutor::CollBatchSendRecvGroupExecutor(
      16           64 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      17           64 :     : CollBatchSendRecvExecutor(dispatcher, topoMatcher)
      18           64 : {}
      19              : 
      20            2 : void CollBatchSendRecvGroupExecutor::CalcIsSuPodAsym(bool isA2MultiModule)
      21              : {
      22              :     // isSuPodAsym 表示A2A3卡数不一致场景或者A3多超节点server数不同场景
      23              :     // 参考alltoallv_direct_fullmesh executor中的判断逻辑
      24            2 :     if (topoAttr_.superPodNum > 1) {
      25              :         isSuPodAsym_
      26            0 :             = (static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode)
      27            0 :                || static_cast<bool>(topoAttr_.multiSuperPodDiffServerNumMode));
      28              :     } else {
      29            6 :         isSuPodAsym_ = (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule)
      30            4 :                        && static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode);
      31              :     }
      32            2 :     HCCL_INFO(
      33              :         "[CalcIsSuPodAsym] superPodNum[%u] isA2MultiModule[%d] isSuPodAsym[%d]", topoAttr_.superPodNum, isA2MultiModule,
      34              :         isSuPodAsym_);
      35            2 : }
      36              : 
      37            2 : HcclResult CollBatchSendRecvGroupExecutor::CalcPingPongHalfSize()
      38              : {
      39            2 :     u32 pingPongSliceNum = GROUP_MAX_CONCURRENT * 2;
      40            2 :     u32 alignSize = HCCL_MIN_SLICE_ALIGN_910B;
      41            2 :     bufferSliceSize_ = algResResp_->cclInputMem.size() / alignSize / pingPongSliceNum * alignSize;
      42              :     // RDMA单流slot大小 = CCLOut / 2:A半区(send, 单流单slot)、B半区(recv, 单流单slot)
      43            2 :     rdmaDataBlockSize_ = algResResp_->cclOutputMem.size() / alignSize / RDMA_CCLOUT_HALF_NUM * alignSize;
      44            2 :     HCCL_INFO(
      45              :         "[CollBatchSendRecvGroupExecutor][CalcPingPongHalfSize] pingPong halfSize[%llu] rdmaDataBlockSize_[%llu]",
      46              :         bufferSliceSize_, rdmaDataBlockSize_);
      47            2 :     return HCCL_SUCCESS;
      48              : }
      49              : 
      50            1 : HcclResult CollBatchSendRecvGroupExecutor::OrganizeSendItemByStream()
      51              : {
      52            1 :     sendQueueBySendstream_.resize(sendStreamNum_);
      53            1 :     HCCL_INFO("[OrganizeSendItemByStream] sendStreamNum_[%u]", sendStreamNum_);
      54            3 :     while (!sendDeque_.empty()) {
      55            2 :         HcclSendRecvItem* curr = sendDeque_.front();
      56            2 :         CHK_PTR_NULL(curr);
      57            2 :         sendQueueBySendstream_[curr->remoteRank % sendStreamNum_].push_back(curr);
      58            2 :         sendDeque_.pop_front();
      59              :     }
      60            1 :     HCCL_INFO("OrganizeSendItemByStream Done!");
      61            1 :     return HCCL_SUCCESS;
      62              : }
      63              : 
      64            2 : HcclResult CollBatchSendRecvGroupExecutor::OrganizeRecvItemByStream()
      65              : {
      66            2 :     recvQueueByRecvstream_.resize(recvStreamNum_);
      67            2 :     HCCL_INFO("[OrganizeRecvItemByStream] recvStreamNum_[%u]", recvStreamNum_);
      68            4 :     while (!recvDeque_.empty()) {
      69            2 :         HcclSendRecvItem* curr = recvDeque_.front();
      70            2 :         CHK_PTR_NULL(curr);
      71            2 :         recvQueueByRecvstream_[curr->remoteRank % recvStreamNum_].push_back(curr);
      72            2 :         recvDeque_.pop_front();
      73              :     }
      74            2 :     HCCL_INFO("OrganizeRecvItemByStream Done!");
      75            2 :     return HCCL_SUCCESS;
      76              : }
      77              : 
      78            2 : HcclResult CollBatchSendRecvGroupExecutor::CalcPodRange()
      79              : {
      80              :     // Determine pod (server/supernode) membership — same approach as alltoallv_direct_fullmesh
      81            2 :     u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
      82            2 :     u32 rankIdxInPod = INVALID_VALUE_RANKID;
      83            2 :     bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B && !topoAttr_.isSingleMeshAggregation;
      84            2 :     if (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) {
      85            1 :         CHK_RET(topoMatcher_->GetLocalServerRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
      86              :     } else {
      87            1 :         CHK_RET(topoMatcher_->GetLocalSuperPodRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
      88              :     }
      89            2 :     podStartRank_ = topoAttr_.userRank - rankIdxInPod;
      90            2 :     podEndRank_ = podStartRank_ + devNumInlocalPod - 1;
      91            2 :     devNumInlocalPod_ = devNumInlocalPod;
      92            2 :     CalcIsSuPodAsym(isA2MultiModule);
      93            2 :     HCCL_INFO(
      94              :         "[CalcPodRange] userRank[%u] pod[%u-%u] devNumInlocalPod[%u] rankIdxInPod[%u] isSuPodAsym[%d]",
      95              :         topoAttr_.userRank, podStartRank_, podEndRank_, devNumInlocalPod, rankIdxInPod, isSuPodAsym_);
      96            2 :     return HCCL_SUCCESS;
      97              : }
      98              : 
      99           35 : bool CollBatchSendRecvGroupExecutor::IsRemoteRankRdma(u32 remoteRank) const
     100              : {
     101              :     // pod内为SDMA,跨pod为RDMA
     102           35 :     return !(remoteRank >= podStartRank_ && remoteRank <= podEndRank_);
     103              : }
     104              : 
     105            0 : HcclResult CollBatchSendRecvGroupExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
     106              : {
     107            0 :     CHK_RET(CollBatchSendRecvExecutor::CalcCommInfo(opTransport));
     108              : 
     109            0 :     LevelNSubCommTransport& commTransport = opTransport[COMM_COMBINE_ORDER];
     110            0 :     for (u32 subCommIndex = 0; subCommIndex < commTransport.size(); subCommIndex++) {
     111            0 :         for (auto& transportRequest : commTransport[subCommIndex].transportRequests) {
     112            0 :             transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
     113              :         }
     114              :     }
     115            0 :     return HCCL_SUCCESS;
     116              : }
     117              : 
     118            1 : HcclResult CollBatchSendRecvGroupExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algResource)
     119              : {
     120            1 :     HcclUs startut = TIME_NOW();
     121            1 :     HCCL_CONFIG_INFO(HCCL_ALG, "[CollBatchSendRecvGroupExecutor] groupsendrecv starts.");
     122              : 
     123            1 :     sendStreamNum_ = GROUP_MAX_CONCURRENT;
     124            1 :     recvStreamNum_ = GROUP_MAX_CONCURRENT;
     125            1 :     HCCL_INFO("[Orchestrate] sendStreamNum_[%u], recvStreamNum_[%u]", sendStreamNum_, recvStreamNum_);
     126              : 
     127            1 :     algResResp_ = &algResource;
     128            1 :     CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_SIZE_TWO));
     129            1 :     CHK_RET(GetPairWiseList(param.BatchSendRecvDataDes.sendRecvItemsPtr, param.BatchSendRecvDataDes.itemNum));
     130            0 :     CHK_RET(ProcessSelfSendRecvTasks(param.stream));
     131            0 :     if (topoAttr_.userRankSize == 1) {
     132            0 :         HCCL_INFO(
     133              :             "tag[%s] BatchSendRecvGroup Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
     134              :             DURATION_US(TIME_NOW() - startut));
     135            0 :         return HCCL_SUCCESS;
     136              :     }
     137            0 :     CHK_RET(CalcPodRange());
     138            0 :     CHK_RET(CalcPingPongHalfSize()); // ping-pong double buffering + RDMA slot
     139            0 :     CHK_RET(OrganizeSendItemByStream());
     140            0 :     CHK_RET(OrganizeRecvItemByStream());
     141              : 
     142            0 :     CHK_RET(CalcSendSlices());
     143            0 :     CHK_RET(CalcRecvSlices());
     144              : 
     145            0 :     CHK_RET(RunLoop(param));
     146              : 
     147            0 :     HCCL_INFO(
     148              :         "tag[%s] BatchSendRecvGroup Executor orchestrate success, take time [%lld]us.", param.tag.c_str(),
     149              :         DURATION_US(TIME_NOW() - startut));
     150            0 :     return HCCL_SUCCESS;
     151              : }
     152              : 
     153            4 : HcclResult CollBatchSendRecvGroupExecutor::CalcStreamTaskStatus(u32& nonEmptySendStream, u32& nonEmptyRecvStream)
     154              : {
     155            4 :     nonEmptySendStream = 0;
     156            4 :     nonEmptyRecvStream = 0;
     157              :     // 记录各从流是否有任务,供头尾同步只唤醒有任务的从流(循环中不再更新)。
     158            4 :     sendStreamHasTask_.assign(sendStreamNum_, false);
     159            4 :     recvStreamHasTask_.assign(recvStreamNum_, false);
     160           19 :     for (u32 i = 0; i < sendStreamNum_; i++) {
     161           15 :         if (!sendDataSlicesBySendStream_[i].empty()) {
     162            5 :             nonEmptySendStream++;
     163            5 :             sendStreamHasTask_[i] = true;
     164              :         }
     165              :     }
     166           18 :     for (u32 i = 0; i < recvStreamNum_; i++) {
     167           14 :         if (!recvDataSlicesByRecvStream_[i].empty()) {
     168            2 :             nonEmptyRecvStream++;
     169            2 :             recvStreamHasTask_[i] = true;
     170              :         }
     171              :     }
     172            4 :     rdmaSendHasTask_ = !rdmaSendSlices_.empty();
     173            4 :     rdmaRecvHasTask_ = !rdmaRecvSlices_.empty();
     174            4 :     return HCCL_SUCCESS;
     175              : }
     176              : 
     177            1 : HcclResult CollBatchSendRecvGroupExecutor::MainPostSubWait(Stream& mainStream)
     178              : {
     179              :     // 主流只通知有任务的从流开始
     180            3 :     for (u32 i = 0; i < sendStreamNum_; i++) {
     181            2 :         if (!sendStreamHasTask_[i]) {
     182            2 :             continue;
     183              :         }
     184            0 :         CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[i], PROF_STAGE_0));
     185            0 :         CHK_RET(
     186              :             LocalNotify::Wait(algResResp_->slaveStreams[i], dispatcher_, algResResp_->notifiesAux[i], PROF_STAGE_0));
     187            0 :         HCCL_DEBUG("MainPost, Send[%u] Wait", i);
     188              :     }
     189              : 
     190            3 :     for (u32 i = 0; i < recvStreamNum_; i++) {
     191            2 :         if (!recvStreamHasTask_[i]) {
     192            2 :             continue;
     193              :         }
     194            0 :         CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[i + sendStreamNum_], PROF_STAGE_0));
     195            0 :         CHK_RET(LocalNotify::Wait(
     196              :             algResResp_->slaveStreams[i + sendStreamNum_], dispatcher_, algResResp_->notifiesAux[i + sendStreamNum_],
     197              :             PROF_STAGE_0));
     198            0 :         HCCL_DEBUG("MainPost, Recv[%u] Wait", i);
     199              :     }
     200              : 
     201              :     // RDMA专用从流(send/recv各一条)
     202            1 :     u32 rdmaSendIdx = RdmaSendStreamIdx();
     203            1 :     u32 rdmaRecvIdx = RdmaRecvStreamIdx();
     204            1 :     if (rdmaSendHasTask_) {
     205            0 :         CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[rdmaSendIdx], PROF_STAGE_0));
     206            0 :         CHK_RET(LocalNotify::Wait(
     207              :             algResResp_->slaveStreams[rdmaSendIdx], dispatcher_, algResResp_->notifiesAux[rdmaSendIdx], PROF_STAGE_0));
     208            0 :         HCCL_DEBUG("MainPost, RdmaSend Wait");
     209              :     }
     210            1 :     if (rdmaRecvHasTask_) {
     211            0 :         CHK_RET(LocalNotify::Post(mainStream, dispatcher_, algResResp_->notifiesAux[rdmaRecvIdx], PROF_STAGE_0));
     212            0 :         CHK_RET(LocalNotify::Wait(
     213              :             algResResp_->slaveStreams[rdmaRecvIdx], dispatcher_, algResResp_->notifiesAux[rdmaRecvIdx], PROF_STAGE_0));
     214            0 :         HCCL_DEBUG("MainPost, RdmaRecv Wait");
     215              :     }
     216              : 
     217            1 :     return HCCL_SUCCESS;
     218              : }
     219              : 
     220            1 : HcclResult CollBatchSendRecvGroupExecutor::MainWaitSubPost(Stream& mainStream)
     221              : {
     222              :     // 最后主流只等待有任务的从流结束
     223            3 :     for (u32 i = 0; i < sendStreamNum_; i++) {
     224            2 :         if (!sendStreamHasTask_[i]) {
     225            2 :             continue;
     226              :         }
     227            0 :         CHK_RET(
     228              :             LocalNotify::Post(algResResp_->slaveStreams[i], dispatcher_, algResResp_->notifiesMain[i], PROF_STAGE_0));
     229            0 :         CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[i], PROF_STAGE_0));
     230            0 :         HCCL_DEBUG("MainWait, Send[%u] Post", i);
     231              :     }
     232              : 
     233            3 :     for (u32 i = 0; i < recvStreamNum_; i++) {
     234            2 :         if (!recvStreamHasTask_[i]) {
     235            2 :             continue;
     236              :         }
     237            0 :         CHK_RET(LocalNotify::Post(
     238              :             algResResp_->slaveStreams[i + sendStreamNum_], dispatcher_, algResResp_->notifiesMain[i + sendStreamNum_],
     239              :             PROF_STAGE_0));
     240            0 :         CHK_RET(
     241              :             LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[i + sendStreamNum_], PROF_STAGE_0));
     242            0 :         HCCL_DEBUG("MainWait, Recv[%u] Post", i);
     243              :     }
     244              : 
     245              :     // RDMA专用从流(send/recv各一条)
     246            1 :     u32 rdmaSendIdx = RdmaSendStreamIdx();
     247            1 :     u32 rdmaRecvIdx = RdmaRecvStreamIdx();
     248            1 :     if (rdmaSendHasTask_) {
     249            0 :         CHK_RET(LocalNotify::Post(
     250              :             algResResp_->slaveStreams[rdmaSendIdx], dispatcher_, algResResp_->notifiesMain[rdmaSendIdx], PROF_STAGE_0));
     251            0 :         CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[rdmaSendIdx], PROF_STAGE_0));
     252            0 :         HCCL_DEBUG("MainWait, RdmaSend Post");
     253              :     }
     254            1 :     if (rdmaRecvHasTask_) {
     255            0 :         CHK_RET(LocalNotify::Post(
     256              :             algResResp_->slaveStreams[rdmaRecvIdx], dispatcher_, algResResp_->notifiesMain[rdmaRecvIdx], PROF_STAGE_0));
     257            0 :         CHK_RET(LocalNotify::Wait(mainStream, dispatcher_, algResResp_->notifiesMain[rdmaRecvIdx], PROF_STAGE_0));
     258            0 :         HCCL_DEBUG("MainWait, RdmaRecv Post");
     259              :     }
     260              : 
     261            1 :     return HCCL_SUCCESS;
     262              : }
     263              : 
     264              : HcclResult
     265            2 : CollBatchSendRecvGroupExecutor::ProcessPreloadedSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream)
     266              : {
     267            2 :     u32 curPhase = sendCurPhase_[streamIdx];
     268            2 :     u32 loadedRank = sendLoadedRemoteRank_[streamIdx];
     269            2 :     u64 loadedSize = sendLoadedSize_[streamIdx];
     270            2 :     u64 kernelOffset = bufferSliceSize_ * (streamIdx * 2 + curPhase);
     271            2 :     u64 phaseOffset = bufferSliceSize_ * curPhase;
     272              : 
     273            2 :     HCCL_INFO(
     274              :         "[RunTasks] SendStream[%u](loaded) phase[%u] offset[%llu] size[%llu] rank[%u]", streamIdx, curPhase,
     275              :         kernelOffset, loadedSize, loadedRank);
     276              : 
     277              :     // Step A.1: Record — notify remote data ready (TxPrepare + TxData)
     278            2 :     LINK sendTargetLink;
     279            2 :     CHK_RET(GetSendTargetLink(loadedRank, sendTargetLink));
     280            2 :     DeviceMem inCommMem = algResResp_->cclInputMem.range(kernelOffset, loadedSize);
     281            2 :     CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[streamIdx]));
     282            2 :     CHK_RET(sendTargetLink->TxData(
     283              :         UserMemType::OUTPUT_MEM, phaseOffset, inCommMem.ptr(), loadedSize, algResResp_->slaveStreams[streamIdx]));
     284              : 
     285            2 :     sendLoadedSize_[streamIdx] = 0;
     286            2 :     pendingSendCount--;
     287              : 
     288              :     // Step A.2: D2D next slice to OTHER half (only if same rank)
     289            2 :     if (!sendDataSlicesBySendStream_[streamIdx].empty()) {
     290            1 :         SendRecvSlice& nextSlice = sendDataSlicesBySendStream_[streamIdx].front();
     291            1 :         if (nextSlice.remoteRank == loadedRank) {
     292            1 :             u32 nextPhase = 1 - curPhase;
     293            1 :             u64 d2dOffset = bufferSliceSize_ * (streamIdx * 2 + nextPhase);
     294            1 :             DeviceMem d2dCommMem = algResResp_->cclInputMem.range(d2dOffset, nextSlice.size);
     295            1 :             DeviceMem inMem(nextSlice.addr, nextSlice.size);
     296            1 :             HCCL_INFO(
     297              :                 "[RunTasks] SendStream[%u] load next to phase[%u] offset[%llu] size[%llu]", streamIdx, nextPhase,
     298              :                 d2dOffset, nextSlice.size);
     299            1 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, d2dCommMem, inMem, algResResp_->slaveStreams[streamIdx]));
     300            1 :             sendLoadedSize_[streamIdx] = nextSlice.size;
     301            1 :             sendLoadedRemoteRank_[streamIdx] = nextSlice.remoteRank;
     302            1 :             sendCurPhase_[streamIdx] = nextPhase;
     303            1 :             sendDataSlicesBySendStream_[streamIdx].pop_front();
     304            1 :             pendingSendCount++;
     305            1 :             if (sendDataSlicesBySendStream_[streamIdx].empty()) {
     306            1 :                 nonEmptySendStream--;
     307            1 :                 HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
     308              :             }
     309            1 :         }
     310              :     }
     311              : 
     312              :     // Step A.3: Wait for remote ack (TxDone)
     313            2 :     CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[streamIdx]));
     314            2 :     return HCCL_SUCCESS;
     315            2 : }
     316              : 
     317              : HcclResult
     318            2 : CollBatchSendRecvGroupExecutor::ProcessNewRankSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream)
     319              : {
     320            2 :     SendRecvSlice& firstSlice = sendDataSlicesBySendStream_[streamIdx].front();
     321            2 :     u32 newRank = firstSlice.remoteRank;
     322              : 
     323              :     // Step B.1: D2D first slice to half A (phase 0)
     324            2 :     u64 offsetA = bufferSliceSize_ * (streamIdx * 2 + 0);
     325            2 :     DeviceMem commMemA = algResResp_->cclInputMem.range(offsetA, firstSlice.size);
     326            2 :     DeviceMem inMem(firstSlice.addr, firstSlice.size);
     327            2 :     HCCL_INFO(
     328              :         "[RunTasks] SendStream[%u] preload rank[%u] to A offset[%llu] size[%llu]", streamIdx, newRank, offsetA,
     329              :         firstSlice.size);
     330            2 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, commMemA, inMem, algResResp_->slaveStreams[streamIdx]));
     331              : 
     332            2 :     sendDataSlicesBySendStream_[streamIdx].pop_front();
     333              : 
     334              :     // Step B.2: Record from A
     335            2 :     LINK sendTargetLink;
     336            2 :     CHK_RET(GetSendTargetLink(newRank, sendTargetLink));
     337            2 :     CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[streamIdx]));
     338            2 :     CHK_RET(sendTargetLink->TxData(
     339              :         UserMemType::OUTPUT_MEM, 0, commMemA.ptr(), firstSlice.size, algResResp_->slaveStreams[streamIdx]));
     340              : 
     341              :     // Step B.3: D2D next to B (only if same rank)
     342            2 :     sendLoadedSize_[streamIdx] = 0;
     343            2 :     sendLoadedRemoteRank_[streamIdx] = newRank;
     344            2 :     sendCurPhase_[streamIdx] = 0;
     345            2 :     if (!sendDataSlicesBySendStream_[streamIdx].empty()) {
     346            1 :         SendRecvSlice& nextSlice = sendDataSlicesBySendStream_[streamIdx].front();
     347            1 :         if (nextSlice.remoteRank == newRank) {
     348            1 :             u64 offsetB = bufferSliceSize_ * (streamIdx * 2 + 1);
     349            1 :             DeviceMem commMemB = algResResp_->cclInputMem.range(offsetB, nextSlice.size);
     350            1 :             DeviceMem inMemB(nextSlice.addr, nextSlice.size);
     351            1 :             HCCL_INFO(
     352              :                 "[RunTasks] SendStream[%u] load next to B offset[%llu] size[%llu]", streamIdx, offsetB, nextSlice.size);
     353            1 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, commMemB, inMemB, algResResp_->slaveStreams[streamIdx]));
     354            1 :             sendLoadedSize_[streamIdx] = nextSlice.size;
     355            1 :             sendLoadedRemoteRank_[streamIdx] = nextSlice.remoteRank;
     356            1 :             sendCurPhase_[streamIdx] = 1;
     357            1 :             sendDataSlicesBySendStream_[streamIdx].pop_front();
     358            1 :             pendingSendCount++;
     359            1 :             if (sendDataSlicesBySendStream_[streamIdx].empty()) {
     360            1 :                 nonEmptySendStream--;
     361            1 :                 HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
     362              :             }
     363            1 :         }
     364              :     } else {
     365            1 :         nonEmptySendStream--;
     366            1 :         HCCL_INFO("[RunTasks] nonEmptySendStream[%u]", nonEmptySendStream);
     367              :     }
     368              : 
     369              :     // Step B.4: Wait for remote ack (TxDone)
     370            2 :     CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[streamIdx]));
     371            2 :     return HCCL_SUCCESS;
     372            2 : }
     373              : 
     374            2 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRecvSlice(u32 streamIdx, u32& nonEmptyRecvStream)
     375              : {
     376            2 :     SendRecvSlice& slice = recvDataSlicesByRecvStream_[streamIdx].front();
     377              : 
     378              :     // Reset phase to 0 when encountering a new rank
     379            2 :     if (slice.remoteRank != recvCurRemoteRank_[streamIdx]) {
     380            2 :         recvCurPhase_[streamIdx] = 0;
     381            2 :         recvCurRemoteRank_[streamIdx] = slice.remoteRank;
     382              :     }
     383              : 
     384            2 :     u32 curPhase = recvCurPhase_[streamIdx];
     385            2 :     u64 offset = bufferSliceSize_ * (streamIdx * 2 + curPhase);
     386            2 :     u64 phaseOffset = bufferSliceSize_ * (curPhase + (topoAttr_.userRank % sendStreamNum_) * 2);
     387            2 :     HCCL_INFO(
     388              :         "[RunTasks] RecvStream[%u] phase[%u] offset[%llu] phaseOffset[%llu] size[%llu] rank[%u]", streamIdx, curPhase,
     389              :         offset, phaseOffset, slice.size, slice.remoteRank);
     390              : 
     391            2 :     LINK recvTargetLink;
     392            2 :     CHK_RET(GetRecvTargetLink(slice.remoteRank, recvTargetLink));
     393              : 
     394            2 :     CHK_RET(recvTargetLink->RxPrepare(algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
     395            2 :     DeviceMem outMem(slice.addr, slice.size);
     396            2 :     CHK_RET(recvTargetLink->RxData(
     397              :         UserMemType::INPUT_MEM, phaseOffset, outMem.ptr(), slice.size,
     398              :         algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
     399            2 :     HCCL_INFO("[RunTasks] RecvStream[%u] direct, outMem ptr[%p], size[%llu]", streamIdx, outMem.ptr(), outMem.size());
     400            2 :     CHK_RET(recvTargetLink->RxDone(algResResp_->slaveStreams[streamIdx + sendStreamNum_]));
     401              : 
     402              :     // Toggle phase for next slice of same rank
     403            2 :     recvCurPhase_[streamIdx] = 1 - curPhase;
     404            2 :     recvDataSlicesByRecvStream_[streamIdx].pop_front();
     405            2 :     if (recvDataSlicesByRecvStream_[streamIdx].empty()) {
     406            2 :         nonEmptyRecvStream--;
     407            2 :         HCCL_INFO("[RunTasks] nonEmptyRecvStream[%u]", nonEmptyRecvStream);
     408              :     }
     409            2 :     return HCCL_SUCCESS;
     410            2 : }
     411              : 
     412            0 : HcclResult CollBatchSendRecvGroupExecutor::RunTasks(OpParam& param)
     413              : {
     414            0 :     u32 nonEmptySendStream = 0;
     415            0 :     u32 nonEmptyRecvStream = 0;
     416            0 :     CHK_RET(CalcStreamTaskStatus(nonEmptySendStream, nonEmptyRecvStream));
     417              : 
     418            0 :     CHK_RET(MainPostSubWait(param.stream));
     419              : 
     420              :     // Initialize ping-pong state (SDMA only; RDMA不做ping-pong)
     421            0 :     sendCurPhase_.resize(sendStreamNum_, 0);
     422            0 :     sendLoadedSize_.resize(sendStreamNum_, 0);
     423            0 :     sendLoadedRemoteRank_.resize(sendStreamNum_, 0);
     424            0 :     recvCurPhase_.resize(recvStreamNum_, 0);
     425            0 :     recvCurRemoteRank_.resize(recvStreamNum_, 0);
     426              : 
     427            0 :     u32 pendingSendCount = 0;
     428              : 
     429            0 :     while (pendingSendCount > 0 || nonEmptySendStream > 0 || nonEmptyRecvStream > 0 || !rdmaSendSlices_.empty()
     430            0 :            || !rdmaRecvSlices_.empty()) {
     431            0 :         HCCL_INFO(
     432              :             "[RunTasks] pending[%u] sendStream[%u] recvStream[%u] rdmaSend[%zu] rdmaRecv[%zu]", pendingSendCount,
     433              :             nonEmptySendStream, nonEmptyRecvStream, rdmaSendSlices_.size(), rdmaRecvSlices_.size());
     434              : 
     435            0 :         for (u32 i = 0; i < sendStreamNum_; i++) {
     436            0 :             if (sendLoadedSize_[i] > 0) {
     437            0 :                 CHK_RET(ProcessPreloadedSendSlice(i, pendingSendCount, nonEmptySendStream));
     438            0 :             } else if (!sendDataSlicesBySendStream_[i].empty()) {
     439            0 :                 CHK_RET(ProcessNewRankSendSlice(i, pendingSendCount, nonEmptySendStream));
     440              :             }
     441              :         }
     442              : 
     443            0 :         for (u32 i = 0; i < recvStreamNum_; i++) {
     444            0 :             if (recvDataSlicesByRecvStream_[i].empty()) {
     445            0 :                 continue;
     446              :             }
     447              :             // SDMA任务:走ping-pong
     448            0 :             CHK_RET(ProcessRecvSlice(i, nonEmptyRecvStream));
     449              :         }
     450              : 
     451            0 :         if (!rdmaSendSlices_.empty()) {
     452            0 :             CHK_RET(ProcessRdmaSendSlice());
     453              :         }
     454            0 :         if (!rdmaRecvSlices_.empty()) {
     455            0 :             CHK_RET(ProcessRdmaRecvSlice());
     456              :         }
     457              : 
     458            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
     459              :     }
     460            0 :     CHK_RET(MainWaitSubPost(param.stream));
     461            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
     462            0 :     return HCCL_SUCCESS;
     463              : }
     464              : 
     465            1 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRdmaSendSlice()
     466              : {
     467            1 :     SendRecvSlice& slice = rdmaSendSlices_.front();
     468              :     // RDMA send使用CCLOut A半区(单流,整半区作为单slot):send scratch offset = 0
     469            1 :     const u64 sendScratchOffset = 0;
     470            1 :     DeviceMem sendScratchMem = algResResp_->cclOutputMem.range(sendScratchOffset, slice.size);
     471            1 :     DeviceMem inMem(slice.addr, slice.size);
     472            1 :     HCCL_INFO(
     473              :         "[RunTasks] RdmaSend D2D user[%p] -> CCLOut_A[%p] offset[%llu] size[%llu] rank[%u]", inMem.ptr(),
     474              :         sendScratchMem.ptr(), sendScratchOffset, slice.size, slice.remoteRank);
     475            1 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, sendScratchMem, inMem, algResResp_->slaveStreams[RdmaSendStreamIdx()]));
     476              : 
     477              :     // Tx: 本地CCLOut_A(单slot) -> 远端CCLOut_B(单slot)。
     478              :     // 远端RDMA recv为单流,其recv scratch offset = rdmaDataBlockSize_(B半区基址)。
     479            1 :     const u64 remoteDstOffset = rdmaDataBlockSize_;
     480            1 :     LINK sendTargetLink;
     481            1 :     CHK_RET(GetSendTargetLink(slice.remoteRank, sendTargetLink));
     482            1 :     CHK_RET(sendTargetLink->TxPrepare(algResResp_->slaveStreams[RdmaSendStreamIdx()]));
     483            1 :     CHK_RET(sendTargetLink->TxData(
     484              :         UserMemType::OUTPUT_MEM, remoteDstOffset, sendScratchMem.ptr(), slice.size,
     485              :         algResResp_->slaveStreams[RdmaSendStreamIdx()]));
     486            1 :     CHK_RET(sendTargetLink->TxDone(algResResp_->slaveStreams[RdmaSendStreamIdx()]));
     487              : 
     488            1 :     rdmaSendSlices_.pop_front();
     489            1 :     return HCCL_SUCCESS;
     490            1 : }
     491              : 
     492            0 : HcclResult CollBatchSendRecvGroupExecutor::ProcessRdmaRecvSlice()
     493              : {
     494            0 :     SendRecvSlice& slice = rdmaRecvSlices_.front();
     495              :     // RDMA recv使用CCLOut B半区(单流,整半区作为单slot):recv scratch offset = rdmaDataBlockSize_(B半区基址)。
     496              :     // 与发送方写入的远端offset一致。
     497            0 :     const u64 recvScratchOffset = rdmaDataBlockSize_;
     498            0 :     DeviceMem recvScratchMem = algResResp_->cclOutputMem.range(recvScratchOffset, slice.size);
     499              : 
     500            0 :     LINK recvTargetLink;
     501            0 :     CHK_RET(GetRecvTargetLink(slice.remoteRank, recvTargetLink));
     502            0 :     CHK_RET(recvTargetLink->RxPrepare(algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
     503            0 :     CHK_RET(recvTargetLink->RxData(
     504              :         UserMemType::OUTPUT_MEM, recvScratchOffset, recvScratchMem.ptr(), slice.size,
     505              :         algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
     506            0 :     CHK_RET(recvTargetLink->RxDone(algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
     507              : 
     508              :     // D2D: CCLOut_B -> user output
     509            0 :     DeviceMem outMem(slice.addr, slice.size);
     510            0 :     HCCL_INFO(
     511              :         "[RunTasks] RdmaRecv D2D CCLOut_B[%p] offset[%llu] -> user[%p] size[%llu] rank[%u]", recvScratchMem.ptr(),
     512              :         recvScratchOffset, outMem.ptr(), slice.size, slice.remoteRank);
     513            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outMem, recvScratchMem, algResResp_->slaveStreams[RdmaRecvStreamIdx()]));
     514              : 
     515            0 :     rdmaRecvSlices_.pop_front();
     516            0 :     return HCCL_SUCCESS;
     517            0 : }
     518              : 
     519            1 : HcclResult CollBatchSendRecvGroupExecutor::SetNormalModeIfDeviceDirect()
     520              : {
     521            3 :     for (const auto& q : sendDataSlicesBySendStream_) {
     522            3 :         for (const auto& slice : q) {
     523            1 :             LINK targetLink;
     524            1 :             CHK_RET(GetSendTargetLink(slice.remoteRank, targetLink));
     525            1 :             if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
     526            0 :                 CHK_RET(SetNormalMode(dispatcher_));
     527            0 :                 HCCL_INFO("[CollBatchSendRecvGroupExecutor]Send Set dispatcher NormalMode");
     528            0 :                 return HCCL_SUCCESS;
     529              :             }
     530            1 :         }
     531              :     }
     532              : 
     533            1 :     for (const auto& q : recvDataSlicesByRecvStream_) {
     534            0 :         for (const auto& slice : q) {
     535            0 :             LINK targetLink;
     536            0 :             CHK_RET(GetRecvTargetLink(slice.remoteRank, targetLink));
     537            0 :             if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
     538            0 :                 CHK_RET(SetNormalMode(dispatcher_));
     539            0 :                 HCCL_INFO("[CollBatchSendRecvGroupExecutor]Recv Set NormalMode dispatcher");
     540            0 :                 return HCCL_SUCCESS;
     541              :             }
     542            0 :         }
     543              :     }
     544              : 
     545            1 :     for (const auto& slice : rdmaSendSlices_) {
     546            0 :         LINK targetLink;
     547            0 :         CHK_RET(GetSendTargetLink(slice.remoteRank, targetLink));
     548            0 :         if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
     549            0 :             CHK_RET(SetNormalMode(dispatcher_));
     550            0 :             HCCL_INFO("[CollBatchSendRecvGroupExecutor]RdmaSend Set NormalMode dispatcher");
     551            0 :             return HCCL_SUCCESS;
     552              :         }
     553            0 :     }
     554              : 
     555            1 :     for (const auto& slice : rdmaRecvSlices_) {
     556            0 :         LINK targetLink;
     557            0 :         CHK_RET(GetRecvTargetLink(slice.remoteRank, targetLink));
     558            0 :         if (targetLink->GetTransportType() == TransportType::TRANS_TYPE_DEVICE_DIRECT) {
     559            0 :             CHK_RET(SetNormalMode(dispatcher_));
     560            0 :             HCCL_INFO("[CollBatchSendRecvGroupExecutor]RdmaRecv Set NormalMode dispatcher");
     561            0 :             return HCCL_SUCCESS;
     562              :         }
     563            0 :     }
     564            1 :     return HCCL_SUCCESS;
     565              : }
     566              : 
     567            0 : HcclResult CollBatchSendRecvGroupExecutor::RunLoop(OpParam& param)
     568              : {
     569            0 :     if (static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())) {
     570            0 :         auto meta = HcclOpMetaInfo::GetOneForBatchSendRecv();
     571            0 :         CHK_RET(InitTask(dispatcher_, param.stream, meta.isEnableCache, meta.GetCacheKey()));
     572              :         // 多流子图前后需加空拷贝
     573            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(
     574              :             algResResp_->cclInputMem, algResResp_->cclOutputMem, param.stream, dispatcher_));
     575              :     }
     576            0 :     CHK_RET(SetNormalModeIfDeviceDirect());
     577            0 :     CHK_RET(RunTasks(param));
     578            0 :     if (static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())) {
     579              :         // 多流子图前后需加空拷贝
     580            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(
     581              :             algResResp_->cclInputMem, algResResp_->cclOutputMem, param.stream, dispatcher_));
     582            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
     583            0 :         HCCL_INFO("LaunchTaskExtend!");
     584              :     }
     585            0 :     return HCCL_SUCCESS;
     586              : }
     587              : 
     588            3 : HcclResult CollBatchSendRecvGroupExecutor::CalcSendSlices()
     589              : {
     590              :     // SDMA slice按remoteRank % streamNum分发到sendDataSlicesBySendStream_(CCLIn, ping-pong);
     591              :     // RDMA slice按rank分组到rdmaByRank,随后按对称场景规则排序到rdmaSendSlices_(CCLOut A半区, 无ping-pong)。
     592            3 :     sendDataSlicesBySendStream_.resize(sendStreamNum_);
     593            3 :     std::map<u32, std::deque<SendRecvSlice>> rdmaByRank;
     594           15 :     for (u32 i = 0; i < sendStreamNum_; i++) {
     595           12 :         const auto& sendQueueInner = sendQueueBySendstream_[i];
     596           14 :         for (u32 j = 0; j < sendQueueInner.size(); j++) {
     597            2 :             HcclSendRecvItem* sendRecvItem = sendQueueInner[j];
     598            2 :             u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
     599            2 :             bool isRdma = IsRemoteRankRdma(sendRecvItem->remoteRank);
     600            2 :             u64 maxCountPerLoop = isRdma ? (rdmaDataBlockSize_ / unitSize) : CalcSendLoopMaxCount(unitSize);
     601            4 :             while (sendRecvItem->count > 0) {
     602            2 :                 u8* curInputPtr = static_cast<u8*>(sendRecvItem->buf);
     603            2 :                 CHK_PTR_NULL(curInputPtr);
     604            2 :                 u64 curCount = (sendRecvItem->count > maxCountPerLoop) ? maxCountPerLoop : sendRecvItem->count;
     605            2 :                 u64 curSize = curCount * unitSize;
     606            2 :                 SendRecvSlice slice(curInputPtr, curSize, sendRecvItem->remoteRank, isRdma);
     607            2 :                 if (isRdma) {
     608            1 :                     rdmaByRank[sendRecvItem->remoteRank].push_back(slice);
     609              :                 } else {
     610            1 :                     sendDataSlicesBySendStream_[i].push_back(slice);
     611              :                 }
     612            2 :                 sendRecvItem->count -= curCount;
     613            2 :                 sendRecvItem->buf = static_cast<u8*>(sendRecvItem->buf) + curSize;
     614              :             }
     615              :         }
     616              :     }
     617              :     // RDMA按对称场景规则排序(send前向递增,跳过本pod与无任务对端)
     618            3 :     OrderRdmaSlices(true, rdmaByRank, rdmaSendSlices_);
     619              : 
     620           15 :     for (u32 i = 0; i < sendStreamNum_; i++) {
     621           12 :         const auto& sendQueueInner = sendDataSlicesBySendStream_[i];
     622           13 :         for (const auto& slice : sendQueueInner) {
     623            1 :             HCCL_INFO(
     624              :                 "[CalcSendSlices] sendstream[%u] addr[%p] size[%llu] rank[%u] isRdma[%d]", i, slice.addr, slice.size,
     625              :                 slice.remoteRank, slice.isRdma);
     626              :         }
     627              :     }
     628            4 :     for (const auto& slice : rdmaSendSlices_) {
     629            1 :         HCCL_INFO("[CalcSendSlices] rdmaSend addr[%p] size[%llu] rank[%u]", slice.addr, slice.size, slice.remoteRank);
     630              :     }
     631            3 :     return HCCL_SUCCESS;
     632            3 : }
     633              : 
     634            3 : HcclResult CollBatchSendRecvGroupExecutor::CalcRecvSlices()
     635              : {
     636              :     // SDMA slice按remoteRank % streamNum分发到recvDataSlicesByRecvStream_(CCLIn, ping-pong);
     637              :     // RDMA slice按rank分组到rdmaByRank,随后按对称场景规则排序到rdmaRecvSlices_(CCLOut B半区, 无ping-pong)。
     638            3 :     recvDataSlicesByRecvStream_.resize(recvStreamNum_);
     639            3 :     std::map<u32, std::deque<SendRecvSlice>> rdmaByRank;
     640           15 :     for (u32 i = 0; i < recvStreamNum_; i++) {
     641           12 :         const auto& recvQueueInner = recvQueueByRecvstream_[i];
     642           14 :         for (u32 j = 0; j < recvQueueInner.size(); j++) {
     643            2 :             HcclSendRecvItem* sendRecvItem = recvQueueInner[j];
     644            2 :             u32 unitSize = SIZE_TABLE[sendRecvItem->dataType];
     645            2 :             bool isRdma = IsRemoteRankRdma(sendRecvItem->remoteRank);
     646            2 :             u64 maxCountPerLoop = isRdma ? (rdmaDataBlockSize_ / unitSize) : CalcRecvLoopMaxCount(unitSize);
     647            4 :             while (sendRecvItem->count > 0) {
     648            2 :                 u8* curOutputPtr = static_cast<u8*>(sendRecvItem->buf);
     649            2 :                 CHK_PTR_NULL(curOutputPtr);
     650            2 :                 u64 curCount = (sendRecvItem->count > maxCountPerLoop) ? maxCountPerLoop : sendRecvItem->count;
     651            2 :                 u64 curSize = curCount * unitSize;
     652            2 :                 SendRecvSlice slice(curOutputPtr, curSize, sendRecvItem->remoteRank, isRdma);
     653            2 :                 if (isRdma) {
     654            1 :                     rdmaByRank[sendRecvItem->remoteRank].push_back(slice);
     655              :                 } else {
     656            1 :                     recvDataSlicesByRecvStream_[i].push_back(slice);
     657              :                 }
     658            2 :                 sendRecvItem->count -= curCount;
     659            2 :                 sendRecvItem->buf = static_cast<u8*>(sendRecvItem->buf) + curSize;
     660              :             }
     661              :         }
     662              :     }
     663              :     // RDMA按对称场景规则排序(recv后向递减,跳过本pod与无任务对端)
     664            3 :     OrderRdmaSlices(false, rdmaByRank, rdmaRecvSlices_);
     665              : 
     666           15 :     for (u32 i = 0; i < recvStreamNum_; i++) {
     667           12 :         const auto& recvQueueInner = recvDataSlicesByRecvStream_[i];
     668           13 :         for (const auto& slice : recvQueueInner) {
     669            1 :             HCCL_INFO(
     670              :                 "[CalcRecvSlices] recvstream[%u] addr[%p] size[%llu] rank[%u] isRdma[%d]", i, slice.addr, slice.size,
     671              :                 slice.remoteRank, slice.isRdma);
     672              :         }
     673              :     }
     674            4 :     for (const auto& slice : rdmaRecvSlices_) {
     675            1 :         HCCL_INFO("[CalcRecvSlices] rdmaRecv addr[%p] size[%llu] rank[%u]", slice.addr, slice.size, slice.remoteRank);
     676              :     }
     677            3 :     return HCCL_SUCCESS;
     678            3 : }
     679              : 
     680           64 : u32 CollBatchSendRecvGroupExecutor::GetNextDstRank(u32& curDstRank)
     681              : {
     682              :     // 对称场景send方向:沿rank id递增环绕遍历,跳过本pod。移植自alltoallv_direct_fullmesh。
     683           64 :     if (curDstRank >= topoAttr_.userRankSize) {
     684            5 :         curDstRank = curDstRank % topoAttr_.userRankSize;
     685              :     }
     686           64 :     if (curDstRank == podStartRank_) {
     687            8 :         curDstRank += devNumInlocalPod_;
     688              :     }
     689           64 :     curDstRank = curDstRank % topoAttr_.userRankSize;
     690           64 :     return curDstRank++;
     691              : }
     692              : 
     693           42 : u32 CollBatchSendRecvGroupExecutor::GetPreSrcRank(u32& curSrcRank)
     694              : {
     695              :     // 对称场景recv方向:沿rank id递减环绕遍历,跳过本pod。移植自alltoallv_direct_fullmesh。
     696           42 :     if (curSrcRank == podStartRank_ + devNumInlocalPod_ - 1) {
     697            2 :         curSrcRank = (curSrcRank + topoAttr_.userRankSize - devNumInlocalPod_) % topoAttr_.userRankSize;
     698              :     }
     699           42 :     if (curSrcRank == 0) {
     700            5 :         curSrcRank = topoAttr_.userRankSize - 1;
     701            5 :         return 0;
     702              :     }
     703           37 :     return curSrcRank--;
     704              : }
     705              : 
     706           11 : void CollBatchSendRecvGroupExecutor::OrderRdmaSlices(
     707              :     bool isSend, const std::map<u32, std::deque<SendRecvSlice>>& byRank, std::deque<SendRecvSlice>& out)
     708              : {
     709              :     // 跨pod对端总数 = userRankSize - devNumInlocalPod。遍历每个候选rank一次。
     710           11 :     u32 totalRdmaRankNum = topoAttr_.userRankSize - devNumInlocalPod_;
     711           11 :     u32 curRank = INVALID_VALUE_RANKID;
     712           11 :     if (isSuPodAsym_) {
     713              :         // 非对称场景:send/recv使用相同遍历顺序(均前向递增GetNextDstRank)。
     714              :         // 起点为本pod之外的第一个rank(从0开始扫描,取首个不在[podStartRank_, podEndRank_]内的rank)。
     715            0 :         for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
     716            0 :             if (i < podStartRank_ || i > podEndRank_) {
     717            0 :                 curRank = i;
     718            0 :                 break;
     719              :             }
     720              :         }
     721            0 :         HCCL_INFO(
     722              :             "[OrderRdmaSlices] %s asym startRank[%u] totalRdmaRankNum[%u]", isSend ? "send" : "recv", curRank,
     723              :             totalRdmaRankNum);
     724            0 :         for (u32 i = 0; i < totalRdmaRankNum; i++) {
     725            0 :             u32 rank = GetNextDstRank(curRank);
     726            0 :             auto it = byRank.find(rank);
     727            0 :             if (it == byRank.end()) {
     728            0 :                 HCCL_INFO("[OrderRdmaSlices] %s asym skip rank[%u] (no task)", isSend ? "send" : "recv", rank);
     729            0 :                 continue;
     730              :             }
     731            0 :             for (const auto& s : it->second) {
     732            0 :                 out.push_back(s);
     733              :             }
     734            0 :             HCCL_INFO(
     735              :                 "[OrderRdmaSlices] %s asym append rank[%u] sliceNum[%zu]", isSend ? "send" : "recv", rank,
     736              :                 it->second.size());
     737              :         }
     738              :     } else {
     739              :         // 对称场景:send前向递增(GetNextDstRank),recv后向递减(GetPreSrcRank)。
     740              :         // send起点取"下一个pod中相同pod内位置的rank";recv起点取"上一个pod中相同pod内位置的rank"。
     741           11 :         curRank = isSend ? (topoAttr_.userRank + devNumInlocalPod_) % topoAttr_.userRankSize :
     742            4 :                            (topoAttr_.userRank + topoAttr_.userRankSize - devNumInlocalPod_) % topoAttr_.userRankSize;
     743           11 :         HCCL_INFO(
     744              :             "[OrderRdmaSlices] %s sym startRank[%u] totalRdmaRankNum[%u]", isSend ? "send" : "recv", curRank,
     745              :             totalRdmaRankNum);
     746          103 :         for (u32 i = 0; i < totalRdmaRankNum; i++) {
     747           92 :             u32 rank = isSend ? GetNextDstRank(curRank) : GetPreSrcRank(curRank);
     748           92 :             auto it = byRank.find(rank);
     749           92 :             if (it == byRank.end()) {
     750           80 :                 HCCL_INFO("[OrderRdmaSlices] %s sym skip rank[%u] (no task)", isSend ? "send" : "recv", rank);
     751           80 :                 continue;
     752              :             }
     753           25 :             for (const auto& s : it->second) {
     754           13 :                 out.push_back(s);
     755              :             }
     756           12 :             HCCL_INFO(
     757              :                 "[OrderRdmaSlices] %s sym append rank[%u] sliceNum[%zu]", isSend ? "send" : "recv", rank,
     758              :                 it->second.size());
     759              :         }
     760              :     }
     761           11 : }
     762              : 
     763            3 : u64 CollBatchSendRecvGroupExecutor::CalcSendLoopMaxCount(const u32 unitSize) const
     764              : {
     765              :     // 中转内存单次最多能够接受的input count
     766            3 :     u64 maxCountPerLoop = bufferSliceSize_ / unitSize;
     767            3 :     HCCL_INFO(
     768              :         "[CollBatchSendRecvGroupExecutor][CalcSendLoopMaxCount]"
     769              :         "using default maxCountPerLoop[%llu] as CCLBuffSize / unitSize.",
     770              :         maxCountPerLoop);
     771            3 :     return maxCountPerLoop;
     772              : }
     773              : 
     774            3 : u64 CollBatchSendRecvGroupExecutor::CalcRecvLoopMaxCount(const u32 unitSize) const
     775              : {
     776              :     // 中转内存单次最多能够接受的output count
     777            3 :     u64 maxCountPerLoop = bufferSliceSize_ / unitSize;
     778            3 :     HCCL_INFO(
     779              :         "[CollBatchSendRecvGroupExecutor][CalcRecvLoopMaxCount]"
     780              :         "using default maxCountPerLoop[%llu] as CCLBuffSize / unitSize.",
     781              :         maxCountPerLoop);
     782            3 :     return maxCountPerLoop;
     783              : }
     784              : 
     785            1 : HcclResult CollBatchSendRecvGroupExecutor::CalcStreamNum(u32& streamNum)
     786              : {
     787            1 :     if (topoAttr_.userRankSize == 1) {
     788            0 :         sendStreamNum_ = 0;
     789            0 :         recvStreamNum_ = 0;
     790            0 :         streamNum = 0;
     791            0 :         HCCL_INFO("[CollBatchSendRecvGroupExecutor] Only one rank, do not need substream, streamNum[%u]", streamNum);
     792            0 :         return HCCL_SUCCESS;
     793              :     }
     794            1 :     sendStreamNum_ = GROUP_MAX_CONCURRENT;
     795            1 :     recvStreamNum_ = GROUP_MAX_CONCURRENT;
     796              :     // SDMA占sendStreamNum_+recvStreamNum_条从流;RDMA单独占RDMA_STREAM_NUM条(1 send + 1 recv)。
     797            1 :     streamNum = sendStreamNum_ + recvStreamNum_ + RDMA_STREAM_NUM;
     798            1 :     HCCL_INFO("[CollBatchSendRecvGroupExecutor][CalcStreamNum] tag_[%s], streamNum[%u].", tag_.c_str(), streamNum);
     799            1 :     return HCCL_SUCCESS;
     800              : }
     801              : 
     802              : REGISTER_EXEC("BatchSendRecvGroup", BatchSendRecvGroupExecutor, CollBatchSendRecvGroupExecutor);
     803              : } // namespace hccl
        

Generated by: LCOV version 2.0-1