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

Generated by: LCOV version 2.0-1