LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - aligned_reduce_scatter_double_ring.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 4.1 % 390 16
Test Date: 2026-08-04 10:52:23 Functions: 10.0 % 30 3

            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 "aligned_reduce_scatter_double_ring.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15           32 : AlignedReduceScatterDoubleRing::AlignedReduceScatterDoubleRing(
      16           32 :     const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
      17              : {
      18           32 : }
      19              : 
      20           32 : AlignedReduceScatterDoubleRing::~AlignedReduceScatterDoubleRing()
      21              : {
      22           32 : }
      23              : 
      24           32 : HcclResult AlignedReduceScatterDoubleRing::Prepare(DeviceMem &inputMem, DeviceMem &outputMem,
      25              :     DeviceMem &scratchMem, const u64 count, const HcclDataType dataType, const Stream &stream,
      26              :     const std::vector<std::vector<Slice>> &multRingsSlices, const HcclReduceOp reductionOp, const u32 root,
      27              :     const u64 baseOffset, const bool disableDMAReduce, const u64 reduceAttrBitMap, const HcomCollOpInfo *opInfo,
      28              :     const u32 userRank, std::vector<Stream> &subStreams, const std::vector<std::shared_ptr<LocalNotify>> &mainSignals,
      29              :     const std::vector<std::shared_ptr<LocalNotify>> &subSignals, const std::vector<std::vector<u32>> &ringsOrders,
      30              :     const std::vector<std::vector<Slice>> &userMemInputSlicesOfDoubleRing)
      31              : {
      32           32 :     reduceAttr_ = reduceAttrBitMap;
      33           32 :     opInfo_ = opInfo;
      34           32 :     userRank_ = userRank;
      35           32 :     subStreams_ = subStreams;
      36           32 :     mainSignals_ = mainSignals;
      37           32 :     subSignals_ = subSignals;
      38           32 :     ringsOrders_ = ringsOrders;
      39           32 :     userMemInputSlicesOfDoubleRing_ = userMemInputSlicesOfDoubleRing;
      40           32 :     return AlgTemplateBase::Prepare(inputMem, outputMem, scratchMem, count, dataType, stream, multRingsSlices,
      41           32 :         reductionOp, root, baseOffset, disableDMAReduce);
      42              : }
      43              : 
      44              : // reduce scatter ring direct算法的函数入口
      45            0 : HcclResult AlignedReduceScatterDoubleRing::RunAsync(const u32 rank, const u32 rankSize,
      46              :                                                        const std::vector<LINK> &links)
      47              : {
      48              :     // 基本的检查
      49            0 :     CHK_RET(CheckParameters(rank, rankSize, links));
      50              : 
      51              :     // 判断rank_size == 1的情况,并拷贝
      52            0 :     if (rankSize == 1) {
      53            0 :         CHK_RET(OneRankMemcpy());
      54            0 :         return HCCL_SUCCESS;
      55              :     }
      56              :     // 收集本地mem信息
      57            0 :     CHK_RET(InitSenderReducer());
      58              : 
      59              :     // 收集邻居信息
      60            0 :     CHK_RET(GetInitializedNeighborLinks(rank, rankSize, links));
      61              : 
      62              :     // 填充slice_
      63            0 :     CHK_RET(SetSlices(rank, rankSize));
      64              : 
      65              :     // 运行reduce-scatter, ring算法
      66            0 :     CHK_RET(RunReduceScatter(rank, rankSize));
      67              : 
      68            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, stream_, subStreams_));
      69              : 
      70            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished: rank[%u] end", rank);
      71            0 :     return HCCL_SUCCESS;
      72              : }
      73              : 
      74            0 : HcclResult AlignedReduceScatterDoubleRing::CheckParameters(const u32 rank, const u32 rankSize,
      75              :                                                               const std::vector<LINK> &links)
      76              : {
      77            0 :     CHK_PTR_NULL(opInfo_);
      78            0 :     CHK_RET(CheckConcurrentDirectParameters(rank, rankSize, links));
      79              :     // 判断subStreams数量是否正确
      80            0 :     CHK_PRT_RET(
      81              :         subStreams_.size() < 1,
      82              :         HCCL_ERROR("[AlignedReduceScatterDoubleRing] subStreams size[%u] is less than 1", subStreams_.size()),
      83              :         HCCL_E_PARA);
      84            0 :     for (auto &s : subStreams_) {
      85            0 :         CHK_PTR_NULL(s.ptr());
      86              :     }
      87              :     // 判断mainSignals数量是否正确
      88            0 :     CHK_PRT_RET(
      89              :         mainSignals_.size() < 1,
      90              :         HCCL_ERROR("[AlignedReduceScatterDoubleRing] mainSignals size[%u] is less than 1", mainSignals_.size()),
      91              :         HCCL_E_PARA);
      92              :     // 判断subSignals数量是否正确
      93            0 :     CHK_PRT_RET(
      94              :         subSignals_.size() < 1,
      95              :         HCCL_ERROR("[AlignedReduceScatterDoubleRing] subSignals size[%u] is less than 1", subSignals_.size()),
      96              :         HCCL_E_PARA);
      97              :     // 判断ringsOrder数量是否正确
      98            0 :     for (u32 ringIndex = 0; ringIndex < ringsOrders_.size(); ringIndex++) {
      99            0 :         CHK_PRT_RET(ringsOrders_[ringIndex].size() != rankSize,
     100              :             HCCL_ERROR("[AlignedReduceScatterDoubleRing] ringsOrders[%u] size[%u] is not equal to rank size[%u]",
     101              :                 ringIndex, ringsOrders_[ringIndex].size(), rankSize),
     102              :             HCCL_E_PARA);
     103              :     }
     104              :     // 判断userMemInputSlices数量是否正确
     105            0 :     for (u32 ringIndex = 0; ringIndex < userMemInputSlicesOfDoubleRing_.size(); ringIndex++) {
     106            0 :         CHK_PRT_RET(userMemInputSlicesOfDoubleRing_[ringIndex].size() % rankSize != 0,
     107              :             HCCL_ERROR("[AlignedReduceScatterDoubleRing] userMemInputSlicesOfDoubleRing[%u] size[%u] can not divided by size[%u]",
     108              :                 ringIndex, userMemInputSlicesOfDoubleRing_[ringIndex].size(), rankSize),
     109              :             HCCL_E_PARA);
     110              :     }
     111            0 :     u32 mainSliceSize = multRingsSlices_[ALIGNED_MAIN_RING_INDEX].size() / rankSize;
     112            0 :     u32 subSliceSize = multRingsSlices_[ALIGNED_SUB_RING_INDEX].size() / rankSize;
     113            0 :     CHK_PRT_RET(mainSliceSize != subSliceSize,
     114              :         HCCL_ERROR("[AlignedReduceScatterDoubleRing] mainSliceSize[%u] is not equal to subSliceSize[%u].",
     115              :             mainSliceSize, subSliceSize),
     116              :         HCCL_E_PARA);
     117            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished to CheckParameters");
     118            0 :     return HCCL_SUCCESS;
     119              : }
     120              : 
     121            0 : HcclResult AlignedReduceScatterDoubleRing::OneRankMemcpy()
     122              : {
     123            0 :     CHK_RET(MainRecordSub()); // 主流通知从流开始通信
     124            0 :     CHK_RET(SubWaitMain());   // 从流等待主流通知
     125            0 :     for (u32 ringIndex = 0; ringIndex < multRingsSlices_.size(); ringIndex++) {
     126            0 :         for (u32 sliceIdx = 0; sliceIdx < multRingsSlices_[ringIndex].size(); sliceIdx++) {
     127            0 :             const Slice &srcSlice = userMemInputSlicesOfDoubleRing_[ringIndex][sliceIdx];
     128            0 :             const Slice &dstSlice = multRingsSlices_[ringIndex][sliceIdx];
     129            0 :             DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcSlice.offset, srcSlice.size);
     130            0 :             DeviceMem dst;
     131            0 :             if (opInfo_->outputAddr != nullptr) {
     132              :                 // opInfo_->outputAddr != nullptr指示要将输出发送至user output
     133            0 :                 u64 stepOffset = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][0]].offset;
     134            0 :                 HCCL_DEBUG("Memcpy operation: stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at userMemOut_",
     135              :                     userRank_, stepOffset, dstSlice.size);
     136            0 :                 dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + stepOffset, dstSlice.size);
     137              :             } else {
     138              :                 // opInfo_->outputAddr == nullptr指示要将输出发送至CCL buffer
     139            0 :                 HCCL_DEBUG("Memcpy operation: stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at outputMem_",
     140              :                     userRank_, dstSlice.offset, dstSlice.size);
     141            0 :                 dst = outputMem_.range(dstSlice.offset, dstSlice.size);
     142              :             }
     143            0 :             if (ringIndex == 1) {
     144            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     145              :             } else {
     146            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
     147              :             }
     148            0 :         }
     149            0 :         HCCL_DEBUG("[AlignedReduceScatterDoubleRing][OneRankMemcpy] ringIndex[%u] Memcpy success", ringIndex);
     150              :     }
     151            0 :     CHK_RET(SubRecordMain()); // 从流通知主流通信完成
     152            0 :     CHK_RET(MainWaitSub());   // 主流等待从流通知
     153            0 :     return HCCL_SUCCESS;
     154              : }
     155              : 
     156            0 : HcclResult AlignedReduceScatterDoubleRing::InitSenderReducer()
     157              : {
     158              :     // 创建reducer & sender
     159            0 :     senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
     160            0 :     CHK_SMART_PTR_NULL(senderInfo_);
     161              : 
     162            0 :     reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
     163            0 :     CHK_SMART_PTR_NULL(reducerInfo_);
     164            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished to InitSenderReducer");
     165            0 :     return HCCL_SUCCESS;
     166              : }
     167              : 
     168            0 : HcclResult AlignedReduceScatterDoubleRing::GetInitializedNeighborLinks(const u32 rank, const u32 rankSize,
     169              :                                                                           const std::vector<LINK> &links)
     170              : {
     171              :     // 收集左邻居信息
     172            0 :     leftLink_ = links[(rank + rankSize - 1) % rankSize];
     173            0 :     CHK_SMART_PTR_NULL(leftLink_);
     174              : 
     175              :     // 收集右邻居信息
     176            0 :     rightLink_ = links[(rank + 1) % rankSize];
     177            0 :     CHK_SMART_PTR_NULL(rightLink_);
     178            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished to GetInitializedNeighborLinks");
     179            0 :     return HCCL_SUCCESS;
     180              : }
     181              : 
     182            0 : HcclResult AlignedReduceScatterDoubleRing::SetSlices(const u32 rank, const u32 rankSize)
     183              : {
     184            0 :     for (u32 ringIndex = 0; ringIndex < multRingsSlices_.size(); ringIndex++) {
     185            0 :         if (multRingsSlices_[ringIndex].size() == 0) {
     186            0 :             multRingsSlices_[ringIndex].resize(rankSize);
     187              : 
     188              :             // 生成std::vector<Slice> multRingsSlices_[ringIndex]
     189            0 :             u64 sliceSize = count_ * SIZE_TABLE[dataType_];
     190              : 
     191            0 :             for (u32 i = 0; i < rankSize; i++) {
     192            0 :                 multRingsSlices_[ringIndex][i].size = sliceSize;
     193              :                 // 用于DMA消减过程中,消除src与dst不对位的风险
     194            0 :                 multRingsSlices_[ringIndex][i].offset = RoundUpWithDivisor(i * sliceSize, HCCL_MIN_SLICE_ALIGN);
     195              : 
     196            0 :                 HCCL_DEBUG("multRingsSlices_[%u], rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]",
     197              :                     ringIndex, rank, i, multRingsSlices_[ringIndex][i].offset, i,
     198              :                         multRingsSlices_[ringIndex][i].size);
     199              :             }
     200              :         }
     201            0 :         for (u32 i = 0; i < multRingsSlices_[ringIndex].size(); i++) {
     202            0 :             HCCL_DEBUG(
     203              :                 "[AlignedReduceScatterDoubleRing][SetSlices] multRingsSlices_[%u], rank[%u], "
     204              :                 "slices[%u].offset=[%llu], slices[%u].size=[%llu]",
     205              :                 ringIndex, rank, i, multRingsSlices_[ringIndex][i].offset, i, multRingsSlices_[ringIndex][i].size);
     206              :         }
     207              :         // 最后一步搬到userMemOut_的offset, 不同的ring环offset不一样
     208              :         u64 toUserMemOffset;
     209            0 :         if (ringIndex == 0) {
     210            0 :             toUserMemOffset = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][0]].offset;
     211              :         } else {
     212            0 :             const auto &prevRingSlice = multRingsSlices_[ringIndex - 1][ringsOrders_[ringIndex - 1][rank]];
     213            0 :             const auto &slice = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][rank]];
     214            0 :             toUserMemOffset = slice.offset - prevRingSlice.offset;
     215              :         }
     216            0 :         HCCL_DEBUG("[AlignedReduceScatterDoubleRing][SetSlices] rank[%u], ring[%u], toUserMemOffset[%u]", rank,
     217              :             ringIndex, toUserMemOffset);
     218            0 :         lastStepOffsets_.emplace_back(toUserMemOffset);
     219              :     }
     220            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished to SetSlices");
     221            0 :     return HCCL_SUCCESS;
     222              : }
     223              : 
     224            0 : HcclResult AlignedReduceScatterDoubleRing::PrepareInitSlices(const u32 rankSize,
     225              :     u64 ringIndex, u32 discontinuousSliceSize, u32 discontinuousSliceIdx, u32 initSlice0Idx, u32 initSlice1Idx,
     226              :     DeviceMem &dstInit, DeviceMem &srcInit, DeviceMem &dstSubInit, DeviceMem &srcSubInit)
     227              : {
     228              :     // 第-1步,片内将部分数据从userIn搬到cclIn
     229            0 :     const Slice &srcInitSlice0 = userMemInputSlicesOfDoubleRing_[ringIndex][initSlice0Idx * discontinuousSliceSize + discontinuousSliceIdx];
     230              :     srcInit
     231            0 :         = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice0.offset, srcInitSlice0.size);
     232            0 :     const Slice &dstInitSlice0 = multRingsSlices_[ringIndex][initSlice0Idx * discontinuousSliceSize + discontinuousSliceIdx];
     233            0 :     dstInit    = inputMem_.range(dstInitSlice0.offset, dstInitSlice0.size);
     234              : 
     235            0 :     const Slice &srcInitSlice1 = userMemInputSlicesOfDoubleRing_[ringIndex][initSlice1Idx * discontinuousSliceSize + discontinuousSliceIdx];
     236              :     srcSubInit
     237            0 :         = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice1.offset, srcInitSlice1.size);
     238            0 :     const Slice &dstInitSlice1 = multRingsSlices_[ringIndex][initSlice1Idx * discontinuousSliceSize + discontinuousSliceIdx];
     239            0 :     dstSubInit       = inputMem_.range(dstInitSlice1.offset, dstInitSlice1.size);
     240              :     // 第-1步并发
     241            0 :     if (rankSize == TWO_RANK_SIZE && opInfo_->outputAddr != nullptr) {
     242            0 :         HCCL_DEBUG(
     243              :             "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     244              :             "userMemInput to offset[%llu], size[%llu] on userMemOut_",
     245              :             userRank_, srcInitSlice1.offset, srcInitSlice1.size, lastStepOffsets_[ringIndex], dstInitSlice1.size);
     246            0 :         dstInit = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
     247            0 :                                     dstInitSlice1.size);
     248              :     } else {
     249            0 :         HCCL_DEBUG(
     250              :             "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     251              :             "userMemInput to offset[%llu], size[%llu] on CCL",
     252              :             userRank_, srcInitSlice1.offset, srcInitSlice1.size, dstInitSlice1.offset, dstInitSlice1.size);
     253              :     }
     254            0 :     HCCL_DEBUG("Memcpy operation: step[-1] stream[sub] src rank[%u] starts to copy(rcv) offset[%llu], "
     255              :         " size[%llu] on userMemInput to offset[%llu], size[%llu] on CCL",
     256              :         userRank_, srcInitSlice0.offset, srcInitSlice0.size, dstInitSlice0.offset, dstInitSlice0.size);
     257            0 :     return HCCL_SUCCESS;
     258              : }
     259              : 
     260            0 : HcclResult AlignedReduceScatterDoubleRing::MemcpyInitSlicesOnMainStreams(
     261              :     u64 ringIndex, DeviceMem &dstInit, DeviceMem &srcInit)
     262              : {
     263            0 :     if (ringIndex == 1) {
     264            0 :         CHK_RET(MainWaitSub());
     265            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     266            0 :         CHK_RET(MainRecordSub());
     267            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, stream_));
     268              :     } else {
     269            0 :         CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     270            0 :         CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     271            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, subStreams_[0]));
     272              :     }
     273            0 :     return HCCL_SUCCESS;
     274              : }
     275              : 
     276            0 : HcclResult AlignedReduceScatterDoubleRing::MemcpyInitSlices(
     277              :     u64 ringIndex, DeviceMem &dstInit, DeviceMem &srcInit, DeviceMem &dstSubInit, DeviceMem &srcSubInit)
     278              : {
     279            0 :     CHK_RET(MemcpyInitSlicesOnMainStreams(ringIndex, dstInit, srcInit));
     280            0 :     if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB && (!disableDMAReduce_)) {
     281            0 :         HCCL_DEBUG("[AlignedReduceScatterDoubleRing][MemcpyInitSlices] no graph mode");
     282            0 :         CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
     283            0 :         CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
     284            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstSubInit, srcSubInit, subStreams_[ringIndex + 1]));
     285              :     } else {
     286            0 :         HCCL_DEBUG("[AlignedReduceScatterDoubleRing][MemcpyInitSlices] graph mode");
     287            0 :         CHK_RET(MemcpyInitSlicesOnMainStreams(ringIndex, dstSubInit, srcSubInit));
     288              :     }
     289            0 :     return HCCL_SUCCESS;
     290              : }
     291              : 
     292            0 : HcclResult AlignedReduceScatterDoubleRing::RunInitStep(const u32 rank, const u32 rankSize)
     293              : {
     294              :     //主环初始indexes
     295            0 :     u32 initSlice0Idx    = (rankSize - rank - 1 + rankSize) % rankSize;
     296            0 :     u32 initSlice1Idx    = (rankSize - rank - DMA_REDUCE_TWO_OFFSET + rankSize) % rankSize;
     297              :     // 从环初始indexes
     298            0 :     u32 subInitSlice0Idx     = (rank + rankSize - 1) % rankSize;
     299            0 :     u32 subInitSlice1Idx     = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     300            0 :     u32 discontinuousSliceSize = multRingsSlices_[ALIGNED_SUB_RING_INDEX].size() / rankSize;
     301            0 :     DeviceMem dstInit;
     302            0 :     DeviceMem srcInit;
     303            0 :     DeviceMem dstSubInit;
     304            0 :     DeviceMem srcSubInit;
     305            0 :     DeviceMem subDstInit;
     306            0 :     DeviceMem subSrcInit;
     307            0 :     DeviceMem subDstSubInit;
     308            0 :     DeviceMem subSrcSubInit;
     309            0 :     for (u32 discontinuousSliceIdx = 0; discontinuousSliceIdx < discontinuousSliceSize; discontinuousSliceIdx++) {
     310            0 :         CHK_RET(PrepareInitSlices(rankSize, ALIGNED_SUB_RING_INDEX,
     311              :             discontinuousSliceSize, discontinuousSliceIdx, subInitSlice0Idx, subInitSlice1Idx,
     312              :             subDstInit, subSrcInit, subDstSubInit, subSrcSubInit));
     313            0 :         CHK_RET(PrepareInitSlices(rankSize, ALIGNED_MAIN_RING_INDEX,
     314              :             discontinuousSliceSize, discontinuousSliceIdx, initSlice0Idx, initSlice1Idx,
     315              :             dstInit, srcInit, dstSubInit, srcSubInit));
     316            0 :         HCCL_DEBUG("Memcpy operation: step[-1] starts on ring[%u]", ALIGNED_SUB_RING_INDEX);
     317            0 :         CHK_RET(MemcpyInitSlices(ALIGNED_SUB_RING_INDEX, subDstInit, subSrcInit, subDstSubInit, subSrcSubInit));
     318            0 :         HCCL_DEBUG("Memcpy operation: step[-1] starts on ring[%u]", ALIGNED_MAIN_RING_INDEX);
     319            0 :         CHK_RET(MemcpyInitSlices(ALIGNED_MAIN_RING_INDEX, dstInit, srcInit, dstSubInit, srcSubInit));
     320              :     }
     321            0 :     return HCCL_SUCCESS;
     322            0 : }
     323              : 
     324            0 : HcclResult AlignedReduceScatterDoubleRing::PrepareRunMainStream(u32 ringIndex, Stream &stream,
     325              :     LINK &preLink, LINK &nextLink)
     326              : {
     327            0 :     HCCL_DEBUG("AlignedReduceScatterDoubleRing PrepareRunMainStream start");
     328            0 :     if (ringIndex == 1) {
     329            0 :         stream = stream_;
     330            0 :         preLink = rightLink_;
     331            0 :         nextLink = leftLink_;
     332              :     } else {
     333            0 :         stream = subStreams_[0];
     334            0 :         preLink = leftLink_;
     335            0 :         nextLink = rightLink_;
     336              :     }
     337            0 :     HCCL_DEBUG("AlignedReduceScatterDoubleRing PrepareRunMainStream end");
     338            0 :     return HCCL_SUCCESS;
     339              : }
     340              : 
     341            0 : HcclResult AlignedReduceScatterDoubleRing::PreSync(const u32 ringIndex)
     342              : {
     343            0 :     if (ringIndex == 1) {
     344            0 :         CHK_RET(MainWaitSub());
     345            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     346            0 :         CHK_RET(MainRecordSub());
     347              :     } else {
     348            0 :         CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     349            0 :         CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     350              :     }
     351            0 :     return HCCL_SUCCESS;
     352              : }
     353              : 
     354            0 : HcclResult AlignedReduceScatterDoubleRing::PrepareDeviceMems(
     355              :     const u32 step, const u32 ringIndex, const u32 rankSize,
     356              :     const u32 txSliceIdx, const u32 rxSliceIdx, const u32 subSliceIdx,
     357              :     std::vector<SenderMemoryInfo> &txReduceMems, std::vector<ReducerMemoryInfo> &rxReduceMems,
     358              :     std::vector<DeviceMem> &localSrcMems, std::vector<DeviceMem> &localDstMems)
     359              : {
     360            0 :     u32 sliceSize = multRingsSlices_[ringIndex].size() / rankSize;
     361            0 :     for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     362            0 :         const Slice &rxSlice = multRingsSlices_[ringIndex][rxSliceIdx * sliceSize + sliceIdx];
     363            0 :         const Slice &cclSlice = multRingsSlices_[ringIndex][subSliceIdx * sliceSize + sliceIdx];
     364            0 :         const Slice &txSlice = multRingsSlices_[ringIndex][txSliceIdx * sliceSize + sliceIdx];
     365            0 :         const Slice &subSlice = userMemInputSlicesOfDoubleRing_[ringIndex][subSliceIdx * sliceSize + sliceIdx];
     366              :         // PrepareReduceDeviceMems
     367              :         // Ack
     368            0 :         DeviceMem dst;
     369            0 :         if (step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
     370            0 :             HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     371              :                 "at userMemOut_", step, userRank_, lastStepOffsets_[ringIndex], rxSlice.size);
     372            0 :             dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
     373            0 :                 rxSlice.size);
     374              :         } else {
     375            0 :             HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     376              :                 "at inputMem_",
     377              :                 step, userRank_, rxSlice.offset, rxSlice.size);
     378            0 :             dst = inputMem_.range(rxSlice.offset, rxSlice.size);
     379              :         }
     380              :         // 在inline reduce场景, 需要利用scratchMem_暂存
     381            0 :         DeviceMem srcMemTemp = scratchMem_.range(rxSlice.offset, rxSlice.size);
     382            0 :         DeviceMem srcMem     = inputMem_.range(txSlice.offset, txSlice.size);
     383            0 :         HCCL_DEBUG("Reduce operation: step[%u] stream[main], receiver starts to rcv offset[%llu], size[%llu]",
     384              :             step, txSlice.offset, txSlice.size);
     385            0 :         rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlice.offset, dst, dst, srcMemTemp});
     386            0 :         txReduceMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSlice.offset, srcMem});
     387              : 
     388              :         // PrepareLocalCopyDeviceMems
     389            0 :         DeviceMem localSrt;
     390            0 :         DeviceMem localDst;
     391            0 :         if (step == rankSize - DMA_REDUCE_TWO_OFFSET) {
     392              :             // do nothing
     393            0 :         } else if (step == rankSize - DMA_REDUCE_THREE_OFFSET && opInfo_->outputAddr != nullptr) {
     394            0 :             HCCL_DEBUG("Memcpy operation: step[%u] subStream[%u], src rank[%u] sends offset[%llu], size[%llu], "
     395              :                 "dst rank[%u] starts to rcv offset[%llu], size[%llu], "
     396              :                 "from userMemIn_ to userMemOut_", step, ringIndex + 1, userRank_, subSlice.offset, subSlice.size,
     397              :                 userRank_, lastStepOffsets_[ringIndex], subSlice.size);
     398            0 :             localSrt = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSlice.offset,
     399            0 :                 subSlice.size);
     400            0 :             localDst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
     401            0 :                 subSlice.size);
     402              :         } else {
     403            0 :             HCCL_DEBUG("Memcpy operation: step[%u] subStream[%u], src rank[%u] sends offset[%llu], size[%llu], "
     404              :                 "dst rank[%u] starts to rcv offset[%llu], size[%llu], "
     405              :                 "from userMemIn_ to inputMem_", step, ringIndex + 1, userRank_, subSlice.offset, subSlice.size,
     406              :                 userRank_, cclSlice.offset, cclSlice.size);
     407            0 :             localSrt = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSlice.offset,
     408            0 :                 subSlice.size);
     409            0 :             localDst = inputMem_.range(cclSlice.offset, cclSlice.size);
     410              :         }
     411            0 :         localSrcMems.emplace_back(localSrt);
     412            0 :         localDstMems.emplace_back(localDst);
     413            0 :     }
     414            0 :     return HCCL_SUCCESS;
     415              : }
     416              : 
     417            0 : HcclResult AlignedReduceScatterDoubleRing::RxAsyncMemcpy(
     418              :     const u32 ringIndex, RxMemoryInfo& mem, Stream &stream, const LINK &link)
     419              : {
     420              :     // PreSync
     421            0 :     CHK_RET(PreSync(ringIndex));
     422            0 :     CHK_PTR_NULL(mem.dst);
     423            0 :     void *srcMemPtr = nullptr;
     424            0 :     CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
     425              : 
     426            0 :     DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + mem.srcOffset, mem.len);
     427            0 :     DeviceMem dstDevMem(static_cast<s8 *>(mem.dst), mem.len);
     428            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstDevMem, srcDevMem,
     429              :         stream, link->GetRemoteRank(), link->GetLinkType()));
     430            0 :     return HCCL_SUCCESS;
     431            0 : }
     432              : 
     433            0 : HcclResult AlignedReduceScatterDoubleRing::ReducerRun(const u32 ringIndex, const HcclDispatcher dispatcher,
     434              :     const LINK &link,
     435              :     ReducerMemoryInfo &reduceMem, Stream &stream)
     436              : {
     437            0 :     CHK_PTR_NULL(stream.ptr());
     438            0 :     bool isSpInlineReduce = link->IsSpInlineReduce();
     439            0 :     HcclResult ret = HCCL_SUCCESS;
     440            0 :     if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_))) {
     441            0 :         void *remoteMem = nullptr;
     442            0 :         CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     443            0 :         const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     444            0 :         CHK_RET(PreSync(ringIndex));
     445            0 :         CHK_RET(
     446              :             HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
     447              :             dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
     448              :             link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
     449              : 
     450            0 :         if (reduceMem.localsrc != reduceMem.localdst) {
     451            0 :             ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     452            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     453              :                 HCCL_ERROR("[AlignedReduceScatterDoubleRing][Run]memcpy_async localSrc[%p] localDst[%p] failed",
     454              :                 reduceMem.localsrc.ptr(), reduceMem.localdst.ptr()),
     455              :                 ret);
     456              :         }
     457            0 :     } else {
     458            0 :         RxMemoryInfo rxMem = RxMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     459            0 :                 reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() };
     460              : 
     461            0 :         u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
     462            0 :         DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
     463            0 :         RxWithReduceMemoryInfo rxWithReduceMem = RxWithReduceMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     464            0 :                 reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size(), reduceSrc.ptr(), reduceMem.localdst.ptr(),
     465            0 :                 dataCount };
     466            0 :         CHK_RET(RxAsyncMemcpy(ringIndex, rxMem, stream, link));
     467            0 :         RxWithReduceMemoryInfo &rxReduceMem = rxWithReduceMem;
     468            0 :         if (ringIndex == ALIGNED_SUB_RING_INDEX) {
     469            0 :             CHK_PRT_RET(stream != subStreams_[0],
     470              :                 HCCL_ERROR("[%s] subStreams_[0] should be used for ringIndex=%d", __func__, ALIGNED_SUB_RING_INDEX), HCCL_E_INTERNAL);
     471            0 :             CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     472            0 :             CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[0], profilerInput_.stage));
     473              :         }
     474            0 :         CHK_RET(HcclReduceAsync(dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_,
     475              :             reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
     476              :             reduceAttr_));
     477            0 :         if (ringIndex == ALIGNED_SUB_RING_INDEX) {
     478            0 :             CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[0], profilerInput_.stage));
     479            0 :             CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     480              :         }
     481            0 :     }
     482            0 :     return HCCL_SUCCESS;
     483              : }
     484              : 
     485            0 : HcclResult AlignedReduceScatterDoubleRing::LocalMemcpy(const u32 step, const u32 rankSize, const u32 ringIndex,
     486              :     DeviceMem &localSrcMem, DeviceMem &localDstMem)
     487              : {
     488              :     // 通过校验流数判断是单算子模式还是图模式
     489            0 :     if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB && (!disableDMAReduce_)) {
     490            0 :         CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
     491            0 :         CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
     492            0 :         if (localSrcMem != localDstMem && step != rankSize - DMA_REDUCE_TWO_OFFSET) {
     493            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, subStreams_[ringIndex + 1]));
     494              :         }
     495              :     } else {
     496              :         // 图模式
     497            0 :         CHK_RET(PreSync(ringIndex));
     498            0 :         if (localSrcMem != localDstMem && step != rankSize - DMA_REDUCE_TWO_OFFSET) {
     499            0 :             if (ringIndex == 1) {
     500            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, stream_));
     501              :             } else {
     502            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, subStreams_[0]));
     503              :             }
     504              :         }
     505              :     }
     506            0 :     return HCCL_SUCCESS;
     507              : }
     508              : 
     509            0 : HcclResult AlignedReduceScatterDoubleRing::RunSubStream(
     510              :     const u32 step, const u32 rankSize, u32 ringIndex,
     511              :     std::vector<DeviceMem> &localSrcMems, std::vector<DeviceMem> &localDstMems)
     512              : {
     513            0 :     for (u32 sliceIdx = 0; sliceIdx < localSrcMems.size(); sliceIdx++) {
     514            0 :         CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
     515            0 :         CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
     516            0 :         if (step != rankSize - DMA_REDUCE_TWO_OFFSET) {
     517            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMems[sliceIdx], localSrcMems[sliceIdx], subStreams_[ringIndex + 1]));
     518              :         }
     519              :     }
     520            0 :     return HCCL_SUCCESS;
     521              : }
     522              : 
     523            0 : HcclResult AlignedReduceScatterDoubleRing::RunAllStreams(const u32 step, const u32 rankSize,
     524              :     std::vector<SenderMemoryInfo> &mainTxReduceMems, std::vector<ReducerMemoryInfo> &mainRxReduceMems,
     525              :     std::vector<SenderMemoryInfo> &subTxReduceMems, std::vector<ReducerMemoryInfo> &subRxReduceMems,
     526              :     std::vector<DeviceMem> &mainLocalSrcMems, std::vector<DeviceMem> &mainLocalDstMems,
     527              :     std::vector<DeviceMem> &subLocalSrcMems, std::vector<DeviceMem> &subLocalDstMems)
     528              : {
     529              :     (void)subTxReduceMems;
     530              :     (void)mainTxReduceMems;
     531            0 :     Stream mainStream;
     532            0 :     LINK mainPreLink;
     533            0 :     LINK mainNextLink;
     534            0 :     Stream subStream;
     535            0 :     LINK subPreLink;
     536            0 :     LINK subNextLink;
     537            0 :     CHK_RET(PrepareRunMainStream(ALIGNED_MAIN_RING_INDEX, mainStream, mainPreLink, mainNextLink));
     538            0 :     HCCL_DEBUG("Reduce: step[%u] ring[%u], src rank[%u] starts to send slice to dst rank[%u]",
     539              :         step, ALIGNED_MAIN_RING_INDEX, mainPreLink->GetRemoteRank(), mainNextLink->GetRemoteRank());
     540            0 :     CHK_RET(PrepareRunMainStream(ALIGNED_SUB_RING_INDEX, subStream, subPreLink, subNextLink));
     541            0 :     HCCL_DEBUG("Reduce: step[%u] ring[%u], src rank[%u] starts to send slice to dst rank[%u]",
     542              :         step, ALIGNED_SUB_RING_INDEX, subPreLink->GetRemoteRank(), subNextLink->GetRemoteRank());
     543              : 
     544            0 :     CHK_RET(mainNextLink->TxAck(mainStream));
     545            0 :     CHK_RET(mainPreLink->RxAck(mainStream));
     546            0 :     CHK_RET(subNextLink->TxAck(subStream));
     547            0 :     CHK_RET(subPreLink->RxAck(subStream));
     548              : 
     549            0 :     u32 sliceSize = multRingsSlices_[ALIGNED_MAIN_RING_INDEX].size() / rankSize;
     550            0 :     for (u32 memIdx = 0; memIdx < sliceSize; memIdx++) {
     551            0 :         CHK_RET(ReducerRun(ALIGNED_MAIN_RING_INDEX, dispatcher_, mainPreLink, mainRxReduceMems[memIdx], mainStream));
     552            0 :         CHK_RET(ReducerRun(ALIGNED_SUB_RING_INDEX, dispatcher_, subPreLink, subRxReduceMems[memIdx], subStream));
     553            0 :         CHK_RET(LocalMemcpy(step, rankSize, ALIGNED_MAIN_RING_INDEX, mainLocalSrcMems[memIdx], mainLocalDstMems[memIdx]));
     554            0 :         CHK_RET(LocalMemcpy(step, rankSize, ALIGNED_SUB_RING_INDEX, subLocalSrcMems[memIdx], subLocalDstMems[memIdx]));
     555              :     }
     556            0 :     CHK_RET(mainPreLink->TxDataSignal(mainStream));
     557            0 :     CHK_RET(mainNextLink->RxDataSignal(mainStream));
     558            0 :     CHK_RET(subPreLink->TxDataSignal(subStream));
     559            0 :     CHK_RET(subNextLink->RxDataSignal(subStream));
     560            0 :     return HCCL_SUCCESS;
     561            0 : }
     562              : 
     563            0 : HcclResult AlignedReduceScatterDoubleRing::PreRunStreams(
     564              :     const u32 step, const u32 rankSize,
     565              :     const u32 txSliceIdxMain, const u32 rxSliceIdxMain, const u32 subSliceIdxMain,
     566              :     const u32 txSliceIdxSub, const u32 rxSliceIdxSub, const u32 subSliceIdxSub,
     567              :     std::vector<SenderMemoryInfo> &txReduceMemsMain,
     568              :     std::vector<ReducerMemoryInfo> &rxReduceMemsMain,
     569              :     std::vector<SenderMemoryInfo> &txReduceMemsSub,
     570              :     std::vector<ReducerMemoryInfo> &rxReduceMemsSub,
     571              :     std::vector<DeviceMem> &localSrcMemsMain, std::vector<DeviceMem> &localDstMemsMain,
     572              :     std::vector<DeviceMem> &localSrcMemsSub, std::vector<DeviceMem> &localDstMemsSub)
     573              : {
     574            0 :     CHK_RET(PrepareDeviceMems(step, ALIGNED_MAIN_RING_INDEX, rankSize,
     575              :         txSliceIdxMain, rxSliceIdxMain, subSliceIdxMain,
     576              :         txReduceMemsMain, rxReduceMemsMain,
     577              :         localSrcMemsMain, localDstMemsMain));
     578            0 :     CHK_RET(PrepareDeviceMems(step, ALIGNED_SUB_RING_INDEX, rankSize,
     579              :         txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub,
     580              :         txReduceMemsSub, rxReduceMemsSub,
     581              :         localSrcMemsSub, localDstMemsSub));
     582            0 :     return HCCL_SUCCESS;
     583              : }
     584              : 
     585            0 : HcclResult AlignedReduceScatterDoubleRing::RunReduceScatter(const u32 rank, const u32 rankSize)
     586              : {
     587            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing starts, the input param rank[%u]", rank);
     588              :     // 空拷贝用于后续操作附着
     589            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     590              :     // 主环主流通知从环主流开始通信
     591            0 :     CHK_RET(MainRecordSub());
     592              :     // 从环主流等待主环主流通知
     593            0 :     CHK_RET(SubWaitMain());
     594            0 :     if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) {
     595            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     596            0 :         CHK_RET(ExecEmptyTasks());
     597            0 :         CHK_RET(RunInitStep(rank, rankSize));
     598              :     }
     599            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     600            0 :     CHK_RET(ExecEmptyTasks());
     601              : 
     602              :     // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3, subSliceIdx = 1
     603              :     // 从环初始indexes
     604            0 :     u32 txSliceIdxSub  = (rank + rankSize - 1) % rankSize;
     605            0 :     u32 rxSliceIdxSub  = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     606            0 :     u32 subSliceIdxSub = (rank + rankSize - DMA_REDUCE_THREE_OFFSET) % rankSize;
     607            0 :     HCCL_DEBUG("[RunReduceScatter]txSliceIdxSub is [%u], rxSliceIdxSub is [%u], subSliceIdxSub is [%u]",
     608              :         txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub);
     609              :     // 主环初始indexes
     610            0 :     u32 txSliceIdxMain  = (rankSize - rank - 1 + rankSize) % rankSize;
     611            0 :     u32 rxSliceIdxMain  = (rankSize - rank - DMA_REDUCE_TWO_OFFSET + rankSize) % rankSize;
     612            0 :     u32 subSliceIdxMain = (rankSize - rank - DMA_REDUCE_THREE_OFFSET + rankSize) % rankSize;
     613              : 
     614            0 :     for (u32 step = 0; step < rankSize - 1; step++) {
     615              :         // 并发
     616            0 :         std::vector<SenderMemoryInfo> txReduceMemsMain;
     617            0 :         std::vector<ReducerMemoryInfo> rxReduceMemsMain;
     618            0 :         std::vector<SenderMemoryInfo> txReduceMemsSub;
     619            0 :         std::vector<ReducerMemoryInfo> rxReduceMemsSub;
     620            0 :         std::vector<DeviceMem> localDstMemsMain;
     621            0 :         std::vector<DeviceMem> localSrcMemsMain;
     622            0 :         std::vector<DeviceMem> localSrcMemsSub;
     623            0 :         std::vector<DeviceMem> localDstMemsSub;
     624            0 :         CHK_RET(PreRunStreams(step, rankSize,
     625              :             txSliceIdxMain, rxSliceIdxMain, subSliceIdxMain,
     626              :             txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub,
     627              :             txReduceMemsMain, rxReduceMemsMain, txReduceMemsSub, rxReduceMemsSub,
     628              :             localSrcMemsMain, localDstMemsMain, localSrcMemsSub, localDstMemsSub));
     629            0 :         CHK_RET(RunAllStreams(step, rankSize, txReduceMemsMain, rxReduceMemsMain, txReduceMemsSub, rxReduceMemsSub,
     630              :             localSrcMemsMain, localDstMemsMain, localSrcMemsSub, localDstMemsSub));
     631              :         // 更新索引
     632            0 :         txSliceIdxSub  = (txSliceIdxSub + rankSize - 1) % rankSize;
     633            0 :         rxSliceIdxSub  = (rxSliceIdxSub + rankSize - 1) % rankSize;
     634            0 :         subSliceIdxSub = (subSliceIdxSub + rankSize - 1) % rankSize;
     635            0 :         txSliceIdxMain  = (txSliceIdxMain + rankSize - 1) % rankSize;
     636            0 :         rxSliceIdxMain  = (rxSliceIdxMain + rankSize - 1) % rankSize;
     637            0 :         subSliceIdxMain = (subSliceIdxMain + rankSize - 1) % rankSize;
     638            0 :     }
     639              :     // 从环主流通知主环主流通信完成
     640            0 :     CHK_RET(SubRecordMain());
     641              :     // 主环主流等待从环主流通知
     642            0 :     CHK_RET(MainWaitSub());
     643            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     644            0 :     CHK_RET(ExecEmptyTasks());
     645            0 :     HCCL_INFO("AlignedReduceScatterDoubleRing finished to RunReduceScatter");
     646            0 :     return HCCL_SUCCESS;
     647              : }
     648              : 
     649            0 : HcclResult AlignedReduceScatterDoubleRing::GetActiveSubstreamNum(u32 &activeSubstreamNum)
     650              : {
     651            0 :     activeSubstreamNum = subStreams_.size();
     652            0 :     if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB &&
     653            0 :         disableDMAReduce_) {
     654            0 :         if (subStreams_.size() <= 2) {
     655            0 :             HCCL_ERROR("[AlignedReduceScatterDoubleRing][GetActiveSubstreamNum]subStreams_.size()[%zu] <= 2",
     656              :                 subStreams_.size());
     657            0 :             return HCCL_E_PARA;
     658              :         }
     659            0 :         activeSubstreamNum = subStreams_.size() - 2;
     660              :     }
     661            0 :     return HCCL_SUCCESS;
     662              : }
     663              : 
     664              : 
     665            0 : HcclResult AlignedReduceScatterDoubleRing::ExecEmptyTasks()
     666              : {
     667            0 :     u32 activeSubstreamNum = 0;
     668            0 :     CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
     669            0 :     for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
     670            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, subStreams_[signalIndex], dispatcher_));
     671              :     }
     672            0 :     return HCCL_SUCCESS;
     673              : }
     674              : 
     675              : // 主流通知从流干活
     676            0 : HcclResult AlignedReduceScatterDoubleRing::MainRecordSub()
     677              : {
     678            0 :     u32 activeSubstreamNum = 0;
     679            0 :     CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
     680            0 :     for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
     681            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[signalIndex],
     682              :             profilerInput_.stage));
     683              :     }
     684            0 :     return HCCL_SUCCESS;
     685              : }
     686              : // 从流等待主流
     687            0 : HcclResult AlignedReduceScatterDoubleRing::SubWaitMain()
     688              : {
     689            0 :     u32 activeSubstreamNum = 0;
     690            0 :     CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
     691            0 :     for (u32 streamIndex = 0; streamIndex < activeSubstreamNum; streamIndex++) {
     692            0 :         CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, subSignals_[streamIndex],
     693              :             profilerInput_.stage));
     694              :     }
     695            0 :     return HCCL_SUCCESS;
     696              : }
     697              : // 主流等待从流
     698            0 : HcclResult AlignedReduceScatterDoubleRing::MainWaitSub()
     699              : {
     700            0 :     u32 activeSubstreamNum = 0;
     701            0 :     CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
     702            0 :     for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
     703            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[signalIndex], profilerInput_.stage));
     704              :     }
     705            0 :     return HCCL_SUCCESS;
     706              : }
     707              : // 从流告诉主流活干完了
     708            0 : HcclResult AlignedReduceScatterDoubleRing::SubRecordMain()
     709              : {
     710            0 :     u32 activeSubstreamNum = 0;
     711            0 :     CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
     712            0 :     for (u32 streamIndex = 0; streamIndex < activeSubstreamNum; streamIndex++) {
     713            0 :         CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, mainSignals_[streamIndex],
     714              :             profilerInput_.stage));
     715              :     }
     716            0 :     return HCCL_SUCCESS;
     717              : }
     718              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_DB_RING, AlignedReduceScatterDoubleRing);
     719              : } // namespace hccl
        

Generated by: LCOV version 2.0-1