LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_reduce_scatter - coll_reduce_scatter_ring_zerocopy_exchange_pipeline_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 322 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 21 0

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "coll_reduce_scatter_ring_zerocopy_exchange_pipeline_executor.h"
      12              : 
      13              : namespace hccl {
      14              : 
      15            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::CollReduceScatterRingZerocopyExchangePipelineExecutor(
      16            0 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      17            0 :     : CollReduceScatterExecutor(dispatcher, topoMatcher)
      18              : {
      19            0 :     CCLMemSlice_ = false;
      20            0 :     DMAReduceFlag_ = true;   // 设为true,以禁用RunLoop中的本地拷贝
      21            0 :     desc_.isZeroCopy = true; // 执行RunLoop的KernelRunInterServer分支
      22            0 :     desc_.deterministic = 1;
      23            0 :     desc_.level1SupportedAlgos = {
      24              :         AlgTypeLevel1::ALG_LEVEL1_RING,
      25              :         AlgTypeLevel1::ALG_LEVEL1_NHR,
      26              :         AlgTypeLevel1::ALG_LEVEL1_NB,
      27            0 :     };
      28            0 :     desc_.level2SupportedAlgos = {AlgTypeLevel2::ALG_LEVEL2_PIPELINE};
      29            0 : }
      30              : 
      31            0 : void CollReduceScatterRingZerocopyExchangePipelineExecutor::ParseParam(const OpParam& param)
      32              : {
      33            0 :     tag_ = param.tag;
      34            0 :     root_ = param.root;
      35            0 :     aicpuUnfoldMode_ = param.aicpuUnfoldMode;
      36            0 :     opType_ = param.opType;
      37              : 
      38            0 :     u32 unitSize = SIZE_TABLE[param.DataDes.dataType];
      39            0 :     totalSize_ = topoAttr_.userRankSize * param.DataDes.count * unitSize;
      40            0 : }
      41              : 
      42            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcStreamNum(u32& streamNum)
      43              : {
      44              :     // level0 需要的stream数,double ring需要2条,single ring需要1条直接用主流
      45            0 :     u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE : 0;
      46              :     // level1 用NHR等ring算法,需要1条stream。但level0与level1串行,直接用主流
      47              :     // level2 用单ring,level2与level0/level1并行,需要1条额外的流
      48            0 :     totalStreamNum += 1;
      49            0 :     streamNum = totalStreamNum;
      50            0 :     HCCL_INFO("[CalcStreamNum] tag[%s] streamNum[%u] topoType_[%d]", tag_.c_str(), streamNum, topoType_);
      51              : 
      52            0 :     return HCCL_SUCCESS;
      53              : }
      54              : 
      55              : HcclResult
      56            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      57              : {
      58            0 :     HCCL_INFO(
      59              :         "[CalcCommInfo] tag[%s] algoLevel0[%d] algoLevel1[%d] algoLevel2[%d]", tag_.c_str(), algType_.algoLevel0,
      60              :         algType_.algoLevel1, algType_.algoLevel2);
      61              : 
      62            0 :     TransportMemType inputType = TransportMemType::CCL_INPUT;
      63            0 :     TransportMemType outputType = TransportMemType::CCL_OUTPUT;
      64            0 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
      65            0 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
      66            0 :     CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
      67            0 :     CHK_RET(CalcExchangeCommInfo(opTransport));
      68            0 :     return HCCL_SUCCESS;
      69              : }
      70              : 
      71            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcLevel0CommInfo(
      72              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      73              : {
      74            0 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
      75            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
      76            0 :     LevelNSubCommTransport& commTransportLevel0 = opTransport[COMM_LEVEL0];
      77            0 :     for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
      78            0 :         commTransportLevel0[subCommIndex].isZeroCopy = true;
      79              :     }
      80            0 :     return HCCL_SUCCESS;
      81            0 : }
      82              : 
      83            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcExchangeCommInfo(
      84              :     std::vector<LevelNSubCommTransport>& opTransport)
      85              : {
      86            0 :     std::set<u32> commTargetUserRankSet;
      87            0 :     u32 remoteRankSend = 0;
      88            0 :     u32 remoteRankRecv = 0;
      89              : 
      90            0 :     CHK_RET(CalExchangeRemoteRank(remoteRankSend, remoteRankRecv));
      91            0 :     HCCL_INFO(
      92              :         "[CalcExchangeCommInfo] tag[%s] userRank[%u] remoteRankSend[%u] remoteRankRecv[%u]", tag_.c_str(),
      93              :         topoAttr_.userRank, remoteRankSend, remoteRankRecv);
      94            0 :     commTargetUserRankSet.insert(remoteRankSend);
      95            0 :     commTargetUserRankSet.insert(remoteRankRecv);
      96              :     CommParaInfo commParaInfo(
      97              :         COMM_COMBINE_ORDER, CommType::COMM_TAG_PARTIAL_MESH_COMBINED, INVALID_VALUE_RANKID, INVALID_VALUE_RANKID, false,
      98            0 :         false, commTargetUserRankSet);
      99              : 
     100            0 :     TransportMemType inputType = TransportMemType::CCL_INPUT;
     101            0 :     TransportMemType outputType = TransportMemType::CCL_OUTPUT;
     102              : 
     103            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaInfo, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
     104            0 :     LevelNSubCommTransport& commTransport = opTransport[COMM_COMBINE_ORDER];
     105            0 :     for (u32 subCommIndex = 0; subCommIndex < commTransport.size(); subCommIndex++) {
     106            0 :         for (auto& transportRequest : commTransport[subCommIndex].transportRequests) {
     107            0 :             transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
     108              :         }
     109              :     }
     110            0 :     return HCCL_SUCCESS;
     111            0 : }
     112              : 
     113              : HcclResult
     114            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::CalExchangeRemoteRank(u32& remoteRankSend, u32& remoteRankRecv)
     115              : {
     116            0 :     u32 l2Size = topoAttr_.superPodNum;
     117            0 :     CHK_PRT_RET(l2Size == 0, HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level2RankSize is 0"), HCCL_E_PARA);
     118            0 :     u32 l1Size = topoAttr_.serverNum / l2Size;
     119            0 :     CHK_PRT_RET(l1Size == 0, HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level1RankSize is 0"), HCCL_E_PARA);
     120            0 :     u32 l0Size = topoAttr_.userRankSize / l2Size / l1Size;
     121            0 :     CHK_PRT_RET(l0Size == 0, HCCL_ERROR("[CalExchangeRemoteRank] invalid rank size, level0RankSize is 0"), HCCL_E_PARA);
     122              : 
     123              :     // 根据rankId计算出坐标(i, j, k)
     124            0 :     u32 l2Index = topoAttr_.userRank / l1Size / l0Size;
     125            0 :     u32 l1Index = (topoAttr_.userRank % (l1Size * l0Size)) / l0Size;
     126            0 :     u32 l0Index = topoAttr_.userRank % l0Size;
     127              : 
     128              :     // 计算本端将要发送数据的目标rank
     129            0 :     remoteRankSend = l2Index * l1Size * l0Size + l0Index * l1Size + l1Index;
     130              : 
     131              :     // 计算本端将要接收数据的目标rank
     132            0 :     u32 r = l1Index * l0Size + l0Index; // 超节点内相对rankid
     133            0 :     l0Index = r / l1Size;
     134            0 :     l1Index = r % l1Size;
     135            0 :     remoteRankRecv = l2Index * l1Size * l0Size + l1Index * l0Size + l0Index;
     136            0 :     return HCCL_SUCCESS;
     137              : }
     138              : 
     139            0 : u64 CollReduceScatterRingZerocopyExchangePipelineExecutor::CalcLoopMaxCount(const u32 unitSize)
     140              : {
     141            0 :     u64 maxCountPerLoop
     142            0 :         = ((inCCLbufferSize_ / topoAttr_.serverNum / HCCL_MIN_SLICE_ALIGN) * HCCL_MIN_SLICE_ALIGN) / unitSize;
     143            0 :     return maxCountPerLoop;
     144              : }
     145              : 
     146              : HcclResult
     147            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::KernelRunIntraServerPre(const OpParam& param, ExecMem& execMem)
     148              : {
     149              :     (void)execMem;
     150            0 :     CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, unitSize_));
     151            0 :     CHK_RET(GetCommRankInfoNormal(
     152              :         level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, false));
     153            0 :     CHK_RET(CalExchangeRemoteRank(exchangeRemoteRankSend_, exchangeRemoteRankRecv_));
     154              : 
     155            0 :     HCCL_INFO(
     156              :         "[KernelRunIntraServerPre] rank[%u:%u,%u,%u], rankSize[%u, %u, %u] exchange remoteRank[send:%u Recv:%u]",
     157              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, level2RankSize_, level1RankSize_, level0RankSize_,
     158              :         exchangeRemoteRankSend_, exchangeRemoteRankRecv_);
     159            0 :     return HCCL_SUCCESS;
     160              : }
     161              : 
     162              : HcclResult
     163            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::KernelRunInterServer(const OpParam& param, ExecMem& execMem)
     164              : {
     165            0 :     curSize_ = execMem.count * unitSize_;
     166            0 :     HCCL_INFO(
     167              :         "[CollReduceScatterRingZerocopyExchangePipelineExecutor] run start, rank[%u:%u,%u,%u], curSize_[%llu]",
     168              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, curSize_);
     169              : 
     170            0 :     for (u32 step = 0; step < level2RankSize_; step++) {
     171            0 :         if (!intraServerDone_) {
     172              :             // 只有第一个loop才需要执行节点内RS
     173            0 :             CHK_RET(RunIntraServer(param, execMem, step));
     174              :         }
     175              : 
     176              :         // 准备节点间RS的数据,user in搬运到ccl in
     177            0 :         CHK_RET(RunInterServerPreProcess(param, execMem, step));
     178              :         // 超节点内、节点间通信执行RS,编排在主流上
     179            0 :         if (level1RankSize_ > 1) {
     180              :             // 节点间RS完成后数据在ccl in
     181            0 :             CHK_RET(RunInterServer(param, execMem, step));
     182              :         }
     183              :         // 数据最终在ccl out
     184            0 :         CHK_RET(RunInterServerPostProcess(param, execMem, step));
     185              : 
     186              :         // 从steep 1开始要进行reduce,将本轮超节点间获取的数据与本轮超节点内的数据进行reduce
     187            0 :         if ((step > 0) && (level2RankSize_ > 1)) {
     188            0 :             CHK_RET(RunSuperPodPostSync(param));
     189              :             // 超节点间通信 与 超节点内通信 都完成后,本地进行reduce操作
     190            0 :             CHK_RET(RunSuperPodAndInterServerPostProcess(param, execMem, step));
     191              :         }
     192              : 
     193            0 :         if (step < (level2RankSize_ - 1)) {
     194              :             // 超节点间通信, 编排在最后一个slaveStreams上
     195            0 :             CHK_RET(RunSuperPodPreSync(param));
     196            0 :             CHK_RET(RunSuperPod(param, execMem, step + 1));
     197              :         }
     198              :     }
     199              : 
     200              :     // 将最终数据从ccl out搬到user out
     201            0 :     CHK_RET(RunFinallyProcess(param, execMem));
     202              : 
     203            0 :     intraServerDone_ = true;
     204            0 :     HCCL_INFO(
     205              :         "[CollReduceScatterRingZerocopyExchangePipelineExecutor] run success, rank[%u:%u,%u,%u]", topoAttr_.userRank,
     206              :         level2Rank_, level1Rank_, level0Rank_);
     207            0 :     return HCCL_SUCCESS;
     208              : }
     209              : 
     210            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodPreSync(const OpParam& param)
     211              : {
     212            0 :     Stream stream = param.stream;
     213            0 :     Stream slaveStream = algResResp_->slaveStreams.back();
     214              :     // 主流RS完成后,通知超节点间通信开始
     215            0 :     CHK_RET(LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux.back(), INVALID_VALUE_STAGE));
     216              :     // 从流等待超节点内RS完成
     217            0 :     CHK_RET(LocalNotify::Wait(slaveStream, dispatcher_, algResResp_->notifiesAux.back(), INVALID_VALUE_STAGE));
     218            0 :     return HCCL_SUCCESS;
     219            0 : }
     220              : 
     221            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodPostSync(const OpParam& param)
     222              : {
     223            0 :     Stream stream = param.stream;
     224            0 :     Stream slaveStream = algResResp_->slaveStreams.back();
     225              :     // 从流通知主流,超节点间数据搬运完成
     226            0 :     CHK_RET(LocalNotify::Post(slaveStream, dispatcher_, algResResp_->notifiesMain.back(), INVALID_VALUE_STAGE));
     227              :     // 主流等待超节点通信完成
     228            0 :     CHK_RET(LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain.back(), INVALID_VALUE_STAGE));
     229            0 :     return HCCL_SUCCESS;
     230            0 : }
     231              : 
     232            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunIntraServer(
     233              :     const OpParam& param, const ExecMem& execMem, u32 step)
     234              : {
     235              :     (void)execMem;
     236              :     // 计算slice信息, 将user in分成level2RankSize_块, 每个step处理一块blockIndex, 每个block需要分成level0RankSize_片
     237            0 :     u64 level0Count = param.DataDes.count * level1RankSize_;
     238            0 :     u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     239            0 :     u64 sliceSize = level0Count * unitSize_;
     240            0 :     u64 blockOffset = blockIndex * sliceSize * level0RankSize_;
     241              : 
     242            0 :     HCCL_DEBUG(
     243              :         "[RunIntraServer] rank[%u:%u,%u,%u] step[%u] blockIndex[%u], level0Count[%llu]", topoAttr_.userRank,
     244              :         level2Rank_, level1Rank_, level0Rank_, step, blockIndex, level0Count);
     245              : 
     246            0 :     std::vector<Slice> dataSegsSlice(level0RankSize_);
     247            0 :     for (u32 i = 0; i < level0RankSize_; i++) {
     248            0 :         dataSegsSlice[i].offset = blockOffset + sliceSize * i; // 相对于param.inputPtr偏移
     249            0 :         dataSegsSlice[i].size = sliceSize;
     250              :     }
     251            0 :     std::vector<std::vector<Slice>> multRingsUserMemSlice = {dataSegsSlice};
     252              : 
     253              :     // 算法编排
     254            0 :     if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
     255            0 :         CHK_RET(MultiRingReduceScatter(
     256              :             param.tag, algResResp_->paramInputMem, algResResp_->paramInputMem, level0Count, param.DataDes.dataType,
     257              :             param.reduceType, multRingsUserMemSlice, param.stream, PROF_STAGE_1, 0, nullptr, multRingsUserMemSlice));
     258              :     } else {
     259            0 :         CHK_PRT_RET(
     260              :             topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
     261              :             HCCL_ERROR("[RunIntraServer] unknown topoType: %u", topoType_), HCCL_E_NOT_SUPPORT);
     262            0 :         CHK_RET(SemiRingReduceScatter(
     263              :             param.tag, algResResp_->paramInputMem, algResResp_->paramInputMem, level0Count, param.DataDes.dataType,
     264              :             param.reduceType, multRingsUserMemSlice, param.stream, PROF_STAGE_1, 0, nullptr, multRingsUserMemSlice));
     265              :     }
     266              : 
     267            0 :     return HCCL_SUCCESS;
     268            0 : }
     269              : 
     270            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::SemiRingReduceScatter(
     271              :     const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
     272              :     const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream,
     273              :     s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
     274              :     const std::vector<std::vector<Slice>> multRingsUserMemSlice)
     275              : {
     276              :     (void)tag;
     277              :     (void)multRingsSliceZero;
     278              :     (void)baseOffset;
     279              :     (void)opInfo;
     280            0 :     HCCL_DEBUG(
     281              :         "[SemiRingReduceScatter] starts, rank[%u:%u,%u,%u]", topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_);
     282              : 
     283            0 :     CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
     284            0 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     285              : 
     286              :     // 此处计算reduceAttr计算,outputmem使用的是scratchmem
     287            0 :     u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
     288              :     // 执行
     289            0 :     std::unique_ptr<AlgTemplateBase> executor = AlgTemplateRegistry::Instance().GetAlgTemplate(
     290            0 :         TemplateType::TEMPLATE_REDUCESCATTER_UNIFIED_MARCH, dispatcher_);
     291            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_UNIFIED_MARCH in COMM_LEVEL0", __func__);
     292            0 :     CHK_SMART_PTR_NULL(executor);
     293              : 
     294            0 :     CHK_RET(executor->Prepare(
     295              :         stream, level0CommInfo, algResResp_->paramInputMem, algResResp_->paramOutputMem, inputMem, outputMem, count,
     296              :         algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux, dataType, reductionOp,
     297              :         multRingsUserMemSlice, reduceAttr));
     298              : 
     299            0 :     HcclResult ret = executor->RegisterProfiler(
     300              :         ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
     301            0 :             + (level0CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank,
     302              :         profStage, HCCL_EXEC_STEP_NOT_SET, stream);
     303            0 :     CHK_PRT_RET(
     304              :         ret != HCCL_SUCCESS, HCCL_ERROR("[SemiRingReduceScatter] Double ring ReduceScatter failed,return[%d]", ret),
     305              :         ret);
     306              : 
     307            0 :     CHK_RET(executor->RunAsync());
     308              : 
     309            0 :     HCCL_DEBUG(
     310              :         "[SemiRingReduceScatter] run success, rank[%u:%u,%u,%u]", topoAttr_.userRank, level2Rank_, level1Rank_,
     311              :         level0Rank_);
     312            0 :     return ret;
     313            0 : }
     314              : 
     315            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServerPreProcess(
     316              :     const OpParam& param, const ExecMem& execMem, u32 step)
     317              : {
     318              :     // 数据准备,将节点内RS的结果从user in搬到ccl in
     319            0 :     u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     320            0 :     u32 cclSliceIndex = blockIndex * level1RankSize_;
     321            0 :     u32 usrInSliceIndex = blockIndex * level1RankSize_ * level0RankSize_ + level1RankSize_ * level0Rank_;
     322            0 :     Stream stream = param.stream;
     323              : 
     324            0 :     HCCL_DEBUG(
     325              :         "[RunInterServerPreProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] sliceIndex[%u, %u]", topoAttr_.userRank,
     326              :         level2Rank_, level1Rank_, level0Rank_, step, blockIndex, cclSliceIndex, usrInSliceIndex);
     327              :     // 本地 user in -> ccl in
     328            0 :     for (u32 i = 0; i < level1RankSize_; i++) {
     329            0 :         u64 ccInOffset = (cclSliceIndex + i) * curSize_;
     330            0 :         u64 userInOffset = (usrInSliceIndex + i) * param.DataDes.count * unitSize_; // 相对于execMem.inputPtr偏移
     331            0 :         DeviceMem dstMem = execMem.inputMem.range(ccInOffset, curSize_);
     332            0 :         DeviceMem srcMem = DeviceMem::create(static_cast<u8*>(execMem.inputPtr) + userInOffset, curSize_);
     333            0 :         HCCL_DEBUG(
     334              :             "[RunInterServerPreProcess] rank[%u:%u,%u,%u] step[%u] userInOffset[%llu] -> ccInOffset[%llu]",
     335              :             topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, userInOffset, ccInOffset);
     336            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     337            0 :     }
     338            0 :     return HCCL_SUCCESS;
     339            0 : }
     340              : 
     341              : HcclResult
     342            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServer(const OpParam& param, ExecMem& execMem, u32 step)
     343              : {
     344              :     // 计算slice信息,也就是在ccl in的偏移
     345            0 :     std::vector<Slice> level1DataSegsSlice(level1RankSize_);
     346            0 :     u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     347            0 :     u32 sliceIndex = blockIndex * level1RankSize_;
     348              : 
     349            0 :     HCCL_DEBUG(
     350              :         "[RunInterServer] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] sliceStart[%u] sliceCnt[%u]", topoAttr_.userRank,
     351              :         level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndex, level1RankSize_);
     352              : 
     353            0 :     for (u32 i = 0; i < level1RankSize_; i++) {
     354            0 :         level1DataSegsSlice[i].offset = (sliceIndex + i) * curSize_;
     355            0 :         level1DataSegsSlice[i].size = curSize_;
     356              :     }
     357              : 
     358            0 :     u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, param.DataDes.dataType, param.reduceType);
     359            0 :     std::unique_ptr<AlgTemplateBase> level1TempAlg;
     360            0 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     361              :         level1TempAlg
     362            0 :             = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     363            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
     364            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     365            0 :         CHK_RET(level1TempAlg->Prepare(reduceAttr));
     366            0 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     367              :         level1TempAlg
     368            0 :             = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     369            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
     370            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     371            0 :         CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
     372            0 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     373              :         level1TempAlg
     374            0 :             = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     375            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
     376            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     377            0 :         CHK_RET(level1TempAlg->Prepare(reduceAttr));
     378              :     }
     379            0 :     CHK_SMART_PTR_NULL(level1TempAlg);
     380              : 
     381              :     // 执行算法编排, 主流上执行,只会使用ccl in,执行完成后数据在ccl in
     382            0 :     CHK_RET(CheckCommSize(COMM_LEVEL1, level0Rank_ + 1));
     383            0 :     SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, level0Rank_);
     384            0 :     CHK_RET(level1TempAlg->Prepare(
     385              :         execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count, param.DataDes.dataType, param.stream,
     386              :         param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSegsSlice));
     387            0 :     CHK_RET(level1TempAlg->RegisterProfiler(
     388              :         (level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_, PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET,
     389              :         param.stream));
     390            0 :     CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
     391              : 
     392            0 :     return HCCL_SUCCESS;
     393            0 : }
     394              : 
     395            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::ExchangeData(
     396              :     const OpParam& param, const ExecMem& execMem, u32 step, u32 remoteRankSend, u32 remoteRankRecv)
     397              : {
     398              :     // 获取通信对端的link
     399            0 :     LINK sendLink;
     400            0 :     LINK recvLink;
     401            0 :     CHK_RET(GetTransportForExchange(remoteRankSend, sendLink));
     402            0 :     CHK_RET(GetTransportForExchange(remoteRankRecv, recvLink));
     403            0 :     CHK_PTR_NULL(sendLink);
     404            0 :     CHK_PTR_NULL(recvLink);
     405              : 
     406              :     // 当通信对端恰好是同server的邻居时,复用Level0的建链,其注册的内存是UserMem
     407              :     // 否则,在CommCombineOrder上建链,其注册内存是ccl buf
     408            0 :     Stream stream = param.stream;
     409            0 :     u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     410            0 :     u32 sliceIndexSnd = blockIndex * level1RankSize_ + level1Rank_; // 要发送的数据块在本地ccl in的位置
     411            0 :     u32 sliceIndexCclOut = blockIndex * level1RankSize_;
     412              : 
     413            0 :     HCCL_DEBUG(
     414              :         "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] send blockIndex[%u] sliceIndex[%u] cclout[%u]",
     415              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexSnd, sliceIndexCclOut);
     416              : 
     417            0 :     bool remoteSndl0Neighbor = IsLevel0Neighbor(remoteRankSend, level0RankSize_);
     418            0 :     bool remoteRcvl0Neighbor = IsLevel0Neighbor(remoteRankRecv, level0RankSize_);
     419            0 :     if (remoteSndl0Neighbor) {
     420              :         // 先本地 ccl in -> user in
     421            0 :         u32 usrInSliceIndex
     422            0 :             = blockIndex * level1RankSize_ * level0RankSize_ + level1RankSize_ * level0Rank_ + level1Rank_;
     423            0 :         u64 userInOffset = usrInSliceIndex * param.DataDes.count * unitSize_; // 相对于param.inputPtr偏移
     424            0 :         DeviceMem srcMem = execMem.inputMem.range(sliceIndexSnd * curSize_, curSize_);
     425            0 :         DeviceMem dstMem = DeviceMem::create(static_cast<u8*>(param.inputPtr) + userInOffset, curSize_);
     426            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     427            0 :         HCCL_DEBUG(
     428              :             "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] ci[%u]->ui[%u]", topoAttr_.userRank,
     429              :             level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexSnd, usrInSliceIndex);
     430              : 
     431              :         // user in send to remote user out
     432            0 :         CHK_RET(recvLink->TxAck(stream));
     433            0 :         CHK_RET(sendLink->RxAck(stream));
     434            0 :         CHK_RET(sendLink->TxAsync(
     435              :             UserMemType::OUTPUT_MEM, 0, static_cast<u8*>(param.inputPtr) + userInOffset, curSize_, stream));
     436            0 :     } else {
     437              :         // ccl in send to remote ccl out
     438            0 :         CHK_RET(recvLink->TxAck(stream));
     439            0 :         CHK_RET(sendLink->RxAck(stream));
     440            0 :         CHK_RET(sendLink->TxAsync(
     441              :             UserMemType::OUTPUT_MEM, sliceIndexCclOut * curSize_,
     442              :             static_cast<u8*>(execMem.inputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, stream));
     443              :     }
     444              : 
     445            0 :     u32 remoteL1Rank = (remoteRankRecv % (level1RankSize_ * level0RankSize_)) / level0RankSize_;
     446            0 :     u32 sliceIndexRcv = blockIndex * level1RankSize_ + remoteL1Rank; // 要接收的数据块在对端ccl in的位置
     447            0 :     if (remoteRcvl0Neighbor) {
     448            0 :         u32 usrInSliceIndexPeer
     449            0 :             = blockIndex * level1RankSize_ * level0RankSize_ + level1Rank_ * level0RankSize_ + level0Rank_;
     450            0 :         u64 userInOffsetPeer = usrInSliceIndexPeer * param.DataDes.count * unitSize_; // 相对于param.inputPtr偏移
     451            0 :         CHK_RET(recvLink->RxAsync(UserMemType::INPUT_MEM, userInOffsetPeer, execMem.outputPtr, curSize_, stream));
     452              :     } else {
     453            0 :         CHK_RET(recvLink->RxAsync(
     454              :             UserMemType::INPUT_MEM, sliceIndexRcv * curSize_,
     455              :             static_cast<u8*>(execMem.outputMem.ptr()) + sliceIndexCclOut * curSize_, curSize_, stream));
     456            0 :         CHK_RET(recvLink->PostFinAck(stream));
     457              :     }
     458              : 
     459            0 :     if (!remoteSndl0Neighbor) {
     460            0 :         CHK_RET(sendLink->WaitFinAck(stream));
     461              :     }
     462              : 
     463              :     // 交换数据的两端之间Barrier,确认收发完成
     464            0 :     CHK_RET(recvLink->TxAck(stream));
     465            0 :     CHK_RET(sendLink->RxAck(stream));
     466            0 :     CHK_RET(sendLink->TxDataSignal(stream));
     467            0 :     CHK_RET(recvLink->RxDataSignal(stream));
     468              : 
     469            0 :     if (remoteRcvl0Neighbor) {
     470              :         // 本地 user out -> ccl out
     471            0 :         DeviceMem srcMem = DeviceMem::create(static_cast<u8*>(execMem.outputPtr), curSize_);
     472            0 :         DeviceMem dstMem = execMem.outputMem.range(sliceIndexCclOut * curSize_, curSize_);
     473            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     474            0 :         HCCL_DEBUG(
     475              :             "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] uo->co[%u]", topoAttr_.userRank,
     476              :             level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexCclOut);
     477            0 :     }
     478            0 :     HCCL_DEBUG(
     479              :         "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] recv blockIndex[%u] sliceIndex[%u] cclout[%u]",
     480              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndexRcv, sliceIndexCclOut);
     481            0 :     return HCCL_SUCCESS;
     482            0 : }
     483              : 
     484            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunInterServerPostProcess(
     485              :     const OpParam& param, const ExecMem& execMem, u32 step)
     486              : {
     487              :     // 超节点内数据交换
     488            0 :     u32 remoteRankSend = exchangeRemoteRankSend_;
     489            0 :     u32 remoteRankRecv = exchangeRemoteRankRecv_;
     490              : 
     491            0 :     HCCL_DEBUG(
     492              :         "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] remoteRankSend[%u] remoteRankRecv[%u]",
     493              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, remoteRankSend, remoteRankRecv);
     494            0 :     if (remoteRankSend == topoAttr_.userRank && remoteRankRecv == topoAttr_.userRank) { // 不需要交换数据
     495              :         // 本地 ccl in -> ccl out
     496            0 :         Stream stream = param.stream;
     497            0 :         u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     498            0 :         u32 srcSliceIndex = blockIndex * level1RankSize_ + level1Rank_;
     499            0 :         u32 dstSliceIndex = blockIndex * level1RankSize_;
     500            0 :         DeviceMem srcMem = execMem.inputMem.range(srcSliceIndex * curSize_, curSize_);
     501            0 :         DeviceMem dstMem = execMem.outputMem.range(dstSliceIndex * curSize_, curSize_);
     502            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     503            0 :         HCCL_DEBUG(
     504              :             "[RunInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] blockIndex[%u] ci[%u]->co[%u]", topoAttr_.userRank,
     505              :             level2Rank_, level1Rank_, level0Rank_, step, blockIndex, srcSliceIndex, dstSliceIndex);
     506            0 :         return HCCL_SUCCESS;
     507            0 :     }
     508              : 
     509            0 :     return ExchangeData(param, execMem, step, remoteRankSend, remoteRankRecv);
     510              : }
     511              : 
     512            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPod(
     513              :     const OpParam& param, const ExecMem& execMem, u32 step)
     514              : {
     515              :     (void)param;
     516            0 :     Stream slaveStream = algResResp_->slaveStreams.back();
     517              :     // 发送前回RS好的数据
     518            0 :     u32 blockIndexSnd = (level2Rank_ + level2RankSize_ - step) % level2RankSize_;
     519            0 :     u32 sliceIndexSnd = blockIndexSnd * level1RankSize_; // 要发送的数据处于本地的哪个slice
     520              : 
     521              :     // 接受上一超节点发来的数据
     522            0 :     u32 blockIndexRcv = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     523            0 :     u32 sliceIndexRcv = blockIndexRcv * level1RankSize_; // 要接收的数据处于对端的哪个slice
     524              : 
     525            0 :     u32 preRank = (level2Rank_ + level2RankSize_ - 1) % level2RankSize_;
     526            0 :     u32 nextRank = (level2Rank_ + 1) % level2RankSize_;
     527            0 :     CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     528            0 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     529            0 :     LINK sendLink = level0CommInfo.links[nextRank];
     530            0 :     LINK recvLink = level0CommInfo.links[preRank];
     531            0 :     CHK_PTR_NULL(sendLink);
     532            0 :     CHK_PTR_NULL(recvLink);
     533              : 
     534              :     // 将数据发给nextRank前回的ccl in范围
     535            0 :     u32 remoteBlockIndexSndTo = (nextRank + level2RankSize_ - step) % level2RankSize_;
     536            0 :     u32 remoteSliceIndexSndTo = remoteBlockIndexSndTo * level1RankSize_; // 对端在哪个slice收对应的数据
     537            0 :     HCCL_DEBUG(
     538              :         "[RunSuperPod] rank[%u:%u,%u,%u] step[%u] send blockIndex[%u] sliceIndex[%u]->[%u]", topoAttr_.userRank,
     539              :         level2Rank_, level1Rank_, level0Rank_, step, blockIndexSnd, sliceIndexSnd, remoteSliceIndexSndTo);
     540            0 :     HCCL_DEBUG(
     541              :         "[RunSuperPod] rank[%u:%u,%u,%u] step[%u] recv blockIndex[%u] sliceIndex[%u]<-[%u]", topoAttr_.userRank,
     542              :         level2Rank_, level1Rank_, level0Rank_, step, blockIndexRcv, sliceIndexSnd, sliceIndexRcv);
     543              : 
     544            0 :     CHK_RET(recvLink->TxAck(slaveStream));
     545            0 :     CHK_RET(sendLink->RxAck(slaveStream));
     546              :     // 建链时其注册内存是ccl in与ccl out
     547              :     // ccl out send to remote ccl in
     548            0 :     CHK_RET(sendLink->TxAsync(
     549              :         UserMemType::INPUT_MEM, remoteSliceIndexSndTo * curSize_,
     550              :         static_cast<s8*>(execMem.outputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, slaveStream));
     551            0 :     CHK_RET(recvLink->RxAsync(
     552              :         UserMemType::OUTPUT_MEM, sliceIndexRcv * curSize_,
     553              :         static_cast<s8*>(execMem.inputMem.ptr()) + sliceIndexSnd * curSize_, curSize_, slaveStream));
     554            0 :     CHK_RET(recvLink->PostFinAck(slaveStream));
     555            0 :     CHK_RET(sendLink->WaitFinAck(slaveStream));
     556              : 
     557              :     // 交换数据的两端之间Barrier,确认收发完成
     558            0 :     CHK_RET(recvLink->TxAck(slaveStream));
     559            0 :     CHK_RET(sendLink->RxAck(slaveStream));
     560            0 :     CHK_RET(sendLink->TxDataSignal(slaveStream));
     561            0 :     CHK_RET(recvLink->RxDataSignal(slaveStream));
     562            0 :     return HCCL_SUCCESS;
     563            0 : }
     564              : 
     565            0 : HcclResult CollReduceScatterRingZerocopyExchangePipelineExecutor::RunSuperPodAndInterServerPostProcess(
     566              :     const OpParam& param, const ExecMem& execMem, u32 step)
     567              : {
     568              :     // ccl in -> ccl out执行reduce
     569            0 :     u32 blockIndexPreStep = (level2Rank_ + level2RankSize_ - step) % level2RankSize_;
     570            0 :     u32 blockIndex = (level2Rank_ + level2RankSize_ - (step + 1)) % level2RankSize_;
     571            0 :     u32 sliceIndex = blockIndex * level1RankSize_;
     572            0 :     u64 dstOffset = sliceIndex * curSize_;
     573            0 :     u64 srcOffset = blockIndexPreStep * level1RankSize_ * curSize_;
     574            0 :     HCCL_DEBUG(
     575              :         "[RunSuperPodAndInterServerPostProcess] rank[%u:%u,%u,%u] step[%u] reduce blockIndex[%u] sliceIndex[%u]",
     576              :         topoAttr_.userRank, level2Rank_, level1Rank_, level0Rank_, step, blockIndex, sliceIndex);
     577              : 
     578            0 :     Stream stream = param.stream;
     579            0 :     CHK_RET(HcclReduceAsync(
     580              :         dispatcher_, static_cast<s8*>(execMem.inputMem.ptr()) + srcOffset, execMem.count, param.DataDes.dataType,
     581              :         param.reduceType, stream, static_cast<s8*>(execMem.outputMem.ptr()) + dstOffset, topoAttr_.userRank,
     582              :         LinkType::LINK_RESERVED, INLINE_REDUCE_BIT));
     583            0 :     return HCCL_SUCCESS;
     584            0 : }
     585              : 
     586              : HcclResult
     587            0 : CollReduceScatterRingZerocopyExchangePipelineExecutor::RunFinallyProcess(const OpParam& param, const ExecMem& execMem)
     588              : {
     589            0 :     HCCL_DEBUG("[RunFinallyProcess] rank[%u:%u,%u,%u] ccl out -> user out");
     590            0 :     u32 sliceIndex = level2Rank_ * level1RankSize_;
     591            0 :     u64 offset = sliceIndex * curSize_;
     592            0 :     DeviceMem srcMem = execMem.outputMem.range(offset, curSize_);
     593            0 :     DeviceMem dstMem = DeviceMem::create(static_cast<u8*>(execMem.outputPtr), curSize_);
     594            0 :     Stream stream = param.stream;
     595            0 :     return HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream);
     596            0 : }
     597              : 
     598              : REGISTER_EXEC(
     599              :     "ReduceScatterRingZerocopyExchangePipelineExecutor", ReduceScatterRingZerocopyExchangePipeline,
     600              :     CollReduceScatterRingZerocopyExchangePipelineExecutor);
     601              : } // namespace hccl
        

Generated by: LCOV version 2.0-1