LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_nhr.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 321 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 22 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.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : 
      16            0 : ReduceScatterNHR::ReduceScatterNHR(const HcclDispatcher dispatcher) : NHRBase(dispatcher) {}
      17              : 
      18            0 : ReduceScatterNHR::~ReduceScatterNHR() {}
      19              : 
      20            0 : HcclResult ReduceScatterNHR::Prepare(u64 reduceAttrBitMap, bool needMerge)
      21              : {
      22            0 :     reduceAttr_ = reduceAttrBitMap;
      23            0 :     isNeedMerge = needMerge;
      24            0 :     return HCCL_SUCCESS;
      25              : }
      26              : 
      27            0 : HcclResult ReduceScatterNHR::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
      28              : {
      29              :     // 基本的检查
      30            0 :     CHK_RET(SimpleCheck(rank, rankSize, links));
      31            0 :     HCCL_INFO(
      32              :         "[ReduceScatterNHR][RunAsync] rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
      33              :         inputMem_.ptr(), outputMem_.ptr(), count_);
      34              : 
      35            0 :     if (isNeedMerge == true) {
      36              :         // 获取tree映射,存储到类对象的成员变量中
      37            0 :         GetSliceMap(rankSize);
      38              :     }
      39              : 
      40              :     // 判断rank_size == 1
      41            0 :     if (rankSize == 1) {
      42            0 :         if (inputMem_ != outputMem_) {
      43            0 :             return HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
      44              :         }
      45            0 :         return HCCL_SUCCESS;
      46              :     }
      47              : 
      48            0 :     u32 unitSize = DataUnitSize(dataType_);
      49            0 :     CHK_PRT_RET(
      50              :         unitSize == 0, HCCL_ERROR("[ReduceScatterNHR][RunAsync] rank[%u] unit data size is zero", rank),
      51              :         HCCL_E_INTERNAL);
      52              : 
      53            0 :     std::vector<Slice> outputSlices(slices_);
      54              : 
      55              :     // 处理和检查Slices
      56            0 :     if (slices_.size() == 0) {
      57            0 :         slices_.resize(rankSize);
      58            0 :         outputSlices.resize(rankSize);
      59              : 
      60              :         // 生成std::vector<Slice> slices_
      61            0 :         u64 sliceSize = count_ * unitSize;
      62              : 
      63            0 :         for (u32 i = 0; i < rankSize; i++) {
      64            0 :             slices_[i].size = sliceSize;
      65            0 :             slices_[i].offset = (i * sliceSize);
      66              : 
      67            0 :             outputSlices[i].size = sliceSize;
      68            0 :             outputSlices[i].offset = (inputMem_.size() > outputMem_.size()) ? 0 : (i * sliceSize);
      69            0 :             HCCL_DEBUG(
      70              :                 "[ReduceScatterNHR][RunAsync] rank[%u], slices[%u].offset=[%llu] slices[%u].size=[%llu] "
      71              :                 "outputSlices[%u].offset=[%llu], outputSlices[%u].size=[%llu] count_[%llu] unitSize[%llu]",
      72              :                 rank, i, slices_[i].offset, i, slices_[i].size, i, outputSlices[i].offset, i, outputSlices[i].size,
      73              :                 count_, unitSize);
      74              :         }
      75              :     }
      76              : 
      77            0 :     CHK_RET(CheckSlices(slices_, rankSize));
      78              : 
      79              :     // 创建reducer & sender
      80            0 :     senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
      81            0 :     CHK_SMART_PTR_NULL(senderInfo_);
      82              : 
      83            0 :     reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
      84            0 :     CHK_SMART_PTR_NULL(reducerInfo_);
      85              : 
      86            0 :     if (sliceMap_.size() != rankSize) {
      87            0 :         GetRankMapping(rankSize, true); // 没有初始化过,说明不是由allreduce或者bcast调入,需要保序
      88              :     }
      89              : 
      90              :     // 运行reduce-scatter, NHR 算法
      91            0 :     CHK_RET(RunReduceScatterNHR(rank, rankSize, links, slices_, outputSlices));
      92              : 
      93            0 :     HCCL_INFO("[ReduceScatterNHR][RunAsync] ReduceScatterNHR finished: rank[%u] end", rank);
      94            0 :     return HCCL_SUCCESS;
      95            0 : }
      96              : 
      97            0 : void ReduceScatterNHR::GetSliceMap(const u32 rankSize)
      98              : {
      99            0 :     std::vector<u32> tree;
     100            0 :     for (u32 i = 0; i < rankSize; i++) {
     101            0 :         tree.push_back(i);
     102              :     }
     103              : 
     104              :     // 其他的再进行计算
     105            0 :     std::vector<u32> tmp(rankSize);
     106            0 :     u32 nSteps = 0;
     107            0 :     for (u32 tmp = rankSize - 1; tmp != 0; tmp >>= 1, nSteps++) {
     108              :     }
     109              : 
     110            0 :     u32 len = rankSize;
     111              : 
     112            0 :     for (u32 step = 0; step < nSteps; step++) {
     113            0 :         u32 nSlices = (rankSize - 1 + (1 << step)) / (1 << (step + 1));
     114            0 :         if (nSlices <= 1) {
     115            0 :             break;
     116              :         }
     117              : 
     118            0 :         bool endFlag = false;
     119              : 
     120            0 :         for (u32 part = 0; part * len < rankSize; part++) {
     121            0 :             u32 start = part * len;
     122            0 :             u32 end = std::min(start + len, rankSize);
     123            0 :             Reorder(start, end, len, tree, tmp);
     124              : 
     125            0 :             if (((end - start) & 1) == 1) {
     126            0 :                 endFlag = true;
     127              :             }
     128              :         }
     129              : 
     130            0 :         for (u32 i = 0; i < rankSize; i++) {
     131            0 :             tree[i] = tmp[i];
     132              :         }
     133              : 
     134            0 :         if (endFlag) {
     135            0 :             break;
     136              :         }
     137              : 
     138            0 :         len >>= 1;
     139              :     }
     140              : 
     141              :     // 因为取的是tree中rank的idx,所以直接返回反向的映射
     142            0 :     sliceMap_.resize(rankSize);
     143            0 :     for (u32 i = 0; i < rankSize; i++) {
     144            0 :         sliceMap_[tree[i]] = i;
     145              :     }
     146              : 
     147            0 :     return;
     148            0 : }
     149              : 
     150            0 : void ReduceScatterNHR::Reorder(u32 start, u32 end, u32 len, std::vector<u32>& tree, std::vector<u32>& tmp)
     151              : {
     152            0 :     const u32 idxTwo = 2;
     153              : 
     154            0 :     for (u32 i = start; i < end; i++) {
     155            0 :         u32 offset = i - start;
     156            0 :         if ((offset & 1) == 0) {
     157            0 :             tmp[start + offset / idxTwo] = tree[i];
     158              :         } else {
     159            0 :             tmp[start + (offset + len) / idxTwo] = tree[i];
     160              :         }
     161              :     }
     162            0 : }
     163              : 
     164            0 : HcclResult ReduceScatterNHR::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
     165              : {
     166              :     // 判断stream, dispatcher是否为空
     167            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     168            0 :     CHK_PTR_NULL(stream_.ptr());
     169              : 
     170              :     // 检查memory
     171            0 :     CHK_PRT_RET(
     172              :         !outputMem_ || !inputMem_,
     173              :         HCCL_ERROR("[ReduceScatterNHR][RunAsync] rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
     174              : 
     175              :     // 判断links数量是否正确
     176            0 :     CHK_PRT_RET(
     177              :         links.size() < rankSize,
     178              :         HCCL_ERROR(
     179              :             "[ReduceScatterNHR][RunAsync] rank[%u] link size[%llu] is "
     180              :             "less than rank size[%u]",
     181              :             rank, links.size(), rankSize),
     182              :         HCCL_E_INTERNAL);
     183            0 :     return HCCL_SUCCESS;
     184              : }
     185              : 
     186            0 : HcclResult ReduceScatterNHR::CheckSlices(const std::vector<Slice>& checkSlices, const u32 rankSize)
     187              : {
     188            0 :     CHK_PRT_RET(
     189              :         checkSlices.size() % rankSize != 0,
     190              :         HCCL_ERROR(
     191              :             "[ReduceScatterNHR][RunAsync] slices.size[%u] should be divided by rankSize[%u]", checkSlices.size(),
     192              :             rankSize),
     193              :         HCCL_E_INTERNAL);
     194            0 :     return HCCL_SUCCESS;
     195              : }
     196              : 
     197            0 : HcclResult ReduceScatterNHR::InlineReducer(const LINK& linkLeft, const std::vector<ReducerMemoryInfo>& rxReduceMems)
     198              : {
     199            0 :     HcclResult ret = HCCL_SUCCESS;
     200            0 :     void* remoteMem = nullptr;
     201            0 :     CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     202            0 :     for (ReducerMemoryInfo reduceMem : rxReduceMems) {
     203            0 :         const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     204            0 :         CHK_RET(HcclReduceAsync(
     205              :             dispatcher_, static_cast<s8*>(remoteMem) + reduceMem.remoteMemOffset, dataBytes / SIZE_TABLE[dataType_],
     206              :             dataType_, reductionOp_, stream_, reduceMem.localsrc.ptr(), linkLeft->GetRemoteRank(),
     207              :             linkLeft->GetLinkType(), INLINE_REDUCE_BIT));
     208              : 
     209            0 :         if (reduceMem.localsrc != reduceMem.localdst) {
     210            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, reduceMem.localdst, reduceMem.localsrc, stream_);
     211            0 :             CHK_PRT_RET(
     212              :                 ret != HCCL_SUCCESS,
     213              :                 HCCL_ERROR(
     214              :                     "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     215              :                     reduceMem.localdst.ptr()),
     216              :                 ret);
     217              :         }
     218            0 :     }
     219            0 :     return HCCL_SUCCESS;
     220              : }
     221              : 
     222              : HcclResult
     223            0 : ReduceScatterNHR::InlineReduceRx(const LINK& linkLeft, std::vector<Slice>& rxSlices, std::vector<Slice>& rxSlicestemp)
     224              : {
     225            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     226            0 :     for (u64 i = 0; i < rxSlices.size(); i++) {
     227            0 :         DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
     228            0 :         DeviceMem srcMemTemp = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
     229            0 :         HCCL_DEBUG(
     230              :             "[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
     231              :             "offset[%llu] size[%llu] ",
     232              :             rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
     233            0 :         rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlices[i].offset, dstMem, dstMem, srcMemTemp});
     234            0 :     }
     235            0 :     CHK_RET(InlineReducer(linkLeft, rxReduceMems));
     236            0 :     return HCCL_SUCCESS;
     237            0 : }
     238              : 
     239            0 : HcclResult ReduceScatterNHR::InlineReduceRxLastStep(
     240              :     const LINK& linkLeft, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     241              :     const std::vector<Slice>& outputSlices)
     242              : {
     243            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     244            0 :     for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
     245            0 :         u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
     246            0 :         DeviceMem dstMem = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     247            0 :         DeviceMem srcMem = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size);
     248            0 :         DeviceMem tmpMem = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     249            0 :         HCCL_DEBUG(
     250              :             "[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
     251              :             "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]",
     252              :             rxSliceIdx, inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     253              : 
     254            0 :         rxReduceMems.emplace_back(
     255            0 :             ReducerMemoryInfo{baseOffset_ + inputSlices[rxSliceIdx].offset, srcMem, dstMem, tmpMem});
     256            0 :     }
     257            0 :     CHK_RET(InlineReducer(linkLeft, rxReduceMems));
     258            0 :     return HCCL_SUCCESS;
     259            0 : }
     260              : 
     261              : HcclResult
     262            0 : ReduceScatterNHR::TbeReduceRx(const LINK& linkLeft, std::vector<Slice>& rxSlices, std::vector<Slice>& rxSlicestemp)
     263              : {
     264            0 :     void* srcMemPtr = nullptr;
     265            0 :     CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &srcMemPtr));
     266            0 :     std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     267            0 :     for (u64 i = 0; i < rxSlices.size(); i++) {
     268            0 :         DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
     269            0 :         DeviceMem srcMem(static_cast<s8*>(srcMemPtr) + baseOffset_ + rxSlices[i].offset, rxSlices[i].size);
     270            0 :         DeviceMem dstMemScratch = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
     271            0 :         u64 dataCount = dstMem.size() / SIZE_TABLE[dataType_];
     272            0 :         HCCL_DEBUG(
     273              :             "[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
     274              :             "offset[%llu] size[%llu] ",
     275              :             rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
     276            0 :         CHK_RET(HcclD2DMemcpyAsync(
     277              :             dispatcher_, dstMemScratch, srcMem, stream_,
     278              :             linkLeft->GetRemoteRank(), // left的inputMem拷到本端的scratchMem
     279              :             linkLeft->GetLinkType()));
     280            0 :         rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{
     281            0 :             UserMemType::INPUT_MEM, baseOffset_ + rxSlices[i].offset, dstMemScratch.ptr(), dstMemScratch.size(),
     282            0 :             dstMemScratch.ptr(), dstMem.ptr(), dataCount});
     283            0 :     }
     284            0 :     for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
     285            0 :         CHK_RET(HcclReduceAsync(
     286              :             dispatcher_, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount,
     287              :             dataType_, // 本端scratchMem localReduce到 本端inputMem
     288              :             reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
     289              :     }
     290            0 :     return HCCL_SUCCESS;
     291            0 : }
     292              : 
     293            0 : HcclResult ReduceScatterNHR::TbeReduceRxLastStep(
     294              :     const LINK& linkLeft, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     295              :     const std::vector<Slice>& outputSlices)
     296              : {
     297            0 :     void* srcMemPtr = nullptr;
     298            0 :     CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &srcMemPtr));
     299            0 :     std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     300            0 :     for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
     301            0 :         u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
     302              :         DeviceMem srcMemRemote(
     303            0 :             static_cast<s8*>(srcMemPtr) + baseOffset_ + inputSlices[rxSliceIdx].offset,
     304            0 :             inputSlices[rxSliceIdx].size); // 对端inputMem
     305              :         DeviceMem dstMem
     306            0 :             = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size); // 本端outputMem
     307              :         DeviceMem srcMem
     308            0 :             = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size); // 本端inputMem
     309              :         DeviceMem tmpMem
     310            0 :             = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size); // 本端scratchMem
     311            0 :         u64 dataCount = dstMem.size() / SIZE_TABLE[dataType_];
     312            0 :         HCCL_DEBUG(
     313              :             "[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
     314              :             "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]",
     315              :             rxSliceIdx, inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     316            0 :         CHK_RET(HcclD2DMemcpyAsync(
     317              :             dispatcher_, tmpMem, srcMemRemote, stream_, linkLeft->GetRemoteRank(), // left的inputMem拷到本端的scratchMem
     318              :             linkLeft->GetLinkType()));
     319            0 :         DeviceMem reduceSrc = (srcMem == dstMem) ? tmpMem : srcMem;
     320            0 :         rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{
     321            0 :             UserMemType::INPUT_MEM, baseOffset_ + inputSlices[rxSliceIdx].offset, tmpMem.ptr(), tmpMem.size(),
     322            0 :             reduceSrc.ptr(), dstMem.ptr(), dataCount});
     323            0 :     }
     324            0 :     for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
     325            0 :         CHK_RET(HcclReduceAsync(
     326              :             dispatcher_, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount,
     327              :             dataType_, // 本端inputMem localReduce到 本端outputMem(之前拷到本端scratch的数据呢?)
     328              :             reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
     329              :     }
     330            0 :     return HCCL_SUCCESS;
     331            0 : }
     332              : 
     333            0 : HcclResult ReduceScatterNHR::RunDestReducerLastStep(
     334              :     const LINK& linkLeft, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     335              :     const std::vector<Slice>& outputSlices)
     336              : {
     337            0 :     HcclResult ret = HCCL_SUCCESS;
     338            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     339            0 :     for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
     340            0 :         u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
     341            0 :         DeviceMem dstMem = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     342            0 :         DeviceMem srcMem = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size);
     343            0 :         DeviceMem tmpMem = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     344            0 :         HCCL_DEBUG(
     345              :             "[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
     346              :             "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]",
     347              :             rxSliceIdx, inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
     348              : 
     349            0 :         rxReduceMems.emplace_back(
     350            0 :             ReducerMemoryInfo{baseOffset_ + inputSlices[rxSliceIdx].offset, srcMem, dstMem, tmpMem});
     351            0 :     }
     352              : 
     353            0 :     ret = reducerInfo_->run(dispatcher_, linkLeft, rxReduceMems, stream_);
     354            0 :     return ret;
     355            0 : }
     356              : 
     357            0 : HcclResult ReduceScatterNHR::GetRxSlices(
     358              :     std::vector<Slice>& rxSlices, std::vector<Slice>& rxSlicestemp, InterServerAlgoStep& stepInfo,
     359              :     const std::vector<Slice>& inputSlices, const std::vector<Slice>& outputSlices)
     360              : {
     361            0 :     for (u32 i = 0; i < stepInfo.nSlices; i++) {
     362            0 :         rxSlices.push_back(inputSlices[stepInfo.rxSliceIdxs[i]]);
     363            0 :         rxSlicestemp.push_back(outputSlices[stepInfo.rxSliceIdxs[i]]);
     364            0 :         HCCL_DEBUG(
     365              :             "[ReduceScatterNHR][RunDestReducer] i[%u] rxSliceIndex[%u] rx offset[%llu] size[%llu]", i,
     366              :             stepInfo.rxSliceIdxs[i], outputSlices[stepInfo.rxSliceIdxs[i]].offset,
     367              :             outputSlices[stepInfo.rxSliceIdxs[i]].size);
     368              :     }
     369              : 
     370            0 :     HCCL_DEBUG(
     371              :         "[ReduceScatterNHR][RunDestReducer] rxslices size [%u], rxslices temp size [%u]", rxSlices.size(),
     372              :         rxSlicestemp.size());
     373              : 
     374              :     // 合并连续slices
     375            0 :     MergeSlices(rxSlices);
     376            0 :     MergeSlices(rxSlicestemp);
     377            0 :     HCCL_DEBUG(
     378              :         "[ReduceScatterNHR][RunDestReducer] merged rxslices size [%u], merged rxslices temp size [%u]", rxSlices.size(),
     379              :         rxSlicestemp.size());
     380            0 :     return HCCL_SUCCESS;
     381              : }
     382              : 
     383            0 : HcclResult ReduceScatterNHR::SdmaReducer(
     384              :     const u32 nSteps, const LINK& linkLeft, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     385              :     const std::vector<Slice>& outputSlices)
     386              : {
     387            0 :     HcclResult ret = HCCL_SUCCESS;
     388            0 :     std::vector<Slice> rxSlices;
     389            0 :     std::vector<Slice> rxSlicestemp;
     390            0 :     if ((INLINE_REDUCE_BITMASK & reduceAttr_) == 1) { // InlineReduce
     391            0 :         if (stepInfo.step == (nSteps - 1)) {
     392            0 :             ret = InlineReduceRxLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
     393              :         } else {
     394            0 :             CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
     395            0 :             ret = InlineReduceRx(linkLeft, rxSlices, rxSlicestemp);
     396              :         }
     397              :     } else { // TbeReduce
     398            0 :         if (stepInfo.step == (nSteps - 1)) {
     399            0 :             ret = TbeReduceRxLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
     400              :         } else {
     401            0 :             CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
     402            0 :             ret = TbeReduceRx(linkLeft, rxSlices, rxSlicestemp);
     403              :         }
     404              :     }
     405            0 :     return ret;
     406            0 : }
     407              : 
     408            0 : HcclResult ReduceScatterNHR::RunReduceScatterNHR(
     409              :     const u32 rank, const u32 rankSize, const std::vector<LINK>& links, const std::vector<Slice>& inputSlices,
     410              :     const std::vector<Slice>& outputSlices)
     411              : {
     412            0 :     bool bRetSize = (inputSlices.size() < rankSize);
     413            0 :     CHK_PRT_RET(
     414              :         bRetSize,
     415              :         HCCL_ERROR(
     416              :             "[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] inputslice size[%llu] is less "
     417              :             "than rank size[%u]",
     418              :             rank, outputSlices.size(), rankSize),
     419              :         HCCL_E_INTERNAL);
     420              : 
     421            0 :     bRetSize = (outputSlices.size() < rankSize);
     422            0 :     CHK_PRT_RET(
     423              :         bRetSize,
     424              :         HCCL_ERROR(
     425              :             "[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] outputslice size[%llu] is less "
     426              :             "than rank size[%u]",
     427              :             rank, outputSlices.size(), rankSize),
     428              :         HCCL_E_INTERNAL);
     429              : 
     430            0 :     HcclResult ret = HCCL_SUCCESS;
     431              : 
     432              :     // 计算通信步数
     433            0 :     u32 nSteps = GetStepNumInterServer(rankSize);
     434              : 
     435              :     // 逐步编排任务
     436            0 :     for (u32 step = 0; step < nSteps; step++) {
     437            0 :         InterServerAlgoStep stepInfo;
     438            0 :         GetStepInfo(step, nSteps, rank, rankSize, stepInfo);
     439              : 
     440              :         // 链的关系没有变化,区别的是发送的slice编号,因为重排tree不影响每棵树节点间的连接关系
     441            0 :         LINK linkLeft = links[stepInfo.fromRank];
     442            0 :         CHK_SMART_PTR_NULL(linkLeft);
     443              : 
     444            0 :         LINK linkRight = links[stepInfo.toRank];
     445            0 :         CHK_SMART_PTR_NULL(linkRight);
     446              : 
     447              :         // 当前每个数据块发送一次ACK、reduce一次、同步一次
     448            0 :         HCCL_DEBUG(
     449              :             "[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] rankSize[%u] from[%u] to[%u] step[%u] nSteps[%u] "
     450              :             "nSlices[%u]",
     451              :             rank, rankSize, stepInfo.fromRank, stepInfo.toRank, step, nSteps, stepInfo.nSlices);
     452              : 
     453            0 :         if (linkLeft->IsSpInlineReduce() && linkRight->IsSpInlineReduce()) { // SDMA
     454            0 :             CHK_RET(linkRight->TxAck(stream_));
     455            0 :             CHK_RET(linkLeft->RxAck(stream_));
     456            0 :             CHK_RET(SdmaReducer(nSteps, linkLeft, stepInfo, inputSlices, outputSlices));
     457            0 :             CHK_RET(linkLeft->TxDataSignal(stream_));  // 告知left我读完了
     458            0 :             CHK_RET(linkRight->RxDataSignal(stream_)); // 等right读完
     459              :         } else {                                       // RDMA
     460            0 :             CHK_RET(linkLeft->TxAck(stream_));
     461            0 :             CHK_RET(linkRight->RxAck(stream_));
     462              :             // tx
     463            0 :             ret = RunSourceSender(linkRight, stepInfo, inputSlices, outputSlices);
     464            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] Tx failed"), ret);
     465              : 
     466              :             // rx
     467            0 :             if (step == (nSteps - 1)) {
     468            0 :                 ret = RunDestReducerLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
     469              :             } else {
     470            0 :                 ret = RunDestReducer(linkLeft, stepInfo, inputSlices, outputSlices);
     471              :             }
     472              : 
     473            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] Rx failed"), ret);
     474            0 :             ret = linkLeft->PostFinAck(stream_);
     475            0 :             CHK_PRT_RET(
     476              :                 ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] PostFinAck failed"), ret);
     477              : 
     478            0 :             ret = linkRight->WaitFinAck(stream_);
     479            0 :             CHK_PRT_RET(
     480              :                 ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] WaitFinAck failed"), ret);
     481              : 
     482            0 :             if (barrierSwitchOn_) {
     483            0 :                 CHK_RET(ExecuteBarrier(linkLeft, linkRight));
     484              :             }
     485              :         }
     486            0 :     }
     487            0 :     return HCCL_SUCCESS;
     488              : }
     489              : 
     490            0 : HcclResult ReduceScatterNHR::RunSourceSender(
     491              :     const LINK& link, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     492              :     const std::vector<Slice>& outputSlices)
     493              : {
     494            0 :     std::vector<Slice> txSlices;
     495            0 :     std::vector<Slice> txSlicestemp;
     496            0 :     for (u32 i = 0; i < stepInfo.nSlices; i++) {
     497            0 :         txSlices.push_back(inputSlices[stepInfo.txSliceIdxs[i]]);
     498            0 :         txSlicestemp.push_back(outputSlices[stepInfo.txSliceIdxs[i]]);
     499            0 :         HCCL_DEBUG(
     500              :             "[ReduceScatterNHR][RunSourceSender] i[%u] txSliceIndex[%u] tx data offset[%llu] size[%llu]", i,
     501              :             stepInfo.txSliceIdxs[i], outputSlices[stepInfo.txSliceIdxs[i]].offset,
     502              :             outputSlices[stepInfo.txSliceIdxs[i]].size);
     503              :     }
     504            0 :     HCCL_DEBUG(
     505              :         "[ReduceScatterNHR][RunSourceSender] txSlices size [%u], txSlices temp size [%u]", txSlices.size(),
     506              :         txSlicestemp.size());
     507              : 
     508              :     // 合并连续slices
     509            0 :     MergeSlices(txSlices);
     510            0 :     MergeSlices(txSlicestemp);
     511            0 :     HCCL_DEBUG(
     512              :         "[ReduceScatterNHR][RunSourceSender] merged txSlices size [%u], merged txSlices temp size [%u]",
     513              :         txSlices.size(), txSlicestemp.size());
     514              : 
     515            0 :     std::vector<SenderMemoryInfo> txMems;
     516            0 :     for (u64 i = 0; i < txSlices.size(); i++) {
     517            0 :         DeviceMem srcMem = inputMem_.range(txSlices[i].offset, txSlices[i].size);
     518            0 :         HCCL_DEBUG(
     519              :             "[ReduceScatterNHR][RunSourceSender] send inputmem range[%llu], size[%llu] tx dstmem offset[%llu]",
     520              :             txSlices[i].offset, txSlices[i].size, txSlicestemp[i].offset);
     521            0 :         txMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSlicestemp[i].offset, srcMem});
     522            0 :     }
     523              : 
     524            0 :     CHK_RET(senderInfo_->run(link, txMems, stream_));
     525            0 :     return HCCL_SUCCESS;
     526            0 : }
     527              : 
     528            0 : HcclResult ReduceScatterNHR::RunDestReducer(
     529              :     const LINK& link, InterServerAlgoStep& stepInfo, const std::vector<Slice>& inputSlices,
     530              :     const std::vector<Slice>& outputSlices)
     531              : {
     532            0 :     std::vector<Slice> rxSlices;
     533            0 :     std::vector<Slice> rxSlicestemp;
     534            0 :     CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
     535              : 
     536            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     537            0 :     for (u64 i = 0; i < rxSlices.size(); i++) {
     538            0 :         DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
     539            0 :         DeviceMem srcMemTemp = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
     540            0 :         HCCL_DEBUG(
     541              :             "[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
     542              :             "offset[%llu] size[%llu] ",
     543              :             rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
     544            0 :         rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlices[i].offset, dstMem, dstMem, srcMemTemp});
     545            0 :     }
     546              : 
     547            0 :     CHK_RET(reducerInfo_->run(dispatcher_, link, rxReduceMems, stream_));
     548            0 :     return HCCL_SUCCESS;
     549            0 : }
     550              : 
     551              : // NHR每步的算法描述原理函数
     552            0 : HcclResult ReduceScatterNHR::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep& stepInfo)
     553              : {
     554              :     (void)nSteps;
     555            0 :     stepInfo.txSliceIdxs.clear();
     556            0 :     stepInfo.rxSliceIdxs.clear();
     557            0 :     u32 sliceSize = slices_.size() / rankSize;
     558            0 :     stepInfo.step = step;
     559            0 :     stepInfo.myRank = rank;
     560              : 
     561              :     // 计算通信对象
     562            0 :     u32 deltaRank = 1 << step;
     563            0 :     u32 sendTo = (rank + rankSize - deltaRank) % rankSize;
     564            0 :     u32 recvFrom = (rank + deltaRank) % rankSize;
     565              : 
     566              :     // 数据份数和数据编号增量
     567            0 :     u32 nSlices = (rankSize - 1 + (1 << step)) / (1 << (step + 1));
     568            0 :     u32 deltaSliceIndex = 1 << (step + 1);
     569            0 :     u32 txSliceIdx = sendTo; // 第一片rank
     570            0 :     u32 rxSliceIdx = rank;
     571              : 
     572            0 :     for (u32 i = 0; i < nSlices; i++) {
     573            0 :         for (u32 j = 0; j < sliceSize; j++) {
     574            0 :             u32 targetTxSliceIdx = sliceMap_[txSliceIdx];
     575            0 :             stepInfo.txSliceIdxs.push_back(targetTxSliceIdx * sliceSize + j);
     576              : 
     577            0 :             u32 targetRxSliceIdx = sliceMap_[rxSliceIdx];
     578            0 :             stepInfo.rxSliceIdxs.push_back(targetRxSliceIdx * sliceSize + j);
     579              : 
     580            0 :             HCCL_DEBUG(
     581              :                 "[ReduceScatterNHR][GetStepInfo] i[%u] txSliceIdx[%u]->targetTxSliceIdx[%u] rxSliceIdx[%u]->"
     582              :                 "targetRxSliceIdx[%u]",
     583              :                 i, txSliceIdx, targetTxSliceIdx, rxSliceIdx, targetRxSliceIdx);
     584              :         }
     585            0 :         txSliceIdx = (txSliceIdx + rankSize - deltaSliceIndex) % rankSize;
     586            0 :         rxSliceIdx = (rxSliceIdx + rankSize - deltaSliceIndex) % rankSize;
     587              :     }
     588              : 
     589            0 :     stepInfo.nSlices = nSlices * sliceSize;
     590            0 :     stepInfo.toRank = sendTo;
     591            0 :     stepInfo.fromRank = recvFrom;
     592            0 :     return HCCL_SUCCESS;
     593              : }
     594              : 
     595            0 : HcclResult ReduceScatterNHR::GetNslbAdjInfo(
     596              :     const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
     597              : {
     598            0 :     if (rankSize == 1) {
     599            0 :         return HCCL_SUCCESS;
     600              :     }
     601            0 :     if (links.size() < rankSize) {
     602            0 :         return HCCL_SUCCESS;
     603              :     }
     604            0 :     u32 nSteps = 0;
     605            0 :     for (u32 temp = rankSize - 1; temp != 0; temp >>= 1, ++nSteps) {
     606              :     }
     607              : 
     608            0 :     for (u32 step = 0; step < nSteps; step++) {
     609            0 :         u32 deltaRank = 1 << step;
     610            0 :         u32 sendTo = (rank + rankSize - deltaRank) % rankSize;
     611              :         ;
     612            0 :         LINK linkRight = links[sendTo];
     613            0 :         CHK_SMART_PTR_NULL(linkRight);
     614              : 
     615            0 :         NslbDpAdjInfo adjInfoStep = {};
     616            0 :         adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
     617            0 :         adjInfoStep.phaseId = step + 1;
     618            0 :         adjInfoStep.rev = 0;
     619            0 :         nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
     620            0 :     }
     621            0 :     nslbAdjInfo.dstRankNum = nSteps;
     622            0 :     return HCCL_SUCCESS;
     623              : }
     624              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_NHR, ReduceScatterNHR);
     625              : } // namespace hccl
        

Generated by: LCOV version 2.0-1