LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_to_all - coll_all_to_all_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 15.7 % 274 43
Test Date: 2026-07-28 12:11:00 Functions: 26.3 % 19 5

            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_all_to_all_executor.h"
      12              : #include "device_capacity.h"
      13              : 
      14              : namespace hccl {
      15              : 
      16            5 : CollAlltoAllExecutor::CollAlltoAllExecutor(const HcclDispatcher dispatcher,
      17            5 :                                            std::unique_ptr<TopoMatcher> &topoMatcher)
      18            5 :     : CollNativeExecutorBase(dispatcher, topoMatcher)
      19              : {
      20            5 : }
      21              : 
      22            0 : HcclResult CollAlltoAllExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algRes)
      23              : {
      24            0 :     HcclUs startut = TIME_NOW();
      25            0 :     tag_ = param.tag;
      26            0 :     algResResp_ = &algRes;
      27            0 :     AlltoAllVParam_ = param;
      28            0 :     ExecMem execMem;
      29            0 :     execMem.count = 0;
      30            0 :     execMem.inputPtr = param.inputPtr;
      31            0 :     execMem.outputPtr = param.outputPtr;
      32              : 
      33            0 :     HcclResult ret = HCCL_SUCCESS;
      34            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      35            0 :         execMem.inputMem = algRes.cclInputMem;
      36            0 :         execMem.outputMem = algRes.cclOutputMem;
      37            0 :         execMem.scratchMem = algRes.scratchMem;
      38              : 
      39            0 :         auto opMeta = GetOpMeta(param.opType, algRes.paramInputMem.size());   // override
      40            0 :         CHK_RET(InitTask(dispatcher_, param.stream, opMeta.isEnableCache, opMeta.GetCacheKey()));
      41            0 :         bool massTasks = HasMassTasks(allMeshAggregationSendRecvInfo_);
      42            0 :         if (massTasks) {
      43            0 :             CHK_RET(SetNormalMode(dispatcher_));
      44              :         }
      45            0 :         ret = KernelRun(param, execMem);
      46              :     } else {
      47            0 :         execMem.inputMem = algRes.paramInputMem;
      48            0 :         execMem.outputMem = algRes.paramOutputMem;
      49            0 :         execMem.scratchMem = algRes.scratchMem;
      50            0 :         ret = KernelRun(param, execMem);
      51              :     }
      52            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
      53              :         HCCL_ERROR("[CollAlltoAllExecutor][Orchestrate]errNo[0x%016llx]executor run failed",
      54              :             HCCL_ERROR_CODE(ret)), ret);
      55              : 
      56              :     // Enforce task launch at the end of Orchestrate
      57              :     // 注意: 不要删除这里的强制launch, 否则会导致aicpu cache功能问题
      58            0 :     HCCL_INFO("%s: enforce task launch at the end of Orchestrate", __func__);
      59            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
      60              : 
      61            0 :     HCCL_INFO("tag[%s], AlltoAll executor orchestrate success, take time [%lld]us.",
      62              :         param.tag.c_str(), DURATION_US(TIME_NOW() - startut));
      63            0 :     return HCCL_SUCCESS;
      64            0 : }
      65              : 
      66            0 : HcclResult CollAlltoAllExecutor::GetAdjInfo(AlgResourceResponse& algRes, AdjInfo& adjInfo)
      67              : {
      68            0 :     algResResp_ = &algRes;
      69            0 :     SubCommInfo levelCommInfo = {0};
      70            0 :     AdjInfo nslbAdjInfo = {0};
      71            0 :     u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
      72              : 
      73            0 :     if (Getlevel1CommRank(levelCommInfo) != HCCL_SUCCESS) {
      74            0 :         return HCCL_SUCCESS;
      75              :     }
      76            0 :     u32 localRank= levelCommInfo.localRank;
      77            0 :     u32 localRankSize = levelCommInfo.localRankSize;
      78              : 
      79            0 :     std::unique_ptr<AlgTemplateBase> levelTempAlg;
      80            0 :     if (SelectTempAlg(levelTempAlg, localRankSize) != HCCL_SUCCESS) {
      81            0 :         return HCCL_SUCCESS;
      82              :     }
      83            0 :     GetDevNumInlocalPod(devNumInlocalPod);
      84            0 :     if (devNumInlocalPod == INVALID_VALUE_RANKSIZE) {
      85            0 :         HCCL_INFO("[GetAdjInfo-nslbdp] devNumInlocalPod == INVALID_VALUE_RANKSIZE.");
      86            0 :         return HCCL_SUCCESS;
      87              :     }
      88              : 
      89            0 :     nslbAdjInfo.dstRankNum = devNumInlocalPod;
      90            0 :     CHK_RET(levelTempAlg->GetNslbAdjInfo(localRank, localRankSize, levelCommInfo.links, nslbAdjInfo));
      91              : 
      92            0 :     adjInfo.dstRankNum = nslbAdjInfo.dstRankNum;
      93            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] adjInfo.dstRankNum[%u].", adjInfo.dstRankNum);
      94              :     
      95            0 :     for (size_t i = 0; i < nslbAdjInfo.nsAdjInfo.size(); i++) {
      96            0 :         NslbDpAdjInfo dpAdjInfo = {0};
      97            0 :         dpAdjInfo.dstLocalRankId = nslbAdjInfo.nsAdjInfo[i].dstLocalRankId;
      98            0 :         dpAdjInfo.phaseId = nslbAdjInfo.nsAdjInfo[i].phaseId;
      99            0 :         dpAdjInfo.rev = 0;
     100            0 :         adjInfo.nsAdjInfo.push_back(dpAdjInfo); 
     101            0 :         HCCL_INFO("[nslbdp]GetAdjInfo dstLocalRankId[%u], phaseId[%u].",
     102              :                    nslbAdjInfo.nsAdjInfo[i].dstLocalRankId, nslbAdjInfo.nsAdjInfo[i].phaseId);
     103              :     }
     104            0 :     return HCCL_SUCCESS;
     105            0 : }
     106              : 
     107              : // override----------------------资源计算接口----------------------
     108            5 : HcclResult CollAlltoAllExecutor::CalcResRequest(const OpParam& param, AlgResourceRequest& resourceRequest)
     109              : {
     110            5 :     (void)ParseParam(param);
     111              : 
     112            5 :     u64 scratchMemSize = 0U;
     113            5 :     u32 streamNum = 0U;
     114            5 :     u32 notifyNum = 0U;
     115            5 :     u64 aivBufferRequest = 0U;
     116              :     std::vector<LevelNSubCommTransport> opTransport {
     117            0 :         std::vector<LevelNSubCommTransport>(static_cast<u32>(COMM_LEVEL_RESERVED))
     118            5 :     };
     119              : 
     120              :     // AICPU aicpuUnfold展开模式下临时强制OP_BASE,使整个资源计算路径与AICPU侧一致
     121              :     // CalcScratchMemSize走OP_BASE分支正确计算scratch
     122              :     // CalcCommInfo走OP_BASE分支设outputMemType为CCL_OUTPUT而非SCRATCH
     123              :     // 避免transport因scratch未分配而拿到nullptr
     124              :     // AIV executor有独立的资源计算逻辑,不需要force OP_BASE
     125            5 :     const bool needForceOpBase = param.aicpuUnfoldMode && !param.isZeroCopy && !desc_.isAivMode;
     126            5 :     HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest] aicpuUnfoldMode[%d] isZeroCopy[%d] "
     127              :         "needForceOpBase[%d] workflowMode[%d] tag[%s]",
     128              :         param.aicpuUnfoldMode, param.isZeroCopy, needForceOpBase,
     129              :         workflowMode_, param.tag.c_str());
     130            5 :     const HcclWorkflowMode savedWorkflowMode = workflowMode_;
     131            5 :     if (needForceOpBase) {
     132            0 :         HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest] aicpuUnfoldMode force OpBase, "
     133              :             "originalWorkflowMode[%d], tag[%s]",
     134              :             savedWorkflowMode, param.tag.c_str());
     135            0 :         workflowMode_ = HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
     136              :     }
     137              : 
     138            5 :     CHK_RET(CalcScratchMemSize(scratchMemSize));
     139            5 :     CHK_RET(CalcStreamNum(streamNum));
     140            5 :     CHK_RET(CalcNotifyNum(streamNum, notifyNum));
     141            5 :     CHK_RET(CalcAivBufferRequest(aivBufferRequest));
     142            5 :     CHK_RET(CalcCommInfo(opTransport));
     143              : 
     144            5 :     if (needForceOpBase) {
     145            0 :         HCCL_DEBUG("[CollAlltoAllExecutor][CalcResRequest] restore workflowMode "
     146              :             "after resource calc, scratchMemSize[%llu]", scratchMemSize);
     147            0 :         workflowMode_ = savedWorkflowMode;
     148              :     }
     149              : 
     150            5 :     CHK_RET(BuildResourceRequest(scratchMemSize, streamNum, notifyNum, aivBufferRequest, opTransport, resourceRequest));
     151            5 :     HCCL_INFO("[CollAlltoAllExecutor][%s] streamNum[%u], notifyNum[%u], sctrachMemSize[%llu], aivBufferRequest[%llu]",
     152              :         __func__, resourceRequest.streamNum, resourceRequest.notifyNum, resourceRequest.scratchMemSize,
     153              :         resourceRequest.aivBufferRequest);
     154              :     // 打印建链诉求
     155           85 :     for (u32 levelIndex = 0; levelIndex < COMM_LEVEL_RESERVED; levelIndex++) {
     156           80 :         LevelNSubCommTransport &levelTransport = resourceRequest.opTransport[levelIndex];
     157           80 :         u32 ringSize = levelTransport.size();
     158           85 :         for (u32 ringIndex = 0; ringIndex < ringSize; ringIndex++) {
     159            5 :             SingleSubCommTransport &subCommTransport = levelTransport[ringIndex];
     160            5 :             u32 rankSize = subCommTransport.transportRequests.size();
     161           10 :             for (u32 rankIndex = 0; rankIndex < rankSize; rankIndex++) {
     162            5 :                 if (subCommTransport.transportRequests[rankIndex].isValid == true) {
     163            0 :                     HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest]" \
     164              :                         "levelIndex[%u], ringIndex[%u], rankIndex[%u], userRank[%u], remoteRank[%u]" \
     165              :                         "isUsedRdma[%d]",
     166              :                         levelIndex, ringIndex, rankIndex, subCommTransport.transportRequests[rankIndex].localUserRank,
     167              :                         subCommTransport.transportRequests[rankIndex].remoteUserRank,
     168              :                         subCommTransport.transportRequests[rankIndex].isUsedRdma);
     169              :                 }
     170              :             }
     171              :         }
     172              :     }
     173            5 :     CHK_RET(CheckNeedCreateVirtualLinks(resourceRequest));
     174            5 :     HCCL_DEBUG("[%s] process success", __func__);
     175            5 :     return HCCL_SUCCESS;
     176            5 : }
     177              : 
     178            5 : HcclResult CollAlltoAllExecutor::CheckNeedCreateVirtualLinks(AlgResourceRequest &resourceRequest)
     179              : {
     180            5 :     return HCCL_SUCCESS;
     181              : }
     182              : 
     183            0 : HcclResult CollAlltoAllExecutor::SetExcutorExtraInfo(const std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, u64 cclbufferSize)
     184              : {
     185            0 :     allMeshAggregationSendRecvInfo_.clear();
     186            0 :     allMeshAggregationSendRecvInfo_ = allMeshAggregationSendRecvInfo;
     187            0 :     UpdateAlltoAllZCopyMode(allMeshAggregationSendRecvInfo_, cclbufferSize);
     188            0 :     HCCL_DEBUG("[%s] allMeshAggregationSendRecvInfo_ size[%u]", __func__, allMeshAggregationSendRecvInfo_.size());
     189              : 
     190            0 :     return HCCL_SUCCESS;
     191              : }
     192              : 
     193            0 : void CollAlltoAllExecutor::UpdateAlltoAllZCopyMode(std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, u64 cclbufferSize)
     194              : {
     195            0 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     196            0 :         u64 maxSendSize = 0;
     197            0 :         u64 maxRecvSize = 0;
     198            0 :         for (auto &sendRecvInfo : allMeshAggregationSendRecvInfo) {
     199            0 :             for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
     200            0 :                 u64 curSendSize = sendRecvInfo.sendLength[i] + sendRecvInfo.sendOffset[i];
     201            0 :                 maxSendSize = std::max(maxSendSize, curSendSize);
     202            0 :                 u64 curRecvSize = sendRecvInfo.recvLength[i] + sendRecvInfo.recvOffset[i];
     203            0 :                 maxRecvSize = std::max(maxRecvSize, curRecvSize);
     204              :             }
     205              :         }
     206            0 :         bool isAlltoAllZCopyMode = (maxSendSize <= cclbufferSize) &&
     207            0 :                                    (maxRecvSize <= cclbufferSize);
     208            0 :         if (isAlltoAllZCopyMode) {
     209            0 :             isAlltoAllZCopyMode_ = true;
     210              :         }
     211            0 :         HCCL_INFO("[CollAlltoAllExecutor][UpdateAlltoAllZCopyMode] maxSendSize[%llu], maxRecvSize[%llu], "\
     212              :             "cclBufferSize[%llu]", maxSendSize, maxRecvSize, cclbufferSize);
     213              :     } else {
     214              :         // 图模式走ZCopy实现
     215            0 :         isAlltoAllZCopyMode_ = true;
     216              :     }
     217            0 :     HCCL_DEBUG("UpdateAlltoAllZCopyMode isAlltoAllZCopyMode_[%d]", isAlltoAllZCopyMode_);
     218            0 : }
     219              : 
     220            0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationSendInfo(const AlltoAllUserRankInfo &userRankInfo,
     221              :     const SendRecvInfo &mySendRecvInfo, const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo,
     222              :     u32 rankInMeshAggregation, u32 infoIndex, OneSendRecvAddrInfo &curSendInfo, u32 meshAggregationRankSize,
     223              :     const bool &isSingleMesh)
     224              : {
     225            0 :     if (infoIndex >= mySendRecvInfo.sendOffset.size() || infoIndex >= mySendRecvInfo.sendLength.size()) {
     226            0 :         HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid infoIndex[%u]", infoIndex);
     227            0 :         return;
     228              :     }
     229            0 :     curSendInfo.localOffset = mySendRecvInfo.sendOffset[infoIndex];
     230            0 :     curSendInfo.localLength = mySendRecvInfo.sendLength[infoIndex];
     231            0 :     u64 remoteOffset = 0;
     232              : 
     233            0 :     if (isSingleMesh) {
     234            0 :         remoteOffset = myMeshAggregationSendRecvInfo[infoIndex].recvOffset[userRankInfo.userRank];
     235              :     } else {
     236            0 :         for (u32 j = infoIndex % meshAggregationRankSize; j <= infoIndex; j += meshAggregationRankSize) {
     237            0 :             for (u32 k = 0; k < meshAggregationRankSize; k++) {
     238            0 :                 if (j == infoIndex && k == rankInMeshAggregation) {
     239            0 :                     break;
     240              :                 }
     241            0 :                 if (k < myMeshAggregationSendRecvInfo.size() && j <
     242            0 :                     myMeshAggregationSendRecvInfo[k].sendLength.size()) {
     243            0 :                     remoteOffset += myMeshAggregationSendRecvInfo[k].sendLength[j];
     244              :                 } else {
     245            0 :                     HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] invalid MeshAggregationSendRecvInfo size[%zu]",
     246              :                         myMeshAggregationSendRecvInfo.size());
     247            0 :                     return;
     248              :                 }
     249              :             }
     250              :         }
     251              :     }
     252              : 
     253            0 :     curSendInfo.remoteOffset = remoteOffset;
     254            0 :     curSendInfo.remoteLength = curSendInfo.localLength;
     255            0 :     HCCL_DEBUG("[CalcIntraMeshAggregationSendInfo] localOffset[%llu], localLength[%llu], "\
     256              :         "remoteOffset[%llu], remoteLength[%llu]", curSendInfo.localOffset,
     257              :         curSendInfo.localLength, curSendInfo.remoteOffset, curSendInfo.remoteLength);
     258              : }
     259              : 
     260            0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationRecvInfoInMeshAggregation(u32 rankIndex, u32 infoIndex,
     261              :     const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo, u64 &localOffset, u32 &offsetCounter,
     262              :     u64 &localLength, u64 &remoteOffset, u32 meshAggregationRankSize)
     263              : {
     264              :     // 这里的判断在外部已经保证了,为了应对coverity sc
     265            0 :     if (myMeshAggregationSendRecvInfo.size() < meshAggregationRankSize) {
     266            0 :         HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid myMeshAggregationSendRecvInfo[%zu]",
     267              :             myMeshAggregationSendRecvInfo.size());
     268            0 :         return;
     269              :     }
     270            0 :     if (myMeshAggregationSendRecvInfo[0].sendLength.size() == 0 ||
     271            0 :         myMeshAggregationSendRecvInfo[0].sendOffset.size() == 0) {
     272            0 :         HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid sendLength size[%zu] or sendOffset size[%zu]",
     273              :             myMeshAggregationSendRecvInfo[0].sendLength.size(), myMeshAggregationSendRecvInfo[0].sendOffset.size());
     274            0 :         return;
     275              :     }
     276            0 :     for (u32 k = 0; k < meshAggregationRankSize; k++) {
     277            0 :         if (infoIndex == 0) {
     278            0 :             localOffset = 0;
     279            0 :             localLength = myMeshAggregationSendRecvInfo[k].sendLength[rankIndex];
     280            0 :             remoteOffset = myMeshAggregationSendRecvInfo[k].sendOffset[rankIndex];
     281            0 :             break;
     282              :         }
     283              : 
     284            0 :         localOffset += myMeshAggregationSendRecvInfo[k].sendLength[rankIndex];
     285            0 :         offsetCounter++;
     286            0 :         if (offsetCounter == infoIndex) {
     287            0 :             if (k == meshAggregationRankSize - 1) {
     288            0 :                 localLength = myMeshAggregationSendRecvInfo[0].sendLength[rankIndex + meshAggregationRankSize];
     289            0 :                 remoteOffset = myMeshAggregationSendRecvInfo[0].sendOffset[rankIndex + meshAggregationRankSize];
     290              :             } else {
     291            0 :                 localLength = myMeshAggregationSendRecvInfo[k + 1].sendLength[rankIndex];
     292            0 :                 remoteOffset = myMeshAggregationSendRecvInfo[k + 1].sendOffset[rankIndex];
     293              :             }
     294            0 :             break;
     295              :         }
     296              :     }
     297            0 :     HCCL_DEBUG("[%s] process success", __func__);
     298              : }
     299              : 
     300            0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationRecvInfo(const AlltoAllUserRankInfo &userRankInfo,
     301              :     const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo, u32 infoIndex, OneSendRecvAddrInfo &curRecvInfo,
     302              :     u32 meshAggregationRankSize, const bool &isSingleMesh)
     303              : {
     304            0 :     u64 localOffset = 0, localLength = 0, remoteLength = 0, remoteOffset = 0;
     305            0 :     u32 offsetCounter = 0;
     306              : 
     307            0 :     if (isSingleMesh) {
     308            0 :         localOffset = myMeshAggregationSendRecvInfo[userRankInfo.userRank].recvOffset[infoIndex];
     309            0 :         localLength = myMeshAggregationSendRecvInfo[userRankInfo.userRank].recvLength[infoIndex];
     310            0 :         remoteLength = myMeshAggregationSendRecvInfo[infoIndex].sendLength[userRankInfo.userRank];
     311            0 :         remoteOffset = myMeshAggregationSendRecvInfo[infoIndex].sendOffset[userRankInfo.userRank];
     312              :     } else {
     313            0 :         for (u32 j = userRankInfo.userRank % meshAggregationRankSize; j < userRankInfo.userRankSize;
     314            0 :             j += meshAggregationRankSize) {
     315            0 :             CalcIntraMeshAggregationRecvInfoInMeshAggregation(j, infoIndex, myMeshAggregationSendRecvInfo, localOffset,
     316              :                 offsetCounter, localLength, remoteOffset, meshAggregationRankSize);
     317            0 :             if (offsetCounter == infoIndex || infoIndex == 0) {
     318              :                 break;
     319              :             }
     320              :         }
     321            0 :         remoteLength = localLength;
     322              :     }
     323            0 :     curRecvInfo.localOffset = localOffset;
     324            0 :     curRecvInfo.localLength = localLength;
     325              : 
     326            0 :     curRecvInfo.remoteOffset = remoteOffset;
     327            0 :     curRecvInfo.remoteLength = remoteLength;
     328            0 :     HCCL_DEBUG("[CalcIntraMeshAggregationRecvInfo] localOffset[%llu], localLength[%llu], "\
     329              :         "remoteOffset[%llu], remoteLength[%llu]", localOffset, localLength, remoteOffset, remoteLength);
     330            0 : }
     331              : 
     332            0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationAlltoAllMemInfo(const AlltoAllUserRankInfo &userRankInfo,
     333              :     const std::vector<SendRecvInfo> &allSendRecvInfo,
     334              :     std::map<u32, std::list<OneSendRecvAddrInfo>> &sendAddrInfosIntra,
     335              :     std::map<u32, std::list<OneSendRecvAddrInfo>> &recvAddrInfosIntra, u32 meshAggregationRankSize,
     336              :     const bool &isSingleMesh)
     337              : {
     338            0 :     sendAddrInfosIntra.clear();
     339            0 :     recvAddrInfosIntra.clear();
     340            0 :     if (allSendRecvInfo.size() != userRankInfo.userRankSize) {
     341            0 :         HCCL_ERROR("Invalid All send recv info size[%zu], should be[%u]", allSendRecvInfo.size(),
     342              :             userRankInfo.userRankSize);
     343            0 :         return;
     344              :     }
     345            0 :     SendRecvInfo mySendRecvInfo = allSendRecvInfo[userRankInfo.userRank];
     346            0 :     u32 rankInMeshAggregation = userRankInfo.userRank % meshAggregationRankSize;
     347            0 :     u32 cluserIndex = userRankInfo.userRank / meshAggregationRankSize;
     348            0 :     auto itBegin = allSendRecvInfo.begin();
     349            0 :     auto itEnd = allSendRecvInfo.begin();
     350            0 :     std::advance(itBegin, cluserIndex * meshAggregationRankSize);
     351            0 :     std::advance(itEnd, (cluserIndex + 1) * meshAggregationRankSize);
     352            0 :     std::vector<SendRecvInfo> myMeshAggregationSendRecvInfo(itBegin, itEnd);
     353              : 
     354            0 :     for (u32 i = 0; i < userRankInfo.userRankSize; i++) {
     355              :         // sendInfo 的计算
     356              :         OneSendRecvAddrInfo curSendInfo;
     357            0 :         u32 remoteRankInMeshAggregation = i % meshAggregationRankSize;
     358            0 :         CalcIntraMeshAggregationSendInfo(userRankInfo, mySendRecvInfo, myMeshAggregationSendRecvInfo,
     359              :             rankInMeshAggregation, i, curSendInfo, meshAggregationRankSize, isSingleMesh);
     360            0 :         sendAddrInfosIntra[remoteRankInMeshAggregation].push_back(curSendInfo);
     361              : 
     362              :         // recvInfo 的计算
     363              :         OneSendRecvAddrInfo curRecvInfo;
     364            0 :         CalcIntraMeshAggregationRecvInfo(userRankInfo, myMeshAggregationSendRecvInfo, i,
     365              :             curRecvInfo, meshAggregationRankSize, isSingleMesh);
     366            0 :         recvAddrInfosIntra[remoteRankInMeshAggregation].push_back(curRecvInfo);
     367              :     }
     368            0 : }
     369              : 
     370            0 : HcclOpMetaInfo CollAlltoAllExecutor::GetOpMeta(HcclCMDType opType, const u64 size)
     371              : {
     372            0 :     bool hugeData = size > SDMA_SEND_MAX_SIZE;
     373            0 :     HcclOpMetaInfoDef opMeta;
     374              : 
     375            0 :     if (isAlltoAllZCopyMode_) {
     376              :         /* zcopy拆分4GB以上SDMA任务前,准备好子图不复用标志 */
     377            0 :         if (opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
     378            0 :             opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, size, hugeData);
     379              :         } else {
     380            0 :             opMeta = HcclOpMetaInfo::GetOneForAllToAllVC(CopyPattern::ZCOPY, size, hugeData);
     381              :         }
     382              :     } else {
     383              :         /* bcopy每次重新生成子图 */
     384            0 :         if (opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
     385            0 :             opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::BCOPY, size, false);
     386              :         } else {
     387            0 :             opMeta = HcclOpMetaInfo::GetOneForAllToAllVC(CopyPattern::BCOPY, size, false);
     388              :         }
     389              :     }
     390              : 
     391            0 :     return opMeta;
     392              : }
     393              : 
     394            0 : u64 CollAlltoAllExecutor::CalAlltoAllVScratchMemSize(u64 &workSpaceMemSize)
     395              : {
     396            0 :     u64 scratchMemSize = 0U;
     397            0 :     if (workSpaceMemSize == 0) {
     398            0 :         scratchMemSize = TINY_MEM_SIZE;
     399            0 :         HCCL_DEBUG("[CalAlltoAllVScratchMemSize] workSpaceMemSize==0, use TINY_MEM_SIZE[%llu]",
     400              :             TINY_MEM_SIZE);
     401              :     } else {
     402            0 :         if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     403            0 :             scratchMemSize = std::max(std::max(workSpaceMemSize, inCCLbufferSize_), TINY_MEM_SIZE);
     404            0 :             HCCL_DEBUG("[CalAlltoAllVScratchMemSize] OpBase mode, workSpaceMemSize[%llu], "
     405              :                 "inCCLbufferSize_[%llu], scratchMemSize[%llu]",
     406              :                 workSpaceMemSize, inCCLbufferSize_, scratchMemSize);
     407              :         } else {
     408            0 :             scratchMemSize = workSpaceMemSize;
     409            0 :             HCCL_DEBUG("[CalAlltoAllVScratchMemSize] non-OpBase mode, workSpaceMemSize[%llu], "
     410              :                 "scratchMemSize[%llu]", workSpaceMemSize, scratchMemSize);
     411              :         }
     412              :     }
     413            0 :     return scratchMemSize;
     414              : }
     415              : 
     416            0 : bool CollAlltoAllExecutor::HasMassTasks(std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo)
     417              : {
     418            0 :     if (isAlltoAllZCopyMode_) {
     419            0 :         return false;
     420              :     }
     421              : 
     422            0 :     u64 maxSendTimes = 0;
     423            0 :     u64 maxRecvTimes = 0;
     424            0 :     const u64 cclBufferSize = algResResp_->cclInputMem.size();
     425            0 :     for (auto &sendRecvInfo : allMeshAggregationSendRecvInfo) {
     426            0 :         u64 sendTimes = 0;
     427            0 :         u64 recvTimes = 0;
     428            0 :         for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
     429            0 :             sendTimes += (sendRecvInfo.sendLength[i] + cclBufferSize - 1) / cclBufferSize;
     430            0 :             recvTimes += (sendRecvInfo.recvLength[i] + cclBufferSize - 1) / cclBufferSize;
     431              :         }
     432            0 :         maxSendTimes = (maxSendTimes > sendTimes) ? maxSendTimes : sendTimes;
     433            0 :         maxRecvTimes = (maxRecvTimes > recvTimes) ? maxRecvTimes : recvTimes;
     434              :     }
     435            0 :     const u64 massThreshold = 65535; //  65535: 单个ffts+任务中,最多承载64K个task
     436            0 :     const u64 maxTasksPerStep = 10;  // BCOPY中每次和远端通信最多消耗task数
     437            0 :     const u64 maxTasksBaseCost = 50; // BCOPY中除每步和远端通信外,最多消耗的task数
     438            0 :     u64 maxTasks = (maxSendTimes + maxRecvTimes) * maxTasksPerStep + maxTasksBaseCost;
     439            0 :     HCCL_DEBUG("[AlltoAll] bcopy maxSendTimes[%llu], maxRecvTimes[%llu], maxTasks[%llu], hasMassTask[%u]",
     440              :         maxSendTimes, maxRecvTimes, maxTasks, (maxTasks > massThreshold));
     441            0 :     return (maxTasks > massThreshold);
     442              : }
     443              : 
     444            5 : HcclResult CollAlltoAllExecutor::SetVirtualDispatcher(const HcclDispatcher virtualDispatcher)
     445              : {
     446            5 :     vDispatcher_ = virtualDispatcher;
     447            5 :     return HCCL_SUCCESS;
     448              : }
     449              : 
     450            1 : HcclResult CollAlltoAllExecutor::CheckNeedRecreateComm(u64 lastScratchMemSize, bool& needRecreateAlltoallComm)
     451              : {
     452            1 :     needRecreateAlltoallComm = false;
     453            1 :     return HCCL_SUCCESS;
     454              : }
     455              : 
     456            0 : HcclResult CollAlltoAllExecutor::RunAlltoAllTemplate(const std::unique_ptr<AlgTemplateBase> &executor,
     457              :     const SubCommInfo &commInfo)
     458              : {
     459            0 :     HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.links);
     460            0 :     CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunAlltoAllTemplate]" \
     461              :         "group has been destroyed. Break!"), ret);
     462            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     463              :         HCCL_ERROR("[CollAlltoAllExecutor][RunAlltoAllTemplate]run executor rank[%u] rank size[%u] failed",
     464              :         commInfo.localRank, commInfo.localRankSize), ret);
     465            0 :     return HCCL_SUCCESS;
     466              : }
     467              : 
     468            0 : HcclResult CollAlltoAllExecutor::RunAlltoAllVTemplateStaged(const std::unique_ptr<AlgTemplateBase> &executor,
     469              :     const SubCommInfo &commInfo)
     470              : {
     471            0 :     HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.links);
     472            0 :     CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunAlltoAllVTemplateStaged]" \
     473              :         "group has been destroyed. Break!"), ret);
     474            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     475              :         HCCL_ERROR("[CollAlltoAllExecutor][RunAlltoAllVTemplateStaged]run executor rank[%u] rank size[%u] failed",
     476              :         commInfo.localRank, commInfo.localRankSize), ret);
     477            0 :     return HCCL_SUCCESS;
     478              : }
     479              : 
     480              : // deprecated
     481            0 : HcclResult CollAlltoAllExecutor::RunTemplateWithVirtualLink(const std::unique_ptr<AlgTemplateBase> &executor,
     482              :     const SubCommInfo &commInfo)
     483              : {
     484            0 :     HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.virtualLinks);
     485            0 :     CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunTemplateWithVirtualLink]" \
     486              :         "group has been destroyed. Break!"), ret);
     487            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     488              :         HCCL_ERROR("[CollAlltoAllExecutor][RunTemplateWithVirtualLink]run executor rank[%u] rank size[%u] failed",
     489              :         commInfo.localRank, commInfo.localRankSize), ret);
     490            0 :     return HCCL_SUCCESS;
     491              : }
     492              : 
     493              : } // namespace hccl
        

Generated by: LCOV version 2.0-1