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

Generated by: LCOV version 2.0-1