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

Generated by: LCOV version 2.0-1