LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_scatter - coll_scatter_ring_for_910_93_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 171 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 12 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_scatter_ring_for_910_93_executor.h"
      12              : 
      13              : namespace hccl {
      14              : 
      15            0 : CollScatterRingFor91093Executor::CollScatterRingFor91093Executor(
      16            0 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      17            0 :     : CollScatterExecutor(dispatcher, topoMatcher)
      18              : {
      19            0 :     DMAReduceFlag_ = workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
      20            0 : }
      21              : 
      22            0 : HcclResult CollScatterRingFor91093Executor::CalcStreamNum(u32& streamNum)
      23              : {
      24            0 :     u32 totalStreamNum
      25            0 :         = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
      26              :                                                              LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
      27              :     // scatter在910_93场景仅支持单算子模式,已有mainstream需要-1
      28            0 :     streamNum = totalStreamNum - 1;
      29            0 :     HCCL_INFO("[CollScatterRingFor91093Executor][CalcStreamNum] tag[%s] streamNum[%u]", tag_.c_str(), streamNum);
      30            0 :     return HCCL_SUCCESS;
      31              : }
      32              : 
      33            0 : HcclResult CollScatterRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      34              : {
      35            0 :     TransportMemType inputType = TransportMemType::RESERVED;
      36            0 :     TransportMemType outputType = TransportMemType::RESERVED;
      37            0 :     CHK_RET(CalcTransportMemType(inputType, outputType));
      38            0 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
      39            0 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
      40            0 :     CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
      41            0 :     return HCCL_SUCCESS;
      42              : }
      43              : 
      44            0 : HcclResult CollScatterRingFor91093Executor::CalcLevel0CommInfo(
      45              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      46              : {
      47            0 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
      48            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
      49            0 :     return HCCL_SUCCESS;
      50            0 : }
      51              : 
      52            0 : HcclResult CollScatterRingFor91093Executor::CalcLevel1CommInfo(
      53              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      54              : {
      55            0 :     CommParaInfo commParaLevel1(COMM_LEVEL1, CommType::COMM_TAG_MAX);
      56            0 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
      57            0 :         commParaLevel1.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
      58            0 :         HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
      59            0 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
      60            0 :         commParaLevel1.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
      61            0 :         HCCL_INFO("[%s]Calc NBCommInfo", __func__);
      62              :     } else {
      63            0 :         commParaLevel1.commType = CommType::COMM_TAG_RING_INNER;
      64            0 :         HCCL_INFO("[%s]Calc RingCommInfo", __func__);
      65              :     }
      66            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel1, opTransport[COMM_LEVEL1], inputType, outputType));
      67              : 
      68            0 :     return HCCL_SUCCESS;
      69            0 : }
      70              : 
      71            0 : HcclResult CollScatterRingFor91093Executor::CalcLevel2CommInfo(
      72              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      73              : {
      74              :     // 910_93 level2当前仅支持nhr、nb、ring算法
      75            0 :     CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX, root_);
      76              : 
      77            0 :     if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
      78            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
      79            0 :         HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
      80            0 :     } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
      81            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
      82            0 :         HCCL_INFO("[%s]Calc NBCommInfo", __func__);
      83              :     } else {
      84            0 :         commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
      85            0 :         HCCL_INFO("[%s]Calc RingCommInfo", __func__);
      86              :     }
      87              : 
      88            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
      89            0 :     return HCCL_SUCCESS;
      90            0 : }
      91              : 
      92            0 : HcclResult CollScatterRingFor91093Executor::KernelRun(const OpParam& param, ExecMem& execMem)
      93              : {
      94            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] starts.", __func__);
      95            0 :     Stream& stream = const_cast<Stream&>(param.stream);
      96              : 
      97            0 :     CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize_));
      98              : 
      99            0 :     CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
     100            0 :     level0CommInfo_ = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     101              : 
     102            0 :     commIndex_ = level0CommInfo_.localRank;
     103              : 
     104            0 :     CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex_ + 1));
     105            0 :     level1CommInfo_ = GetSubCommInfo(COMM_LEVEL1, commIndex_);
     106              : 
     107            0 :     CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     108            0 :     level2CommInfo_ = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     109              : 
     110            0 :     CHK_RET(KernelRunLevel2(param, execMem, stream));
     111            0 :     CHK_RET(KernelRunLevel1(param, execMem, stream));
     112            0 :     CHK_RET(KernelRunLevel0(param, execMem, stream));
     113              : 
     114            0 :     if (!DMAReduceFlag_) {
     115              :         DeviceMem srcMem = execMem.inputMem.range(
     116            0 :             serverSliceOffset_ + execMem.outputMem.size() * commIndex_, execMem.count * perDataSize_);
     117            0 :         CHK_SMART_PTR_NULL(srcMem.ptr());
     118            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, stream));
     119            0 :     }
     120            0 :     HCCL_INFO("scatter ring run success");
     121            0 :     return HCCL_SUCCESS;
     122              : }
     123              : 
     124              : /* ***********超节点间scatter*********** */
     125            0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel2(const OpParam& param, ExecMem& execMem, Stream& stream)
     126              : {
     127            0 :     u32 level2RankSize = level2CommInfo_.localRankSize;
     128            0 :     u32 level2Rank = level2CommInfo_.localRank;
     129            0 :     subUserRankRootSupperPod_ = topoMatcher_->GetSubRootWithSuperPod(topoAttr_.userRank, param.root);
     130              : 
     131            0 :     if (level2RankSize > 1 && subUserRankRootSupperPod_ == topoAttr_.userRank) {
     132            0 :         u32 planeRootSupperPod = 0;
     133            0 :         CHK_RET(GetRankByUserRank(COMM_LEVEL2, COMM_INDEX_0, param.root, planeRootSupperPod));
     134            0 :         std::unique_ptr<AlgTemplateBase> level2TempAlg;
     135            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     136              :             level2TempAlg
     137            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
     138            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NB in COMM_LEVEL2", __func__);
     139            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     140              :             level2TempAlg
     141            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
     142            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NHR in COMM_LEVEL2", __func__);
     143              :         } else {
     144              :             level2TempAlg
     145            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
     146            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_RING in COMM_LEVEL2", __func__);
     147              :         }
     148              : 
     149            0 :         CHK_SMART_PTR_NULL(level2TempAlg);
     150              : 
     151            0 :         u64 level2Count = execMem.inputMem.size() / perDataSize_;
     152            0 :         CHK_RET(level2TempAlg->Prepare(
     153              :             execMem.inputMem, execMem.inputMem, execMem.scratchMem, level2Count, param.DataDes.dataType, stream,
     154              :             HCCL_REDUCE_RESERVED, planeRootSupperPod, std::vector<Slice>(0)));
     155            0 :         CHK_RET(level2TempAlg->RegisterProfiler(
     156              :             (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank, PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET,
     157              :             stream));
     158            0 :         CHK_RET(RunTemplate(level2TempAlg, level2CommInfo_));
     159            0 :     }
     160            0 :     return HCCL_SUCCESS;
     161              : }
     162              : 
     163              : /* ***********节点间scatter*********** */
     164            0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel1(const OpParam& param, ExecMem& execMem, Stream& stream)
     165              : {
     166            0 :     u32 level2RankSize = level2CommInfo_.localRankSize;
     167            0 :     u32 level2Rank = level2CommInfo_.localRank;
     168            0 :     u32 level1RankSize = level1CommInfo_.localRankSize;
     169            0 :     u32 level1Rank = level1CommInfo_.localRank;
     170            0 :     HCCL_DEBUG("level1RankSize:%u level1Rank:%u", level1RankSize, level1Rank);
     171              : 
     172            0 :     u64 level1SliceSize = execMem.inputMem.size() / level2RankSize;
     173            0 :     u64 level1SliceCount = level1SliceSize / perDataSize_;
     174            0 :     level1SliceOffset_ = level1SliceSize * level2Rank;
     175              : 
     176            0 :     CHK_RET(topoMatcher_->GetSubRootForScatter(subUserRankRootSupperPod_, subRoot_));
     177            0 :     CHK_PRT_RET(
     178              :         subRoot_ == INVALID_VALUE_RANKID,
     179              :         HCCL_ERROR(
     180              :             "[CollScatterRingFor91093Executor][KernelRun]GetSubRootForScatter failed, "
     181              :             "userRank[%u], root[%u], subRoot[%u]",
     182              :             topoAttr_.userRank, param.root, subRoot_),
     183              :         HCCL_E_INTERNAL);
     184            0 :     HCCL_DEBUG(
     185              :         "[CollScatterRingFor91093Executor][KernelRun]GetSubRootForScatter, userRank[%u], root[%u], subRoot[%u]",
     186              :         topoAttr_.userRank, param.root, subRoot_);
     187              : 
     188            0 :     if (level1RankSize > 1 && subRoot_ == topoAttr_.userRank) {
     189            0 :         u32 rootRankLevel1 = 0;
     190            0 :         CHK_RET(GetRankByUserRank(COMM_LEVEL1, commIndex_, subUserRankRootSupperPod_, rootRankLevel1));
     191              : 
     192            0 :         std::unique_ptr<AlgTemplateBase> level1TempAlg;
     193            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     194              :             level1TempAlg
     195            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
     196            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NB in COMM_LEVEL1", __func__);
     197            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     198              :             level1TempAlg
     199            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
     200            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NHR in COMM_LEVEL1", __func__);
     201              :         } else {
     202              :             level1TempAlg
     203            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
     204            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_RING in COMM_LEVEL1", __func__);
     205              :         }
     206            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     207              : 
     208            0 :         DeviceMem level1InputMem = execMem.inputMem.range(level1SliceOffset_, level1SliceSize);
     209            0 :         CHK_SMART_PTR_NULL(level1InputMem.ptr());
     210              : 
     211            0 :         CHK_RET(level1TempAlg->Prepare(
     212              :             level1InputMem, level1InputMem, level1InputMem, level1SliceCount, param.DataDes.dataType, stream,
     213              :             HCCL_REDUCE_RESERVED, rootRankLevel1, std::vector<Slice>(0), level1SliceOffset_));
     214            0 :         CHK_RET(level1TempAlg->RegisterProfiler(
     215              :             (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank, PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET,
     216              :             stream));
     217            0 :         CHK_RET(RunTemplate(level1TempAlg, level1CommInfo_));
     218            0 :     }
     219            0 :     return HCCL_SUCCESS;
     220              : }
     221              : 
     222              : /* ***********节点内scatter*********** */
     223            0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel0(const OpParam& param, ExecMem& execMem, Stream& stream)
     224              : {
     225              :     // 每个server分配的slice大小
     226            0 :     u32 level0RankSize = level0CommInfo_.localRankSize;
     227            0 :     u32 level2RankSize = level2CommInfo_.localRankSize;
     228            0 :     u32 level1RankSize = level1CommInfo_.localRankSize;
     229            0 :     u32 level1Rank = level1CommInfo_.localRank;
     230              : 
     231            0 :     u64 serverSliceSize = execMem.inputMem.size() / (level1RankSize * level2RankSize);
     232            0 :     serverSliceOffset_ = serverSliceSize * level1Rank + level1SliceOffset_;
     233            0 :     HCCL_DEBUG(
     234              :         "inputMem.size()=%llu, commLevel0->RankSize()=%u, serverSliceSize=%llu, serverSliceOffset=%llu "
     235              :         "commIndex=%u commLevel1[commIndex]->rank=%u",
     236              :         execMem.inputMem.size(), level0RankSize, serverSliceSize, serverSliceOffset_, commIndex_, level1Rank);
     237              : 
     238            0 :     DeviceMem scatterRingInput = execMem.inputMem.range(serverSliceOffset_, serverSliceSize);
     239            0 :     CHK_SMART_PTR_NULL(scatterRingInput);
     240              : 
     241              :     // 将根节点数据切分成level0RankSize份
     242            0 :     std::vector<Slice> dataSegsSlice;             // 数据分成ranksize份,每份的起始偏移和大小
     243            0 :     std::vector<std::vector<Slice>> mulRingSlice; // 每个stream使用的数据基于用户buffer的偏移
     244              :     // 根据数据量算每个环上数据的偏移和大小
     245            0 :     CHK_RET(PrepareDataSlice(execMem.count, perDataSize_, level0RankSize, dataSegsSlice));
     246              : 
     247              :     u32 ringNum;
     248            0 :     if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
     249            0 :         ringNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
     250            0 :         mulRingSlice = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
     251              :     } else {
     252            0 :         ringNum = LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
     253            0 :         mulRingSlice.push_back(dataSegsSlice);
     254              :     }
     255            0 :     CHK_PRT_RET(
     256              :         mulRingSlice.size() != ringNum,
     257              :         HCCL_ERROR(
     258              :             "[CollScatterRingFor91093Executor][KernelRunLevel0]ringNum[%u] != mulRingSlice size[%zu]", ringNum,
     259              :             mulRingSlice.size()),
     260              :         HCCL_E_INTERNAL);
     261            0 :     HCCL_INFO("scatter ring/scatter ring direct: using multiring algo inner-server.");
     262            0 :     HcomCollOpInfo* scatterOpInfoPtr = nullptr;
     263            0 :     HcomCollOpInfo scatterOpInfo
     264            0 :         = {"", nullptr, execMem.outputPtr, param.DataDes.count, param.DataDes.dataType, subRoot_, param.reduceType, 0};
     265            0 :     if (DMAReduceFlag_) {
     266            0 :         scatterOpInfoPtr = &scatterOpInfo;
     267              :     }
     268            0 :     CHK_RET(MultiRingScatter(
     269              :         param.tag, scatterRingInput, scatterRingInput, execMem.count, param.DataDes.dataType, mulRingSlice, subRoot_,
     270              :         stream, scatterOpInfoPtr, serverSliceOffset_));
     271            0 :     return HCCL_SUCCESS;
     272            0 : }
     273            0 : HcclResult CollScatterRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     274              : {
     275            0 :     if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
     276            0 :         return HCCL_E_UNAVAIL;
     277              :     }
     278            0 :     level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     279              : 
     280            0 :     return HCCL_SUCCESS;
     281              : }
     282              : 
     283              : HcclResult
     284            0 : CollScatterRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
     285              : {
     286            0 :     if (level1RankSize > 1) {
     287            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     288              :             level1TempAlg
     289            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
     290            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     291            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     292              :             level1TempAlg
     293            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
     294            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     295              :         } else {
     296              :             level1TempAlg
     297            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
     298            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     299              :         }
     300            0 :         return HCCL_SUCCESS;
     301              :     }
     302            0 :     return HCCL_E_UNAVAIL;
     303              : }
     304              : REGISTER_EXEC("ScatterRingFor91093Executor", ScatterRingFor91093, CollScatterRingFor91093Executor);
     305              : } // namespace hccl
        

Generated by: LCOV version 2.0-1