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

Generated by: LCOV version 2.0-1