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

Generated by: LCOV version 2.0-1