LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_to_all - coll_all_to_all_v_direct_fullmesh_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 264 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 19 0

            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_v_direct_fullmesh_executor.h"
      12              : 
      13              : namespace hccl {
      14              : 
      15            0 : CollRunAlltoAllDirectFullmesh::CollRunAlltoAllDirectFullmesh(const HcclDispatcher dispatcher,
      16            0 :                                                    std::unique_ptr<TopoMatcher> &topoMatcher)
      17            0 :     : CollAlltoAllExecutor(dispatcher, topoMatcher)
      18              : {
      19            0 : }
      20              : 
      21            0 : HcclResult CollRunAlltoAllDirectFullmesh::Orchestrate(OpParam& param, AlgResourceResponse& algRes)
      22              : {
      23            0 :     HcclUs startut = TIME_NOW();
      24            0 :     HcclResult ret = HCCL_SUCCESS;
      25            0 :     tag_ = param.tag;
      26            0 :     algResResp_ = &algRes;
      27            0 :     AlltoAllVParam_ = param;
      28              : 
      29            0 :     ExecMem execMem;
      30            0 :     execMem.count = 0;
      31            0 :     execMem.inputPtr = param.inputPtr;
      32            0 :     execMem.outputPtr = param.outputPtr;
      33            0 :     execMem.inputMem = algRes.cclInputMem;
      34            0 :     execMem.outputMem = algRes.cclOutputMem;
      35            0 :     ret = KernelRun(param, execMem);
      36              : 
      37            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
      38              :         HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][Orchestrate]errNo[0x%016llx]executor run failed",
      39              :             HCCL_ERROR_CODE(ret)), ret);
      40              :     
      41              :     // Enforce task launch at the end of Orchestrate
      42              :     // 注意: 不要删除这里的强制launch, 否则会导致aicpu cache功能问题
      43            0 :     HCCL_INFO("%s: enforce task launch at the end of Orchestrate", __func__);
      44            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
      45              : 
      46            0 :     HCCL_INFO("tag[%s], AlltoAllDirectFullmesh tempAlg orchestrate success, take time [%lld]us.",
      47              :         param.tag.c_str(), DURATION_US(TIME_NOW() - startut));
      48            0 :     return HCCL_SUCCESS;
      49            0 : }
      50              : 
      51            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetAdjInfo(AlgResourceResponse& algRes, AdjInfo& adjInfo)
      52              : {
      53            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] GetAdjInfo.");
      54            0 :     algResResp_ = &algRes;
      55            0 :     SubCommInfo levelCommInfo = {0};
      56            0 :     AdjInfo nslbAdjInfo = {0};
      57            0 :     u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
      58              : 
      59            0 :     u32 localRank= topoAttr_.userRank;
      60            0 :     u32 localRankSize = topoAttr_.userRankSize;
      61              : 
      62            0 :     std::unique_ptr<AlgTemplateBase> levelTempAlg;
      63              : 
      64            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] SelectTempAlg.");
      65            0 :     levelTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
      66            0 :             TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
      67            0 :     CHK_SMART_PTR_NULL(levelTempAlg);
      68            0 :     u32 rankIdxInPod = INVALID_VALUE_RANKID;
      69            0 :     CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
      70              : 
      71            0 :     if (devNumInlocalPod == INVALID_VALUE_RANKSIZE) {
      72            0 :         HCCL_INFO("[GetAdjInfo-nslbdp] devNumInlocalPod == INVALID_VALUE_RANKSIZE.");
      73            0 :         return HCCL_SUCCESS;
      74              :     }
      75              : 
      76            0 :     nslbAdjInfo.dstRankNum = devNumInlocalPod;
      77            0 :     CHK_RET(levelTempAlg->GetNslbAdjInfo(localRank, localRankSize, levelCommInfo.links, nslbAdjInfo));
      78              : 
      79            0 :     adjInfo.dstRankNum = nslbAdjInfo.dstRankNum;
      80            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] adjInfo.dstRankNum[%u].", adjInfo.dstRankNum);
      81              :     
      82            0 :     for (size_t i = 0; i < nslbAdjInfo.nsAdjInfo.size(); i++) {
      83            0 :         NslbDpAdjInfo dpAdjInfo = {0};
      84            0 :         dpAdjInfo.dstLocalRankId = nslbAdjInfo.nsAdjInfo[i].dstLocalRankId;
      85            0 :         dpAdjInfo.phaseId = nslbAdjInfo.nsAdjInfo[i].phaseId;
      86            0 :         dpAdjInfo.rev = 0;
      87            0 :         adjInfo.nsAdjInfo.push_back(dpAdjInfo); 
      88            0 :         HCCL_INFO("[nslbdp]GetAdjInfo dstLocalRankId[%u], phaseId[%u].",
      89              :                    nslbAdjInfo.nsAdjInfo[i].dstLocalRankId, nslbAdjInfo.nsAdjInfo[i].phaseId);
      90              :     }
      91            0 :     return HCCL_SUCCESS;
      92            0 : }
      93              : 
      94            0 : HcclResult CollRunAlltoAllDirectFullmesh::MarkNeedAlltoallvCache() {
      95            0 :     needAlltoallvCache_ = true;
      96            0 :     HCCL_INFO("[CollRunAlltoAllDirectFullmesh][MarkNeedAlltoallvCache] set needAlltoallvCache_[%u]"\
      97              :         "for alltoallv aicpu cache", needAlltoallvCache_);
      98            0 :     return HCCL_SUCCESS;
      99              : }
     100              : 
     101            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetHcclOffsetDstRanksMap(std::unordered_map<uint64_t,
     102              :     std::vector<uint32_t>>& hcclOffsetDstRanksMap) const {
     103            0 :     hcclOffsetDstRanksMap.clear();
     104            0 :     hcclOffsetDstRanksMap = hcclOffsetDstRanksMap_; // Deep copy
     105              : 
     106            0 :     return HCCL_SUCCESS;
     107              : }
     108              : 
     109            0 : HcclOpMetaInfo CollRunAlltoAllDirectFullmesh::GetOpMeta(HcclCMDType opType, const u64 size)
     110              : {
     111              :     (void)opType;
     112            0 :     HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, size, true);
     113            0 :     return opMeta;
     114              : }
     115              : 
     116            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSDMAGroupInfo(const u32 userRank,
     117              :     u32& devNumInlocalPod, u32& rankIdxInPod)
     118              : {
     119              :     (void) userRank;
     120            0 :     bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B &&
     121            0 :                             !topoAttr_.isSingleMeshAggregation;
     122            0 :     if (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) {
     123            0 :         CHK_RET(topoMatcher_->GetLocalServerRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
     124              :     } else {
     125            0 :         CHK_RET(topoMatcher_->GetLocalSuperPodRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
     126              :     }
     127            0 :     CHK_PRT_RET(devNumInlocalPod == INVALID_VALUE_RANKSIZE,
     128              :         HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][GetLocalSDMAGroupInfo]get local superPod total ranksize failed."),
     129              :         HCCL_E_PARA);
     130            0 :     return HCCL_SUCCESS;
     131              : }
     132              : 
     133            0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcStreamNum(u32& streamNum)
     134              : {
     135              :     // 每个超节点内的卡数
     136            0 :     u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
     137            0 :     u32 rankIdxInPod = INVALID_VALUE_RANKID;
     138            0 :     CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
     139              : 
     140              :     // 单超节点场景需要的从流数量
     141            0 :     streamNum = (devNumInlocalPod > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
     142            0 :         (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE * RANK_SET_COMPUTE_CONST) : (devNumInlocalPod * RANK_SET_COMPUTE_CONST);
     143              : 
     144              :     // 多超节点场景下,RDMA会设置独立的并发度
     145            0 :     if ((topoAttr_.userRankSize - devNumInlocalPod) > 0) {
     146            0 :         streamNum += 1; // 一条从流专门用来管理超节点间的RDMA通信
     147            0 :         u32 totalRdmaRankNum = topoAttr_.userRankSize - devNumInlocalPod;
     148            0 :         streamNum += (totalRdmaRankNum > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
     149            0 :             (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) : (totalRdmaRankNum);
     150              :     }
     151              : 
     152            0 :     HCCL_INFO("[CollRunAlltoAllDirectFullmesh][CalcStreamNum] tag[%s] streamNum[%u]",
     153              :         tag_.c_str(), streamNum);
     154            0 :     return HCCL_SUCCESS;
     155              : }
     156              : 
     157              : // level0-level1 打平fullmesh
     158              : // 超节点内建SDMA链路;超节点间建RDMA链路
     159            0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcLevel0CommInfo(TransportMemType inputType, TransportMemType outputType,
     160              :     std::vector<LevelNSubCommTransport>& opTransport)
     161              : {
     162            0 :     CommParaInfo commCombinePara(COMM_COMBINE_ORDER, CommType::COMM_TAG_MESH);
     163            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commCombinePara, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
     164              : 
     165            0 :     LevelNSubCommTransport &commTransportLevel0 = opTransport[COMM_COMBINE_ORDER];
     166            0 :     for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
     167            0 :         for (auto &transportRequest : commTransportLevel0[subCommIndex].transportRequests) {
     168            0 :             transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
     169              :         }
     170              :     }
     171            0 :     return HCCL_SUCCESS;
     172            0 : }
     173              : 
     174            0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcTransportMemType(TransportMemType &inputType, TransportMemType &outputType)
     175              : {
     176            0 :     inputType = TransportMemType::CCL_INPUT;
     177            0 :     outputType = TransportMemType::CCL_OUTPUT;
     178              : 
     179            0 :     HCCL_INFO("[CollRunAlltoAllDirectFullmesh][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
     180              :         tag_.c_str(), inputType, outputType);
     181            0 :     return HCCL_SUCCESS;
     182              : }
     183              : 
     184            0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
     185              : {
     186            0 :     TransportMemType inputType = TransportMemType::RESERVED;
     187            0 :     TransportMemType outputType = TransportMemType::RESERVED;
     188              : 
     189            0 :     CHK_RET(CalcTransportMemType(inputType, outputType));
     190              :     // level0 - level1 全连接通信域
     191            0 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
     192            0 :     return HCCL_SUCCESS;
     193              : }
     194              : 
     195            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoallV(const OpParam &param)
     196              : {
     197              :     // 注意: 如果send/recv info的计算逻辑发生变化, 需要同步修改framework下的IsSmallDataAlltoallv()函数
     198            0 :     for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
     199            0 :         u64 curSendCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCounts) + j);
     200            0 :         u64 curSendDispls = *(static_cast<const u64 *>(param.All2AllDataDes.sdispls) + j);
     201            0 :         localSendRecvInfo_.sendCounts[j] = curSendCounts;
     202            0 :         localSendRecvInfo_.sendDispls[j] = curSendDispls;
     203            0 :         localSendRecvInfo_.sendLength[j] = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
     204            0 :         localSendRecvInfo_.sendOffset[j] = curSendDispls * SIZE_TABLE[param.All2AllDataDes.sendType];
     205              : 
     206            0 :         u64 curRecvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.recvCounts) + j);
     207            0 :         u64 curRecvDispls = *(static_cast<const u64 *>(param.All2AllDataDes.rdispls) + j);
     208            0 :         localSendRecvInfo_.recvCounts[j] = curRecvCounts;
     209            0 :         localSendRecvInfo_.recvDispls[j] = curRecvDispls;
     210            0 :         localSendRecvInfo_.recvLength[j] = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
     211            0 :         localSendRecvInfo_.recvOffset[j] = curRecvDispls * SIZE_TABLE[param.All2AllDataDes.recvType];
     212              : 
     213            0 :         HCCL_DEBUG("GetLocalSendRecvInfoforAlltoallV rank[%u], sendCounts[%llu], sendDispls[%llu] "\
     214              :             "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
     215              :             localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
     216              :             localSendRecvInfo_.recvDispls[j]);
     217              :     }
     218            0 :     return HCCL_SUCCESS;
     219              : }
     220              : 
     221            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoall(const OpParam &param)
     222              : {
     223            0 :     u64 curSendDispls = 0;
     224            0 :     u64 curSendOffset = 0;
     225            0 :     u64 curRecvDispls = 0;
     226            0 :     u64 curRecvOffset = 0;
     227            0 :     for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
     228            0 :         u64 curSendCounts = param.All2AllDataDes.sendCount;
     229            0 :         u64 curSendLength = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
     230            0 :         localSendRecvInfo_.sendCounts[j] = curSendCounts;
     231            0 :         localSendRecvInfo_.sendDispls[j] = curSendDispls;
     232            0 :         localSendRecvInfo_.sendLength[j] = curSendLength;
     233            0 :         localSendRecvInfo_.sendOffset[j] = curSendOffset;
     234            0 :         curSendDispls += curSendCounts;
     235            0 :         curSendOffset += curSendLength;
     236              : 
     237            0 :         u64 curRecvCounts = param.All2AllDataDes.sendCount;
     238            0 :         u64 curRecvLength = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
     239            0 :         localSendRecvInfo_.recvCounts[j] = curRecvCounts;
     240            0 :         localSendRecvInfo_.recvDispls[j] = curRecvDispls;
     241            0 :         localSendRecvInfo_.recvLength[j] = curRecvLength;
     242            0 :         localSendRecvInfo_.recvOffset[j] = curRecvOffset;
     243            0 :         curRecvDispls += curRecvCounts;
     244            0 :         curRecvOffset += curRecvLength;
     245            0 :         HCCL_DEBUG("GetLocalSendRecvInfoforAlltoAll rank[%u], sendCounts[%llu], sendDispls[%llu] "\
     246              :             "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
     247              :             localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
     248              :             localSendRecvInfo_.recvDispls[j]);
     249              :     }
     250            0 :     return HCCL_SUCCESS;
     251              : }
     252              : 
     253            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoallVC(const OpParam &param)
     254              : {
     255            0 :     u64 curSendDispls = 0;
     256            0 :     u64 curSendOffset = 0;
     257            0 :     u64 curRecvDispls = 0;
     258            0 :     u64 curRecvOffset = 0;
     259            0 :     u64 rankSize = topoAttr_.userRankSize;
     260            0 :     u64 usrRank = topoAttr_.userRank;
     261            0 :     for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
     262            0 :         u64 curSendCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank * rankSize + j);
     263            0 :         u64 curSendLength = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
     264            0 :         localSendRecvInfo_.sendCounts[j] = curSendCounts;
     265            0 :         localSendRecvInfo_.sendDispls[j] = curSendDispls;
     266            0 :         localSendRecvInfo_.sendLength[j] = curSendLength;
     267            0 :         localSendRecvInfo_.sendOffset[j] = curSendOffset;
     268            0 :         curSendDispls += curSendCounts;
     269            0 :         curSendOffset += curSendLength;
     270              : 
     271            0 :         u64 curRecvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank + rankSize * j);
     272            0 :         u64 curRecvLength = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
     273            0 :         localSendRecvInfo_.recvCounts[j] = curRecvCounts;
     274            0 :         localSendRecvInfo_.recvDispls[j] = curRecvDispls;
     275            0 :         localSendRecvInfo_.recvLength[j] = curRecvLength;
     276            0 :         localSendRecvInfo_.recvOffset[j] = curRecvOffset;
     277            0 :         curRecvDispls += curRecvCounts;
     278            0 :         curRecvOffset += curRecvLength;
     279            0 :         HCCL_DEBUG("GetLocalSendRecvInfoforAlltoallVC rank[%u], sendCounts[%llu], sendDispls[%llu] "\
     280              :             "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
     281              :             localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
     282              :             localSendRecvInfo_.recvDispls[j]);
     283              :     }
     284            0 :     return HCCL_SUCCESS;
     285              : }
     286              : 
     287            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetAlltoAllvTmpRankSendRecvInfo(const OpParam &param)
     288              : {
     289            0 :     localSendRecvInfo_.sendCounts.resize(topoAttr_.userRankSize, 0);
     290            0 :     localSendRecvInfo_.sendDispls.resize(topoAttr_.userRankSize, 0);
     291            0 :     localSendRecvInfo_.sendLength.resize(topoAttr_.userRankSize, 0);
     292            0 :     localSendRecvInfo_.sendOffset.resize(topoAttr_.userRankSize, 0);
     293              : 
     294            0 :     localSendRecvInfo_.recvCounts.resize(topoAttr_.userRankSize, 0);
     295            0 :     localSendRecvInfo_.recvDispls.resize(topoAttr_.userRankSize, 0);
     296            0 :     localSendRecvInfo_.recvLength.resize(topoAttr_.userRankSize, 0);
     297            0 :     localSendRecvInfo_.recvOffset.resize(topoAttr_.userRankSize, 0);
     298            0 :     if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
     299            0 :         CHK_RET(GetLocalSendRecvInfoforAlltoallV(param));
     300            0 :     } else if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALL) {
     301            0 :         CHK_RET(GetLocalSendRecvInfoforAlltoall(param));
     302            0 :     } else if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALLVC) {
     303            0 :         CHK_RET(GetLocalSendRecvInfoforAlltoallVC(param));
     304              :     } else {
     305            0 :         HCCL_ERROR("Only support optype AllToAll , AllToAllV and AllToAllVC !");
     306              :     }
     307            0 :     return HCCL_SUCCESS;
     308              : }
     309              : 
     310            0 : HcclResult CollRunAlltoAllDirectFullmesh::KernelRun(const OpParam &param, ExecMem &execMem)
     311              : {
     312            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] AllToAll fullmesh start.", __func__);
     313              : 
     314              :     // 准备数据
     315            0 :     CHK_RET(ActiveSlaveStreams(param.stream));
     316            0 :     CHK_RET(GetAlltoAllvTmpRankSendRecvInfo(param));
     317              : 
     318              :     // 获取当前超节点内总卡数
     319            0 :     u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
     320            0 :     u32 rankIdxInPod = INVALID_VALUE_RANKID;
     321            0 :     CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
     322              : 
     323              :     // 获取通信域
     324            0 :     CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_INDEX_0 + 1));
     325            0 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_COMBINE_ORDER, COMM_INDEX_0);
     326            0 :     bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B &&
     327            0 :                             !topoAttr_.isSingleMeshAggregation;
     328              :     // isSuPodAsym 表示A2A3卡数不一致场景或者A3多超节点server数不同场景
     329            0 :     bool isSuPodAsym = false;
     330            0 :     if (topoAttr_.superPodNum > 1) {
     331            0 :         isSuPodAsym = (static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode) || static_cast<bool>(topoAttr_.multiSuperPodDiffServerNumMode));
     332              :     } else {
     333            0 :         isSuPodAsym = (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) && static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode);
     334              :     }
     335              : 
     336              :     // 执行
     337              :     // 注意: 如果使用了非AlltoAllVDirectFullMesh的算法模板, 需要同步修改framework中的NeedOpUnfoldCache()函数
     338            0 :     std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     339            0 :         TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
     340            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH in COMM_COMBINE_ORDER", __func__);
     341            0 :     CHK_SMART_PTR_NULL(tempAlg);
     342              : 
     343            0 :     PrepareData prepareData;
     344            0 :     prepareData.stream = param.stream;
     345            0 :     prepareData.userRank = topoAttr_.userRank;
     346            0 :     prepareData.userRankSize = topoAttr_.userRankSize;
     347            0 :     prepareData.linksPtr = &level0CommInfo.links;
     348            0 :     prepareData.localSendRecvInfoPtr = &localSendRecvInfo_;
     349            0 :     prepareData.devNumInlocalPod = devNumInlocalPod;
     350            0 :     prepareData.rankIdxInPod = rankIdxInPod;
     351              : 
     352            0 :     prepareData.inputMem = algResResp_->paramInputMem;
     353            0 :     prepareData.outputMem = algResResp_->paramOutputMem;
     354            0 :     prepareData.cclInMem = execMem.inputMem;
     355            0 :     prepareData.cclOutMem = execMem.outputMem;
     356            0 :     prepareData.workMode = workflowMode_;
     357            0 :     prepareData.subStreamsPtr = &algResResp_->slaveStreams;
     358            0 :     prepareData.signalPtr = &algResResp_->notifiesMain;
     359            0 :     prepareData.signalAuxPtr = &algResResp_->notifiesAux;
     360            0 :     prepareData.isSuPodAsym = isSuPodAsym;
     361            0 :     prepareData.opType = param.opType;
     362            0 :     prepareData.algOpContext = algOpContext_;
     363              : 
     364              :     // 如果使能alltoallv aicpu cache
     365            0 :     if (needAlltoallvCache_) {
     366              :         // 注意: 一定是alltoallv类算子才有可能设置needAlltoallvCache_, 让alltoallv temp alg感知cache并保存算法中间结果
     367            0 :         CHK_PRT_RET(!(param.opType == HcclCMDType::HCCL_CMD_ALLTOALLV || param.opType == HcclCMDType::HCCL_CMD_ALLTOALLVC),
     368              :             HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][KernelRun] needAlltoallvCache_[%u] opType[%u]",
     369              :                 needAlltoallvCache_, param.opType),
     370              :             HCCL_E_INTERNAL);
     371              : 
     372              :         // 使能alltoallv temp alg感知alltoallv aicpu cache
     373            0 :         prepareData.needAlltoallvCache = true;
     374              :     } else {
     375            0 :         prepareData.needAlltoallvCache = false;
     376              :     }
     377              : 
     378            0 :     CHK_RET(tempAlg->Prepare(prepareData));
     379              : 
     380            0 :     CHK_RET(tempAlg->RunAsync());
     381              : 
     382            0 :     if (needAlltoallvCache_) {
     383              :         // 在tempAlg被销毁前保存hcclOffset-dstRank之间的mapping信息
     384              :         // 注意: CollRunAlltoAllDirectFullmesh executor使用的一定是AlltoAllVDirectFullMesh template
     385            0 :         hcclOffsetDstRanksMap_.clear();
     386            0 :         HCCL_INFO("[CollRunAlltoAllDirectFullmesh][KernelRun] get hcclOffset-dstRanks mapping for AlltoAllVDirectFullMesh");
     387            0 :         CHK_RET(tempAlg->GetHcclOffsetDstRanksMap(hcclOffsetDstRanksMap_));
     388              :     }
     389              : 
     390            0 :     HCCL_INFO("[CollRunAlltoAllDirectFullmesh] executor run success.");
     391            0 :     if (algOpContext_.opRetryHandler.isPostSync == true) {
     392            0 :         OpParam postSyncParam = param;
     393            0 :         if ((*prepareData.subStreamsPtr).size() == 0) {
     394            0 :             CHK_RET(PostSyncWithoutSubstream(postSyncParam, execMem));
     395              :         } else {
     396            0 :             PrepareData postSyncPrepareData = prepareData;
     397            0 :             CHK_RET(PostSyncWithSubstream(postSyncParam, execMem, postSyncPrepareData));
     398            0 :         }
     399            0 :     }
     400            0 :     return HCCL_SUCCESS;
     401            0 : }
     402            0 : HcclResult CollRunAlltoAllDirectFullmesh::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     403              : {
     404            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] Getlevel1CommRank userRank[%u]--userRankSize[%u].",topoAttr_.userRank, topoAttr_.userRankSize);
     405            0 :     level1CommInfo.localRank = topoAttr_.userRank;
     406            0 :     level1CommInfo.localRankSize = topoAttr_.userRankSize;
     407            0 :     return HCCL_SUCCESS;
     408              : }
     409              : 
     410            0 : HcclResult CollRunAlltoAllDirectFullmesh::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
     411              : {
     412              :     (void) level1RankSize;
     413            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] SelectTempAlg.");
     414            0 :     level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     415            0 :             TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
     416            0 :     CHK_SMART_PTR_NULL(level1TempAlg);
     417              : 
     418            0 :     return HCCL_SUCCESS;
     419              : }
     420              : 
     421            0 : HcclResult CollRunAlltoAllDirectFullmesh::GetDevNumInlocalPod(u32& devNumInlocalPod)
     422              : {
     423            0 :     HCCL_INFO("[GetAdjInfo-nslbdp] GetDevNumInlocalPod.");
     424              :     // 获取当前超节点内总卡数
     425            0 :     u32 rankIdxInPod = INVALID_VALUE_RANKID;
     426            0 :     CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
     427              : 
     428            0 :     return HCCL_SUCCESS;
     429              : }
     430              : REGISTER_EXEC("RunAlltoAllDirectFullmesh", AlltoAllVDirectFullMesh, CollRunAlltoAllDirectFullmesh);
     431              : } // namespace hccl
        

Generated by: LCOV version 2.0-1