LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_reduce - all_reduce_multi_deter_pipeline.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 378 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 24 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 "all_reduce_multi_deter_pipeline.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : AllReduceMultiDeterPipeline::AllReduceMultiDeterPipeline(const HcclDispatcher dispatcher)
      16            0 :     : MultiDeterPipeline(dispatcher)
      17            0 : {}
      18              : 
      19            0 : AllReduceMultiDeterPipeline::~AllReduceMultiDeterPipeline() {}
      20              : 
      21            0 : HcclResult AllReduceMultiDeterPipeline::GetRemoteCclbufferDeviceMem(
      22              :     u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
      23              : {
      24            0 :     void* remoteMemPtr = nullptr;
      25            0 :     CHK_RET(link->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
      26            0 :     u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
      27            0 :     u64 size = slices_[inputSliceIndex].size;
      28            0 :     u64 offset = slices_[outputSliceIndex].offset;
      29            0 :     u8* intraSrcAddr = beginAddrU8 + offset;
      30            0 :     remoteMem = DeviceMem::create(intraSrcAddr, size);
      31            0 :     if (remoteMem.ptr() == nullptr) {
      32            0 :         HCCL_ERROR(
      33              :             "[%s] offset + size = [%llu] > cclBufferSize[%llu] > cclBufferSize", __func__, offset + size,
      34              :             outCclBuffer_.size());
      35            0 :         return HCCL_E_MEMORY;
      36              :     }
      37            0 :     HCCL_DEBUG(
      38              :         "[%s] rank[%u], beginAddr[%p], offset[%llu](slices_[outputSliceIndex].offset), "
      39              :         "curSize[%llu], totalBufferSize[%llu]",
      40              :         __func__, outputSliceIndex, remoteMem.ptr(), offset, size, outCclBuffer_.size());
      41            0 :     return HCCL_SUCCESS;
      42              : }
      43              : 
      44              : HcclResult
      45            0 : AllReduceMultiDeterPipeline::GetLocalInCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
      46              : {
      47            0 :     u64 size = ifUseLastSize ? lastSize_ : slices_[rankIdInAllRanks].size;
      48            0 :     u64 offset = slices_[rankIdInAllRanks].offset;
      49            0 :     localMem = inCclBuffer_.range(offset, size);
      50            0 :     if (localMem.ptr() == nullptr) {
      51            0 :         HCCL_ERROR(
      52              :             "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
      53              :             inCclBuffer_.size());
      54            0 :         return HCCL_E_MEMORY;
      55              :     }
      56            0 :     HCCL_DEBUG(
      57              :         "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu]", __func__, rankIdInAllRanks,
      58              :         localMem.ptr(), offset, size, inCclBuffer_.size());
      59            0 :     return HCCL_SUCCESS;
      60              : }
      61              : 
      62              : // reduce scatter RDMA、SDMA 发送后 取本地 sliceOffset = slices_[rankIdInAllRanks].offset偏移处存放地址
      63              : // 什么时候取小内存额外进行判断
      64              : // allgather 都是发到相同内存块,所以不需要额外判断是否为小内存
      65            0 : HcclResult AllReduceMultiDeterPipeline::GetLocalOutCclbufferDeviceMem(
      66              :     u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
      67              : {
      68            0 :     u64 size = ifUseLastSize ? lastSize_ : slices_[rankIdInAllRanks].size;
      69            0 :     u64 offset = slices_[rankIdInAllRanks].offset;
      70            0 :     localMem = outCclBuffer_.range(offset, size);
      71            0 :     if (localMem.ptr() == nullptr) {
      72            0 :         HCCL_ERROR(
      73              :             "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
      74              :             outCclBuffer_.size());
      75            0 :         return HCCL_E_MEMORY;
      76              :     }
      77            0 :     HCCL_DEBUG(
      78              :         "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu]", __func__, rankIdInAllRanks,
      79              :         localMem.ptr(), offset, size, outCclBuffer_.size());
      80            0 :     return HCCL_SUCCESS;
      81              : }
      82              : 
      83            0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool isUserIn)
      84              : {
      85            0 :     u8* beginAddrU8 = isUserIn ? static_cast<u8*>(usrInMemPtr_) : static_cast<u8*>(usrOutMemPtr_);
      86            0 :     u64 offset = slices_[rankIdInAllRanks].offset;
      87            0 :     u64 size = slices_[rankIdInAllRanks].size;
      88            0 :     u8* intraSrcAddr = beginAddrU8 + offset; // 不用 + offset_,因为usrInMem_已经加过了
      89            0 :     localMem = DeviceMem::create(intraSrcAddr, size);
      90            0 :     if (localMem.ptr() == nullptr) {
      91            0 :         HCCL_ERROR(
      92              :             "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
      93              :             outCclBuffer_.size());
      94            0 :         return HCCL_E_MEMORY;
      95              :     }
      96            0 :     HCCL_DEBUG(
      97              :         "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu] isUserIn[%u]", __func__,
      98              :         rankIdInAllRanks, localMem.ptr(), offset, size, outCclBuffer_.size(), isUserIn);
      99            0 :     return HCCL_SUCCESS;
     100              : }
     101              : 
     102            0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
     103              : {
     104            0 :     CHK_RET(GetLocalUserDeviceMem(rankIdInAllRanks, localMem, true));
     105            0 :     return HCCL_SUCCESS;
     106              : }
     107              : 
     108            0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserOutDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
     109              : {
     110            0 :     CHK_RET(GetLocalUserDeviceMem(rankIdInAllRanks, localMem, false));
     111            0 :     return HCCL_SUCCESS;
     112              : }
     113              : 
     114            0 : HcclResult AllReduceMultiDeterPipeline::RunLocalCopy()
     115              : {
     116            0 :     if (intraRankId_ != intraRankSize_ - 1) {
     117            0 :         HCCL_DEBUG("[%s] intra-card no need to copy userRank[%u], intraRankId_[%u]", __func__, userRank_, intraRankId_);
     118            0 :         return HCCL_SUCCESS;
     119              :     }
     120              : 
     121            0 :     DeviceMem userIn;
     122            0 :     CHK_RET(GetLocalUserInDeviceMem(userRank_, userIn));
     123            0 :     DeviceMem cclbuffer;
     124            0 :     CHK_RET(GetLocalOutCclbufferDeviceMem(userRank_, cclbuffer, false));
     125              :     // 使用主流搬迁卡内数据
     126            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, cclbuffer, userIn, mainStream_));
     127            0 :     HCCL_DEBUG(
     128              :         "[%s] intra-card copy data from userInMem to number[%u] cclbuffer[%p] size[%llu]", __func__, userRank_,
     129              :         cclbuffer.ptr(), curSize_);
     130            0 :     return HCCL_SUCCESS;
     131            0 : }
     132              : 
     133              : // 机内alltoall full mesh收集数据, #step表示pairwise的第step步
     134            0 : HcclResult AllReduceMultiDeterPipeline::RunIntraAlltoallPreSync(u32 step)
     135              : {
     136            0 :     HCCL_DEBUG("[%s] intra-server alltoall begin, step[%u]", __func__, step);
     137              :     // alltoall需要准备跨机要的reduce数据 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块
     138              :     // 输入是input的第(R1+1)%S1组数据(往后1)
     139              :     // 每个rank机内只需拷贝intraRankSize_ - 1次
     140            0 :     HCCL_DEBUG(
     141              :         "[%s] intra-server SDMA send begin, [serverId, intraRankId] = [%u, %u]", __func__, serverId_, intraRankId_);
     142            0 :     for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
     143            0 :         HCCL_DEBUG(
     144              :             "[%s] intra-server SDMA send begin, userRank[%u] step[%u] pro[%u/%u]", __func__, userRank_, i, i + 1,
     145              :             intraRankSize_ - 1);
     146              :         // 从机内rankId为recvIntraRankId收集数据,也发给机内rankId为sendIntraRankId数据
     147            0 :         u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
     148            0 :         LINK sendIntraLink = intraLinks_[sendIntraRankId];
     149            0 :         CHK_RET(sendIntraLink->TxAck(subStreams_[i]));
     150            0 :         CHK_RET(sendIntraLink->RxAck(subStreams_[i]));
     151            0 :     }
     152              :     // 增加主从流同步,目的是让SDMA同时进行
     153            0 :     CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     154            0 :     CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     155            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
     156            0 :     CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     157            0 :     CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     158            0 :     return HCCL_SUCCESS;
     159              : }
     160              : 
     161            0 : HcclResult AllReduceMultiDeterPipeline::BatchPostNotifyForStreams(
     162              :     const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
     163              : {
     164            0 :     if (useMainStream) {
     165            0 :         HCCL_DEBUG("[%s] use mainStrem, skip notify wait", __func__);
     166            0 :         return HCCL_SUCCESS;
     167              :     }
     168            0 :     for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
     169            0 :         if (streamTasks[s].empty())
     170            0 :             continue; // 无任务的流跳过
     171            0 :         u32 streamIdx = reduceStreamBegin_ + s;
     172            0 :         if (reduceMainStreamIdx_ == streamIdx) {
     173            0 :             continue;
     174              :         }
     175            0 :         if (isStartPhase) {
     176            0 :             CHK_RET(AlgTemplateBase::ExecEmptyTask(
     177              :                 inCclBuffer_, outCclBuffer_, subStreams_[reduceMainStreamIdx_], dispatcher_));
     178            0 :             CHK_RET(LocalNotify::Post(
     179              :                 subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
     180            0 :             CHK_RET(LocalNotify::Wait(
     181              :                 subStreams_[streamIdx], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
     182            0 :             HCCL_DEBUG("[%s] stream[%u] start phase notify done", __func__, streamIdx);
     183              :         } else {
     184            0 :             CHK_RET(LocalNotify::Post(
     185              :                 subStreams_[streamIdx], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
     186            0 :             CHK_RET(LocalNotify::Wait(
     187              :                 subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
     188            0 :             CHK_RET(AlgTemplateBase::ExecEmptyTask(
     189              :                 inCclBuffer_, outCclBuffer_, subStreams_[reduceMainStreamIdx_], dispatcher_));
     190            0 :             HCCL_DEBUG("[%s] stream[%u] sync phase notify done", __func__, streamIdx);
     191              :         }
     192              :     }
     193            0 :     return HCCL_SUCCESS;
     194              : }
     195              : 
     196            0 : bool AllReduceMultiDeterPipeline::IfUseLastSize(u32 step, u32 sendServerId)
     197              : {
     198              :     // 第0步的最后一块rank,使用小块内存
     199            0 :     if (step == 0 && userRank_ == userRankSize_ - 1) {
     200            0 :         return true;
     201              :     }
     202              :     // 其他步骤,接收数据的rank是机内最后一个且allreduce后的结果是发给最后一个server
     203            0 :     if (step != 0 && (sendServerId == serverSize_ - 1) && (intraRankId_ == intraRankSize_ - 1)) {
     204            0 :         return true;
     205              :     }
     206            0 :     return false;
     207              : }
     208              : 
     209              : // 机内localreduce首先按序收集所有内存块,接着二分归并reduce,最多使用4条流并行
     210            0 : HcclResult AllReduceMultiDeterPipeline::RunIntraLocalReduce(u32 step)
     211              : {
     212            0 :     HCCL_DEBUG("[%s] inter-server local reduce begin, step[%u]", __func__, step);
     213            0 :     u32 recvServerId = GetPreServerIdByStep(step);  // 从上一个收
     214            0 :     u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个
     215            0 :     std::vector<DeviceMem> reduceMem;
     216            0 :     std::vector<bool> isReduceBlock;
     217            0 :     isReduceBlock.resize(intraRankSize_);
     218            0 :     reduceMem.resize(intraRankSize_);
     219            0 :     u32 retIndex = 0;
     220              :     // 机内,最后rank的规约结果放在倒数第2块,其他放在倒数第1块
     221              :     // 机内第0步,规约结果放在第userRank块cclbuffer
     222            0 :     if (intraRankId_ == intraRankSize_ - 1) {
     223            0 :         retIndex = intraRankSize_ - SECOND_TO_LAST;
     224              :     } else {
     225            0 :         retIndex = intraRankSize_ - 1;
     226              :     }
     227            0 :     bool ifUseLastSize = IfUseLastSize(step, sendServerId);
     228            0 :     u32 userInIdx = GetRankIdx(sendServerId, intraRankId_);
     229              :     // 最后一块留给allreduce,idx为第sendServerId大块内存的第idx小块
     230            0 :     u32 idx = 0;
     231            0 :     for (u32 i = 0; i < intraRankSize_; ++i) {
     232            0 :         u32 outCCLbufferIdx = 0;
     233              :         // i == intraRankId_时,需要取userIn内存,
     234            0 :         if (i == intraRankId_) {
     235              :             // 第0步,机内最后一个rank取cclbuffer,因为localcopy时将该内存搬到了第userRank_块cclbuffer
     236            0 :             if (step == 0 && i == intraRankSize_ - 1) {
     237            0 :                 isReduceBlock[i] = true;
     238            0 :                 outCCLbufferIdx = userRank_;
     239            0 :                 DeviceMem cclbufferIntraMem;
     240            0 :                 CHK_RET(GetLocalOutCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, ifUseLastSize));
     241            0 :                 reduceMem[i] = std::move(cclbufferIntraMem);
     242            0 :                 HCCL_DEBUG(
     243              :                     "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i,
     244              :                     outCCLbufferIdx);
     245            0 :                 retIndex = i;
     246            0 :             } else {
     247            0 :                 isReduceBlock[i] = false;
     248            0 :                 DeviceMem usrInIntraMem;
     249            0 :                 CHK_RET(GetLocalUserInDeviceMem(userInIdx, usrInIntraMem));
     250            0 :                 reduceMem[i] = std::move(usrInIntraMem);
     251            0 :                 HCCL_DEBUG(
     252              :                     "[%s] inter-server local reduce, NO.%u reduceMem stores userInIdx[%u]", __func__, i, userInIdx);
     253            0 :             }
     254            0 :             continue;
     255            0 :         }
     256              :         // 其他情况:统一处理CCLBuffer,
     257            0 :         isReduceBlock[i] = true;
     258            0 :         outCCLbufferIdx = GetRankIdx(recvServerId, idx);
     259            0 :         DeviceMem cclbufferIntraMem;
     260            0 :         CHK_RET(GetLocalOutCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, ifUseLastSize));
     261            0 :         reduceMem[i] = std::move(cclbufferIntraMem);
     262            0 :         HCCL_DEBUG(
     263              :             "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i, outCCLbufferIdx);
     264            0 :         idx++;
     265              :         // step 0, localreduce到userRank_块内存上
     266            0 :         if (step == 0 && outCCLbufferIdx == userRank_) {
     267            0 :             retIndex = i;
     268              :         }
     269            0 :     }
     270            0 :     HCCL_DEBUG(
     271              :         "[%s] intra-server local reduce, retIndex[%u], intraRankId[%u], intraRankSize[%u]", __func__, retIndex,
     272              :         intraRankId_, intraRankSize_);
     273            0 :     CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, false));
     274            0 :     HCCL_INFO("[%s] intra-server step[%u] run local reduce success", __func__, step);
     275            0 :     return HCCL_SUCCESS;
     276            0 : }
     277              : 
     278            0 : HcclResult AllReduceMultiDeterPipeline::RunInterSend(u32 step)
     279              : {
     280            0 :     HCCL_DEBUG("[%s] inter-server RDMA write begin, step[%u]", __func__, step);
     281              :     // 使用主流进行rdma
     282            0 :     u32 recvServerId = GetPreServerIdByStep(step);  // 从上一个收
     283            0 :     u32 sendServerId = GetNextServerIdByStep(step); // 发给下一个
     284            0 :     LINK recvInterLink = serverLinks_[recvServerId];
     285            0 :     LINK sendInterLink = serverLinks_[sendServerId];
     286              :     // 跨机 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块的第1块
     287              :     // 跨机 输入是input的第(R1+1)%S1组数据(往后1)的第2块
     288            0 :     u32 reduceRetIndex = intraRankSize_ - SECOND_TO_LAST;
     289            0 :     u32 recvCCLbufferIdx = GetRankIdx(recvServerId, 0);
     290            0 :     u32 sendCCLbufferIdx = GetRankIdx(recvServerId, reduceRetIndex);
     291              :     // 跨机写对端
     292            0 :     u32 sendRankId = GetRankIdx(sendServerId, intraRankId_); // 发送到sendRankId
     293            0 :     u32 recvRankId = GetRankIdx(recvServerId, intraRankId_); // 从recvRankId接收
     294            0 :     u32 remoteRecvCCLbufferIdx = GetRankIdx(serverId_, 0);   // 接收端:收的位置
     295              :     // 2机又从serverId_收也从serverId_发
     296            0 :     u32 remoteSendCCLbufferIdx = serverSize_ == MIN_SERVER_NUM ?
     297            0 :                                      GetRankIdx(serverId_, reduceRetIndex) :
     298            0 :                                      GetRankIdx(sendServerId, reduceRetIndex); // 发送端:发的位置
     299            0 :     DeviceMem recvMem;
     300            0 :     DeviceMem sendMem;
     301            0 :     bool ifSendToLastServer = IfUseLastSize(step, sendServerId);
     302              :     // 最后一个rank则只收小块内存
     303            0 :     CHK_RET(GetLocalOutCclbufferDeviceMem(recvCCLbufferIdx, recvMem, isLastRank_));
     304              :     // 如果是发给最后一个rank则只发小块内存
     305            0 :     CHK_RET(GetLocalOutCclbufferDeviceMem(sendCCLbufferIdx, sendMem, ifSendToLastServer));
     306              : 
     307            0 :     HCCL_DEBUG(
     308              :         "[%s] inter-server RDMA write begin, cclbufferIdx: [%u] send to [%u], [%u] recv from [%u]", __func__,
     309              :         sendCCLbufferIdx, remoteRecvCCLbufferIdx, recvCCLbufferIdx, remoteSendCCLbufferIdx);
     310            0 :     HCCL_DEBUG(
     311              :         "[%s] inter-server RDMA write begin, rankId: [%u] send to [%u], [%u] recv from [%u]", __func__, userRank_,
     312              :         sendRankId, userRank_, recvRankId);
     313            0 :     HCCL_DEBUG(
     314              :         "[%s] inter-server RDMA write begin, if use last small mem? : isLastRank[%u], ifSendToLastServer[%u]", __func__,
     315              :         isLastRank_, ifSendToLastServer);
     316            0 :     if (recvInterLink->IsSpInlineReduce() && sendInterLink->IsSpInlineReduce()) {
     317            0 :         CHK_RET(sendInterLink->TxAck(mainStream_));
     318            0 :         CHK_RET(recvInterLink->RxAck(mainStream_));
     319            0 :         DeviceMem dstMem = std::move(recvMem);
     320            0 :         void* remoteMemPtr = nullptr;
     321            0 :         CHK_RET(recvInterLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
     322            0 :         u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
     323            0 :         u64 size = slices_[recvCCLbufferIdx].size;
     324            0 :         u64 offset = slices_[remoteSendCCLbufferIdx].offset;
     325            0 :         u8* intraSrcAddr = beginAddrU8 + offset;
     326            0 :         DeviceMem srcMem = DeviceMem::create(intraSrcAddr, isLastRank_ ? lastSize_ : size);
     327            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_, recvRankId, recvInterLink->GetLinkType()));
     328            0 :         CHK_RET(sendInterLink->TxDataSignal(mainStream_));
     329            0 :         CHK_RET(recvInterLink->RxDataSignal(mainStream_));
     330            0 :     } else {
     331            0 :         CHK_RET(recvInterLink->TxAck(mainStream_));
     332            0 :         CHK_RET(sendInterLink->RxAck(mainStream_));
     333              : 
     334            0 :         u64 size = slices_[sendCCLbufferIdx].size; // 没有特殊情况发送接收内存大小都是一样
     335              :         // 发送的size取本地发的数据大小,如果是发给最后一个rank则只发小块内存
     336            0 :         CHK_RET(sendInterLink->TxAsync(
     337              :             UserMemType::OUTPUT_MEM, slices_[remoteRecvCCLbufferIdx].offset, sendMem.ptr(),
     338              :             ifSendToLastServer ? lastSize_ : size, mainStream_));
     339              :         // 接收的size取远端发的数据大小, 如果是最后一个rank则只收小块内存
     340            0 :         CHK_RET(recvInterLink->RxAsync(
     341              :             UserMemType::OUTPUT_MEM, slices_[remoteSendCCLbufferIdx].offset, recvMem.ptr(),
     342              :             isLastRank_ ? lastSize_ : size, mainStream_));
     343            0 :         CHK_RET(recvInterLink->PostFinAck(mainStream_));
     344            0 :         CHK_RET(sendInterLink->WaitFinAck(mainStream_));
     345              :     }
     346            0 :     HCCL_INFO("[%s] inter-server step[%u] run RDMA send success", __func__, step);
     347            0 :     return HCCL_SUCCESS;
     348            0 : }
     349              : 
     350            0 : HcclResult AllReduceMultiDeterPipeline::RunFinalReduce()
     351              : {
     352            0 :     HCCL_DEBUG("[%s] intra-server final reduce begin", __func__);
     353            0 :     std::vector<DeviceMem> reduceMem;
     354            0 :     std::vector<bool> isReduceBlock;
     355            0 :     isReduceBlock.resize(serverSize_);
     356            0 :     reduceMem.resize(serverSize_);
     357              : 
     358            0 :     u32 retIndex = serverId_;
     359              :     // userRank_ == userRankSize_ - 1时取小内存进行最后一次reduce
     360            0 :     bool ifUseLastSize = isLastRank_;
     361            0 :     HCCL_DEBUG(
     362              :         "[%s] intra-server retIndex[%u], interRankSize[%u], ifUseLastSize[%u]", __func__, retIndex, serverSize_,
     363              :         ifUseLastSize);
     364              :     // 收集每个机子的数据进行最后的reduce
     365            0 :     for (u32 i = 0; i < serverSize_; ++i) {
     366            0 :         DeviceMem cclbufferIntraMem;
     367            0 :         u32 cclbufferIdx = 0;
     368            0 :         if (i == serverId_) {
     369            0 :             cclbufferIdx = userRank_;
     370              :         } else {
     371            0 :             cclbufferIdx = GetRankIdx(i, 0);
     372              :         }
     373              :         // 所有local reduce数据为第serverId大块的第0块
     374            0 :         isReduceBlock[i] = true;
     375            0 :         CHK_RET(GetLocalOutCclbufferDeviceMem(cclbufferIdx, cclbufferIntraMem, ifUseLastSize));
     376            0 :         reduceMem[i] = std::move(cclbufferIntraMem);
     377            0 :         HCCL_DEBUG(
     378              :             "[%s] inter-server final local reduce, NO.%u reduceMem stores cclbufferIdx[%u]", __func__, i, cclbufferIdx);
     379            0 :     }
     380            0 :     CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, true)); // final reduce使用主流进行操作
     381            0 :     HCCL_INFO("[%s] intra-server run final local reduce success", __func__);
     382            0 :     return HCCL_SUCCESS;
     383            0 : }
     384              : 
     385            0 : HcclResult AllReduceMultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase)
     386              : {
     387            0 :     if (isStartPhase) {
     388            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
     389            0 :         CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     390            0 :         CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     391            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
     392              :     } else {
     393            0 :         CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     394            0 :         CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
     395            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
     396            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
     397              :     }
     398            0 :     return HCCL_SUCCESS;
     399              : }
     400              : 
     401            0 : HcclResult AllReduceMultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase)
     402              : {
     403            0 :     if (isStartPhase) {
     404            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
     405            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[reduceMainStreamIdx_], -1));
     406            0 :         CHK_RET(LocalNotify::Wait(
     407              :             subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[reduceMainStreamIdx_],
     408              :             INVALID_VALUE_STAGE));
     409            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
     410              :     } else {
     411            0 :         CHK_RET(LocalNotify::Post(
     412              :             subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], -1));
     413            0 :         CHK_RET(
     414              :             LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], INVALID_VALUE_STAGE));
     415            0 :         CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
     416            0 :         HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
     417              :     }
     418            0 :     return HCCL_SUCCESS;
     419              : }
     420              : 
     421              : HcclResult
     422            0 : AllReduceMultiDeterPipeline::RunAllGatherInterServer(u32 step, const LINK& prevInterLink, const LINK& nextInterLink)
     423              : {
     424            0 :     HCCL_INFO(
     425              :         "[%s] inter-server allgather run, userRank[%u], step[%u/%u]", __func__, userRank_, step,
     426              :         serverSize_ - STEP_OFFSET_TWO);
     427            0 :     CHK_RET(prevInterLink->TxAck(mainStream_));
     428            0 :     CHK_RET(nextInterLink->RxAck(mainStream_));
     429            0 :     u32 rxDMAMemSliceId = (serverSize_ + step) % PARITY_BASE;
     430            0 :     u32 txDMAMemSliceId = (serverSize_ + step - 1) % PARITY_BASE;
     431            0 :     UserMemType srcMemType = txDMAMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM;
     432            0 :     UserMemType dstMemType = rxDMAMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM;
     433            0 :     u32 txSliceId = ((serverId_ + step) % serverSize_) * intraRankSize_ + intraRankId_;
     434            0 :     u32 txDataSize = slices_[txSliceId].size;
     435            0 :     DeviceMem txlocalMem;
     436            0 :     if (txDMAMemSliceId == serverSizeParity_) {
     437            0 :         CHK_RET(GetLocalOutCclbufferDeviceMem(txSliceId, txlocalMem, false));
     438              :     } else {
     439            0 :         CHK_RET(GetLocalInCclbufferDeviceMem(txSliceId, txlocalMem, false));
     440              :     }
     441            0 :     CHK_RET(nextInterLink->TxAsync(dstMemType, slices_[txSliceId].offset, txlocalMem.ptr(), txDataSize, mainStream_));
     442              : 
     443            0 :     u32 rxSliceId = ((serverId_ + step + 1) % serverSize_) * intraRankSize_ + intraRankId_;
     444            0 :     DeviceMem rxLocalMem;
     445            0 :     if (rxDMAMemSliceId == serverSizeParity_) {
     446            0 :         CHK_RET(GetLocalOutCclbufferDeviceMem(rxSliceId, rxLocalMem, false));
     447              :     } else {
     448            0 :         CHK_RET(GetLocalInCclbufferDeviceMem(rxSliceId, rxLocalMem, false));
     449              :     }
     450            0 :     u64 rxDataSize = slices_[rxSliceId].size;
     451            0 :     CHK_RET(prevInterLink->RxAsync(srcMemType, slices_[rxSliceId].offset, rxLocalMem.ptr(), rxDataSize, mainStream_));
     452            0 :     HCCL_DEBUG(
     453              :         "[%s] step[%u], txId[%u], rxId[%u], srcMemType[%u], dstMemType[%u]", __func__, step, txDMAMemSliceId,
     454              :         rxDMAMemSliceId, srcMemType, dstMemType);
     455            0 :     HCCL_DEBUG(
     456              :         "[%s] txlocalMem: ptr[%p], size[%llu], rxLocalMem: ptr[%p], size[%llu]", __func__, txlocalMem.ptr(),
     457              :         txlocalMem.size(), rxLocalMem.ptr(), rxLocalMem.size());
     458            0 :     HCCL_DEBUG(
     459              :         "[%s] send txlocalMem to txSliceId[%llu], recv rxLocalMem from rxSliceId[%llu]", __func__, txSliceId,
     460              :         rxSliceId);
     461            0 :     HCCL_INFO("[%s] inter-server allgather success", __func__);
     462            0 :     return HCCL_SUCCESS;
     463            0 : }
     464              : 
     465            0 : HcclResult AllReduceMultiDeterPipeline::RunAllGatherIntraServer(u32 step)
     466              : {
     467            0 :     HCCL_INFO("[%s] intra-server allgather run, userRank[%u], step[%u/%u]", __func__, userRank_, step, serverSize_ - 1);
     468            0 :     u32 dmaMemSliceId = (serverSize_ + step - 1) % PARITY_BASE;
     469            0 :     for (u32 i = 1; i < intraRankSize_; i++) {
     470            0 :         u32 remIntraRankId = (intraRankId_ + i) % intraRankSize_;
     471            0 :         CHK_RET(intraLinks_[remIntraRankId]->TxAck(subStreams_[i - 1]));
     472            0 :         CHK_RET(intraLinks_[remIntraRankId]->RxAck(subStreams_[i - 1]));
     473            0 :         void* remoteMemPtr = nullptr;
     474            0 :         CHK_RET(intraLinks_[remIntraRankId]->GetRemoteMem(
     475              :             dmaMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM, &remoteMemPtr));
     476            0 :         u32 remoteCclbufferId = ((serverId_ + step) % serverSize_) * intraRankSize_ + remIntraRankId;
     477              :         DeviceMem src = DeviceMem::create(
     478            0 :             static_cast<u8*>(remoteMemPtr) + slices_[remoteCclbufferId].offset, slices_[remoteCclbufferId].size);
     479              :         DeviceMem dst = DeviceMem::create(
     480            0 :             static_cast<u8*>(usrOutMemPtr_) + slices_[remoteCclbufferId].offset, slices_[remoteCclbufferId].size);
     481            0 :         CHK_RET(HcclD2DMemcpyAsync(
     482              :             dispatcher_, dst, src, subStreams_[i - 1], intraLinks_[remIntraRankId]->GetRemoteRank(),
     483              :             intraLinks_[remIntraRankId]->GetLinkType()));
     484            0 :         CHK_RET(intraLinks_[remIntraRankId]->TxDataSignal(subStreams_[i - 1]));
     485            0 :         CHK_RET(intraLinks_[remIntraRankId]->RxDataSignal(subStreams_[i - 1]));
     486            0 :     }
     487            0 :     HCCL_INFO("[%s] intra-server allgather success", __func__);
     488            0 :     return HCCL_SUCCESS;
     489              : }
     490              : 
     491            0 : HcclResult AllReduceMultiDeterPipeline::RunAsyncAllgatherPipeline()
     492              : {
     493            0 :     HCCL_INFO("[%s] begin, userRank[%u]", __func__, userRank_);
     494              :     //  机间 ring algo 逆时针,从后往前
     495            0 :     u32 prevInterRankId = GetNextServerIdByStep(1);
     496            0 :     u32 nextInterRankId = GetPreServerIdByStep(1);
     497            0 :     LINK prevInterLink = serverLinks_[prevInterRankId];
     498            0 :     LINK nextInterLink = serverLinks_[nextInterRankId];
     499            0 :     for (u32 step = 0; step < serverSize_; step++) {
     500            0 :         HCCL_INFO("[%s] allgather pipeline, userRank[%u], step[%u/%u]", __func__, userRank_, step, serverSize_ - 1);
     501            0 :         CHK_RET(MainRecordSub(0, subStreamNum_));
     502            0 :         CHK_RET(SubWaitMain(0, subStreamNum_));
     503            0 :         if (step < serverSize_ - 1) {
     504            0 :             CHK_RET(RunAllGatherInterServer(step, prevInterLink, nextInterLink));
     505            0 :             CHK_RET(prevInterLink->PostFinAck(mainStream_));
     506            0 :             CHK_RET(nextInterLink->WaitFinAck(mainStream_));
     507              :             // inter的最后一步需要barrier确保数据发完
     508            0 :             if (step == serverSize_ - STEP_OFFSET_TWO) {
     509            0 :                 CHK_RET(ExecuteBarrier(prevInterLink, nextInterLink, mainStream_));
     510              :             }
     511              :         }
     512            0 :         CHK_RET(RunAllGatherIntraServer(step));
     513            0 :         CHK_RET(SubRecordMain(0, subStreamNum_));
     514            0 :         CHK_RET(MainWaitSub(0, subStreamNum_));
     515            0 :         u32 cclbufferFlag = (serverSize_ + step - 1) % PARITY_BASE;
     516            0 :         u32 sliceId = ((serverId_ + step) % serverSize_) * intraRankSize_ + intraRankId_;
     517            0 :         DeviceMem srcMem;
     518            0 :         if (cclbufferFlag == serverSizeParity_) {
     519            0 :             CHK_RET(GetLocalOutCclbufferDeviceMem(sliceId, srcMem, false));
     520              :         } else {
     521            0 :             CHK_RET(GetLocalInCclbufferDeviceMem(sliceId, srcMem, false));
     522              :         }
     523            0 :         DeviceMem dstMem;
     524            0 :         CHK_RET(GetLocalUserOutDeviceMem(sliceId, dstMem));
     525            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
     526            0 :         HCCL_DEBUG(
     527              :             "[%s] step[%u], cclbufferFlag[%u], sliceId[%u], cclBufferSrcMem: ptr[%p], size[%llu]", __func__, step,
     528              :             cclbufferFlag, sliceId, srcMem.ptr(), srcMem.size());
     529            0 :     }
     530            0 :     HCCL_INFO("[%s] end, userRank[%u]", __func__, userRank_);
     531            0 :     return HCCL_SUCCESS;
     532            0 : }
     533              : 
     534              : // 实现为确定性reduce scatter pipeline + all gather pipeline
     535            0 : HcclResult AllReduceMultiDeterPipeline::RunAsync()
     536              : {
     537            0 :     HCCL_INFO(
     538              :         "[AllReduceMultiDeterPipeline] run begin: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] "
     539              :         "cclBuffer[%p].",
     540              :         userRank_, userRankSize_, usrInMemPtr_, usrOutMemPtr_, outCclBuffer_.ptr());
     541            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     542            0 :     CHK_RET(RunAsyncReduceScatterPipeline());
     543            0 :     CHK_RET(RunAsyncAllgatherPipeline());
     544            0 :     HCCL_INFO("[AllReduceMultiDeterPipeline] AllReduceMultiDeterPipeline success userRank[%u] ", userRank_);
     545            0 :     return HCCL_SUCCESS;
     546              : }
     547              : 
     548              : // 适配新CollExecutor接口
     549            0 : HcclResult AllReduceMultiDeterPipeline::Prepare(
     550              :     HcomCollOpInfo* opInfo, DeviceMem& inBuffer, DeviceMem& outBuffer, const u64 count,
     551              :     const std::vector<Slice>& slices, const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo,
     552              :     Stream& mainStream, std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
     553              :     std::vector<std::shared_ptr<LocalNotify>>& notifySub)
     554              : {
     555              :     // opInfo
     556            0 :     opInfo_ = opInfo;
     557            0 :     dataType_ = opInfo_->dataType;
     558            0 :     unitSize_ = SIZE_TABLE[opInfo_->dataType];
     559            0 :     memSliceSize_ = opInfo_->count * unitSize_; // 一整块rank的内存大小
     560            0 :     usrInMemPtr_ = opInfo_->inputAddr;
     561            0 :     usrOutMemPtr_ = opInfo_->outputAddr;
     562            0 :     reductionOp_ = opInfo_->reduceOp;
     563              : 
     564              :     // stream
     565            0 :     mainStream_ = mainStream;
     566            0 :     subStreams_ = subStream;
     567            0 :     subStreamNum_ = subStreams_.size();
     568            0 :     CHK_RET(PrepareTopoInfo(level0CommInfo, level1CommInfo));
     569            0 :     all2allStreamBegin_ = 0;
     570            0 :     all2allStreamSize_ = intraRankSize_ - 1; // alltoall 从流只需要 intraRankSize_ - 1条
     571            0 :     reduceMainStreamIdx_ = intraRankSize_ - 1;
     572            0 :     reduceStreamBegin_ = intraRankSize_ - 1;
     573            0 :     reduceStreamSize_ = MAX_REDUCE_STREAM_NUM; // reduce 从流只需要MAX_REDUCE_STREAM_NUM条
     574            0 :     HCCL_INFO(
     575              :         "[%s] stream: all2allStreamBegin[%u], size[%u], reduceStreamBegin[%u], size[%u], reduceMainStreamIdx[%u]",
     576              :         __func__, all2allStreamBegin_, all2allStreamSize_, reduceStreamBegin_, reduceStreamSize_, reduceMainStreamIdx_);
     577              : 
     578              :     // streamNotify, size: n
     579            0 :     streamNotifyMain_ = notifyMain;
     580            0 :     if (streamNotifyMain_.size() < intraRankSize_) {
     581            0 :         HCCL_ERROR(
     582              :             "[%s] rank[%u] streamNotifyMain_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
     583              :             streamNotifyMain_.size(), intraRankSize_);
     584            0 :         return HCCL_E_INTERNAL;
     585              :     }
     586            0 :     streamNotifySub_ = notifySub;
     587            0 :     if (streamNotifySub_.size() < intraRankSize_) {
     588            0 :         HCCL_ERROR(
     589              :             "[%s] rank[%u] streamNotifySub_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
     590              :             streamNotifySub_.size(), intraRankSize_);
     591            0 :         return HCCL_E_INTERNAL;
     592              :     }
     593              : 
     594              :     // 此次reduce scatter数据信息
     595            0 :     inCclBuffer_ = inBuffer;
     596            0 :     outCclBuffer_ = outBuffer;
     597            0 :     bufferSize_ = inBuffer.size();
     598            0 :     slices_ = slices;
     599              :     // allreduce count为此次处理的数据总数
     600            0 :     curSize_ = slices_[userRank_].size;
     601            0 :     count_ = slices_[userRank_].size / unitSize_;
     602            0 :     lastSize_ = slices_[userRankSize_ - 1].size;
     603            0 :     isLastRank_ = (userRank_ == userRankSize_ - 1) ? true : false;
     604              :     // serverSize_是偶数,与正常allreduce pipeline中的allgather pipeline流程一样;若为奇数,则颠倒内存
     605            0 :     serverSizeParity_ = (serverSize_ % PARITY_BASE == 0) ? 1 : 0;
     606            0 :     perRankAvgDataSize_ = count * unitSize_ / userRankSize_;
     607            0 :     if (slices_.size() != userRankSize_) {
     608            0 :         HCCL_ERROR("[%s] slices size[%llu] not match userRankSize[%u]", __func__, slices_.size(), userRankSize_);
     609            0 :         return HCCL_E_INTERNAL;
     610              :     }
     611            0 :     HCCL_INFO(
     612              :         "[%s] this time: bufferSize[%u], count[%u], curSize[%u], lastSize[%u], slicesNum[%u] "
     613              :         "serverSizeParity[%u]",
     614              :         __func__, bufferSize_, count_, curSize_, lastSize_, slices_.size(), serverSizeParity_);
     615            0 :     return HCCL_SUCCESS;
     616              : }
     617              : 
     618            0 : u64 AllReduceMultiDeterPipeline::GetLocalReduceSerialThresh() { return perRankAvgDataSize_; }
     619              : 
     620              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_MULTI_DETERMINISTIC_PIPELINE, AllReduceMultiDeterPipeline);
     621              : } // namespace hccl
        

Generated by: LCOV version 2.0-1