LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_nhr_v1.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 156 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 13 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 "reduce_scatter_nhr_v1.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : 
      16            0 : ReduceScatterNHRV1::ReduceScatterNHRV1(const HcclDispatcher dispatcher)
      17            0 :     : NHRV1Base(dispatcher)
      18              : {
      19            0 : }
      20              : 
      21            0 : ReduceScatterNHRV1::~ReduceScatterNHRV1()
      22              : {
      23            0 : }
      24              : 
      25            0 : HcclResult ReduceScatterNHRV1::Prepare(u64 reduceAttrBitMap, HcomCollOpInfo *opInfo)
      26              : {
      27              :     (void)opInfo;
      28            0 :     reduceAttr_ = reduceAttrBitMap;
      29            0 :     return HCCL_SUCCESS;
      30              : }
      31              : 
      32            0 : HcclResult ReduceScatterNHRV1::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      33              : {
      34              :     // 基本的检查
      35            0 :     CHK_RET(SimpleCheck(rank, rankSize, links));
      36            0 :     HCCL_INFO("ReduceScatterNHRV1 run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
      37              :         rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
      38              : 
      39              :     // 判断rank_size == 1
      40            0 :     if (rankSize == 1) {
      41            0 :         if (inputMem_ != outputMem_) {
      42            0 :             return HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
      43              :         }
      44            0 :         return HCCL_SUCCESS;
      45              :     }
      46              : 
      47              :     // 处理和检查Slices
      48            0 :     if (slices_.size() == 0) {
      49            0 :         CHK_RET(SetDefaultSlices(rank, rankSize));
      50              :     }
      51            0 :     CHK_RET(CheckSlices(rankSize));
      52              : 
      53              :     // 获取通信关系
      54            0 :     RingInfo info = GetRingInfo(rankSize);
      55              : 
      56              :     // 垂直方向做Ring
      57            0 :     CHK_RET(RunReduceScatterOnVertical(rank, links, info));
      58              : 
      59              :     // 水平方向做Ring
      60            0 :     CHK_RET(RunReduceScatterOnHorizontal(rank, links, info));
      61              : 
      62              :     // 额外的搬运(从(x,sqrt-1)搬运到(x,sqrt))
      63              :     /* 一个可能的优化点:
      64              :     以8节点为例:    0   1   2
      65              :                     3   4   5
      66              :                     6   7
      67              :     当前的做法是:{0,1}、{3,4}、{6、7}做水平Ring,0/3/6/7拿到各自的那份结果,1拿到1和2的结果,4拿到4和5的结果,
      68              :                  最后1把2的结果发给2,4把5的结果发给5
      69              :     一种可能的更优做法是:以ReduceOp=Sum为例,首先把2和5的数据都置为0,
      70              :                         然后直接{0,1,2}、{3,4,5}、{6,7}做水平Ring,避免不等分Ring和额外的拷贝步骤,但需要调用TBE-asign
      71              :     */
      72            0 :     CHK_RET(RunLastCopyStep(rank, links, info));
      73              : 
      74              :     // 搬运数据到OutputMem
      75            0 :     CHK_RET(RunCopyDataToOutputMem(rank));
      76              : 
      77            0 :     HCCL_INFO("ReduceScatterNHRV1 finished: rank[%u] end", rank);
      78            0 :     return HCCL_SUCCESS;
      79            0 : }
      80              : 
      81            0 : HcclResult ReduceScatterNHRV1::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      82              : {
      83              :     // 判断stream, dispatcher是否为空
      84            0 :     CHK_SMART_PTR_NULL(dispatcher_);
      85            0 :     CHK_PTR_NULL(stream_.ptr());
      86              : 
      87              :     // 检查memory
      88            0 :     CHK_PRT_RET(!outputMem_ || !inputMem_,
      89              :         HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
      90              : 
      91              :     // 判断links数量是否正确
      92            0 :     CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] link size[%llu] is less than "\
      93              :         "rank size[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
      94            0 :     return HCCL_SUCCESS;
      95              : }
      96              : 
      97            0 : HcclResult ReduceScatterNHRV1::SetDefaultSlices(const u32 rank, const u32 rankSize)
      98              : {
      99            0 :     u32 unitSize = DataUnitSize(dataType_);
     100            0 :     CHK_PRT_RET(unitSize == 0,
     101              :         HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] unit data size is zero", rank), HCCL_E_INTERNAL);
     102              : 
     103            0 :     slices_.resize(rankSize);
     104            0 :     u64 sliceSize = count_ * unitSize;
     105            0 :     for (u32 i = 0; i < rankSize; i++) {
     106            0 :         slices_[i].size = sliceSize;
     107            0 :         slices_[i].offset = (i * sliceSize);
     108            0 :         HCCL_DEBUG("rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu] ",
     109              :             rank, i, slices_[i].offset, i, slices_[i].size);
     110              :     }
     111            0 :     return HCCL_SUCCESS;
     112              : }
     113              : 
     114            0 : HcclResult ReduceScatterNHRV1::CheckSlices(const u32 rankSize)
     115              : {
     116            0 :     CHK_PRT_RET(slices_.size() != rankSize,
     117              :         HCCL_ERROR("[ReduceScatterNHRV1]slices.size[%u] should be equal to rankSize[%u]",
     118              :             slices_.size(), rankSize), HCCL_E_INTERNAL);
     119              : 
     120            0 :     for (u32 idx = 1; idx < slices_.size(); idx++) {
     121            0 :         CHK_PRT_RET(slices_[idx-1].offset + slices_[idx-1].size != slices_[idx].offset,
     122              :             HCCL_ERROR("[ReduceScatterNHRV1]only support continuous slices, but get "\
     123              :                 "slices[%u].offset[%u], slices[%u].size[%u], slices[%u].offset[%u]",
     124              :                 idx-1, slices_[idx-1].offset, idx-1, slices_[idx-1].size, idx, slices_[idx].offset), HCCL_E_INTERNAL);
     125              :     }
     126            0 :     return HCCL_SUCCESS;
     127              : }
     128              : 
     129            0 : HcclResult ReduceScatterNHRV1::RunReduceScatterBrokenRing(const u32 rank, const std::vector<LINK> &links,
     130              :     const std::vector<Slice> &slices)
     131              : {
     132            0 :     std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     133            0 :         TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     134            0 :     CHK_SMART_PTR_NULL(tempAlg);
     135            0 :     CHK_RET(tempAlg->Prepare(reduceAttr_));
     136              : 
     137            0 :     if (!barrierSwitchOn_) {
     138            0 :         tempAlg->CloseBarrier();
     139              :     }
     140              : 
     141            0 :     CHK_RET(tempAlg->Prepare(inputMem_, inputMem_, scratchMem_, count_, dataType_,
     142              :         stream_, reductionOp_, root_, slices));
     143              : 
     144            0 :     CHK_RET(tempAlg->RegisterProfiler(
     145              :         profilerInput_.planeID, profilerInput_.stage, profilerInput_.step, stream_));
     146              : 
     147            0 :     return tempAlg->RunAsync(rank, links.size(), links);
     148            0 : }
     149              : 
     150            0 : HcclResult ReduceScatterNHRV1::RunReduceScatterOnVertical(const u32 rank, const std::vector<LINK> &links,
     151              :     const RingInfo &info)
     152              : {
     153            0 :     u32 hIndex = info.GetHIndex(rank);      // 查找自己位于第几列
     154              : 
     155              :     // 构造新的links和slices
     156            0 :     std::vector<LINK> subLinks;
     157            0 :     std::vector<Slice> subSlices;
     158            0 :     u32 sliceIndexOffset = 0;        // slice数量的累计偏移
     159            0 :     u32 hIndexForRing = (hIndex < info.GetRowSize()) ? hIndex : info.GetVIndex(rank); // 属于第几个垂直Ring
     160            0 :     u32 vSizeForRing = info.GetVSizeByHIndex(hIndexForRing);    // 所属垂直Ring的大小
     161            0 :     for (u32 vIdx = 0; vIdx < vSizeForRing; vIdx++) {   // 处理垂直方向上的Rank
     162              :         // 增加link
     163            0 :         u32 rankInRing = info.GetRank(vIdx, hIndexForRing);
     164            0 :         CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Vertical] rank[%u] out of range, "\
     165              :             "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
     166            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] links[%u]=%u", rank, vIdx, rankInRing);
     167            0 :         subLinks.push_back(links[rankInRing]);
     168              : 
     169              :         // 寻找要合并的slice
     170            0 :         u32 nSlices = info.GetHSizeByVIndex(vIdx);
     171            0 :         Slice& headSlice = slices_[sliceIndexOffset];
     172            0 :         Slice& tailSlice = slices_[sliceIndexOffset + nSlices - 1];
     173              : 
     174              :         // 增加slice
     175            0 :         Slice slice;
     176            0 :         slice.offset = headSlice.offset;
     177            0 :         slice.size = tailSlice.offset + tailSlice.size - headSlice.offset;
     178            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
     179              :             rank, vIdx, slice.offset, vIdx, slice.size);
     180            0 :         subSlices.push_back(slice);
     181              : 
     182              :         // 更新偏移
     183            0 :         sliceIndexOffset += nSlices;
     184              :     }
     185              : 
     186              :     // -- 可能还涉及跳跃的一个链接,比如8节点
     187              :     // ---- 0   1   2
     188              :     // ---- 3   4   5
     189              :     // ---- 6   7
     190              :     // -- 两个垂直Ring分别是{0,3,6,2}和{1,4,7,5},而不是{0,3,6}和{1,4,7}
     191            0 :     if (info.GetHSizeByVIndex(hIndexForRing) > info.GetRowSize()) {
     192              :         // 添加link
     193            0 :         u32 rankInRing = info.GetRank(hIndexForRing, info.GetRowSize());
     194            0 :         CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Vertical] rank[%u] out of range, "\
     195              :             "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
     196            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] links[%u]=%u", rank, subLinks.size(), rankInRing);
     197            0 :         subLinks.push_back(links[rankInRing]);
     198              : 
     199              :         // 添加slice
     200            0 :         Slice slice;
     201            0 :         slice.offset = 0;
     202            0 :         slice.size = 0;
     203            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
     204              :             rank, subLinks.size(), slice.offset, subLinks.size(), slice.size);
     205            0 :         subSlices.push_back(slice);
     206              :     }
     207              : 
     208              :     // 长度不足2,直接跳过
     209            0 :     if (subLinks.size() < 2) {
     210            0 :         return HCCL_SUCCESS;
     211              :     }
     212              : 
     213              :     // 计算在垂直Ring中的rank号
     214            0 :     u32 subRank = (hIndex == hIndexForRing) ? info.GetVIndex(rank) : vSizeForRing;
     215            0 :     HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subRank=%u", rank, subRank);
     216              : 
     217              :     // 执行Broken Ring ReduceScatter
     218            0 :     return RunReduceScatterBrokenRing(subRank, subLinks, subSlices);
     219            0 : }
     220              : 
     221            0 : HcclResult ReduceScatterNHRV1::RunReduceScatterOnHorizontal(const u32 rank, const std::vector<LINK> &links,
     222              :     const RingInfo &info)
     223              : {
     224            0 :     u32 hIndex = info.GetHIndex(rank);
     225            0 :     if (hIndex >= info.GetRowSize()) {
     226            0 :         return HCCL_SUCCESS;
     227              :     }
     228              : 
     229              :     // 构造新的links和slices
     230            0 :     u32 vIndex = info.GetVIndex(rank);
     231            0 :     std::vector<LINK> subLinks;
     232            0 :     std::vector<Slice> subSlices;
     233            0 :     for (u32 hIdx = 0; hIdx < info.GetRowSize(); hIdx++) {
     234              :         // 增加link
     235            0 :         u32 rankInRing = info.GetRank(vIndex, hIdx);
     236            0 :         CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Horizontal] rank[%u] out of range, "\
     237              :             "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
     238            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] links[%u]=%u", rank, hIdx, rankInRing);
     239            0 :         subLinks.push_back(links[rankInRing]);
     240              : 
     241              :         // 增肌slice(CheckSlices()已经约束slices_里的Slice都是连续的)
     242              :         // -- 比如8节点在做完水平Ring后
     243              :         // ---- 0   1   2
     244              :         // ---- 3   4   5
     245              :         // ---- 6   7
     246              :         // -- 0/3/6/7节点只拿到自己那份ReduceScatter结果,而1拿到1和2的两份数据、4拿到4和5的两份数据
     247            0 :         u64 sliceSize = slices_[rankInRing].size;
     248            0 :         if (hIdx == info.GetRowSize() - 1 && info.GetHSizeByVIndex(vIndex) > info.GetRowSize()) {
     249            0 :             sliceSize += slices_[rankInRing+1].size;
     250              :         }
     251              : 
     252            0 :         Slice slice;
     253            0 :         slice.offset = slices_[rankInRing].offset;
     254            0 :         slice.size = sliceSize;
     255            0 :         HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
     256              :             rank, hIdx, slice.offset, hIdx, slice.size);
     257            0 :         subSlices.push_back(slice);
     258              :     }
     259              : 
     260              :     // 长度不足2,直接跳过
     261            0 :     if (subLinks.size() < 2) {
     262            0 :         return HCCL_SUCCESS;
     263              :     }
     264              : 
     265              :     // 计算在水平Ring中的rank号
     266            0 :     u32 subRank = hIndex;
     267            0 :     HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] subRank=%u", rank, subRank);
     268              : 
     269              :     // 执行Broken Ring ReduceScatter
     270            0 :     return RunReduceScatterBrokenRing(subRank, subLinks, subSlices);
     271            0 : }
     272              : 
     273            0 : HcclResult ReduceScatterNHRV1::RunLastCopyStep(const u32 rank, const std::vector<LINK> &links, const RingInfo &info)
     274              : {
     275              :     HcclResult ret;
     276              : 
     277            0 :     u32 hIndex = info.GetHIndex(rank);          // 查找自己位于第几列
     278            0 :     u32 vIndex = info.GetVIndex(rank);          // 查找自己位于第几行
     279            0 :     if (hIndex >= info.GetRowSize() - 1 && info.GetHSizeByVIndex(vIndex) > info.GetRowSize()) {
     280            0 :         u32 peerRank = (hIndex == info.GetRowSize() - 1) ? rank + 1 : rank - 1;
     281              : 
     282              :         // 检查指针
     283            0 :         CHK_SMART_PTR_NULL(links[peerRank]);
     284              : 
     285              :         // TxAck
     286            0 :         ret = links[peerRank]->TxAck(stream_);
     287            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     288              :             HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]rank[%u] tx ack from peerank[%u] failed",
     289              :                 rank, peerRank), ret);
     290              : 
     291              :         // RxAck
     292            0 :         ret = links[peerRank]->RxAck(stream_);
     293            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     294              :             HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]rank[%u] rx ack from peerank[%u] failed",
     295              :                 rank, peerRank), ret);
     296              : 
     297            0 :         if (hIndex == info.GetRowSize() - 1) {    // 发数据
     298            0 :             Slice& txSlice = slices_[peerRank];
     299            0 :             DeviceMem srcMem = inputMem_.range(txSlice.offset, txSlice.size);
     300            0 :             HCCL_DEBUG("tx srcMem[%p] range[%llu] size[%llu] ", srcMem.ptr(), txSlice.offset, txSlice.size);
     301            0 :             CHK_RET(ExecuteTxSync(links[peerRank], UserMemType::INPUT_MEM, txSlice.offset + baseOffset_,
     302              :                 srcMem.ptr(), srcMem.size(), stream_));
     303              : 
     304            0 :             ret = links[peerRank]->TxWaitDone(stream_);
     305            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]TxWaitDone failed"), ret);
     306            0 :         } else {        // 收数据
     307            0 :             Slice& rxSlice = slices_[rank];
     308            0 :             DeviceMem dstMem = inputMem_.range(rxSlice.offset, rxSlice.size);
     309            0 :             HCCL_DEBUG("rx dstMem[%p] range[%llu], size[%llu] ",  dstMem.ptr(), rxSlice.offset, rxSlice.size);
     310            0 :             CHK_RET(ExecuteRxSync(links[peerRank], UserMemType::INPUT_MEM, rxSlice.offset + baseOffset_,
     311              :                 dstMem.ptr(), dstMem.size(), stream_));
     312              : 
     313            0 :             ret = links[peerRank]->RxWaitDone(stream_);
     314            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]RxWaitDone failed"), ret);
     315            0 :         }
     316              : 
     317              :         // 如果不Barrier,SDMA结果在数据量超过CCL Buffer之后结果不正确
     318            0 :         CHK_RET(ExecuteBarrier(links[peerRank], stream_));
     319              :     }
     320            0 :     return HCCL_SUCCESS;
     321              : }
     322              : 
     323            0 : HcclResult ReduceScatterNHRV1::RunCopyDataToOutputMem(const u32 rank)
     324              : {
     325            0 :     if (inputMem_ != outputMem_) {
     326            0 :         Slice& srcSlice = slices_[rank];
     327            0 :         DeviceMem dst = outputMem_.range(0, srcSlice.size);
     328            0 :         DeviceMem src = inputMem_.range(srcSlice.offset, srcSlice.size);
     329            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     330            0 :     }
     331            0 :     return HCCL_SUCCESS;
     332              : }
     333              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, ReduceScatterNHRV1);
     334              : }   // ~~ namespace hccl
        

Generated by: LCOV version 2.0-1