LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_executor/ins_alg_executor/all_to_all - ins_all_to_all_sole_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 50.6 % 156 79
Test Date: 2026-08-04 10:52:23 Functions: 11.1 % 63 7

            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 "log.h"
      12              : 
      13              : #include "ins_coll_alg_registry.h"
      14              : #include "ins_all_to_all_sole_executor.h"
      15              : 
      16              : #include "ccu_temp_all_to_all_mesh_1D.h"
      17              : #include "ccu_temp_all_to_all_v_mesh_1D.h"
      18              : #include "ccu_temp_all_to_all_mesh2d.h"
      19              : #include "ccu_temp_all_to_all_v_mesh_2D.h"
      20              : #include "ccu_temp_all_to_all_v_mesh_2Die.h"
      21              : 
      22              : namespace Hccl {
      23              : template <typename AlgTopoMatch, typename InsAlgTemplate>
      24            1 : InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::InsAlltoAllSoleExecutor() : InsCollAlgBase()
      25              : {
      26            1 : }
      27              : 
      28              : template <typename AlgTopoMatch, typename InsAlgTemplate>
      29            2 : InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::~InsAlltoAllSoleExecutor()
      30              : {
      31            2 : }
      32              : 
      33              : template <typename AlgTopoMatch, typename InsAlgTemplate>
      34            1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::InitParams(const CollAlgOperator &op, const CollAlgParams &params)
      35              : {
      36            1 :     opMode_        = params.opMode;
      37            1 :     maxTmpMemSize_ = params.maxTmpMemSize;
      38            1 :     CHK_PRT_RET((maxTmpMemSize_ == 0),
      39              :                 HCCL_ERROR("[InitParams] maxTmpMemSize equals to zero for OPBASE."), HcclResult::HCCL_E_PARA);
      40              : 
      41            1 :     CHK_PRT_RET(GetAlltoAllLocalSendRecvInfo(op, myRank_, rankSize_, localSendRecvInfo_), HCCL_ERROR("[InitParams] unable to init DataInfo."),
      42              :                 HcclResult::HCCL_E_PARA);
      43            1 :     if (op.opType == OpType::ALLTOALL) {
      44            1 :         sendType_ = op.all2AllDataDes.sendType;
      45            1 :         recvType_ = op.all2AllDataDes.recvType;
      46            0 :     } else if (op.opType == OpType::ALLTOALLV) {
      47            0 :         sendType_ = op.all2AllVDataDes.sendType;
      48            0 :         recvType_ = op.all2AllVDataDes.recvType;
      49            0 :     } else if (op.opType == OpType::ALLTOALLVC) {
      50            0 :         sendType_ = op.all2AllVCDataDes.sendType;
      51            0 :         recvType_ = op.all2AllVCDataDes.recvType;
      52            0 :     } else if (op.opType != OpType::HALFALLTOALLV) {
      53            0 :         HCCL_ERROR("[InsAlltoAllSoleExecutor] opType [%s] is invalid.", op.opType.Describe().c_str());
      54            0 :         return HcclResult::HCCL_E_PARA;
      55              :     }
      56            1 :     CHK_PRT_RET(InitOpInfo(op, opType_, redOp_, root_), HCCL_ERROR("[InitParams] unable to init OpInfo."),
      57              :                 HcclResult::HCCL_E_PARA);
      58            1 :     return HcclResult::HCCL_SUCCESS;
      59              : }
      60              : 
      61              : template <typename AlgTopoMatch, typename InsAlgTemplate>
      62            0 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::CalcResOffload(const RankGraph *rankGraph,
      63              :                                                                                   const u64 &dataSize,
      64              :                                                                                   CollOffloadOpResReq &resReq)
      65              : {
      66              :     (void)dataSize;
      67            0 :     resReq.requiredScratchMemSize = 200 * 1024 * 1024; //  200 * 1024*1024 = 200M
      68              : 
      69              :     // Topo Match
      70            0 :     AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
      71            0 :     CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
      72              : 
      73              :     // instantiate a template
      74            0 :     InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
      75              : 
      76            0 :     std::map<u32, u32>rank2PathNumMap;
      77            0 :     HCCL_INFO("[InsAlltoAllSoleExecutor] CalcResOffload SetPathNumMap");
      78            0 :     for(auto rankIdx : virtRanks_){
      79            0 :         if(rankIdx==myRank_){
      80            0 :             continue;
      81              :         }
      82            0 :         std::vector<NetInstance::Path> tmpPaths0 =
      83              :             rankGraph->GetPaths(0, myRank_, rankIdx);
      84            0 :         std::vector<NetInstance::Path> tmpPaths1 =
      85              :             rankGraph->GetPaths(1, myRank_, rankIdx);
      86            0 :         HCCL_INFO("[InsAlltoAllSoleExecutor]tmpPaths0.size() = %zu,tmpPaths1.size() = %zu", tmpPaths0.size(), tmpPaths1.size());
      87            0 :         rank2PathNumMap[rankIdx] = tmpPaths0.size() + tmpPaths1.size();
      88              :     }
      89              :     
      90            0 :     tempAlg.setPathNumMap(rank2PathNumMap);
      91              : 
      92              :     // calculate required insQueues and prepare queue
      93            0 :     AlgTempResReq tempResReq;
      94            0 :     if (enableDetour_) {
      95            0 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor], CalcResOffload with detouring enabled.");
      96            0 :         CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
      97              :     } else {
      98            0 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor], CalcResOffload with detouring disabled.");
      99            0 :         CHK_RET(tempAlg.CalcRes(tempResReq));
     100              :     }
     101              : 
     102            0 :     resReq.requiredSubQueNum = tempResReq.streamNum - 1;
     103              : 
     104            0 :     HCCL_INFO("[InsAlltoAllSoleExecutor][CalcResOffload] requiredSubQueNum[%llu], requiredScratchMemSize[%llu].",
     105              :                resReq.requiredSubQueNum, resReq.requiredScratchMemSize);
     106              : 
     107            0 :     return HcclResult::HCCL_SUCCESS;
     108            0 : }
     109              : 
     110              : template <typename AlgTopoMatch, typename InsAlgTemplate>
     111            1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::CalcRes(const RankGraph *rankGraph,
     112              :                                                                            CollAlgResReq     &algResReq)
     113              : {
     114              :     // Topo Match
     115            1 :     AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
     116            1 :     CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
     117            1 :     algResReq.topoInfo.UpdateSingleLevelTopo(virtRanks_, virtRankMap_, vTopo_);
     118            3 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor][CalcRes]topoInfo.virtRanks[%u], topoInfo.virtRankMap[%u], topoInfo.vTopo[%u].",
     119              :                algResReq.topoInfo.virtRanks.size(), algResReq.topoInfo.virtRankMap.size(), algResReq.topoInfo.vTopo.size());
     120              :     // instantiate a template
     121            1 :     InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
     122              : 
     123            1 :     std::map<u32, u32>rank2PathNumMap;
     124            3 :     HCCL_INFO("[InsAlltoAllSoleExecutor] CalcRes SetPathNumMap");
     125            8 :     for(auto rankIdx : virtRanks_){
     126            4 :         if(rankIdx==myRank_){
     127            1 :             continue;
     128              :         }
     129            3 :         std::vector<NetInstance::Path> tmpPaths0 =
     130              :             rankGraph->GetPaths(0, myRank_, rankIdx);
     131            3 :         std::vector<NetInstance::Path> tmpPaths1 =
     132              :             rankGraph->GetPaths(1, myRank_, rankIdx);
     133            3 :         rank2PathNumMap[rankIdx] = tmpPaths0.size() + tmpPaths1.size();
     134              :     }
     135            1 :     tempAlg.setPathNumMap(rank2PathNumMap);
     136              : 
     137              :     // calculate required insQues and prepare queue
     138            1 :     AlgTempResReq tempResReq;
     139            1 :     if (enableDetour_) {
     140            0 :         HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
     141            0 :         CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
     142              :     } else {
     143            3 :         HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
     144            1 :         CHK_RET(tempAlg.CalcRes(tempResReq));
     145              :     }
     146            1 :     CHK_RET(CalcLinkInfo(myRank_, rankGraph, tempResReq.links, algResReq.levelRankPairs));
     147            1 :     algResReq.primQueueNum= tempResReq.streamNum;
     148            1 :     algResReq.queueNotifys = tempResReq.queNotifys;
     149            3 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], requiredQueNum [%u].", myRank_, algResReq.primQueueNum);
     150            1 :     CHK_RET(CalcResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, algResReq.links));
     151            1 :     return HcclResult::HCCL_SUCCESS;
     152            1 : }
     153              : 
     154              : // dataSize_ as input
     155              : template <typename AlgTopoMatch, typename InsAlgTemplate>
     156            1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::Orchestrate(const RankGraph     *rankGraph,
     157              :                                                                               const CollAlgOperator &op,
     158              :                                                                               const CollAlgParams   &params,
     159              :                                                                               InsQuePtr              insQue)
     160              : {
     161              :     // init and check params
     162            1 :     CHK_RET(Init(op, params, insQue));
     163              :     // Topo Match
     164            1 :     AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
     165            1 :     CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
     166            3 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], [%s].", myRank_, topoMatch.Describe().c_str());
     167              : 
     168              :     // instantiate a template
     169            1 :     InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
     170            1 :     tempAlg.SetDmaMode(dmaMode_);
     171            1 :     tempAlg.SetCollOp(op);  // CCU template需要传递op信息
     172            1 :     tempAlg.SetA2ASendRecvInfo(localSendRecvInfo_);
     173            1 :     tempAlg.SetLoadInfo(params);
     174              : 
     175              :     // calculate required insQues and prepare queue
     176            1 :     AlgTempResReq tempResReq;
     177            1 :     if (enableDetour_) {
     178            0 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
     179            0 :         CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
     180              :     } else {
     181            3 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
     182            1 :         CHK_RET(tempAlg.CalcRes(tempResReq));
     183              :     }
     184              : 
     185            1 :     CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
     186            3 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], template [%s], requiredQue Num [%u].", myRank_,
     187              :                tempAlg.Describe().c_str(), tempResReq.queNum);
     188              : 
     189            1 :     CHK_RET(PrepResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, tempResLinks_));
     190              : 
     191            1 :     CHK_RET(OrchestrateOpbase(tempAlg));
     192              : 
     193            1 :     return HcclResult::HCCL_SUCCESS;
     194            1 : }
     195              : 
     196              : template <typename AlgTopoMatch, typename InsAlgTemplate>
     197            0 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::Orchestrate(const AlgTopoInfo     &topoInfo,
     198              :                                                                                  const CollAlgOperator &op,
     199              :                                                                                  const CollAlgParams   &params,
     200              :                                                                                  ConnectedLinkMgr      *linkMgr,
     201              :                                                                                  InsQuePtr              insQue)
     202              : {
     203            0 :     HCCL_INFO("[InsAlltoAllSoleExecutor] Begin to orchestrate.");
     204              :     // init and check params
     205            0 :     CHK_RET(Init(op, params, insQue));
     206            0 :     dataType_ = op.dataType;
     207              :     // instantiate a template
     208            0 :     if(topoInfo.vTopo.size() == 0) {
     209            0 :         HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank[%d], vTopo size is zero.", myRank_);
     210            0 :         return HcclResult::HCCL_E_PARA;
     211              :     }
     212            0 :     if(topoInfo.virtRankMap.size() == 0) {
     213            0 :         HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank[%d], virtRankMap size is zero.", myRank_);
     214            0 :         return HcclResult::HCCL_E_PARA;
     215              :     }
     216            0 :     InsAlgTemplate tempAlg(myRank_, rankSize_, topoInfo.vTopo[0], topoInfo.virtRankMap[0]);
     217            0 :     tempAlg.SetDmaMode(dmaMode_);
     218            0 :     tempAlg.SetCollOp(op);
     219            0 :     tempAlg.SetA2ASendRecvInfo(localSendRecvInfo_);
     220            0 :     tempAlg.SetLoadInfo(params);
     221            0 :     tempAlg.SetDataType(dataType_);
     222            0 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], Init insAlgTemplate with rankSize [%u] and dmaMode [%s].", myRank_,
     223              :                rankSize_, dmaMode_.Describe().c_str());
     224            0 :     virtRankMap_ = topoInfo.virtRankMap[0];
     225            0 :     virtRanks_ = topoInfo.virtRanks[0];
     226            0 :     std::map<u32, u32>rank2PathNumMap;
     227            0 :     for(u32 rankIdx:virtRanks_){
     228            0 :         auto links0 = linkMgr->GetLinks(0, rankIdx);
     229            0 :         auto links1 = linkMgr->GetLinks(1, rankIdx);
     230            0 :         if(links0.size() + links1.size() != 0){
     231            0 :             rank2PathNumMap[rankIdx] = links0.size() + links1.size();
     232              :         }
     233              :     }
     234            0 :     tempAlg.setPathNumMap(rank2PathNumMap);
     235              : 
     236              :     // calculate required insQues and prepare queue
     237            0 :     AlgTempResReq tempResReq;
     238            0 :     if (enableDetour_) {
     239            0 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
     240            0 :         CHK_RET(tempAlg.CalcResDetour(linkMgr, tempResReq));
     241              :     } else {
     242            0 :         HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
     243            0 :         CHK_RET(tempAlg.CalcRes(tempResReq));
     244              :     }
     245              : 
     246            0 :     CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
     247            0 :     HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], template [%s], requiredQue Num [%u].", myRank_,
     248              :                tempAlg.Describe().c_str(), tempResReq.queNum);
     249              : 
     250            0 :     CHK_RET(PrepResLinks(myRank_, tempResReq.links, linkMgr, tempResLinks_));
     251              : 
     252            0 :     CHK_RET(OrchestrateOpbase(tempAlg));
     253            0 :     HCCL_INFO("[InsAlltoAllSoleExecutor] Orchestrate success.");
     254              : 
     255            0 :     return HcclResult::HCCL_SUCCESS;
     256            0 : }
     257              : 
     258              : template <typename AlgTopoMatch, typename InsAlgTemplate>
     259            1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::OrchestrateOpbase(InsAlgTemplate &tempAlg)
     260              : {
     261            3 :     HCCL_DEBUG("[CollAlgFactory][InsAlltoAllSoleExecutor] AlgTemplate is [%s]", tempAlg.Describe().c_str());
     262            1 :     CHK_PRT_RET(maxTmpMemSize_ == 0,
     263              :                 HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank [%d], Invalid input maxTmpMemSize [%u].", myRank_, maxTmpMemSize_),
     264              :                 HcclResult::HCCL_E_PARA);
     265              : 
     266            1 :     CHK_RET(tempAlg.GetScratchBufferInfo(maxTmpMemSize_, sendType_));
     267              : 
     268            1 :     BuffInfo buffInfo;
     269            1 :     buffInfo.inBuffType     = BufferType::SCRATCH;
     270            1 :     buffInfo.outBuffType    = BufferType::SCRATCH;
     271            1 :     buffInfo.inBuffBaseOff  = 0;
     272            1 :     buffInfo.outBuffBaseOff = maxTmpMemSize_ / 2;  // 占据scratch memory的后半部分,除以2
     273            1 :     RankSliceInfo sliceInfoVec;
     274              : 
     275            1 :     TempFuncs tempFuncs;
     276            1 :     tempFuncs.opMode              = opMode_;
     277            1 :     tempFuncs.enableCounterNotify = IsEnableCounterNotifyByDevType(myRank_, devType_);
     278            1 :     tempFuncs.isForepart          = true; // Usr Buff to CCL Buff required
     279            1 :     tempFuncs.isBottom            = true; // CCL Buff to Usr Buff required
     280            1 :     CHK_RET(tempAlg.Run(tempFuncs, sliceInfoVec, buffInfo, tempResLinks_, requiredQue_));
     281            3 :     HCCL_INFO("[InsAlltoAllSoleExecutor][OrchestrateOpbase] Run templet success.");
     282              : 
     283            1 :     return HcclResult::HCCL_SUCCESS;
     284            1 : }
     285              : 
     286              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, InsAlltoAllMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
     287              :                           InsTempAlltoAllMesh);
     288              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, InsAlltoAllvMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
     289              :                           InsTempAlltoAllMesh);
     290              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLVC, InsAlltoAllvcMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
     291              :                           InsTempAlltoAllMesh);
     292              : #ifndef CCL_KERNEL_AICPU
     293              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, CcuAlltoAllMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
     294              :                         CcuTempAllToAllMesh1D);
     295              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
     296              :                         CcuTempAlltoAllVMesh1D);
     297              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, CcuAlltoAllMesh2D, InsAlltoAllSoleExecutor, TopoMatchConcurrMesh,
     298              :                         CcuTempAlltoAllMesh2D);
     299              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh2D, InsAlltoAllSoleExecutor, TopoMatchConcurrMesh,
     300              :                         CcuTempAlltoAllVMesh2D);
     301              : INS_REGISTER_IMPL_BY_TEMP(OpType::HALFALLTOALLV, CcuHalfAll2AllVMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
     302              :                         CcuTempHalfAllToAllVMesh1D);
     303              : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh2Die, InsAlltoAllSoleExecutor, TopoMatchMesh,
     304              :                         CcuTempAlltoAllVMesh2Die);
     305              : #endif
     306              : 
     307              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1