LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_multi_deter_pipeline.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 261 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 17 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_multi_deter_pipeline.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : ReduceScatterMultiDeterPipeline::ReduceScatterMultiDeterPipeline(const HcclDispatcher dispatcher)
      16            0 :     : MultiDeterPipeline(dispatcher)
      17            0 : {}
      18              : 
      19            0 : ReduceScatterMultiDeterPipeline::~ReduceScatterMultiDeterPipeline() {}
      20              : 
      21            0 : HcclResult ReduceScatterMultiDeterPipeline::GetRemoteCclbufferDeviceMem(
      22              :     u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
      23              : {
      24            0 :     u64 inputSliceOffset = memSliceSize_ * inputSliceIndex + offset_;
      25            0 :     u64 eachOffset = eachRankCclbufferSize_;
      26            0 :     u64 outputSliceOffset = eachOffset * outputSliceIndex;
      27            0 :     u64 outputInSliceOffset = (HCCL_MIN_SLICE_ALIGN_910B + (inputSliceOffset % HCCL_MIN_SLICE_ALIGN_910B)
      28              :                                - (outputSliceOffset % HCCL_MIN_SLICE_ALIGN_910B))
      29            0 :                               % HCCL_MIN_SLICE_ALIGN_910B;
      30            0 :     void* remoteMemPtr = nullptr;
      31            0 :     CHK_RET(link->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
      32            0 :     u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
      33            0 :     u8* intraSrcAddr = beginAddrU8 + outputSliceOffset + outputInSliceOffset;
      34            0 :     remoteMem = DeviceMem::create(intraSrcAddr, curSize_);
      35            0 :     if (remoteMem.ptr() == nullptr) {
      36            0 :         HCCL_ERROR(
      37              :             "[%s] outputSliceOffset + outputInSliceOffset + curSize_ = [%llu] > cclBufferSize[%llu]", __func__,
      38              :             outputSliceOffset + outputInSliceOffset + curSize_, cclBuffer_.size());
      39            0 :         return HCCL_E_MEMORY;
      40              :     }
      41            0 :     HCCL_DEBUG(
      42              :         "[%s] rank[%u], beginAddr[%p], outputSliceOffset[%llu](outputSliceIndex * eachOffset[%llu]), "
      43              :         "outputInSliceOffset[%llu], curSize[%llu], totalBufferSize[%llu]",
      44              :         __func__, outputSliceIndex, remoteMem.ptr(), outputSliceOffset, eachOffset, outputInSliceOffset, curSize_,
      45              :         cclBuffer_.size());
      46            0 :     return HCCL_SUCCESS;
      47              : }
      48              : 
      49              : // RDMA 发送时顶格收,所以不需要128K对齐,故sliceOffset为0
      50              : // SDMA 发送后 取本地 sliceOffset = slices_[rankIdInAllRanks].offset偏移处存放地址
      51              : HcclResult
      52            0 : ReduceScatterMultiDeterPipeline::GetLocalCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, u64 sliceOffset)
      53              : {
      54            0 :     u64 eachOffset = eachRankCclbufferSize_; // 当前轮有效数据大小 + HCCL_MIN_SLICE_ALIGN_910B作为偏移
      55            0 :     u64 rdmaOffset = eachOffset * rankIdInAllRanks;
      56            0 :     u64 sdmaOffset = sliceOffset;
      57            0 :     u64 offset = sliceOffset == 0 ? rdmaOffset : sdmaOffset;
      58            0 :     localMem = cclBuffer_.range(offset, curSize_);
      59            0 :     if (localMem.ptr() == nullptr) {
      60            0 :         HCCL_ERROR(
      61              :             "[%s] get localMem failed, rdmaOffset + curSize_"
      62              :             "= [%llu] or sdmaOffset + curSize = [%llu] > cclBufferSize[%llu]",
      63              :             __func__, rdmaOffset + curSize_, sdmaOffset + curSize_, cclBuffer_.size());
      64            0 :         return HCCL_E_MEMORY;
      65              :     }
      66            0 :     HCCL_DEBUG(
      67              :         "[%s] rank[%u], beginAddr[%p], offset[%llu](rdmaOffset[%u] or sdmaOffset[%llu]), "
      68              :         "curSize[%llu], totalBufferSize[%llu]",
      69              :         __func__, rankIdInAllRanks, localMem.ptr(), sliceOffset, rdmaOffset, sdmaOffset, curSize_, cclBuffer_.size());
      70            0 :     return HCCL_SUCCESS;
      71              : }
      72              : 
      73            0 : HcclResult ReduceScatterMultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
      74              : {
      75            0 :     u8* beginAddrU8 = static_cast<u8*>(usrInMemPtr_);
      76            0 :     u64 eachOffset = memSliceSize_;
      77            0 :     u8* intraSrcAddr = beginAddrU8 + (rankIdInAllRanks * eachOffset); // 不用 + offset_,因为usrInMem_已经加过了
      78            0 :     localMem = DeviceMem::create(intraSrcAddr, curSize_);
      79            0 :     if (localMem.ptr() == nullptr) {
      80            0 :         HCCL_ERROR(
      81              :             "[%s] get localMem failed, rankIdInAllRanks * eachOffset + curSize_ = [%u] is too big", __func__,
      82              :             rankIdInAllRanks * eachOffset + curSize_, cclBuffer_.size());
      83            0 :         return HCCL_E_MEMORY;
      84              :     }
      85            0 :     HCCL_DEBUG(
      86              :         "[%s] ranks[%u], offset[%llu](rankIdInAllRanks * eachOffset[%llu])"
      87              :         "intraSrcAddr[%p], memSliceSize[%llu], usrInMemPtr[%p]",
      88              :         __func__, rankIdInAllRanks, rankIdInAllRanks * eachOffset, eachOffset, intraSrcAddr, memSliceSize_,
      89              :         usrInMemPtr_);
      90            0 :     return HCCL_SUCCESS;
      91              : }
      92              : 
      93            0 : HcclResult ReduceScatterMultiDeterPipeline::RunLocalCopy()
      94              : {
      95              :     // 每张卡将自己的input的第user rank块数据搬到output,例如0A 1B 2C
      96            0 :     DeviceMem userIn;
      97            0 :     CHK_RET(GetLocalUserInDeviceMem(userRank_, userIn));
      98            0 :     DeviceMem userOut = DeviceMem::create(usrOutMemPtr_, curSize_);
      99              :     // 使用主流搬迁卡内数据
     100            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, userOut, userIn, mainStream_));
     101            0 :     HCCL_DEBUG(
     102              :         "[%s] intra-card copy data from [%u] to usrOutMem[%p] size[%llu]", __func__, userRank_, usrOutMemPtr_,
     103              :         curSize_);
     104            0 :     return HCCL_SUCCESS;
     105            0 : }
     106              : 
     107              : // 机内alltoall full mesh收集数据, #step表示pairwise的第step步
     108            0 : HcclResult ReduceScatterMultiDeterPipeline::RunIntraAlltoallPreSync(u32 step)
     109              : {
     110            0 :     HCCL_DEBUG("[%s] intra-server alltoall begin, step[%u]", __func__, step);
     111            0 :     HCCL_DEBUG(
     112              :         "[%s] intra-server SDMA send begin, [serverId, intraRankId] = [%u, %u]", __func__, serverId_, intraRankId_);
     113            0 :     for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
     114            0 :         HCCL_DEBUG(
     115              :             "[%s] intra-server SDMA send begin, userRank[%u] step[%u] pro[%u/%u]", __func__, userRank_, i, i + 1,
     116              :             intraRankSize_ - 1);
     117              :         // 从机内rankId为recvIntraRankId收集数据,也发给机内rankId为sendIntraRankId数据
     118            0 :         u32 recvIntraRankId = GetPreIntraRankIdByStep(i + 1);
     119            0 :         u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
     120            0 :         LINK recvIntraLink = intraLinks_[recvIntraRankId];
     121            0 :         LINK sendIntraLink = intraLinks_[sendIntraRankId];
     122            0 :         CHK_RET(sendIntraLink->TxAck(subStreams_[i]));
     123            0 :         CHK_RET(sendIntraLink->RxAck(subStreams_[i]));
     124            0 :     }
     125              :     // 增加主从流同步,目的是让SDMA同时进行
     126            0 :     CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     127            0 :     CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     128            0 :     CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     129            0 :     CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     130            0 :     return HCCL_SUCCESS;
     131              : }
     132              : 
     133            0 : HcclResult ReduceScatterMultiDeterPipeline::BatchPostNotifyForStreams(
     134              :     const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
     135              : {
     136            0 :     if (useMainStream) {
     137            0 :         HCCL_DEBUG("[%s] use mainStrem, skip notify wait", __func__);
     138            0 :         return HCCL_SUCCESS;
     139              :     }
     140            0 :     for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
     141            0 :         if (streamTasks[s].empty())
     142            0 :             continue; // 无任务的流跳过
     143            0 :         u32 streamIdx = reduceStreamBegin_ + s;
     144            0 :         if (reduceMainStreamIdx_ == streamIdx) {
     145            0 :             continue;
     146              :         }
     147            0 :         if (isStartPhase) {
     148              :             // 启动阶段:主流→子流 通知(Post主流,Wait子流)
     149            0 :             CHK_RET(LocalNotify::Post(
     150              :                 subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
     151            0 :             CHK_RET(LocalNotify::Wait(
     152              :                 subStreams_[streamIdx], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
     153            0 :             HCCL_DEBUG("[%s] stream[%u] start phase notify done", __func__, streamIdx);
     154              :         } else {
     155              :             // 同步阶段:子流→主流 通知(Post子流,Wait主流)
     156            0 :             CHK_RET(LocalNotify::Post(
     157              :                 subStreams_[streamIdx], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
     158            0 :             CHK_RET(LocalNotify::Wait(
     159              :                 subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
     160            0 :             HCCL_DEBUG("[%s] stream[%u] sync phase notify done", __func__, streamIdx);
     161              :         }
     162              :     }
     163            0 :     return HCCL_SUCCESS;
     164              : }
     165              : 
     166              : // 机内localreduce首先按序收集所有内存块,接着二分归并reduce,最多使用4条流并行
     167            0 : HcclResult ReduceScatterMultiDeterPipeline::RunIntraLocalReduce(u32 step)
     168              : {
     169            0 :     HCCL_DEBUG("[%s] inter-server local reduce begin, step[%u]", __func__, step);
     170            0 :     u32 recvServerId = GetPreServerIdByStep(step);  // 从上一个收
     171            0 :     u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个
     172            0 :     std::vector<DeviceMem> reduceMem;
     173            0 :     std::vector<bool> isReduceBlock;
     174            0 :     u32 retIndex = 0;
     175            0 :     isReduceBlock.resize(intraRankSize_);
     176              :     // 机内,最后rank的规约结果放在倒数第2块,其他放在倒数第1块
     177            0 :     if (intraRankId_ == intraRankSize_ - 1) {
     178            0 :         retIndex = intraRankSize_ - SECOND_TO_LAST;
     179              :     } else {
     180            0 :         retIndex = intraRankSize_ - 1;
     181              :     }
     182              :     // 机内第0步,规约到usrOut,即rank所在机间的intraRankId_处
     183            0 :     if (step == 0) {
     184            0 :         retIndex = intraRankId_;
     185              :     }
     186            0 :     HCCL_DEBUG(
     187              :         "[%s] intra-server local reduce, retIndex[%u], intraRankSize[%u] intraRankId[%u]", __func__, retIndex,
     188              :         intraRankSize_, intraRankId_);
     189            0 :     reduceMem.resize(intraRankSize_);
     190            0 :     u32 userInIdx = GetRankIdx(sendServerId, intraRankId_);
     191              :     // 最后一块留给allreduce
     192            0 :     u32 idx = 0;
     193            0 :     const u32 serverId = (step == 0) ? serverId_ : recvServerId;
     194            0 :     for (u32 i = 0; i < intraRankSize_; ++i) {
     195            0 :         if (i == intraRankId_) {
     196            0 :             if (step == 0) {
     197              :                 // step=0:归到usrOut
     198            0 :                 isReduceBlock[i] = true;
     199            0 :                 DeviceMem usrOutInraMem = DeviceMem::create(usrOutMemPtr_, curSize_);
     200            0 :                 reduceMem[i] = std::move(usrOutInraMem);
     201            0 :                 HCCL_DEBUG("[%s] inter-server local reduce, NO.%u reduceMem stores userOut", __func__, i);
     202            0 :             } else {
     203              :                 // step≠0:填充usrIn
     204            0 :                 isReduceBlock[i] = false;
     205            0 :                 DeviceMem usrInIntraMem;
     206            0 :                 CHK_RET(GetLocalUserInDeviceMem(userInIdx, usrInIntraMem));
     207            0 :                 reduceMem[i] = std::move(usrInIntraMem);
     208            0 :                 HCCL_DEBUG(
     209              :                     "[%s] inter-server local reduce, NO.%u reduceMem stores userInIdx[%u],", __func__, i, userInIdx);
     210            0 :             }
     211            0 :             continue;
     212            0 :         }
     213              :         // 其他情况:统一处理CCLBuffer
     214            0 :         isReduceBlock[i] = true;
     215            0 :         const u32 outCCLbufferIdx = GetRankIdx(serverId, idx);
     216            0 :         DeviceMem cclbufferIntraMem;
     217            0 :         CHK_RET(GetLocalCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, slices_[outCCLbufferIdx].offset));
     218            0 :         reduceMem[i] = std::move(cclbufferIntraMem);
     219            0 :         HCCL_DEBUG(
     220              :             "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i, outCCLbufferIdx);
     221            0 :         idx++;
     222            0 :     }
     223            0 :     CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, false));
     224            0 :     HCCL_INFO("[%s] intra-server step[%u] run local reduce success", __func__, step);
     225            0 :     return HCCL_SUCCESS;
     226            0 : }
     227              : 
     228            0 : HcclResult ReduceScatterMultiDeterPipeline::RunInterSend(u32 step)
     229              : {
     230            0 :     HCCL_DEBUG("[%s] inter-server RDMA write begin, step[%u]", __func__, step);
     231              :     // 使用主流进行rdma
     232            0 :     u32 recvServerId = GetPreServerIdByStep(step);  // 从上一个收
     233            0 :     u32 sendServerId = GetNextServerIdByStep(step); // 发给下一个
     234            0 :     LINK recvInterLink = serverLinks_[recvServerId];
     235            0 :     LINK sendInterLink = serverLinks_[sendServerId];
     236              :     // 跨机 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块的第1块
     237              :     // 跨机 输入是input的第(R1+1)%S1组数据(往后1)的倒数第2块
     238            0 :     u32 reduceRetIndex = intraRankSize_ - SECOND_TO_LAST;
     239            0 :     u32 recvCCLbufferIdx = GetRankIdx(recvServerId, 0);              // 本端rank0:收的位置
     240            0 :     u32 sendCCLbufferIdx = GetRankIdx(recvServerId, reduceRetIndex); // 本端rank0:发的位置
     241              :     // 跨机写对端
     242            0 :     u32 sendRankId = GetRankIdx(sendServerId, intraRankId_); // 发送到sendRankId
     243            0 :     u32 recvRankId = GetRankIdx(recvServerId, intraRankId_); // 从recvRankId接收
     244            0 :     u32 remoteRecvCCLbufferIdx = GetRankIdx(serverId_, 0);   // 接收端:收的位置
     245              :     // 2机又从serverId_收也从serverId_发
     246            0 :     u32 remoteSendCCLbufferIdx = serverSize_ == MIN_SERVER_NUM ?
     247            0 :                                      GetRankIdx(serverId_, reduceRetIndex) :
     248            0 :                                      GetRankIdx(sendServerId, reduceRetIndex); // 发送端:发的位置
     249            0 :     DeviceMem recvMem;
     250            0 :     DeviceMem sendMem;
     251            0 :     CHK_RET(GetLocalCclbufferDeviceMem(recvCCLbufferIdx, recvMem, eachRankCclbufferSize_ * recvCCLbufferIdx));
     252            0 :     CHK_RET(GetLocalCclbufferDeviceMem(sendCCLbufferIdx, sendMem, slices_[sendCCLbufferIdx].offset));
     253            0 :     HCCL_DEBUG(
     254              :         "[%s] inter-server RDMA write begin, cclbufferIdx: [%u] send to [%u], [%u] recv from [%u]", __func__,
     255              :         sendCCLbufferIdx, remoteRecvCCLbufferIdx, recvCCLbufferIdx, remoteSendCCLbufferIdx);
     256            0 :     HCCL_DEBUG(
     257              :         "[%s] inter-server RDMA write begin, rankId: [%u] send to [%u], [%u] recv from [%u]", __func__, userRank_,
     258              :         sendRankId, userRank_, recvRankId);
     259              :     // A + X 单机16卡为SDMA读语义
     260            0 :     if (recvInterLink->IsSpInlineReduce() && sendInterLink->IsSpInlineReduce()) {
     261            0 :         CHK_RET(sendInterLink->TxAck(mainStream_));
     262            0 :         CHK_RET(recvInterLink->RxAck(mainStream_));
     263            0 :         DeviceMem dstMem = std::move(recvMem);
     264            0 :         DeviceMem srcMem;
     265            0 :         void* remoteMemPtr = nullptr;
     266            0 :         CHK_RET(recvInterLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr));
     267            0 :         u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
     268            0 :         u8* intraSrcAddr = beginAddrU8 + slices_[remoteSendCCLbufferIdx].offset;
     269            0 :         srcMem = DeviceMem::create(intraSrcAddr, curSize_);
     270            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_, recvRankId, recvInterLink->GetLinkType()));
     271            0 :         CHK_RET(sendInterLink->TxDataSignal(mainStream_));
     272            0 :         CHK_RET(recvInterLink->RxDataSignal(mainStream_));
     273            0 :     } else {
     274            0 :         CHK_RET(recvInterLink->TxAck(mainStream_));
     275            0 :         CHK_RET(sendInterLink->RxAck(mainStream_));
     276              : 
     277            0 :         CHK_RET(sendInterLink->TxAsync(
     278              :             UserMemType::OUTPUT_MEM, remoteRecvCCLbufferIdx * eachRankCclbufferSize_, sendMem.ptr(), curSize_,
     279              :             mainStream_));
     280            0 :         CHK_RET(recvInterLink->RxAsync(
     281              :             UserMemType::OUTPUT_MEM, slices_[remoteSendCCLbufferIdx].offset, recvMem.ptr(), curSize_, mainStream_));
     282              : 
     283            0 :         CHK_RET(recvInterLink->PostFinAck(mainStream_));
     284            0 :         CHK_RET(sendInterLink->WaitFinAck(mainStream_));
     285              :     }
     286            0 :     HCCL_INFO("[%s] inter-server step[%u] run RDMA send success", __func__, step);
     287            0 :     return HCCL_SUCCESS;
     288            0 : }
     289              : 
     290            0 : HcclResult ReduceScatterMultiDeterPipeline::RunFinalReduce()
     291              : {
     292              :     // 主从流同步
     293            0 :     HCCL_DEBUG("[%s] intra-server final reduce begin", __func__);
     294            0 :     std::vector<DeviceMem> reduceMem;
     295            0 :     std::vector<bool> isReduceBlock;
     296            0 :     u32 retIndex = serverId_;
     297            0 :     DeviceMem usrOutInraMem = DeviceMem::create(usrOutMemPtr_, curSize_);
     298            0 :     isReduceBlock.resize(serverSize_);
     299            0 :     reduceMem.resize(serverSize_);
     300              : 
     301            0 :     HCCL_DEBUG("[%s] intra-server retIndex[%u], interRankSize[%u]", __func__, retIndex, serverSize_);
     302              :     // 收集每个机子的数据进行最后的redeuce
     303            0 :     for (u32 i = 0; i < serverSize_; ++i) {
     304            0 :         if (i == serverId_) {
     305            0 :             isReduceBlock[i] = true;
     306            0 :             reduceMem[i] = std::move(usrOutInraMem);
     307            0 :             HCCL_DEBUG("[%s] inter-server final local reduce, NO.%u reduceMem stores userOut", __func__, i);
     308            0 :             continue;
     309              :         }
     310              :         // 所有local reduce数据为第serverId大块的第0块
     311            0 :         isReduceBlock[i] = true;
     312            0 :         u32 cclbufferIdx = GetRankIdx(i, 0);
     313            0 :         DeviceMem cclbufferIntraMem;
     314            0 :         CHK_RET(GetLocalCclbufferDeviceMem(cclbufferIdx, cclbufferIntraMem, 0));
     315            0 :         reduceMem[i] = std::move(cclbufferIntraMem);
     316            0 :         HCCL_DEBUG(
     317              :             "[%s] inter-server final local reduce, NO.%u reduceMem stores cclbufferIdx[%u]", __func__, i, cclbufferIdx);
     318            0 :     }
     319            0 :     CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, true));
     320            0 :     HCCL_INFO("[%s] intra-server run final local reduce success", __func__);
     321            0 :     return HCCL_SUCCESS;
     322            0 : }
     323              : 
     324            0 : HcclResult ReduceScatterMultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase)
     325              : {
     326            0 :     if (isStartPhase) {
     327            0 :         CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     328            0 :         CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     329            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
     330              :     } else {
     331            0 :         CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     332            0 :         CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     333            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
     334              :     }
     335            0 :     return HCCL_SUCCESS;
     336              : }
     337              : 
     338            0 : HcclResult ReduceScatterMultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase)
     339              : {
     340            0 :     if (isStartPhase) {
     341            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[reduceMainStreamIdx_], -1));
     342            0 :         CHK_RET(LocalNotify::Wait(
     343              :             subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[reduceMainStreamIdx_],
     344              :             INVALID_VALUE_STAGE));
     345            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
     346              :     } else {
     347            0 :         CHK_RET(LocalNotify::Post(
     348              :             subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], -1));
     349            0 :         CHK_RET(
     350              :             LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], INVALID_VALUE_STAGE));
     351            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
     352              :     }
     353            0 :     return HCCL_SUCCESS;
     354              : }
     355              : 
     356              : // 每个server内首先要进行alltoall full mesh收集数据,再进行机内local reduce,最后发送给指定server
     357            0 : HcclResult ReduceScatterMultiDeterPipeline::RunAsync()
     358              : {
     359            0 :     CHK_RET(RunAsyncReduceScatterPipeline());
     360            0 :     return HCCL_SUCCESS;
     361              : }
     362              : 
     363              : // 适配新CollExecutor接口
     364            0 : HcclResult ReduceScatterMultiDeterPipeline::Prepare(
     365              :     HcomCollOpInfo* opInfo, DeviceMem& cclBuffer, const u64 count, const u64 offset, const std::vector<Slice>& slices,
     366              :     const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo, Stream& mainStream,
     367              :     std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
     368              :     std::vector<std::shared_ptr<LocalNotify>>& notifySub)
     369              : {
     370              :     // stream
     371            0 :     subStreams_ = subStream;
     372            0 :     mainStream_ = mainStream;
     373            0 :     subStreamNum_ = subStreams_.size();
     374            0 :     CHK_RET(PrepareTopoInfo(level0CommInfo, level1CommInfo));
     375            0 :     all2allStreamBegin_ = 0;
     376            0 :     all2allStreamSize_ = intraRankSize_ - 1; // alltoall 从流只需要 intraRankSize_ - 1条
     377            0 :     reduceStreamBegin_ = intraRankSize_ - 1;
     378            0 :     reduceMainStreamIdx_ = intraRankSize_ - 1;
     379            0 :     reduceStreamSize_ = MAX_REDUCE_STREAM_NUM; // reduce 从流只需要MAX_REDUCE_STREAM_NUM条
     380            0 :     HCCL_INFO(
     381              :         "[%s] stream: all2allStreamBegin[%u], size[%u], reduceStreamBegin[%u], reduceMainStreamIdx[%u], size[%u]",
     382              :         __func__, all2allStreamBegin_, all2allStreamSize_, reduceStreamBegin_, reduceMainStreamIdx_, reduceStreamSize_);
     383              : 
     384              :     // opInfo
     385            0 :     opInfo_ = opInfo;
     386            0 :     reductionOp_ = opInfo_->reduceOp;
     387            0 :     usrInMemPtr_ = opInfo_->inputAddr;
     388            0 :     usrOutMemPtr_ = opInfo_->outputAddr;
     389            0 :     dataType_ = opInfo_->dataType;
     390            0 :     unitSize_ = SIZE_TABLE[opInfo_->dataType];
     391            0 :     memSliceSize_ = opInfo_->count * unitSize_; // 一整块rank的内存大小
     392              : 
     393              :     // streamNotify, size: n
     394            0 :     streamNotifySub_ = notifySub;
     395            0 :     if (streamNotifySub_.size() < intraRankSize_) {
     396            0 :         HCCL_ERROR(
     397              :             "[%s] rank[%u] streamNotifySub_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
     398              :             streamNotifySub_.size(), intraRankSize_);
     399            0 :         return HCCL_E_INTERNAL;
     400              :     }
     401            0 :     streamNotifyMain_ = notifyMain;
     402            0 :     if (streamNotifyMain_.size() < intraRankSize_) {
     403            0 :         HCCL_ERROR(
     404              :             "[%s] rank[%u] streamNotifyMain_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
     405              :             streamNotifyMain_.size(), intraRankSize_);
     406            0 :         return HCCL_E_INTERNAL;
     407              :     }
     408            0 :     HCCL_INFO(
     409              :         "[%s] notify: streamNum[%u], streamNotifySubNum[%u], streamNotifyMainNum[%u]", __func__, subStreams_.size(),
     410              :         streamNotifySub_.size(), streamNotifyMain_.size());
     411              : 
     412              :     // 此次reduce scatter数据信息
     413            0 :     cclBuffer_ = cclBuffer;
     414            0 :     count_ = count;
     415            0 :     curSize_ = count_ * unitSize_;
     416            0 :     bufferSize_ = cclBuffer.size();
     417            0 :     offset_ = offset;
     418            0 :     slices_ = slices;
     419            0 :     eachRankCclbufferSize_ = curSize_ + HCCL_MIN_SLICE_ALIGN_910B;
     420            0 :     if (slices_.size() != userRankSize_) {
     421            0 :         HCCL_ERROR("[%s] slices size[%llu] not match userRankSize[%u]", __func__, slices_.size(), userRankSize_);
     422            0 :         return HCCL_E_INTERNAL;
     423              :     }
     424            0 :     HCCL_INFO(
     425              :         "[%s] this time: bufferSize[%u], count[%u], curSize[%u], offset[%u], slicesNum[%u]", __func__, bufferSize_,
     426              :         count_, curSize_, offset_, slices_.size());
     427            0 :     return HCCL_SUCCESS;
     428              : }
     429              : 
     430            0 : u64 ReduceScatterMultiDeterPipeline::GetLocalReduceSerialThresh() { return curSize_; }
     431              : 
     432              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_MULTI_DETERMINISTIC_PIPELINE, ReduceScatterMultiDeterPipeline);
     433              : } // namespace hccl
        

Generated by: LCOV version 2.0-1