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 % 308 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 29 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(const HcclDispatcher dispatcher)
      16            0 :     : AlgTemplateBase(dispatcher)
      17            0 : {}
      18              : 
      19            0 : ReduceScatterRingConcurrentDirect::~ReduceScatterRingConcurrentDirect() {}
      20              : 
      21            0 : HcclResult ReduceScatterRingConcurrentDirect::Prepare(
      22              :     const u64 reduceAttrBitMap, const HcomCollOpInfo* opInfo, const u32 userRank, std::vector<Stream>& subStreams,
      23              :     const std::vector<std::shared_ptr<LocalNotify>>& mainSignals,
      24              :     const std::vector<std::shared_ptr<LocalNotify>>& subSignals, const std::vector<u32>& ringsOrder,
      25              :     const std::vector<Slice>& userMemInputSlices, bool isSdma)
      26              : {
      27            0 :     reduceAttr_ = reduceAttrBitMap;
      28            0 :     opInfo_ = opInfo;
      29            0 :     userRank_ = userRank;
      30            0 :     subStreams_ = subStreams;
      31            0 :     mainSignals_ = mainSignals;
      32            0 :     subSignals_ = subSignals;
      33            0 :     ringsOrder_ = ringsOrder;
      34            0 :     userSlices_ = userMemInputSlices;
      35            0 :     isSdma_ = isSdma;
      36            0 :     return HCCL_SUCCESS;
      37              : }
      38              : 
      39              : // reduce scatter ring direct算法的函数入口
      40              : HcclResult
      41            0 : ReduceScatterRingConcurrentDirect::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
      42              : {
      43              :     // 基本的检查
      44            0 :     CHK_RET(CheckParameters(rank, rankSize, links));
      45              : 
      46              :     // 判断rank_size == 1的情况,并拷贝
      47            0 :     if (rankSize == 1) {
      48            0 :         CHK_RET(OneRankMemcpy());
      49            0 :         return HCCL_SUCCESS;
      50              :     }
      51            0 :     HCCL_DEBUG("ReduceScatterRingConcurrentDirect starts: rank[%u]", rank);
      52              :     // 收集本地mem信息
      53            0 :     CHK_RET(InitSenderReducer());
      54              : 
      55              :     // 收集邻居信息
      56            0 :     CHK_RET(GetInitializedNeighborLinks(rank, rankSize, links));
      57              : 
      58              :     // 填充slice_
      59            0 :     CHK_RET(SetSlices(rank, rankSize));
      60              : 
      61              :     // 运行reduce-scatter, ring算法
      62            0 :     CHK_RET(RunReduceScatter(rank, rankSize));
      63              : 
      64            0 :     if (barrierSwitchOn_) {
      65              :         // 执行barrier,保证数据收发完成
      66            0 :         CHK_RET(ExecuteBarrier(leftLink_, rightLink_));
      67              :     }
      68              : 
      69            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, stream_, subStreams_));
      70              : 
      71            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished: rank[%u]", rank);
      72            0 :     return HCCL_SUCCESS;
      73              : }
      74              : 
      75              : HcclResult
      76            0 : ReduceScatterRingConcurrentDirect::CheckParameters(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
      77              : {
      78            0 :     CHK_PTR_NULL(opInfo_);
      79            0 :     CHK_RET(CheckConcurrentDirectParameters(rank, rankSize, links));
      80              :     // 判断subStreams数量是否正确
      81            0 :     CHK_PRT_RET(
      82              :         subStreams_.size() < 1,
      83              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subStreams size[%u] is less than 1", subStreams_.size()),
      84              :         HCCL_E_PARA);
      85            0 :     for (auto& s : subStreams_) {
      86            0 :         CHK_PTR_NULL(s.ptr());
      87              :     }
      88              :     // 判断mainSignals数量是否正确
      89            0 :     CHK_PRT_RET(
      90              :         mainSignals_.size() < 1,
      91              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] mainSignals size[%u] is less than 1", mainSignals_.size()),
      92              :         HCCL_E_PARA);
      93              :     // 判断subSignals数量是否正确
      94            0 :     CHK_PRT_RET(
      95              :         subSignals_.size() < 1,
      96              :         HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subSignals size[%u] is less than 1", subSignals_.size()),
      97              :         HCCL_E_PARA);
      98              :     // 判断ringsOrder数量是否正确
      99            0 :     CHK_PRT_RET(
     100              :         ringsOrder_.size() != rankSize,
     101              :         HCCL_ERROR(
     102              :             "[ReduceScatterRingConcurrentDirect] ringsOrder size[%u] is not equal to rank size[%u]", ringsOrder_.size(),
     103              :             rankSize),
     104              :         HCCL_E_PARA);
     105              :     // 判断userMemInputSlices数量是否正确
     106            0 :     CHK_PRT_RET(
     107              :         userSlices_.size() % rankSize != 0,
     108              :         HCCL_ERROR(
     109              :             "[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(
     127              :                 "[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at userMemOut_",
     128              :                 userRank_, stepOffset, dstSlice.size);
     129            0 :             dst = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) + stepOffset, dstSlice.size);
     130              :         } else {
     131              :             // opInfo_->outputAddr == nullptr指示要将输出发送至CCL buffer
     132            0 :             HCCL_DEBUG(
     133              :                 "[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at outputMem_",
     134              :                 userRank_, dstSlice.offset, dstSlice.size);
     135            0 :             dst = outputMem_.range(dstSlice.offset, dstSlice.size);
     136              :         }
     137            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     138            0 :     }
     139            0 :     return HCCL_SUCCESS;
     140              : }
     141              : 
     142            0 : HcclResult ReduceScatterRingConcurrentDirect::InitSenderReducer()
     143              : {
     144              :     // 创建reducer & sender
     145            0 :     senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
     146            0 :     CHK_SMART_PTR_NULL(senderInfo_);
     147              : 
     148            0 :     reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
     149            0 :     CHK_SMART_PTR_NULL(reducerInfo_);
     150            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to InitSenderReducer");
     151            0 :     return HCCL_SUCCESS;
     152              : }
     153              : 
     154            0 : HcclResult ReduceScatterRingConcurrentDirect::GetInitializedNeighborLinks(
     155              :     const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
     156              : {
     157              :     // 收集左邻居信息
     158            0 :     leftLink_ = links[(rank + rankSize - 1) % rankSize];
     159            0 :     CHK_SMART_PTR_NULL(leftLink_);
     160              : 
     161              :     // 收集右邻居信息
     162            0 :     rightLink_ = links[(rank + 1) % rankSize];
     163            0 :     CHK_SMART_PTR_NULL(rightLink_);
     164            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to GetInitializedNeighborLinks");
     165            0 :     return HCCL_SUCCESS;
     166              : }
     167              : 
     168            0 : HcclResult ReduceScatterRingConcurrentDirect::SetSlices(const u32 rank, const u32 rankSize)
     169              : {
     170            0 :     if (slices_.size() == 0) {
     171            0 :         slices_.resize(rankSize);
     172              : 
     173              :         // 生成std::vector<Slice> slices_
     174            0 :         u64 sliceSize = count_ * SIZE_TABLE[dataType_];
     175              :         ;
     176              : 
     177            0 :         for (u32 i = 0; i < rankSize; i++) {
     178            0 :             slices_[i].size = sliceSize;
     179              :             // 用于DMA消减过程中,消除src与dst不对位的风险
     180            0 :             slices_[i].offset = RoundUpWithDivisor(i * sliceSize, HCCL_MIN_SLICE_ALIGN);
     181              : 
     182            0 :             HCCL_DEBUG(
     183              :                 "rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]", rank, i, slices_[i].offset, i,
     184              :                 slices_[i].size);
     185              :         }
     186              :     }
     187            0 :     if (UNLIKELY(HcclCheckLogLevel(DLOG_DEBUG))) {
     188            0 :         for (u32 i = 0; i < slices_.size(); i++) {
     189            0 :             HCCL_DEBUG(
     190              :                 "[ReduceScatterRingConcurrentDirect][SetSlices]rank[%u], slices[%u].offset=[%llu], "
     191              :                 "slices[%u].size=[%llu]",
     192              :                 rank, i, slices_[i].offset, i, slices_[i].size);
     193              :         }
     194              :     }
     195              :     // 最后一步搬到userMemOut_的offset, 不同的ring环offset不一样
     196            0 :     lastStepOffset_ = slices_[ringsOrder_[0]].offset;
     197            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to SetSlices");
     198            0 :     return HCCL_SUCCESS;
     199              : }
     200              : 
     201            0 : HcclResult ReduceScatterRingConcurrentDirect::RunInitStep(const u32 rank, const u32 rankSize)
     202              : {
     203              :     // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3
     204            0 :     u32 initSlice0Idx = 0;
     205            0 :     u32 initSlice1Idx = 0;
     206            0 :     initSlice0Idx = (rank + rankSize - 1) % rankSize;
     207            0 :     initSlice1Idx = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     208            0 :     u32 sliceSize = slices_.size() / rankSize;
     209              : 
     210            0 :     for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     211              :         // 第-1步,片内将部分数据从userIn搬到cclIn
     212            0 :         const Slice& srcInitSlice0 = userSlices_[initSlice0Idx * sliceSize + sliceIdx];
     213              :         DeviceMem srcSubInit
     214            0 :             = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr) + srcInitSlice0.offset, srcInitSlice0.size);
     215            0 :         const Slice& dstInitSlice0 = slices_[initSlice0Idx * sliceSize + sliceIdx];
     216            0 :         DeviceMem dstSubInit = inputMem_.range(dstInitSlice0.offset, dstInitSlice0.size);
     217            0 :         const Slice& srcInitSlice1 = userSlices_[initSlice1Idx * sliceSize + sliceIdx];
     218              :         DeviceMem srcInit
     219            0 :             = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr) + srcInitSlice1.offset, srcInitSlice1.size);
     220            0 :         const Slice& dstInitSlice1 = slices_[initSlice1Idx * sliceSize + sliceIdx];
     221            0 :         DeviceMem dstInit = inputMem_.range(dstInitSlice1.offset, dstInitSlice1.size);
     222              :         // 第-1步并发
     223            0 :         CHK_RET(MainRecordSub()); // 主流通知从流开始通信
     224            0 :         CHK_RET(SubWaitMain());   // 从流等待主流通知
     225            0 :         if (rankSize == TWO_RANK_SIZE && opInfo_->outputAddr != nullptr) {
     226            0 :             HCCL_DEBUG(
     227              :                 "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     228              :                 "userMemInput to offset[%llu], size[%llu] on userMemOut_",
     229              :                 userRank_, srcInitSlice1.offset, srcInitSlice1.size, lastStepOffset_, dstInitSlice1.size);
     230            0 :             dstInit = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) + lastStepOffset_, dstInitSlice1.size);
     231              :         } else {
     232            0 :             HCCL_DEBUG(
     233              :                 "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
     234              :                 "userMemInput to offset[%llu], size[%llu] on CCL",
     235              :                 userRank_, srcInitSlice1.offset, srcInitSlice1.size, dstInitSlice1.offset, dstInitSlice1.size);
     236              :         }
     237            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, stream_));
     238            0 :         HCCL_DEBUG(
     239              :             "Memcpy operation: step[-1] stream[sub] src rank[%u] starts to copy(rcv) offset[%llu], "
     240              :             "size[%llu] on userMemInput to offset[%llu], size[%llu] on CCL",
     241              :             userRank_, srcInitSlice0.offset, srcInitSlice0.size, dstInitSlice0.offset, dstInitSlice0.size);
     242            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstSubInit, srcSubInit, subStreams_[0]));
     243            0 :         CHK_RET(SubRecordMain()); // 从流通知主流通信完成
     244            0 :         CHK_RET(MainWaitSub());   // 主流等待从流通知
     245            0 :     }
     246            0 :     return HCCL_SUCCESS;
     247              : }
     248              : 
     249            0 : HcclResult ReduceScatterRingConcurrentDirect::PreSync()
     250              : {
     251            0 :     CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[0], profilerInput_.stage));
     252            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     253            0 :     CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[0], profilerInput_.stage));
     254            0 :     return HCCL_SUCCESS;
     255              : }
     256              : 
     257            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerSpInlineSlice(
     258              :     const HcclDispatcher dispatcher, const LINK& link, void* remoteMem, ReducerMemoryInfo reduceMem, Stream& stream)
     259              : {
     260            0 :     const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     261            0 :     CHK_RET(HcclReduceAsync(
     262              :         dispatcher, static_cast<s8*>(remoteMem) + reduceMem.remoteMemOffset, dataBytes / SIZE_TABLE[dataType_],
     263              :         dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(), link->GetRemoteRank(), link->GetLinkType(),
     264              :         INLINE_REDUCE_BIT));
     265              : 
     266            0 :     if (reduceMem.localsrc != reduceMem.localdst) {
     267            0 :         HcclResult ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     268            0 :         CHK_PRT_RET(
     269              :             ret != HCCL_SUCCESS,
     270              :             HCCL_ERROR(
     271              :                 "[ReducerRun]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     272              :                 reduceMem.localdst.ptr()),
     273              :             ret);
     274              :     }
     275            0 :     return HCCL_SUCCESS;
     276              : }
     277              : 
     278              : // 仅rdma场景调用:主流一次性PreSync后批量下发inline reduce任务
     279            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunSpInlineReduce(
     280              :     const HcclDispatcher dispatcher, const LINK& link, const std::vector<ReducerMemoryInfo>& reducerMems,
     281              :     Stream& stream)
     282              : {
     283            0 :     CHK_RET(link->RxDataSignal(stream));
     284            0 :     void* remoteMem = nullptr;
     285            0 :     CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     286            0 :     CHK_RET(PreSync());
     287            0 :     for (ReducerMemoryInfo reduceMem : reducerMems) {
     288            0 :         CHK_RET(ReducerSpInlineSlice(dispatcher, link, remoteMem, reduceMem, stream));
     289            0 :     }
     290            0 :     return HCCL_SUCCESS;
     291              : }
     292              : 
     293              : // sdma场景主流单个slice的远端读任务
     294            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerSdmaRemoteReadSlice(
     295              :     const LINK& link, const ReducerMemoryInfo& reduceMem, Stream& stream)
     296              : {
     297            0 :     RxMemoryInfo mem{
     298            0 :         UserMemType::INPUT_MEM, reduceMem.remoteMemOffset, reduceMem.remoteRcvTemp.ptr(),
     299            0 :         reduceMem.remoteRcvTemp.size()};
     300            0 :     CHK_PTR_NULL(mem.dst);
     301            0 :     void* srcMemPtr = nullptr;
     302            0 :     CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
     303              : 
     304            0 :     DeviceMem srcDevMem(static_cast<s8*>(srcMemPtr) + mem.srcOffset, mem.len);
     305            0 :     DeviceMem dstDevMem(static_cast<s8*>(mem.dst), mem.len);
     306            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstDevMem, srcDevMem, stream, link->GetRemoteRank(), link->GetLinkType()));
     307            0 :     return HCCL_SUCCESS;
     308            0 : }
     309              : 
     310              : // 数据接收确认 + 主流批量本地reduce任务
     311            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerLocalReduceSuffix(
     312              :     const HcclDispatcher dispatcher, const LINK& link, const std::vector<ReducerMemoryInfo>& reducerMems,
     313              :     Stream& stream)
     314              : {
     315            0 :     if (link->GetSupportDataReceivedAck()) {
     316            0 :         CHK_RET(link->DataReceivedAck(stream));
     317              :     }
     318            0 :     for (ReducerMemoryInfo reduceMem : reducerMems) {
     319            0 :         u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
     320            0 :         DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
     321            0 :         CHK_RET(HcclReduceAsync(
     322              :             dispatcher, reduceSrc.ptr(), dataCount, dataType_, reductionOp_, stream, reduceMem.localdst.ptr(),
     323              :             INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
     324            0 :     }
     325            0 :     return HCCL_SUCCESS;
     326              : }
     327              : 
     328              : // 仅rdma场景调用:主流一次性PreSync、RxAsync后批量下发本地reduce任务
     329            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunNoSpInlineReduce(
     330              :     const HcclDispatcher dispatcher, const LINK& link, const std::vector<ReducerMemoryInfo>& reducerMems,
     331              :     Stream& stream)
     332              : {
     333            0 :     std::vector<RxMemoryInfo> rxMems;
     334            0 :     for (const ReducerMemoryInfo& reduceMem : reducerMems) {
     335            0 :         rxMems.emplace_back(RxMemoryInfo{
     336            0 :             UserMemType::INPUT_MEM, reduceMem.remoteMemOffset, reduceMem.remoteRcvTemp.ptr(),
     337            0 :             reduceMem.remoteRcvTemp.size()});
     338              :     }
     339            0 :     CHK_RET(PreSync());
     340            0 :     CHK_RET(link->RxAsync(rxMems, stream));
     341            0 :     CHK_RET(ReducerLocalReduceSuffix(dispatcher, link, reducerMems, stream));
     342            0 :     return HCCL_SUCCESS;
     343            0 : }
     344              : 
     345            0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRun(
     346              :     const HcclDispatcher dispatcher, const LINK& link, const std::vector<ReducerMemoryInfo>& reducerMems,
     347              :     Stream& stream)
     348              : {
     349            0 :     CHK_PTR_NULL(stream.ptr());
     350            0 :     bool isSpInlineReduce = link->IsSpInlineReduce();
     351            0 :     if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_))) {
     352            0 :         CHK_RET(ReducerRunSpInlineReduce(dispatcher, link, reducerMems, stream));
     353            0 :     } else {
     354            0 :         CHK_RET(ReducerRunNoSpInlineReduce(dispatcher, link, reducerMems, stream));
     355              :     }
     356            0 :     return HCCL_SUCCESS;
     357              : }
     358              : 
     359            0 : HcclResult ReduceScatterRingConcurrentDirect::RunMainStreamTx(
     360              :     const u32 step, const std::vector<Slice>& txSliceVector, const std::vector<Slice>& rxSliceVector, const u32 rank,
     361              :     const u32 rankSize, std::vector<ReducerMemoryInfo>& rxReduceMems)
     362              : {
     363              :     (void)rank;
     364            0 :     CHK_RET(leftLink_->TxAck(stream_));
     365            0 :     CHK_RET(rightLink_->RxAck(stream_));
     366            0 :     u32 sliceSize = slices_.size() / rankSize;
     367              : 
     368              :     // 通信,如果是最后一步,则做消减拷贝
     369            0 :     std::vector<SenderMemoryInfo> txMems;
     370            0 :     DeviceMem dst;
     371            0 :     for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     372              :         // Ack
     373            0 :         HCCL_DEBUG(
     374              :             "Reduce operation: step[%u] stream[main], src rank[%u] starts to send offset[%llu] size[%llu] "
     375              :             "from leftMem_",
     376              :             step, leftLink_->GetRemoteRank(), rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     377            0 :         if (isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
     378            0 :             HCCL_DEBUG("[RunReduceScatter] sdma DMAReduce step");
     379            0 :             HCCL_DEBUG(
     380              :                 "Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     381              :                 "at userMemOut_",
     382              :                 step, userRank_, lastStepOffset_, rxSliceVector[sliceIdx].size);
     383            0 :             dst = DeviceMem::create(
     384            0 :                 static_cast<u8*>(opInfo_->outputAddr) + lastStepOffset_, rxSliceVector[sliceIdx].size);
     385              :         } else {
     386            0 :             HCCL_DEBUG(
     387              :                 "Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     388              :                 "at inputMem_",
     389              :                 step, userRank_, rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     390            0 :             dst = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     391            0 :             if (!isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
     392            0 :                 HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce step");
     393            0 :                 finalSrc_ = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     394            0 :                 finalDst_ = DeviceMem::create(
     395            0 :                     static_cast<u8*>(opInfo_->outputAddr) + lastStepOffset_, rxSliceVector[sliceIdx].size);
     396              :             }
     397              :         }
     398              :         // 在inline reduce场景, 需要利用scratchMem_暂存
     399            0 :         DeviceMem srcMemTemp = scratchMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
     400            0 :         DeviceMem srcMem = inputMem_.range(txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
     401            0 :         HCCL_DEBUG(
     402              :             "Reduce operation: step[%u] stream[main], senderInfo_ rank[%u] starts to rcv offset[%llu], "
     403              :             "size[%llu]",
     404              :             step, rightLink_->GetRemoteRank(), txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
     405            0 :         rxReduceMems.emplace_back(
     406            0 :             ReducerMemoryInfo{baseOffset_ + rxSliceVector[sliceIdx].offset, dst, dst, srcMemTemp});
     407            0 :         txMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSliceVector[sliceIdx].offset, srcMem});
     408            0 :     }
     409            0 :     CHK_RET(senderInfo_->run(rightLink_, txMems, stream_));
     410            0 :     return HCCL_SUCCESS;
     411            0 : }
     412              : 
     413              : // 仅rdma场景调用:主流Tx前缀 + ReducerRun整段下发
     414            0 : HcclResult ReduceScatterRingConcurrentDirect::RunMainStream(
     415              :     const u32 step, std::vector<Slice> txSliceVector, std::vector<Slice> rxSliceVector, const u32 rank,
     416              :     const u32 rankSize)
     417              : {
     418            0 :     std::vector<ReducerMemoryInfo> rxReduceMems;
     419            0 :     CHK_RET(RunMainStreamTx(step, txSliceVector, rxSliceVector, rank, rankSize, rxReduceMems));
     420            0 :     CHK_RET(ReducerRun(dispatcher_, leftLink_, rxReduceMems, stream_));
     421            0 :     return HCCL_SUCCESS;
     422            0 : }
     423              : 
     424              : // 从流提前下发Post(mainSignals),使主流Wait(mainSignals)可在队列未积压时及时通过
     425            0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStreamPrePost()
     426              : {
     427            0 :     CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     428            0 :     return HCCL_SUCCESS;
     429              : }
     430              : 
     431              : // 仅rdma场景调用:从流Wait(subSignals)后批量下发本步拷贝任务
     432            0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStream(
     433              :     const u32 step, std::vector<Slice> subSliceVector, std::vector<Slice> cclSliceVector, const u32 rank,
     434              :     const u32 rankSize)
     435              : {
     436              :     (void)rank;
     437            0 :     CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     438            0 :     for (u32 sliceIdx = 0; sliceIdx < subSliceVector.size(); sliceIdx++) {
     439            0 :         CHK_RET(RunSubStreamSlice(step, sliceIdx, subSliceVector, cclSliceVector, rankSize));
     440              :     }
     441            0 :     return HCCL_SUCCESS;
     442              : }
     443              : 
     444              : // 从流单个slice的拷贝任务
     445            0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStreamSlice(
     446              :     const u32 step, const u32 sliceIdx, const std::vector<Slice>& subSliceVector,
     447              :     const std::vector<Slice>& cclSliceVector, const u32 rankSize)
     448              : {
     449            0 :     HCCL_DEBUG(
     450              :         "Memcpy operation: step[%u] stream[sub], src rank[%u] starts to send offset[%llu], size[%llu] "
     451              :         "from userMemIn_",
     452              :         step, userRank_, subSliceVector[sliceIdx].offset, subSliceVector[sliceIdx].size);
     453              :     DeviceMem src = DeviceMem::create(
     454            0 :         static_cast<u8*>(opInfo_->inputAddr) + subSliceVector[sliceIdx].offset, subSliceVector[sliceIdx].size);
     455            0 :     DeviceMem dst;
     456            0 :     if (step == rankSize - DMA_REDUCE_TWO_OFFSET) {
     457              :         // do nothing
     458            0 :     } else if (isSdma_ && step == rankSize - DMA_REDUCE_THREE_OFFSET && opInfo_->outputAddr != nullptr) {
     459            0 :         HCCL_DEBUG(
     460              :             "Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     461              :             "to userMemOut_",
     462              :             step, userRank_, lastStepOffset_, subSliceVector[sliceIdx].size);
     463            0 :         dst = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) + lastStepOffset_, subSliceVector[sliceIdx].size);
     464            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
     465            0 :     } else {
     466            0 :         HCCL_DEBUG(
     467              :             "Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
     468              :             "to inputMem_",
     469              :             step, userRank_, cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
     470            0 :         dst = inputMem_.range(cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
     471            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
     472              :     }
     473            0 :     return HCCL_SUCCESS;
     474            0 : }
     475              : 
     476              : // 仅sdma场景调用:RxDataSignal后,主从流逐slice交替下发,最后补本地reduce后缀
     477            0 : HcclResult ReduceScatterRingConcurrentDirect::RunSdmaStepConcurrent(
     478              :     const u32 step, const std::vector<ReducerMemoryInfo>& rxReduceMems, const std::vector<Slice>& subSliceVector,
     479              :     const std::vector<Slice>& cclSliceVector, const u32 rank, const u32 rankSize)
     480              : {
     481              :     (void)rank;
     482            0 :     CHK_RET(leftLink_->RxDataSignal(stream_));
     483            0 :     bool isSpInlineReduce = leftLink_->IsSpInlineReduce() && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_));
     484            0 :     void* remoteMem = nullptr;
     485            0 :     if (isSpInlineReduce) {
     486            0 :         CHK_RET(leftLink_->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     487              :     }
     488              :     // 每个slice按 从流Post(mainSignals) -> 主流Wait/Empty/Post -> 从流Wait(subSignals) -> 从流拷贝 ->
     489              :     // 主流reduce/远端读 的顺序交替下发
     490            0 :     for (u32 sliceIdx = 0; sliceIdx < subSliceVector.size(); sliceIdx++) {
     491            0 :         CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
     492            0 :         CHK_RET(PreSync());
     493            0 :         CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
     494            0 :         CHK_RET(RunSubStreamSlice(step, sliceIdx, subSliceVector, cclSliceVector, rankSize));
     495            0 :         if (isSpInlineReduce) {
     496            0 :             CHK_RET(ReducerSpInlineSlice(dispatcher_, leftLink_, remoteMem, rxReduceMems[sliceIdx], stream_));
     497              :         } else {
     498            0 :             CHK_RET(ReducerSdmaRemoteReadSlice(leftLink_, rxReduceMems[sliceIdx], stream_));
     499              :         }
     500              :     }
     501            0 :     if (!isSpInlineReduce) {
     502            0 :         CHK_RET(ReducerLocalReduceSuffix(dispatcher_, leftLink_, rxReduceMems, stream_));
     503              :     }
     504            0 :     return HCCL_SUCCESS;
     505              : }
     506              : 
     507            0 : HcclResult ReduceScatterRingConcurrentDirect::RunReduceScatter(const u32 rank, const u32 rankSize)
     508              : {
     509            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect starts, the input param rank[%u]", rank);
     510              :     // 空拷贝用于后续操作附着
     511            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     512              : 
     513            0 :     CHK_RET(RunInitStep(rank, rankSize));
     514            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     515            0 :     CHK_RET(MainRecordSub()); // 主流通知从流开始通信
     516            0 :     CHK_RET(SubWaitMain());   // 从流等待主流通知
     517              :     // 空拷贝用于主从流任务并发
     518            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
     519            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, subStreams_[0], dispatcher_));
     520            0 :     u32 sliceSize = slices_.size() / rankSize;
     521              : 
     522              :     // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3, subSliceIdx = 1
     523            0 :     u32 txSliceIdx = (rank + rankSize - 1) % rankSize;
     524            0 :     u32 rxSliceIdx = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
     525            0 :     u32 subSliceIdx = (rank + rankSize - DMA_REDUCE_THREE_OFFSET) % rankSize;
     526            0 :     for (u32 step = 0; step < rankSize - 1; step++) {
     527            0 :         std::vector<Slice> rxSliceVector;
     528            0 :         std::vector<Slice> cclSliceVector;
     529            0 :         std::vector<Slice> txSliceVector;
     530            0 :         std::vector<Slice> subSliceVector;
     531            0 :         for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
     532            0 :             rxSliceVector.push_back(slices_[rxSliceIdx * sliceSize + sliceIdx]);
     533            0 :             cclSliceVector.push_back(slices_[subSliceIdx * sliceSize + sliceIdx]);
     534            0 :             txSliceVector.push_back(slices_[txSliceIdx * sliceSize + sliceIdx]);
     535            0 :             subSliceVector.push_back(userSlices_[subSliceIdx * sliceSize + sliceIdx]);
     536              :         }
     537              : 
     538              :         // dispatcher_aicpu 单条流的任务队列存在上限,队列满后host会阻塞下发,因此主流与从流的任务必须
     539              :         // 交替下发:从流Wait(subSignals)依赖主流Post(subSignals),主流Wait(mainSignals)依赖从流
     540              :         // Post(mainSignals)。若先集中下发某一条流的全部任务,队列被占满后host阻塞,而队列中等待的信号
     541              :         // 又需要另一条流尚未下发的任务来产生,两条流互相死等。以下保证每个Wait与其配对的Post在小窗口
     542              :         // 内先后完成下发,且每条流上的任务序列保持不变。
     543            0 :         if (!isSdma_) {
     544              :             // 从流先Post(mainSignals),主流整段下发完成后,从流再Wait并下发本步拷贝任务
     545            0 :             CHK_RET(RunSubStreamPrePost());
     546              :             // 主流
     547            0 :             CHK_RET(RunMainStream(step, txSliceVector, rxSliceVector, rank, rankSize));
     548              :             // 从流
     549            0 :             CHK_RET(RunSubStream(step, subSliceVector, cclSliceVector, rank, rankSize));
     550              :         } else {
     551              :             // sdma场景:主流Tx前缀下发后,主从流逐slice交替下发
     552            0 :             std::vector<ReducerMemoryInfo> rxReduceMems;
     553            0 :             CHK_RET(RunMainStreamTx(step, txSliceVector, rxSliceVector, rank, rankSize, rxReduceMems));
     554            0 :             CHK_RET(RunSdmaStepConcurrent(step, rxReduceMems, subSliceVector, cclSliceVector, rank, rankSize));
     555            0 :         }
     556              : 
     557              :         // 更新索引
     558            0 :         subSliceIdx = (subSliceIdx + rankSize - 1) % rankSize;
     559            0 :         txSliceIdx = (txSliceIdx + rankSize - 1) % rankSize;
     560            0 :         rxSliceIdx = (rxSliceIdx + rankSize - 1) % rankSize;
     561            0 :     }
     562            0 :     CHK_RET(SubRecordMain()); // 从流通知主流通信完成
     563            0 :     CHK_RET(MainWaitSub());   // 主流等待从流通知
     564            0 :     if (!isSdma_ && opInfo_->outputAddr != nullptr) {
     565            0 :         HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce last step");
     566            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, finalDst_, finalSrc_, stream_));
     567              :     }
     568            0 :     HCCL_INFO("ReduceScatterRingConcurrentDirect finished to RunReduceScatter");
     569            0 :     return HCCL_SUCCESS;
     570              : }
     571              : 
     572              : // 主流通知从流干活
     573            0 : HcclResult ReduceScatterRingConcurrentDirect::MainRecordSub()
     574              : {
     575            0 :     for (u32 signalIndex = 0; signalIndex < subSignals_.size(); signalIndex++) {
     576            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[signalIndex], profilerInput_.stage));
     577              :     }
     578            0 :     return HCCL_SUCCESS;
     579              : }
     580              : // 从流等待主流
     581            0 : HcclResult ReduceScatterRingConcurrentDirect::SubWaitMain()
     582              : {
     583            0 :     for (u32 streamIndex = 0; streamIndex < subSignals_.size(); streamIndex++) {
     584            0 :         CHK_RET(
     585              :             LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, subSignals_[streamIndex], profilerInput_.stage));
     586              :     }
     587            0 :     return HCCL_SUCCESS;
     588              : }
     589              : // 主流等待从流
     590            0 : HcclResult ReduceScatterRingConcurrentDirect::MainWaitSub()
     591              : {
     592            0 :     for (u32 signalIndex = 0; signalIndex < mainSignals_.size(); signalIndex++) {
     593            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[signalIndex], profilerInput_.stage));
     594              :     }
     595            0 :     return HCCL_SUCCESS;
     596              : }
     597              : // 从流告诉主流活干完了
     598            0 : HcclResult ReduceScatterRingConcurrentDirect::SubRecordMain()
     599              : {
     600            0 :     for (u32 streamIndex = 0; streamIndex < mainSignals_.size(); streamIndex++) {
     601            0 :         CHK_RET(
     602              :             LocalNotify::Post(subStreams_[streamIndex], dispatcher_, mainSignals_[streamIndex], profilerInput_.stage));
     603              :     }
     604            0 :     return HCCL_SUCCESS;
     605              : }
     606              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, ReduceScatterRingConcurrentDirect);
     607              : } // namespace hccl
        

Generated by: LCOV version 2.0-1