LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template - asymmetric_hierarchical_concatenate_alg_template_base.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 296 0
Test Date: 2026-07-28 12:11:00 Functions: 0.0 % 31 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 "asymmetric_hierarchical_concatenate_alg_template_base.h"
      12              :  
      13              : #include <iostream>
      14              : #include <fstream>
      15              :  
      16              : namespace hccl {
      17              :  
      18            0 : AHCAlgTemplateBase::AHCAlgTemplateBase(const HcclDispatcher dispatcher)
      19            0 :     : AlgTemplateBase(dispatcher), needTraslateSliceAddr_(false), rankSize_(1), extendFlag_(false)
      20              : {
      21            0 : }
      22              :  
      23            0 : AHCAlgTemplateBase::~AHCAlgTemplateBase()
      24              : {
      25            0 : }
      26              : 
      27            0 : HcclResult AHCAlgTemplateBase::Prepare(u64 reduceAttrBitMap, HcomCollOpInfo *opInfo)
      28              : {
      29            0 :     reduceAttr_ = reduceAttrBitMap;
      30            0 :     return HCCL_SUCCESS;
      31              : }
      32              : 
      33            0 : HcclResult AHCAlgTemplateBase::Prepare(u64 totalCount, const std::vector<std::vector<std::vector<u32>>> &globalSubGroups,
      34              :     std::map<AHCConcOpType, TemplateType> &ahcAlgOption, bool extendFlag, AHCExtendPreparePara extendPara)
      35              : {
      36            0 :     globalSubGroups_ = globalSubGroups;
      37            0 :     totalCount_ = totalCount;
      38            0 :     ahcAlgOption_ = ahcAlgOption;
      39            0 :     extendFlag_ = extendFlag;
      40            0 :     ahcExtendPreparePara_ = extendPara;
      41            0 :     return HCCL_SUCCESS;
      42              : }
      43              : 
      44            0 : HcclResult AHCAlgTemplateBase::DisposeSubGroups(const u32 rank)
      45              : {
      46            0 :     return HCCL_SUCCESS;
      47              : }
      48              :  
      49            0 : HcclResult AHCAlgTemplateBase::CommAHCInfoInit()
      50              : {
      51            0 :     return HCCL_SUCCESS;
      52              : }
      53              : 
      54            0 : HcclResult AHCAlgTemplateBase::PrepareRunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      55              : {
      56            0 :     HcclResult ret = HCCL_SUCCESS;
      57            0 :     CHK_SMART_PTR_NULL(dispatcher_);
      58            0 :     CHK_PTR_NULL(stream_.ptr());
      59            0 :     CHK_PRT_RET(!outputMem_ || !inputMem_,
      60              :         HCCL_ERROR("[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] run_async inputmem or outputmem is null", rank), HCCL_E_PTR);
      61              :  
      62            0 :     HCCL_INFO("AHCAlgTemplateBase run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", \
      63              :               rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
      64              :  
      65            0 :     CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] linksize[%llu] is less "\
      66              :         "than rankSize[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
      67              :  
      68              :     // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
      69            0 :     if (rankSize == 1) {
      70            0 :         if (inputMem_ != outputMem_) {
      71            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
      72            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
      73              :                 HCCL_ERROR("[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] memcpy async failed", rank), ret);
      74              :         }
      75            0 :         return ret;
      76              :     }
      77              :  
      78            0 :     DisposeSubGroups(rank);
      79              :  
      80            0 :     rankSize_ = rankSize;
      81              :  
      82            0 :     CommAHCInfoInit();
      83              :     
      84              :     // 保存物理slice,可能非连续
      85            0 :     physicalSlices_ = slices_;
      86              : 
      87              :     // 检查、并清空逻辑slices_
      88            0 :     if (slices_.size() != 0) {
      89            0 :         HCCL_DEBUG("[AHCAlgTemplateBase][PrepareRunAsync] clear logic slice_");
      90            0 :         slices_.clear();
      91              :     }
      92              :  
      93            0 :     return HCCL_SUCCESS;
      94              : }
      95              : 
      96            0 : HcclResult AHCAlgTemplateBase::GetNslbAdjInfoPro(const u32 rank, const u32 rankSize,
      97              :                                        const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
      98              : {
      99            0 :     HCCL_DEBUG("[NSLB-AHC] entry GetNslbAdjInfoPro");
     100            0 :     DisposeSubGroups(rank);
     101            0 :     CommAHCInfoInit();
     102              :     
     103            0 :     if (rankSize == 1 || links.size() < rankSize) {
     104            0 :         return HCCL_SUCCESS;
     105              :     }
     106            0 :     u32 nSteps  = 0;
     107            0 :     std::vector<u32> dstRanks;
     108            0 :     HCCL_DEBUG("[NSLB-AHC] try to GetNslbDstRanks, rank = %u, ranksize = %u", rank, rankSize);
     109            0 :     CHK_RET(commAHCBaseInfo_->GetNslbDstRanks(rank, dstRanks));
     110            0 :     if (dstRanks.size() == 0 || dstRanks.size() > NSLBDP_MAX_PHASE) {
     111            0 :         HCCL_DEBUG("[NSLB-AHC]  dstRanks size not support");
     112            0 :         return HCCL_SUCCESS;
     113              :     }
     114            0 :     for (u32 nextRank : dstRanks) {
     115            0 :         LINK linkRight = links[nextRank];
     116            0 :         CHK_SMART_PTR_NULL(linkRight);
     117            0 :         NslbDpAdjInfo adjInfoStep = {0, 0, 0};
     118            0 :         adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
     119            0 :         adjInfoStep.phaseId = nSteps + 1;
     120            0 :         adjInfoStep.rev = 0;
     121            0 :         nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
     122            0 :         nSteps ++;
     123            0 :     }
     124            0 :     nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
     125            0 :     return HCCL_SUCCESS;
     126            0 : }
     127              : 
     128              : 
     129            0 : HcclResult AHCAlgTemplateBase::PrepareAlgTemplate(std::unique_ptr<AlgTemplateBase> &tempAlg, const std::vector<Slice> &slices, AHCOpType opType)
     130              : {
     131            0 :     HcclResult ret = HCCL_SUCCESS;
     132            0 :     switch (opType) {
     133            0 :         case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER: {
     134            0 :             ret = tempAlg->Prepare(inputMem_, inputMem_, scratchMem_, count_, dataType_,
     135            0 :                 stream_, reductionOp_, root_, slices, baseOffset_);
     136            0 :             break;
     137              :         }
     138            0 :         case AHCOpType::AHC_OP_TYPE_ALLGATHER: {
     139            0 :             ret = tempAlg->Prepare(outputMem_, outputMem_, scratchMem_, count_, dataType_,
     140            0 :                 stream_, reductionOp_, root_, slices, baseOffset_);
     141            0 :             break;
     142              :         }
     143            0 :         case AHCOpType::AHC_OP_TYPE_ALLREDUCE: {
     144            0 :             ret = tempAlg->Prepare(inputMem_, outputMem_, scratchMem_, count_, dataType_,
     145            0 :                 stream_, reductionOp_, root_, slices, baseOffset_);
     146            0 :             break;
     147              :         }
     148            0 :         case AHCOpType::AHC_OP_TYPE_RESERVED:{
     149              :             // 其他算子不支持,直接返回
     150            0 :             ret = HCCL_E_PARA;
     151            0 :             break;
     152              :         }
     153              :     }
     154              : 
     155            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     156              :         HCCL_ERROR("[AHCAlgTemplateBase][PrepareAlgTemplate] prepare step error"), ret);
     157              :      
     158            0 :     return ret;
     159              : }
     160              :  
     161            0 : HcclResult AHCAlgTemplateBase::MemcpyForSingleOp(const u32 rank, AHCOpType opType)
     162              : {
     163            0 :     HcclResult ret = HCCL_SUCCESS;
     164            0 :     u32 commRank = commAHCBaseInfo_->GetCommRank(rank);
     165            0 :     HCCL_DEBUG("[AHCAlgTemplateBase][MemcpyForSingleOp] rank[%u] commRank[%u]", rank, commRank);
     166            0 :     switch (opType) {
     167            0 :         case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER:{
     168            0 :             u64 srcSize = (inputMem_.size() - commRank * count_ * DataUnitSize(dataType_)) > count_ * DataUnitSize(dataType_) ?
     169            0 :                 count_ * DataUnitSize(dataType_) : (inputMem_.size() - commRank * count_ * DataUnitSize(dataType_)); 
     170            0 :             DeviceMem srcMem = inputMem_.range(commRank * count_ * DataUnitSize(dataType_), srcSize);
     171            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, srcMem, stream_);
     172            0 :             break;
     173            0 :         }
     174            0 :         case AHCOpType::AHC_OP_TYPE_ALLGATHER:{
     175            0 :             u64 dstSize = (outputMem_.size() - commRank * count_ * DataUnitSize(dataType_)) > count_ * DataUnitSize(dataType_) ?
     176            0 :                 count_ * DataUnitSize(dataType_) : (outputMem_.size() - commRank * count_ * DataUnitSize(dataType_));
     177            0 :             DeviceMem dstMem = outputMem_.range(commRank * count_ * DataUnitSize(dataType_), dstSize);
     178            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, inputMem_, stream_);
     179            0 :             break;
     180            0 :         }
     181            0 :         case AHCOpType::AHC_OP_TYPE_ALLREDUCE:{
     182              :             // 使用 RS+AG 实现 AR 时,需要在 RS 完成时进行一次额外的数据搬运
     183            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, inputMem_, outputMem_, stream_);
     184            0 :             break;
     185              :         }
     186            0 :         case AHCOpType::AHC_OP_TYPE_RESERVED:{
     187              :             // 其他算子不支持,无需copy,直接返回
     188            0 :             break;
     189              :         }
     190              :     }
     191            0 :     return ret;
     192              : }
     193              :  
     194            0 : HcclResult AHCAlgTemplateBase::RunInstance(const u32 rank, const std::vector<LINK> &links, std::vector<Slice> &slices,
     195              :         std::unique_ptr<AlgTemplateBase> &tempAlg, AHCOpType opType)
     196              : {
     197            0 :     HcclResult ret = HCCL_SUCCESS;
     198              :  
     199              :     // 判断是否关闭reducescatter的barrier
     200            0 :     if (!barrierSwitchOn_) {
     201            0 :         tempAlg->CloseBarrier();
     202              :     }
     203              : 
     204              :     //地址映射
     205            0 :     if (needTraslateSliceAddr_) {
     206            0 :         ret = commAHCBaseInfo_->TrasLogicSliceToPhysical(slices, physicalSlices_);
     207            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     208              :             HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] optype[%d] translate slice failed", rank, opType), ret);
     209              :     }
     210              : 
     211              :     // 调用算法执行
     212            0 :     ret = PrepareAlgTemplate(tempAlg, slices, opType);
     213            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     214              :         HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] prepare optype[%d] failed", rank, opType), ret);
     215              :   
     216            0 :     ret = tempAlg->RegisterProfiler(
     217            0 :         profilerInput_.planeID, profilerInput_.stage, profilerInput_.step, stream_);
     218            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     219              :         HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] registerProfiler optype[%d] failed", rank, opType), ret);
     220              :     
     221            0 :     ret = tempAlg->RunAsync(rank, links.size(), links);
     222            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     223              :         HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] run optype[%d] failed", rank, opType), ret);
     224              :  
     225            0 :     return ret;
     226              : }
     227              :  
     228            0 : ReduceScatterAHCBase::ReduceScatterAHCBase(const HcclDispatcher dispatcher)
     229            0 :     : AHCAlgTemplateBase(dispatcher)
     230              : {
     231            0 : }
     232              :  
     233            0 : ReduceScatterAHCBase::~ReduceScatterAHCBase()
     234              : {
     235            0 : }
     236              :  
     237            0 : HcclResult ReduceScatterAHCBase::RunAsync(const u32 rank, const u32 rankSize,
     238              :     const std::vector<LINK> &links)
     239              : {
     240            0 :     HCCL_INFO("[ReduceScatterAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
     241              :  
     242            0 :     HcclResult ret = HCCL_SUCCESS;
     243            0 :     ret = PrepareRunAsync(rank, rankSize, links);
     244            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(), outputMem_.size(), count_);
     245              :  
     246              :     // 设置地址翻译标记有效,并计算逻辑的totalsize
     247            0 :     needTraslateSliceAddr_ = true;
     248            0 :     commAHCBaseInfo_->ParseInputSlice(physicalSlices_);
     249              :  
     250            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     251              :         HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_), ret);
     252              :  
     253            0 :     CHK_PRT_RET(rankSize == 1, HCCL_INFO("[ReduceScatterAHCBase][RunAsync] rankSize[%u], do nothing.",
     254              :         rankSize), HCCL_SUCCESS);
     255              :  
     256            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] begin intra rs", rank);
     257              :  
     258              :     // 做组内 reduce-scatter
     259            0 :     ret = RunIntraReduceScatter(rank, links, commAHCBaseInfo_);
     260            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     261              :         "RunIntraReduceScatter  step", rank, count_), ret);
     262              :  
     263            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] end intra rs begin inter", rank);
     264              :  
     265              :     // 做组间 reduce-scatter
     266            0 :     ret = RunInterReduceScatter(rank, links, commAHCBaseInfo_);
     267            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     268              :         "RunInterReduceScatter step", rank, count_), ret);
     269              :  
     270              :     // 对于单独的 Reduce-scatter 算子,在运算结束时进行数据搬运
     271            0 :     if (inputMem_ != outputMem_) {
     272            0 :         ret = MemcpyForSingleOp(rank, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER);
     273            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     274              :             HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] memcpy failed", rank), ret);
     275              :     }
     276              :  
     277            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] end inter rs", rank);
     278              :  
     279            0 :     HCCL_INFO("[ReduceScatterAHCBase][RunAsync] finished: rank[%u]", rank);
     280            0 :     return HCCL_SUCCESS;
     281              : }
     282              : 
     283            0 : HcclResult ReduceScatterAHCBase::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
     284              :                                         const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
     285              : {
     286            0 :     return GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo);
     287              : }
     288              : 
     289            0 : HcclResult ReduceScatterAHCBase::RunIntraReduceScatter(const u32 rank, const std::vector<LINK> &links,
     290              :     const std::unique_ptr<CommAHCBaseInfo> &commAHCBaseInfo)
     291              : {
     292              :     // 获取当前rank的组内rank
     293            0 :     HcclResult ret = HCCL_SUCCESS;
     294            0 :     HCCL_INFO("[ReduceScatterAHC][RunIntraReduceScatter] begin intra ReduceScatter rank[%u] count[%llu]", rank, count_);
     295              :  
     296            0 :     u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
     297              :  
     298              :     // 创建执行算子实列
     299            0 :     std::unique_ptr<AlgTemplateBase> tempAlg;
     300            0 :     commAHCBaseInfo->GetIntraAlgTemplateOpInstance(AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER, tempAlg, dispatcher_, reduceAttr_,
     301            0 :         extendFlag_, ahcExtendPreparePara_);
     302              :  
     303            0 :     std::vector<std::vector<Slice>> intraSlicesVector;
     304            0 :     std::vector<std::vector<LINK>> intraLinksVector;
     305            0 :     CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(rank, DataUnitSize(dataType_), count_, links, intraLinksVector, intraSlicesVector));
     306              :  
     307            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunIntraReduceScatter] run inst rank[%u] intraRank[%u]",
     308              :         rank, intraRank);
     309              :  
     310            0 :     for (u32 i = 0; i < intraLinksVector.size(); i++) {
     311            0 :         std::vector<Slice> intraSlices = intraSlicesVector[i];
     312            0 :         std::vector<LINK> intraLinks = intraLinksVector[i];
     313            0 :         if (intraLinks.size() <= 1 ) {
     314            0 :             continue;
     315              :         }
     316            0 :         CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER));
     317            0 :     }
     318              :  
     319            0 :     HCCL_DEBUG("[ReduceScatterAHCBase][RunIntraReduceScatter] end intra ReduceScatter rank[%u]", rank);
     320              :  
     321            0 :     return ret;
     322            0 : }
     323              :  
     324            0 : AllGatherAHCBase::AllGatherAHCBase(const HcclDispatcher dispatcher)
     325            0 :     : AHCAlgTemplateBase(dispatcher)
     326              : {
     327            0 : }
     328              :  
     329            0 : AllGatherAHCBase::~AllGatherAHCBase()
     330              : {
     331            0 : }
     332              :  
     333            0 : HcclResult AllGatherAHCBase::RunAsync(const u32 rank, const u32 rankSize,
     334              :     const std::vector<LINK> &links)
     335              : {
     336            0 :     HCCL_INFO("[AllGatherAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
     337              :  
     338            0 :     HcclResult ret = HCCL_SUCCESS;
     339            0 :     ret = PrepareRunAsync(rank, rankSize, links);
     340            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(), outputMem_.size(), count_);
     341              : 
     342              :     // 设置地址翻译标记有效,并计算逻辑的totalsize
     343            0 :     needTraslateSliceAddr_ = true;
     344            0 :     commAHCBaseInfo_->ParseInputSlice(physicalSlices_);
     345              :  
     346            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     347              :         HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_), ret);
     348              :  
     349            0 :     CHK_PRT_RET(rankSize == 1, HCCL_INFO("[AllGatherAHCBase][RunAsync] rankSize[%u], do nothing.",
     350              :         rankSize), HCCL_SUCCESS);
     351              :  
     352            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] begin intra ag", rank);
     353              :  
     354              :     // 对于单独的 All-gather 算子,在运算开始时进行数据搬运
     355            0 :     if (inputMem_ != outputMem_) {
     356            0 :         ret = MemcpyForSingleOp(rank, AHCOpType::AHC_OP_TYPE_ALLGATHER);
     357            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     358              :             HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] memcpy failed", rank), ret);
     359              :     }
     360              :  
     361              :     // 做组间 all-gather
     362            0 :     ret = RunInterAllGather(rank, links, commAHCBaseInfo_);
     363            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     364              :         "RunInterAllGather step", rank, count_), ret);
     365              :  
     366            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] end inter ag", rank);
     367              :  
     368              :     // 做组内 allgather
     369            0 :     ret = RunIntraAllGather(rank, links, commAHCBaseInfo_);
     370            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     371              :         "RunIntraAllGather  step", rank, count_), ret);
     372              :  
     373            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] end intra ag begin inter", rank);
     374              :  
     375            0 :     HCCL_INFO("[AllGatherAHCBase][RunAsync] finished: rank[%u]", rank);
     376            0 :     return HCCL_SUCCESS;
     377              : }
     378              : 
     379            0 : HcclResult AllGatherAHCBase::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
     380              :                                         const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
     381              : {
     382            0 :     return GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo);
     383              : }
     384              : 
     385            0 : HcclResult AllGatherAHCBase::RunIntraAllGather(const u32 rank, const std::vector<LINK> &links,
     386              :     const std::unique_ptr<CommAHCBaseInfo> &commAHCBaseInfo)
     387              : {
     388              :     // 获取当前rank的组内rank
     389            0 :     HCCL_INFO("[AllGatherAHCBase][RunIntraAllGather] begin intra AllGather rank[%u]", rank);
     390              :  
     391            0 :     u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
     392              :  
     393              :     // 创建执行算子实列
     394            0 :     std::unique_ptr<AlgTemplateBase> tempAlg;
     395            0 :     commAHCBaseInfo->GetIntraAlgTemplateOpInstance(AHCOpType::AHC_OP_TYPE_ALLGATHER, tempAlg, dispatcher_, reduceAttr_,
     396            0 :         extendFlag_, ahcExtendPreparePara_);
     397              :  
     398            0 :     std::vector<std::vector<Slice>> intraSlicesVector;
     399            0 :     std::vector<std::vector<LINK>> intraLinksVector;
     400            0 :     CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(rank, DataUnitSize(dataType_), count_, links, intraLinksVector, intraSlicesVector));
     401              :  
     402            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunIntraAllGather] run inst rank[%u] intraRank[%u]",
     403              :         rank, intraRank);
     404              :  
     405            0 :     for (u32 i = 0; i < intraLinksVector.size(); i++) {
     406            0 :         std::vector<Slice> intraSlices = intraSlicesVector[i];
     407            0 :         std::vector<LINK> intraLinks = intraLinksVector[i];
     408            0 :         if (intraLinks.size() <= 1) {
     409            0 :             continue;
     410              :         }
     411            0 :         CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_ALLGATHER));
     412            0 :     }
     413              :  
     414            0 :     HCCL_DEBUG("[AllGatherAHCBase][RunIntraAllGather] end intra AllGather rank[%u]", rank);
     415              :  
     416            0 :     return HCCL_SUCCESS;
     417            0 : }
     418              : 
     419            0 : AllReduceAHCBase::AllReduceAHCBase(const HcclDispatcher dispatcher)
     420            0 :     : AHCAlgTemplateBase(dispatcher)
     421              : {
     422            0 : }
     423              :  
     424            0 : AllReduceAHCBase::~AllReduceAHCBase()
     425              : {
     426            0 : }
     427              :  
     428            0 : HcclResult AllReduceAHCBase::RunAsync(const u32 rank, const u32 rankSize,
     429              :     const std::vector<LINK> &links)
     430              : {  
     431            0 :     HCCL_INFO("[AllReduceAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
     432              :  
     433            0 :     HcclResult ret = HCCL_SUCCESS;
     434            0 :     ret = PrepareRunAsync(rank, rankSize, links);
     435            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(), outputMem_.size(), count_);
     436              : 
     437            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     438              :         HCCL_ERROR("[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_), ret);
     439              :  
     440            0 :     CHK_PRT_RET(rankSize == 1, HCCL_INFO("[AllReduceAHCBase][RunAsync] rankSize[%u], do nothing.",
     441              :         rankSize), HCCL_SUCCESS);
     442              :  
     443            0 :     CHK_PRT_RET(count_ == 0, HCCL_INFO("[AllReduceAHCBase][RunAsync] count_[%llu], do nothing.", count_), HCCL_SUCCESS);
     444              : 
     445            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] begin intra rs", rank);
     446              :  
     447            0 :     ret = RunIntraReduceScatter(rank, links, commAHCBaseInfo_);
     448            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     449              :         "RunIntraReduceScatter  step", rank, count_), ret);
     450              :  
     451            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end intra rs begin inter", rank);
     452              :  
     453              :     // 垂直方向做allreduce ring
     454            0 :     ret = RunInterAllReduce(rank, links, commAHCBaseInfo_);
     455            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     456              :         "RunInterAllReduce step", rank, count_), ret);
     457              :  
     458            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end inter begin intra ag", rank);
     459              :  
     460              :     // 水平方向做broken allgather ring
     461            0 :     ret = RunIntraAllGather(rank, links, commAHCBaseInfo_);
     462            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "\
     463              :         "RunIntraAllGather step", rank, count_), ret);
     464              :  
     465            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end intra ag", rank);
     466              :  
     467            0 :     HCCL_INFO("[AllReduceAHCBase][RunAsync] finished: rank[%u]", rank);
     468            0 :     return HCCL_SUCCESS;
     469              : }
     470              :  
     471            0 : HcclResult AllReduceAHCBase::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
     472              :                                         const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
     473              : {
     474              :     //获取reducescatter的部分
     475            0 :     CHK_RET(GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo));
     476              :     //后续模拟all_gather的部分
     477            0 :     HCCL_INFO("[NSLB-AHC]try to get allgather part");
     478            0 :     if(nslbAdjInfo.dstRankNum == 0 || nslbAdjInfo.nsAdjInfo.size() == 0) {
     479            0 :         HCCL_INFO("[NSLB-AHC] get reducescatter part is null");
     480            0 :         return HCCL_SUCCESS;
     481              :     }
     482            0 :     uint16_t nsteps = nslbAdjInfo.nsAdjInfo.size();
     483              : 
     484            0 :     for (size_t index = 0; index < nsteps; index ++) {
     485            0 :         NslbDpAdjInfo adjInfoStep = {0, 0, 0};
     486            0 :         adjInfoStep.dstLocalRankId = nslbAdjInfo.nsAdjInfo[nsteps - index - 1].dstLocalRankId;
     487            0 :         adjInfoStep.phaseId = nsteps + index + 1;
     488            0 :         adjInfoStep.rev = 0;
     489            0 :         nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
     490            0 :         nsteps ++;
     491              :     }
     492            0 :     nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
     493            0 :     HCCL_INFO("[NSLB-AHC]success to get allgather part");
     494            0 :     return HCCL_SUCCESS;
     495              : }
     496              : 
     497            0 : HcclResult AllReduceAHCBase::RunIntraReduceScatter(const u32 rank, const std::vector<LINK> &links,
     498              :     const std::unique_ptr<CommAHCBaseInfo> &commAHCBaseInfo)
     499              : {
     500              :     // 获取当前rank的组内rank
     501            0 :     HCCL_INFO("[AllReduceAHCBase][RunIntraReduceScatter] begin intra ReduceScatter rank[%u]", rank);
     502              :  
     503            0 :     u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
     504              :  
     505              :     // 创建执行算子实列
     506            0 :     std::unique_ptr<AlgTemplateBase> tempAlg;
     507            0 :     commAHCBaseInfo->GetIntraAlgTemplateOpInstance(AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER, tempAlg, dispatcher_, reduceAttr_,
     508            0 :         extendFlag_, ahcExtendPreparePara_);
     509              :  
     510            0 :     std::vector<Slice> intraSlices;
     511            0 :     std::vector<LINK> intraLinks;
     512            0 :     CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(rank, DataUnitSize(dataType_), count_, links, intraLinks, intraSlices));
     513              :  
     514              :     // 长度不足2,直接跳过
     515            0 :     if (intraLinks.size() <= 1) {
     516            0 :         return HCCL_SUCCESS;
     517              :     }
     518              :  
     519            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunIntraReduceScatter] run inst rank[%u] intraRank[%u], IntraSize=%u",
     520              :         rank, intraRank, intraLinks.size());
     521              :  
     522            0 :     CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER));
     523              :  
     524            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunIntraReduceScatter] end intra ReduceScatter rank[%u]", rank);
     525              :  
     526            0 :     return HCCL_SUCCESS;
     527            0 : }
     528              :  
     529            0 : HcclResult AllReduceAHCBase::RunIntraAllGather(const u32 rank, const std::vector<LINK> &links,
     530              :     const std::unique_ptr<CommAHCBaseInfo> &commAHCBaseInfo)
     531              : {
     532            0 :     HCCL_INFO("[AllReduceAHCBase][RunIntraAllGather] begin intra allgather rank[%u]", rank);
     533              :  
     534              :     // 获取当前rank的组内rank
     535            0 :     u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
     536              :  
     537              :     // 创建执行算子实列
     538            0 :     std::unique_ptr<AlgTemplateBase> tempAlg;
     539            0 :     commAHCBaseInfo->GetIntraAlgTemplateOpInstance(AHCOpType::AHC_OP_TYPE_ALLGATHER, tempAlg, dispatcher_, reduceAttr_,
     540            0 :         extendFlag_, ahcExtendPreparePara_);
     541              :  
     542            0 :     std::vector<Slice> intraSlices;
     543            0 :     std::vector<LINK> intraLinks;
     544              :  
     545            0 :     CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(rank, DataUnitSize(dataType_), count_, links, intraLinks, intraSlices));
     546              :  
     547              :     // 长度不足2,直接跳过
     548            0 :     if (intraLinks.size() <= 1) {
     549            0 :         return HCCL_SUCCESS;
     550              :     }
     551              :  
     552            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunIntraAllGather] run inst rank[%u] intraRank[%u], IntraSize=%u",
     553              :         rank, intraRank, intraLinks.size());
     554              :  
     555            0 :     CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_ALLGATHER));
     556              :  
     557            0 :     HCCL_DEBUG("[AllReduceAHCBase][RunIntraAllGather] end intra allgather rank[%u]", rank);
     558            0 :     return HCCL_SUCCESS;
     559            0 : }
     560              :  
     561              : }   // ~~ namespace hccl
        

Generated by: LCOV version 2.0-1