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

Generated by: LCOV version 2.0-1