LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_ring_concurrent_direct.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 283 0
Test Date: 2026-07-28 12:11:00 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_ring_concurrent_direct.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : ReduceScatterRingConcurrentDirect::ReduceScatterRingConcurrentDirect(
      16            0 :     const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
      17              : {
      18            0 : }
      19              : 
      20            0 : ReduceScatterRingConcurrentDirect::~ReduceScatterRingConcurrentDirect()
      21              : {
      22            0 : }
      23              : 
      24            0 : HcclResult ReduceScatterRingConcurrentDirect::Prepare(const u64 reduceAttrBitMap, const HcomCollOpInfo *opInfo,
      25              :                                                       const u32 userRank, std::vector<Stream> &subStreams,
      26              :                                                       const std::vector<std::shared_ptr<LocalNotify>> &mainSignals,
      27              :                                                       const std::vector<std::shared_ptr<LocalNotify>> &subSignals,
      28              :                                                       const std::vector<u32> &ringsOrder,
      29              :                                                       const std::vector<Slice> &userMemInputSlices, bool isSdma)
      30              : {
      31            0 :     reduceAttr_ = reduceAttrBitMap;
      32            0 :     opInfo_ = opInfo;
      33            0 :     userRank_ = userRank;
      34            0 :     subStreams_ = subStreams;
      35            0 :     mainSignals_ = mainSignals;
      36            0 :     subSignals_ = subSignals;
      37            0 :     ringsOrder_ = ringsOrder;
      38            0 :     userSlices_ = userMemInputSlices;
      39            0 :     isSdma_ = isSdma;
      40            0 :     return HCCL_SUCCESS;
      41              : }
      42              : 
      43              : // reduce scatter ring direct算法的函数入口
      44            0 : HcclResult ReduceScatterRingConcurrentDirect::RunAsync(const u32 rank, const u32 rankSize,
      45              :                                                        const std::vector<LINK> &links)
      46              : {
      47              :     // 基本的检查
      48            0 :     CHK_RET(CheckParameters(rank, rankSize, links));
      49              : 
      50              :     // 判断rank_size == 1的情况,并拷贝
      51            0 :     if (rankSize == 1) {
      52            0 :         CHK_RET(OneRankMemcpy());
      53            0 :         return HCCL_SUCCESS;
      54              :     }
      55            0 :     HCCL_DEBUG("ReduceScatterRingConcurrentDirect starts: rank[%u]", rank);
      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 :     if (barrierSwitchOn_) {
      69              :         // 执行barrier,保证数据收发完成
      70            0 :         CHK_RET(ExecuteBarrier(leftLink_, rightLink_));
      71              :     }
      72              : 
      73            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, stream_, subStreams_));
      74              : 
      75            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished: rank[%u]", rank);
      76            0 :     return HCCL_SUCCESS;
      77              : }
      78              : 
      79            0 : HcclResult ReduceScatterRingConcurrentDirect::CheckParameters(const u32 rank, const u32 rankSize,
      80              :                                                               const std::vector<LINK> &links)
      81              : {
      82            0 :     CHK_PTR_NULL(opInfo_);
      83            0 :     CHK_RET(CheckConcurrentDirectParameters(rank, rankSize, links));
      84              :     // 判断subStreams数量是否正确
      85            0 :     CHK_PRT_RET(
      86              :         subStreams_.size() < 1,
      87              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subStreams size[%u] is less than 1", subStreams_.size()),
      88              :         HCCL_E_PARA);
      89            0 :     for (auto &s : subStreams_) {
      90            0 :         CHK_PTR_NULL(s.ptr());
      91              :     }
      92              :     // 判断mainSignals数量是否正确
      93            0 :     CHK_PRT_RET(
      94              :         mainSignals_.size() < 1,
      95              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] mainSignals size[%u] is less than 1", mainSignals_.size()),
      96              :         HCCL_E_PARA);
      97              :     // 判断subSignals数量是否正确
      98            0 :     CHK_PRT_RET(
      99              :         subSignals_.size() < 1,
     100              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subSignals size[%u] is less than 1", subSignals_.size()),
     101              :         HCCL_E_PARA);
     102              :     // 判断ringsOrder数量是否正确
     103            0 :     CHK_PRT_RET(ringsOrder_.size() != rankSize,
     104              :                 HCCL_ERROR("[ReduceScatterRingConcurrentDirect] ringsOrder size[%u] is not equal to rank size[%u]",
     105              :                            ringsOrder_.size(), rankSize),
     106              :                 HCCL_E_PARA);
     107              :     // 判断userMemInputSlices数量是否正确
     108            0 :     CHK_PRT_RET(userSlices_.size() % rankSize != 0,
     109              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] userMemInputSlices size[%u] can not divided by size[%u]",
     110              :                    userSlices_.size(), rankSize),
     111              :         HCCL_E_PARA);
     112            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect CheckParameters success");
     113            0 :     return HCCL_SUCCESS;
     114              : }
     115              : 
     116            0 : HcclResult ReduceScatterRingConcurrentDirect::OneRankMemcpy()
     117              : {
     118            0 :     for (u32 sliceIdx = 0; sliceIdx < slices_.size(); sliceIdx++) {
     119            0 :         const Slice &srcSlice = userSlices_[sliceIdx];
     120            0 :         const Slice &dstSlice = slices_[sliceIdx];
     121            0 :         DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcSlice.offset, srcSlice.size);
     122            0 :         DeviceMem dst;
     123            0 :         if (opInfo_->outputAddr != nullptr) {
     124              :             // opInfo_->outputAddr != nullptr指示要将输出发送至user output
     125            0 :             u64 stepOffset = slices_[ringsOrder_[0]].offset;
     126            0 :             HCCL_DEBUG("[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at userMemOut_",
     127              :                 userRank_, stepOffset, dstSlice.size);
     128            0 :             dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + stepOffset, dstSlice.size);
     129              :         } else {
     130              :             // opInfo_->outputAddr == nullptr指示要将输出发送至CCL buffer
     131            0 :             HCCL_DEBUG("[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at outputMem_",
     132              :                 userRank_, dstSlice.offset, dstSlice.size);
     133            0 :             dst = outputMem_.range(dstSlice.offset, dstSlice.size);
     134              :         }
     135            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     136            0 :     }
     137            0 :     return HCCL_SUCCESS;
     138              : }
     139              : 
     140            0 : HcclResult ReduceScatterRingConcurrentDirect::InitSenderReducer()
     141              : {
     142              :     // 创建reducer & sender
     143            0 :     senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
     144            0 :     CHK_SMART_PTR_NULL(senderInfo_);
     145              : 
     146            0 :     reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
     147            0 :     CHK_SMART_PTR_NULL(reducerInfo_);
     148            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to InitSenderReducer");
     149            0 :     return HCCL_SUCCESS;
     150              : }
     151              : 
     152            0 : HcclResult ReduceScatterRingConcurrentDirect::GetInitializedNeighborLinks(const u32 rank, const u32 rankSize,
     153              :                                                                           const std::vector<LINK> &links)
     154              : {
     155              :     // 收集左邻居信息
     156            0 :     leftLink_ = links[(rank + rankSize - 1) % rankSize];
     157            0 :     CHK_SMART_PTR_NULL(leftLink_);
     158              : 
     159              :     // 收集右邻居信息
     160            0 :     rightLink_ = links[(rank + 1) % rankSize];
     161            0 :     CHK_SMART_PTR_NULL(rightLink_);
     162            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to GetInitializedNeighborLinks");
     163            0 :     return HCCL_SUCCESS;
     164              : }
     165              : 
     166            0 : HcclResult ReduceScatterRingConcurrentDirect::SetSlices(const u32 rank, const u32 rankSize)
     167              : {
     168            0 :     if (slices_.size() == 0) {
     169            0 :         slices_.resize(rankSize);
     170              : 
     171              :         // 生成std::vector<Slice> slices_
     172            0 :         u64 sliceSize = count_ * SIZE_TABLE[dataType_];
     173              :         ;
     174              : 
     175            0 :         for (u32 i = 0; i < rankSize; i++) {
     176            0 :             slices_[i].size = sliceSize;
     177              :             // 用于DMA消减过程中,消除src与dst不对位的风险
     178            0 :             slices_[i].offset = RoundUpWithDivisor(i * sliceSize, HCCL_MIN_SLICE_ALIGN);
     179              : 
     180            0 :             HCCL_DEBUG("rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]", rank, i, slices_[i].offset, i,
     181              :                        slices_[i].size);
     182              :         }
     183              :     }
     184            0 :     if (UNLIKELY(HcclCheckLogLevel(DLOG_DEBUG))) {
     185            0 :         for (u32 i = 0; i < slices_.size(); i++) {
     186            0 :             HCCL_DEBUG(
     187              :                 "[ReduceScatterRingConcurrentDirect][SetSlices]rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]",
     188              :                 rank, i, slices_[i].offset, i, slices_[i].size);
     189              :         }
     190              :     }
     191              :     // 最后一步搬到userMemOut_的offset, 不同的ring环offset不一样
     192            0 :     lastStepOffset_ = slices_[ringsOrder_[0]].offset;
     193            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to SetSlices");
     194            0 :     return HCCL_SUCCESS;
     195              : }
     196              : 
     197            0 : HcclResult ReduceScatterRingConcurrentDirect::RunInitStep(const u32 rank, const u32 rankSize)
     198              : {
     199              :     // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3
     200            0 :     u32 initSlice0Idx = 0;
     201            0 :     u32 initSlice1Idx = 0;
     202            0 :     initSlice0Idx     = (rank + rankSize - 1) % rankSize;
     203            0 :     initSlice1Idx     = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     204            0 :     u32 sliceSize = slices_.size() / rankSize;
     205              : 
     206            0 :     for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     207              :         // 第-1步,片内将部分数据从userIn搬到cclIn
     208            0 :         const Slice &srcInitSlice0 = userSlices_[initSlice0Idx * sliceSize + sliceIdx];
     209              :         DeviceMem    srcSubInit
     210            0 :             = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice0.offset, srcInitSlice0.size);
     211            0 :         const Slice &dstInitSlice0 = slices_[initSlice0Idx * sliceSize + sliceIdx];
     212            0 :         DeviceMem    dstSubInit    = inputMem_.range(dstInitSlice0.offset, dstInitSlice0.size);
     213            0 :         const Slice &srcInitSlice1 = userSlices_[initSlice1Idx * sliceSize + sliceIdx];
     214              :         DeviceMem    srcInit
     215            0 :             = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice1.offset, srcInitSlice1.size);
     216            0 :         const Slice &dstInitSlice1 = slices_[initSlice1Idx * sliceSize + sliceIdx];
     217            0 :         DeviceMem    dstInit       = inputMem_.range(dstInitSlice1.offset, dstInitSlice1.size);
     218              :         // 第-1步并发
     219            0 :         CHK_RET(MainRecordSub()); // 主流通知从流开始通信
     220            0 :         CHK_RET(SubWaitMain());   // 从流等待主流通知
     221            0 :         if (rankSize == TWO_RANK_SIZE && opInfo_->outputAddr != nullptr) {
     222            0 :             HCCL_DEBUG(
     223              :                 "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     224              :                 "userMemInput to offset[%llu], size[%llu] on userMemOut_",
     225              :                 userRank_, srcInitSlice1.offset, srcInitSlice1.size, lastStepOffset_, dstInitSlice1.size);
     226            0 :             dstInit = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
     227            0 :                 dstInitSlice1.size);
     228              :         } else {
     229            0 :             HCCL_DEBUG(
     230              :                 "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     231              :                 "userMemInput to offset[%llu], size[%llu] on CCL",
     232              :                 userRank_, srcInitSlice1.offset, srcInitSlice1.size, dstInitSlice1.offset, dstInitSlice1.size);
     233              :         }
     234            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, stream_));
     235            0 :         HCCL_DEBUG("Memcpy operation: step[-1] stream[sub] src rank[%u] starts to copy(rcv) offset[%llu], "
     236              :             " size[%llu] on userMemInput to offset[%llu], size[%llu] on CCL",
     237              :             userRank_, srcInitSlice0.offset, srcInitSlice0.size, dstInitSlice0.offset, dstInitSlice0.size);
     238            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstSubInit, srcSubInit, subStreams_[0]));
     239            0 :         CHK_RET(SubRecordMain()); // 从流通知主流通信完成
     240            0 :         CHK_RET(MainWaitSub());   // 主流等待从流通知
     241            0 :     }
     242            0 :     return HCCL_SUCCESS;
     243              : }
     244              : 
     245            0 : HcclResult ReduceScatterRingConcurrentDirect::PreSync()
     246              : {
     247            0 :     CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[0], profilerInput_.stage));
     248            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     249            0 :     CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[0], profilerInput_.stage));
     250            0 :     return HCCL_SUCCESS;
     251              : }
     252              : 
     253            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunSpInlineReduce(
     254              :     const HcclDispatcher dispatcher, const LINK &link,
     255              :     const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
     256              : {
     257            0 :     HcclResult ret = HCCL_SUCCESS;
     258            0 :     CHK_RET(link->RxDataSignal(stream));
     259            0 :     void *remoteMem = nullptr;
     260            0 :     CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     261            0 :     if (!isSdma_) {
     262            0 :         CHK_RET(PreSync());
     263              :     }
     264            0 :     for (ReducerMemoryInfo reduceMem : reducerMems) {
     265            0 :         if (isSdma_) {
     266            0 :             CHK_RET(PreSync());
     267              :         }
     268            0 :         const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     269            0 :         CHK_RET(
     270              :             HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
     271              :             dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
     272              :             link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
     273              : 
     274            0 :         if (reduceMem.localsrc != reduceMem.localdst) {
     275            0 :             ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     276            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     277              :                 HCCL_ERROR("[ReducerRun]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     278              :                 reduceMem.localdst.ptr()),
     279              :                 ret);
     280              :         }
     281            0 :     }
     282            0 :     return HCCL_SUCCESS;
     283              : }
     284              : 
     285            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunNoSpInlineReduce(
     286              :     const HcclDispatcher dispatcher, const LINK &link,
     287              :     const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
     288              : {
     289            0 :     std::vector<RxMemoryInfo> rxMems;
     290            0 :     for (const ReducerMemoryInfo &reduceMem : reducerMems) {
     291            0 :         rxMems.emplace_back(RxMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     292            0 :             reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() });
     293              :     }
     294              : 
     295            0 :     std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     296            0 :     for (ReducerMemoryInfo reduceMem : reducerMems) {
     297            0 :         u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
     298            0 :         DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
     299              : 
     300            0 :         rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     301            0 :             reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size(), reduceSrc.ptr(), reduceMem.localdst.ptr(),
     302              :             dataCount });
     303            0 :     }
     304            0 :     if (!isSdma_) {
     305              :         // AnyPath
     306            0 :         CHK_RET(PreSync());
     307            0 :         CHK_RET(link->RxAsync(rxMems, stream));
     308              :     } else {
     309              :         // SDMA
     310            0 :         CHK_RET(link->RxDataSignal(stream));
     311            0 :         for (auto& mem : rxMems) {
     312            0 :             CHK_RET(PreSync());
     313            0 :             CHK_PTR_NULL(mem.dst);
     314            0 :             void *srcMemPtr = nullptr;
     315            0 :             CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
     316              : 
     317            0 :             DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + mem.srcOffset, mem.len);
     318            0 :             DeviceMem dstDevMem(static_cast<s8 *>(mem.dst), mem.len);
     319            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher, dstDevMem, srcDevMem,
     320              :                 stream, link->GetRemoteRank(), link->GetLinkType()));
     321            0 :         }
     322              :     }
     323            0 :     if (link->GetSupportDataReceivedAck()) {
     324            0 :         CHK_RET(link->DataReceivedAck(stream));
     325              :     }
     326            0 :     for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
     327            0 :         CHK_RET(HcclReduceAsync(dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_,
     328              :             reductionOp_, stream, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
     329              :             reduceAttr_));
     330              :     }
     331            0 :     return HCCL_SUCCESS;
     332            0 : }
     333              : 
     334            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRun(const HcclDispatcher dispatcher, const LINK &link,
     335              :     const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
     336              : {
     337            0 :     CHK_PTR_NULL(stream.ptr());
     338            0 :     bool isSpInlineReduce = link->IsSpInlineReduce();
     339            0 :     if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_))) {
     340            0 :         CHK_RET(ReducerRunSpInlineReduce(dispatcher, link, reducerMems, stream));
     341            0 :     } else {
     342            0 :         CHK_RET(ReducerRunNoSpInlineReduce(dispatcher, link, reducerMems, stream));
     343              :     }
     344            0 :     return HCCL_SUCCESS;
     345              : }
     346              : 
     347            0 : HcclResult ReduceScatterRingConcurrentDirect::RunMainStream(const u32 step, std::vector<Slice> txSliceVector,
     348              :     std::vector<Slice> rxSliceVector, const u32 rank, const u32 rankSize)
     349              : {
     350              :     (void) rank;
     351            0 :     CHK_RET(leftLink_->TxAck(stream_));
     352            0 :     CHK_RET(rightLink_->RxAck(stream_));
     353            0 :     u32 sliceSize = slices_.size() / rankSize;
     354              : 
     355              :     // 通信,如果是最后一步,则做消减拷贝
     356            0 :     std::vector<SenderMemoryInfo> txMems;
     357            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     358            0 :     DeviceMem dst;
     359            0 :     for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     360              :         // Ack
     361            0 :         HCCL_DEBUG("Reduce operation: step[%u] stream[main], src rank[%u] starts to send offset[%llu] size[%llu] "
     362              :             "from leftMem_",
     363              :             step, leftLink_->GetRemoteRank(), rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     364            0 :         if (isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
     365            0 :             HCCL_DEBUG("[RunReduceScatter] sdma DMAReduce step");
     366            0 :             HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     367              :                 "at userMemOut_",
     368              :                 step, userRank_, lastStepOffset_, rxSliceVector[sliceIdx].size);
     369            0 :             dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
     370            0 :                 rxSliceVector[sliceIdx].size);
     371              :         } else {
     372            0 :             HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     373              :                 "at inputMem_",
     374              :                 step, userRank_, rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     375            0 :             dst = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     376            0 :             if (!isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
     377            0 :                 HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce step");
     378            0 :                 finalSrc_ = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     379            0 :                 finalDst_ = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
     380            0 :                 rxSliceVector[sliceIdx].size);
     381              :             }
     382              :         }
     383              :         // 在inline reduce场景, 需要利用scratchMem_暂存
     384            0 :         DeviceMem srcMemTemp = scratchMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     385            0 :         DeviceMem srcMem     = inputMem_.range(txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
     386            0 :         HCCL_DEBUG("Reduce operation: step[%u] stream[main], senderInfo_ rank[%u] starts to rcv offset[%llu], "
     387              :             " size[%llu]",
     388              :             step, rightLink_->GetRemoteRank(), txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
     389            0 :         rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSliceVector[sliceIdx].offset,
     390              :             dst, dst, srcMemTemp});
     391            0 :         txMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSliceVector[sliceIdx].offset, srcMem});
     392            0 :     }
     393            0 :     CHK_RET(senderInfo_->run(rightLink_, txMems, stream_));
     394            0 :     CHK_RET(ReducerRun(dispatcher_, leftLink_, rxReduceMems, stream_));
     395            0 :     return HCCL_SUCCESS;
     396            0 : }
     397              : 
     398            0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStream(const u32 step, std::vector<Slice> subSliceVector,
     399              :     std::vector<Slice> cclSliceVector, const u32 rank, const u32 rankSize)
     400              : {
     401              :     (void) rank;
     402            0 :     if (!isSdma_) {
     403            0 :         CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     404            0 :         CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     405              :     }
     406            0 :     for (u32 sliceIdx = 0; sliceIdx < subSliceVector.size(); sliceIdx++) {
     407            0 :         if (isSdma_) {
     408            0 :             CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     409            0 :             CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     410              :         }
     411            0 :         HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], src rank[%u] starts to send offset[%llu], size[%llu] "
     412              :             "from userMemIn_", step, userRank_, subSliceVector[sliceIdx].offset, subSliceVector[sliceIdx].size);
     413            0 :         DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSliceVector[sliceIdx].offset,
     414            0 :             subSliceVector[sliceIdx].size);
     415            0 :         DeviceMem dst;
     416            0 :         if (step == rankSize - DMA_REDUCE_TWO_OFFSET) {
     417              :             // do nothing
     418            0 :         } else if (isSdma_ && step == rankSize - DMA_REDUCE_THREE_OFFSET && opInfo_->outputAddr != nullptr) {
     419            0 :             HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     420              :                 "to userMemOut_", step, userRank_, lastStepOffset_, subSliceVector[sliceIdx].size);
     421            0 :             dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
     422            0 :                 subSliceVector[sliceIdx].size);
     423            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
     424            0 :         } else {
     425            0 :             HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     426              :                 "to inputMem_",
     427              :                 step, userRank_, cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
     428            0 :             dst = inputMem_.range(cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
     429            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
     430              :         }
     431            0 :     }
     432            0 :     return HCCL_SUCCESS;
     433              : }
     434              : 
     435            0 : HcclResult ReduceScatterRingConcurrentDirect::RunReduceScatter(const u32 rank, const u32 rankSize)
     436              : {
     437            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect starts, the input param rank[%u]", rank);
     438              :     // 空拷贝用于后续操作附着
     439            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     440              : 
     441            0 :     CHK_RET(RunInitStep(rank, rankSize));
     442            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     443            0 :     CHK_RET(MainRecordSub()); // 主流通知从流开始通信
     444            0 :     CHK_RET(SubWaitMain());   // 从流等待主流通知
     445              :     // 空拷贝用于主从流任务并发
     446            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     447            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, subStreams_[0], dispatcher_));
     448            0 :     u32 sliceSize = slices_.size() / rankSize;
     449              : 
     450              :     // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3, subSliceIdx = 1
     451            0 :     u32 txSliceIdx  = (rank + rankSize - 1) % rankSize;
     452            0 :     u32 rxSliceIdx  = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     453            0 :     u32 subSliceIdx = (rank + rankSize - DMA_REDUCE_THREE_OFFSET) % rankSize;
     454            0 :     for (u32 step = 0; step < rankSize - 1; step++) {
     455            0 :         std::vector<Slice> rxSliceVector;
     456            0 :         std::vector<Slice> cclSliceVector;
     457            0 :         std::vector<Slice> txSliceVector;
     458            0 :         std::vector<Slice> subSliceVector;
     459            0 :         for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     460            0 :             rxSliceVector.push_back(slices_[rxSliceIdx * sliceSize + sliceIdx]);
     461            0 :             cclSliceVector.push_back(slices_[subSliceIdx * sliceSize + sliceIdx]);
     462            0 :             txSliceVector.push_back(slices_[txSliceIdx * sliceSize + sliceIdx]);
     463            0 :             subSliceVector.push_back(userSlices_[subSliceIdx * sliceSize + sliceIdx]);
     464              :         }
     465              : 
     466              :         // 主流
     467            0 :         CHK_RET(RunMainStream(step, txSliceVector, rxSliceVector, rank, rankSize));
     468              : 
     469              :         // 从流
     470            0 :         CHK_RET(RunSubStream(step, subSliceVector, cclSliceVector, rank, rankSize));
     471              : 
     472              :         // 更新索引
     473            0 :         subSliceIdx = (subSliceIdx + rankSize - 1) % rankSize;
     474            0 :         txSliceIdx  = (txSliceIdx + rankSize - 1) % rankSize;
     475            0 :         rxSliceIdx  = (rxSliceIdx + rankSize - 1) % rankSize;
     476            0 :     }
     477            0 :     CHK_RET(SubRecordMain()); // 从流通知主流通信完成
     478            0 :     CHK_RET(MainWaitSub());   // 主流等待从流通知
     479            0 :     if (!isSdma_ && opInfo_->outputAddr != nullptr) {
     480            0 :         HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce last step");
     481            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, finalDst_, finalSrc_, stream_));
     482              :     }
     483            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to RunReduceScatter");
     484            0 :     return HCCL_SUCCESS;
     485              : }
     486              : 
     487              : // 主流通知从流干活
     488            0 : HcclResult ReduceScatterRingConcurrentDirect::MainRecordSub()
     489              : {
     490            0 :     for (u32 signalIndex = 0; signalIndex < subSignals_.size(); signalIndex++) {
     491            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[signalIndex],
     492              :             profilerInput_.stage));
     493              :     }
     494            0 :     return HCCL_SUCCESS;
     495              : }
     496              : // 从流等待主流
     497            0 : HcclResult ReduceScatterRingConcurrentDirect::SubWaitMain()
     498              : {
     499            0 :     for (u32 streamIndex = 0; streamIndex < subSignals_.size(); streamIndex++) {
     500            0 :         CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, subSignals_[streamIndex],
     501              :             profilerInput_.stage));
     502              :     }
     503            0 :     return HCCL_SUCCESS;
     504              : }
     505              : // 主流等待从流
     506            0 : HcclResult ReduceScatterRingConcurrentDirect::MainWaitSub()
     507              : {
     508            0 :     for (u32 signalIndex = 0; signalIndex < mainSignals_.size(); signalIndex++) {
     509            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[signalIndex], profilerInput_.stage));
     510              :     }
     511            0 :     return HCCL_SUCCESS;
     512              : }
     513              : // 从流告诉主流活干完了
     514            0 : HcclResult ReduceScatterRingConcurrentDirect::SubRecordMain()
     515              : {
     516            0 :     for (u32 streamIndex = 0; streamIndex < mainSignals_.size(); streamIndex++) {
     517            0 :         CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, mainSignals_[streamIndex],
     518              :             profilerInput_.stage));
     519              :     }
     520            0 :     return HCCL_SUCCESS;
     521              : }
     522              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, ReduceScatterRingConcurrentDirect);
     523              : } // namespace hccl
        

Generated by: LCOV version 2.0-1