LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_unified_march.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 13.4 % 202 27
Test Date: 2026-07-28 12:11:00 Functions: 30.8 % 13 4

            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_unified_march.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : static const u32 NEIGHBORS_NUM_TWO = 2; //  2: 邻居数量
      16              : static const u32 NEIGHBORS_NUM_ONE = 1; //  1: 邻居数量
      17              : static const u32 DIVISOR_NUM_TWO = 2;
      18              : 
      19            1 : ReduceScatterUnifiedMarch::ReduceScatterUnifiedMarch(const HcclDispatcher dispatcher)
      20            1 :     : AlgTemplateBase(dispatcher)
      21            1 : {}
      22              : 
      23            2 : ReduceScatterUnifiedMarch::~ReduceScatterUnifiedMarch() {}
      24              : 
      25            1 : HcclResult ReduceScatterUnifiedMarch::Prepare(Stream &mainStream, SubCommInfo &level0CommInfo,
      26              :     DeviceMem &userInput, DeviceMem &userOutput, DeviceMem &usrInMem,
      27              :     DeviceMem &scratchMem, u64 totalCount, std::vector<Stream> &subStreams,
      28              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalMainToSub,
      29              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalSubToMain,
      30              :     const HcclDataType dataType, const HcclReduceOp reductionOp,
      31              :     const std::vector<std::vector<Slice>> &multRingsUserMemSlice, u64 reduceAttrBitMap)
      32              : {
      33            1 :     reduceAttr_ = reduceAttrBitMap;
      34            1 :     mainStream_ = mainStream;
      35            1 :     intraRank_ = level0CommInfo.localRank;
      36            1 :     intraRankSize_ = level0CommInfo.localRankSize;
      37            1 :     CHK_PRT_RET(intraRankSize_ == 0 || (intraRankSize_ % DIVISOR_NUM_TWO != 0),
      38              :         HCCL_ERROR("[ReduceScatterUnifiedMarch][Prepare]intraRankSize_ is zero or not divisible by 2"),
      39              :         HCCL_E_PARA);
      40            1 :     links_ = level0CommInfo.links;
      41              : 
      42            1 :     userInput_ = userInput;
      43            1 :     userOutput_ = userOutput;
      44            1 :     usrInMem_ = usrInMem;
      45            1 :     scratchMem_ = scratchMem;
      46            1 :     HCCL_INFO("userInput_[%p] size[%llu], userOutput_[%p] size[%llu], usrInMem_[%p] size[%llu], scratchMem_[%p] size[%llu]",
      47              :         userInput_.ptr(), userInput_.size(), userOutput_.ptr(), userOutput_.size(), usrInMem_.ptr(), usrInMem_.size(),
      48              :         scratchMem_.ptr(), scratchMem_.size());
      49              : 
      50            1 :     subStreams_ = subStreams;
      51            1 :     meshSignalMainToSub_ = meshSignalMainToSub;
      52            1 :     meshSignalSubToMain_ = meshSignalSubToMain;
      53            1 :     CHK_PRT_RET(subStreams_.size() < NEIGHBORS_NUM_TWO || meshSignalMainToSub_.size() < NEIGHBORS_NUM_TWO ||
      54              :         meshSignalSubToMain_.size() < NEIGHBORS_NUM_TWO,
      55              :         HCCL_ERROR("[AllGatherUnifiedMarch] subStreams_ size[%u] or meshSignalMainToSub_ size[%u] or "\
      56              :         "meshSignalSubToMain_ size[%u] is less than 2",
      57              :         subStreams_.size(), meshSignalMainToSub_.size(), meshSignalSubToMain_.size()), HCCL_E_PARA);
      58              : 
      59            1 :     totalCount_ = totalCount;
      60            1 :     dataType_ = dataType;
      61            1 :     reductionOp_ = reductionOp;
      62            1 :     blockDataByte_ = totalCount_ * SIZE_TABLE[dataType_];
      63            1 :     multRingsUserMemSlice_ = multRingsUserMemSlice;
      64            1 :     CHK_PRT_RET(multRingsUserMemSlice_[0].size() % intraRankSize_ != 0,
      65              :         HCCL_ERROR("[ReduceScatterUnifiedMarch] multRingsUserMemSlice_[0] size[%u] can not be divided by rank size[%u]",
      66              :         multRingsUserMemSlice_[0].size(), intraRankSize_), HCCL_E_PARA);
      67              : 
      68            1 :     return HCCL_SUCCESS;
      69              : }
      70              : 
      71            0 : std::string ReduceScatterUnifiedMarch::GetStreamIndexString()
      72              : {
      73            0 :     std::string res = "";
      74            0 :     for (u32 streamIndex = 0; streamIndex < subStreams_.size(); streamIndex++) {
      75            0 :         res += std::to_string(streamIndex) + ", ";
      76              :     }
      77            0 :     return res;
      78            0 : }
      79              : 
      80              : // 主流通知所有从流
      81            0 : HcclResult ReduceScatterUnifiedMarch::NotifySubStreamStart(u32 streamSize)
      82              : {
      83            0 :     CHK_PRT_RET(streamSize > subStreams_.size() || streamSize > meshSignalSubToMain_.size(),
      84              :         HCCL_ERROR("[ReduceScatterUnifiedMarch][NotifySubStreamStart] streamSize[%u] is out of range"\
      85              :         "subStreams_ size[%zu] or meshSignalSubToMain_ size[%zu]", streamSize, subStreams_.size(),
      86              :         meshSignalSubToMain_.size()), HCCL_E_PARA);    
      87            0 :     for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
      88            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, meshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
      89            0 :         CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, meshSignalSubToMain_[streamIndex],
      90              :             INVALID_VALUE_STAGE));
      91              :     }
      92            0 :     HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifySubStreamStart] intraRank [%u] main stream notify substream [%s]",
      93              :         intraRank_, GetStreamIndexString().c_str());
      94            0 :     return HCCL_SUCCESS;
      95              : }
      96              : 
      97            0 : HcclResult ReduceScatterUnifiedMarch::WaitSubStreamFinish(u32 streamSize)
      98              : {
      99            0 :     CHK_PRT_RET(streamSize > subStreams_.size() || streamSize > meshSignalMainToSub_.size(),
     100              :         HCCL_ERROR("[ReduceScatterUnifiedMarch][WaitSubStreamFinish] streamSize[%u] is out of range"\
     101              :         "subStreams_ size[%zu] or meshSignalMainToSub_ size[%zu]",
     102              :         streamSize, subStreams_.size(), meshSignalMainToSub_.size()), HCCL_E_PARA);
     103            0 :     for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
     104            0 :         CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, meshSignalMainToSub_[streamIndex],
     105              :             INVALID_VALUE_STAGE));
     106            0 :         CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, meshSignalMainToSub_[streamIndex],
     107              :             INVALID_VALUE_STAGE));
     108              :     }
     109            0 :     HCCL_DEBUG("[ReduceScatterUnifiedMarch][WaitSubStreamFinish] intraRank [%u] main stream wait substream [%s]",
     110              :         intraRank_, GetStreamIndexString().c_str());
     111            0 :     return HCCL_SUCCESS;
     112              : }
     113              : 
     114            0 : HcclResult ReduceScatterUnifiedMarch::NotifyNeighborsStart(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
     115              : {
     116              :     // 图模式保持使用Post/Wait接口
     117            0 :     if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     118              :         // notify是否越界由平台侧保证
     119            0 :         for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
     120            0 :             if (neighborRankId == 0) {
     121            0 :                 CHK_RET(nextIntralLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
     122            0 :                 CHK_RET(prevIntraLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait
     123            0 :             } else if (neighborRankId == 1) {
     124            0 :                 CHK_RET(prevIntraLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
     125            0 :                 CHK_RET(nextIntralLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait        
     126              :             }
     127              :         }
     128            0 :         HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done",
     129              :             intraRank_, neighbors);
     130            0 :         return HCCL_SUCCESS;
     131              :     }
     132              : 
     133              :     // 一条流负责一个环
     134            0 :     for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
     135              :         // 交替使用Ack和DataSignal两种notify
     136            0 :         const u32 NOTIFY_IDX_TWO = 2;
     137            0 :         if (neighborRankId == 0) {
     138            0 :             if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
     139            0 :                 CHK_RET(nextIntralLink->TxAck(subStreams_[neighborRankId]));    // AckRecord
     140            0 :                 CHK_RET(prevIntraLink->RxAck(subStreams_[neighborRankId]));     // AckWait
     141              :             } else {
     142            0 :                 CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId]));     // DataRecord
     143            0 :                 CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId]));      // DataWait
     144              :             }
     145            0 :         } else if (neighborRankId == 1) {
     146            0 :             if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
     147            0 :                 CHK_RET(prevIntraLink->TxAck(subStreams_[neighborRankId]));     // AckRecord
     148            0 :                 CHK_RET(nextIntralLink->RxAck(subStreams_[neighborRankId]));    // AckWait
     149              :             } else {
     150            0 :                 CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId]));      // DataRecord
     151            0 :                 CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId]));     // DataWait
     152              :             }
     153              :         }
     154              :     }
     155            0 :     HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done",
     156              :         intraRank_, neighbors);
     157            0 :     return HCCL_SUCCESS;
     158              : }
     159              : 
     160            0 : HcclResult ReduceScatterUnifiedMarch::NotifyNeighborsEnd(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
     161              : {
     162            0 :     for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
     163            0 :         if (neighborRankId == 0) {
     164            0 :             CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
     165            0 :             CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId]));     
     166            0 :         } else if (neighborRankId == 1) {
     167            0 :             CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId]));
     168            0 :             CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId]));   
     169              :         }
     170              :     }
     171            0 :     HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsEnd] intraRank[%u] notifys [%u]neigbhbors reduce done",
     172              :         intraRank_, neighbors);
     173            0 :     return HCCL_SUCCESS;
     174              : }
     175              : 
     176            0 : HcclResult ReduceScatterUnifiedMarch::DoSerialReduce(void* remDMAMemPtr, void* dstAddr, u64 memSize,
     177              :     u64 dataCount, Stream &tmpStream, LINK& tmpLink, u64 remoteOffsetByte)
     178              : {
     179            0 :     for (u32 sliceIdx = 0; sliceIdx < (multRingsUserMemSlice_[0].size() / intraRankSize_); sliceIdx++) {
     180            0 :         DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr) + remoteOffsetByte +
     181            0 :             multRingsUserMemSlice_[0][sliceIdx].offset, memSize);
     182              :         DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(dstAddr) +
     183            0 :             multRingsUserMemSlice_[0][sliceIdx].offset, memSize);
     184              : 
     185            0 :         if ((INLINE_REDUCE_BITMASK & reduceAttr_) == 1) { // inlineReduce
     186            0 :             struct hccl::Transport::Buffer remoteBuf;
     187            0 :             remoteBuf.addr = srcMem.ptr();
     188            0 :             remoteBuf.size = srcMem.size();
     189            0 :             struct hccl::Transport::Buffer localBuf;
     190            0 :             localBuf.addr = dstMem.ptr();
     191            0 :             localBuf.size = dstMem.size();
     192            0 :             HCCL_DEBUG("intralRank[%u] slice[%u] offset[%llu] do inlinereduce with remoteBuf[addr[%p], size[%llu]] and "\
     193              :                 "localBuf[addr[%p], size[%llu]]", intraRank_, sliceIdx, multRingsUserMemSlice_[0][sliceIdx].offset,
     194              :                 remoteBuf.addr, remoteBuf.size, localBuf.addr, localBuf.size);
     195            0 :             CHK_RET(tmpLink->ReadReduceSync(localBuf, remoteBuf, dataType_, reductionOp_, tmpStream));
     196              :         } else { // TBE_reduce
     197              :             // left的inputMem拷到本端的scratchMem
     198            0 :             DeviceMem tempMem = scratchMem_.range(remoteOffsetByte, srcMem.size());
     199            0 :             struct hccl::Transport::Buffer remoteBuf;
     200            0 :             remoteBuf.addr = srcMem.ptr();
     201            0 :             remoteBuf.size = srcMem.size();
     202            0 :             struct hccl::Transport::Buffer localBuf;
     203            0 :             localBuf.addr = tempMem.ptr();
     204            0 :             localBuf.size = tempMem.size();
     205            0 :             HCCL_DEBUG("intralRank[%u] slice[%u] offset[%llu] do SDMA read with remoteBuf[addr[%p], size[%llu]] and "\
     206              :                 "localBuf[addr[%p], size[%llu]]", intraRank_, sliceIdx, multRingsUserMemSlice_[0][sliceIdx].offset,
     207              :                 remoteBuf.addr, remoteBuf.size, localBuf.addr, localBuf.size);
     208            0 :             CHK_RET(tmpLink->ReadSync(localBuf, remoteBuf, tmpStream));
     209            0 :             CHK_RET(HcclReduceAsync(dispatcher_, tempMem.ptr(), dataCount, dataType_, reductionOp_, tmpStream,
     210              :                 dstMem.ptr(), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
     211            0 :         }
     212            0 :     }
     213            0 :     return HCCL_SUCCESS; 
     214              : }
     215              : 
     216            0 : HcclResult ReduceScatterUnifiedMarch::RunSingleSliceRead(u32 ringPrevRank, u32 ringNextRank, u32 step, u32 totalStep)
     217              : {
     218            0 :     LINK prevIntraLink = links_[ringPrevRank];
     219            0 :     CHK_SMART_PTR_NULL(prevIntraLink);
     220            0 :     LINK nextIntralLink = links_[ringNextRank];
     221            0 :     CHK_SMART_PTR_NULL(nextIntralLink);
     222            0 :     u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
     223            0 :     CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, neighbors));
     224              : 
     225              :     //拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
     226            0 :     CHK_RET(WaitSubStreamFinish(neighbors));
     227            0 :     CHK_RET(NotifySubStreamStart(neighbors));
     228              : 
     229              :     // 从前向rank读取数据
     230            0 :     void* preRemDMAMemPtr = nullptr;
     231            0 :     CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::INPUT_MEM, &preRemDMAMemPtr));
     232            0 :     u32 preDataIndex = (intraRank_ + intraRankSize_ - step - totalStep) % intraRankSize_;
     233            0 :     u64 preOffsetByte = preDataIndex * blockDataByte_;
     234            0 :     void* preDstAddr = static_cast<u8 *>(userInput_.ptr()) + preOffsetByte;
     235              : 
     236            0 :     CHK_RET(DoSerialReduce(preRemDMAMemPtr, preDstAddr, blockDataByte_, totalCount_,
     237              :         subStreams_[0], prevIntraLink, preOffsetByte));
     238            0 :     HCCL_INFO("[ReduceScatterUnifiedMarch][RunSingleSliceRead] intralRank [%u] reduce with ringPrevRank [%u] done",
     239              :         intraRank_, ringPrevRank);
     240              : 
     241              :     // 从后向rank读取数据
     242            0 :     if (neighbors > NEIGHBORS_NUM_ONE) {
     243            0 :         void* nextRemDMAMemPtr = nullptr;
     244            0 :         CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::INPUT_MEM, &nextRemDMAMemPtr));
     245            0 :         u32 nextDataIndex = (intraRank_ + totalStep + step) % intraRankSize_;
     246            0 :         u64 nextOffsetByte = nextDataIndex * blockDataByte_;
     247            0 :         void* nextDstAddr = static_cast<u8 *>(userInput_.ptr()) + nextOffsetByte;
     248              : 
     249            0 :         CHK_RET(DoSerialReduce(nextRemDMAMemPtr, nextDstAddr, blockDataByte_, totalCount_,
     250              :             subStreams_[1], nextIntralLink, nextOffsetByte));
     251            0 :         HCCL_INFO("[ReduceScatterUnifiedMarch][RunSingleSliceRead] intralRank [%u]"\
     252              :             "reduce with ringNextRank [%u] done", intraRank_, ringNextRank);
     253              :     }
     254              : 
     255              :     /* 2卡 场景,在最后一步的notifyDone */
     256            0 :     if (step == 0) {
     257            0 :         CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, neighbors));
     258              :     }
     259            0 :     notifyIdx_++;
     260              : 
     261            0 :     return HCCL_SUCCESS;
     262            0 : }
     263              : 
     264            0 : HcclResult ReduceScatterUnifiedMarch::RunHalfSliceRead(u32 ringPrevRank, u32 ringNextRank, u32 step, u32 totalStep)
     265              : {
     266            0 :     LINK prevIntraLink = links_[ringPrevRank];
     267            0 :     CHK_SMART_PTR_NULL(prevIntraLink);
     268            0 :     LINK nextIntralLink = links_[ringNextRank];
     269            0 :     CHK_SMART_PTR_NULL(nextIntralLink);
     270            0 :     CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
     271              : 
     272              :     //拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
     273            0 :     CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
     274            0 :     CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
     275              : 
     276              :     // 从前向rank读取数据
     277            0 :     void* preRemDMAMemPtr = nullptr;
     278            0 :     CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::INPUT_MEM, &preRemDMAMemPtr));
     279            0 :     u32 temIdx = (step == 0) ? totalStep : 0;
     280            0 :     u32 preDataIndex = (intraRank_ + intraRankSize_ - temIdx) % intraRankSize_;
     281              :     // 考虑总数据量不能被整除的情况
     282            0 :     u32 partOneCount = (step != totalStep) ? (totalCount_ / DIVISOR_NUM_TWO) : (totalCount_ - totalCount_ / DIVISOR_NUM_TWO);
     283            0 :     u64 partOneSize = partOneCount * SIZE_TABLE[dataType_];
     284            0 :     u64 preOffsetByte = (step != totalStep) ? (preDataIndex * blockDataByte_) :
     285            0 :         (preDataIndex * blockDataByte_ + totalCount_ / DIVISOR_NUM_TWO * SIZE_TABLE[dataType_]);
     286            0 :     void* preDstAddr = static_cast<u8 *>(userInput_.ptr()) + preOffsetByte;
     287              : 
     288            0 :     CHK_RET(DoSerialReduce(preRemDMAMemPtr, preDstAddr, partOneSize, partOneCount,
     289              :         subStreams_[0], prevIntraLink, preOffsetByte));
     290            0 :     HCCL_INFO("[ReduceScatterUnifiedMarch][RunHalfSliceRead] intralRank [%u] reduce with ringPrevRank [%u] done",
     291              :         intraRank_, ringPrevRank);
     292              : 
     293              :     // 从后向rank读取数据
     294            0 :     void* nextRemDMAMemPtr = nullptr;
     295            0 :     CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::INPUT_MEM, &nextRemDMAMemPtr));
     296            0 :     temIdx = (step == 0) ? totalStep : 0;
     297            0 :     u32 nextDataIndex = (intraRank_ + temIdx) % intraRankSize_;
     298            0 :     u32 partTwoCount = totalCount_ - partOneCount;
     299            0 :     u64 partTwoSize = partTwoCount * SIZE_TABLE[dataType_];
     300            0 :     u64 nextOffsetByte = (step != totalStep) ? (nextDataIndex * blockDataByte_ + partOneSize) :
     301            0 :         nextDataIndex * blockDataByte_;
     302            0 :     void* nextDstAddr = static_cast<u8 *>(userInput_.ptr()) + nextOffsetByte;
     303              : 
     304            0 :     CHK_RET(DoSerialReduce(nextRemDMAMemPtr, nextDstAddr, partTwoSize, partTwoCount,
     305              :         subStreams_[1], nextIntralLink, nextOffsetByte));
     306            0 :     HCCL_INFO("[ReduceScatterUnifiedMarch][RunHalfSliceRead] intralRank [%u] reduce with ringNextRank [%u] done",
     307              :         intraRank_, ringNextRank);
     308              : 
     309              :     /* 4卡及以上的 场景,在最后一步的notifyDone */
     310              :     // 单算子使用Ack/Datasignal接口,必须保证两者交替使用
     311            0 :     if (step == totalStep) {
     312            0 :         const u32 NOTIFY_IDX_TWO = 2;
     313            0 :         if (notifyIdx_ % NOTIFY_IDX_TWO != 0 && GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     314            0 :             notifyIdx_++;
     315            0 :             CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
     316            0 :             CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
     317            0 :             CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
     318              :         }
     319            0 :         CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
     320              :     }
     321            0 :     notifyIdx_++;
     322              : 
     323            0 :     return HCCL_SUCCESS;
     324            0 : }
     325              : 
     326            0 : HcclResult ReduceScatterUnifiedMarch::RunAsync()
     327              : {
     328            0 :     HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForReduceScatter();
     329            0 :     CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
     330              : 
     331              :     // 获取link的收、发
     332            0 :     u32 ringPrevRank = (intraRank_ + intraRankSize_ - 1) % intraRankSize_;
     333            0 :     u32 ringNextRank = (intraRank_ + 1) % intraRankSize_;
     334              : 
     335            0 :     u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
     336            0 :     CHK_RET(NotifySubStreamStart(neighbors));
     337              : 
     338              :     // 计算所需的总步骤
     339            0 :     u32 totalStep = intraRankSize_ / DIVISOR_NUM_TWO + 1;
     340            0 :     u32 step = 0;
     341            0 :     if(totalStep == DIVISOR_NUM_TWO) {
     342            0 :         CHK_RET(RunSingleSliceRead(ringPrevRank, ringNextRank, step, totalStep));
     343              :     } else {
     344              :         // 进行第1步收发
     345            0 :         CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, step, totalStep));
     346              : 
     347              :         // 进行第k步收发
     348            0 :         step++;
     349            0 :         for (; step < totalStep - DIVISOR_NUM_TWO; step++) {
     350            0 :             CHK_RET(RunSingleSliceRead(ringPrevRank, ringNextRank, step, totalStep));
     351              :         }
     352              : 
     353              :         // 进行第totalStep - 1步收发
     354            0 :         CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, step, totalStep));
     355              : 
     356              :         // 进行第totalStep步收发
     357            0 :         CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, totalStep, totalStep));
     358              :     }
     359              : 
     360            0 :     CHK_RET(WaitSubStreamFinish(neighbors));
     361            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, subStreams_));
     362              : 
     363            0 :     HCCL_INFO("[ReduceScatterUnifiedMarch][RunAsync] finished.");
     364            0 :     return HCCL_SUCCESS;
     365              : }
     366              : 
     367              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_UNIFIED_MARCH, ReduceScatterUnifiedMarch);
     368              : } // namespace hccl
        

Generated by: LCOV version 2.0-1