LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_alltoall - alltoall_symmetric_memory.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 249 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 25 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 "alltoall_symmetric_memory.h"
      12              : 
      13              : namespace hccl {
      14            0 : AlltoAllFullMeshSymmetricMemory::AlltoAllFullMeshSymmetricMemory(const HcclDispatcher dispatcher)
      15            0 :     : AlgTemplateBase(dispatcher)
      16              : {
      17            0 : }
      18              : 
      19            0 : AlltoAllFullMeshSymmetricMemory::~AlltoAllFullMeshSymmetricMemory() {}
      20              : 
      21            0 : HcclResult AlltoAllFullMeshSymmetricMemory::GenerateSubStreamInfo(const std::vector<Stream> &subStreams,
      22              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalMainToSub,
      23              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalSubToMain)
      24              : {
      25            0 :     u32 totalSubstreamSize = sdmaConcurrentNum_;
      26            0 :     if (subStreams.size() < totalSubstreamSize || meshSignalMainToSub.size() < totalSubstreamSize ||
      27            0 :         meshSignalSubToMain.size() < totalSubstreamSize) {
      28            0 :         HCCL_ERROR("[AlltoAllFullMeshSymmetricMemory][GenerateSubStreamInfo]subStreamsSize[%zu], meshSignalMainToSubSize[%zu]"\
      29              :             "meshSignalSubToMainSize[%zu] is smaller than totalSubstreamSize[%u]",subStreams.size(),
      30              :             meshSignalMainToSub.size(), meshSignalSubToMain.size(), totalSubstreamSize);
      31            0 :         return HCCL_E_PARA;
      32              :     }
      33            0 :     CHK_PRT_RET(links_.size() < userRankSize_, HCCL_ERROR("[AlltoAllFullMeshSymmetricMemory][GenerateSubStreamInfo]"\
      34              :         "links_.size()[%zu] is smaller than userRankSize_[%u].", links_.size(), userRankSize_),
      35              :         HCCL_E_PARA);
      36            0 :     HCCL_DEBUG("subStreams.size[%zu], meshSignalMainToSub.size[%zu], links_.size[%zu]",
      37              :         subStreams.size(), meshSignalMainToSub.size(), links_.size());
      38            0 :     for (u32 sdmaIndex = 0; sdmaIndex < sdmaConcurrentNum_; sdmaIndex++) {
      39            0 :         sdmaSubStream_.push_back(subStreams[sdmaIndex]);
      40            0 :         sdmaMeshSignalMainToSub_.push_back(meshSignalMainToSub[sdmaIndex]);
      41            0 :         sdmaMeshSignalSubToMain_.push_back(meshSignalSubToMain[sdmaIndex]);
      42              :     }
      43            0 :     return HCCL_SUCCESS;
      44              : }
      45              : 
      46            0 : HcclResult AlltoAllFullMeshSymmetricMemory::Prepare(PrepareData &param)
      47              : {
      48            0 :     mainStream_ = param.stream;
      49            0 :     userRank_ = param.userRank;
      50            0 :     userRankSize_ = param.userRankSize;
      51            0 :     links_ = *param.linksPtr;
      52            0 :     sendRecvInfoPtr_ = param.sendRecvInfoPtr;
      53            0 :     devNumInlocalPod_ = param.devNumInlocalPod;
      54            0 :     rankIdxInPod_ = param.rankIdxInPod;
      55            0 :     opType_ = param.opType;
      56            0 :     algOpContext_ = param.algOpContext;
      57              : 
      58            0 :     podStartRank_ = userRank_ - rankIdxInPod_;
      59            0 :     podEndRank_ = podStartRank_ + devNumInlocalPod_ - 1;
      60            0 :     sdmaConcurrentNum_ = (devNumInlocalPod_ > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
      61            0 :         (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) : (devNumInlocalPod_);
      62              : 
      63            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory]devNumInlocalPod_[%u], userRankSize_[%u] podStartRank_[%u]" \
      64              :         "podEndRank_[%u], sdmaConcurrentNum_[%u]",
      65              :         devNumInlocalPod_, userRankSize_, podStartRank_, podEndRank_, sdmaConcurrentNum_);
      66              : 
      67            0 :     CHK_PRT_RET(userRankSize_ == 0, HCCL_ERROR("[AlltoAllFullMeshSymmetricMemory][Prepare]userRankSize_ is zero."),
      68              :         HCCL_E_PARA);
      69              : 
      70            0 :     userInput_ = param.inputMem;
      71            0 :     userOutput_ = param.outputMem;
      72            0 :     workMode_ = param.workMode;
      73              : 
      74            0 :     CHK_RET(GenerateSubStreamInfo(*param.subStreamsPtr, *param.signalPtr, *param.signalAuxPtr));
      75            0 :     return HCCL_SUCCESS;
      76              : }
      77              : 
      78            0 : std::string AlltoAllFullMeshSymmetricMemory::GetStreamIndexString()
      79              : {
      80            0 :     std::string res = "";
      81            0 :     for (auto& info : subStreamReadInfo_) {
      82            0 :         u32 destRank = info.first;
      83            0 :         u32 streamIndex = destRank % sdmaConcurrentNum_;
      84            0 :         res += std::to_string(streamIndex) + ", ";
      85              :     }
      86            0 :     return res;
      87            0 : }
      88              : 
      89            0 : void AlltoAllFullMeshSymmetricMemory::UpdateCurrRankRecvInfo(u32 roundIdx, u32 side, u32 destRank,
      90              :     ReadDataBlock& readInfo)
      91              : {
      92            0 :     const ZCopySendRecvInfo& sendRecvInfo = *sendRecvInfoPtr_;
      93            0 :     u64 recvLen = sendRecvInfo.localRecvLength[destRank];
      94            0 :     u64 userOutOffset = sendRecvInfo.localRecvOffset[destRank];
      95            0 :     u64 remoteUserInOffset = sendRecvInfo.remoteSendOffset[destRank];
      96            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
      97              :         "recvLen[%lu] remoteUserInOffset[%llu] userOutOffset[%llu]",
      98              :         userRank_, destRank, recvLen, remoteUserInOffset, userOutOffset);
      99            0 :     readInfo = {recvLen, remoteUserInOffset, userOutOffset};
     100            0 : }
     101              : 
     102            0 : void AlltoAllFullMeshSymmetricMemory::UpdateSendRecvInfo(u32 roundIdx,
     103              :     std::unordered_map<u32, ReadDataBlock> &subStreamReadInfo,
     104              :     const std::vector<std::vector<std::pair<u32,u32>>> &partialCommRankSet)
     105              : {
     106            0 :     for (u32 side = 0; side < partialCommRankSet.size(); side++) {
     107            0 :         for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
     108            0 :             u32 readRemoteRank = partialCommRankSet[side][j].first;
     109            0 :             if (readRemoteRank == userRank_) {
     110            0 :                 continue;
     111              :             }
     112              :             ReadDataBlock readInfo;
     113            0 :             UpdateCurrRankRecvInfo(roundIdx, side, readRemoteRank, readInfo);
     114              : 
     115            0 :             subStreamReadInfo[readRemoteRank] = readInfo;
     116              :         }
     117              :     }
     118            0 : }
     119              : 
     120            0 : void AlltoAllFullMeshSymmetricMemory::UpdateRemoteRankSet(u32 roundIdx, u32 groupRankSize)
     121              : {
     122            0 :     if (sdmaConcurrentNum_ == 1) {
     123            0 :         UpdatePartialCommunicationRankSetPairWise(roundIdx, groupRankSize);
     124              :     } else {
     125            0 :         UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
     126              :     }
     127            0 : }
     128              : 
     129            0 : void AlltoAllFullMeshSymmetricMemory::UpdatePartialCommunicationRankSetPairWise(u32 roundIdx, u32 groupRankSize)
     130              : {
     131            0 :     partialCommRankSet_.clear();
     132            0 :     partialCommRankSet_.resize(1);
     133            0 :     for (u32 i = roundIdx * sdmaConcurrentNum_; i < (roundIdx * sdmaConcurrentNum_ + groupRankSize); i++) {
     134            0 :         u32 readRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
     135            0 :         u32 sendRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
     136            0 :         partialCommRankSet_[0].push_back(std::make_pair(readRemoteRank, sendRemoteRank));
     137            0 :         HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSetPairWise] userRank [%u] i[%u]" \
     138              :             "readRemoteRank[%u] writeRemoteRank[%u]", userRank_, i, readRemoteRank, sendRemoteRank);
     139              :     }
     140            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSetPairWise] partialCommRankSet_ size[%zu]",
     141              :         partialCommRankSet_[0].size());
     142            0 : }
     143              : 
     144            0 : void AlltoAllFullMeshSymmetricMemory::UpdatePartialCommunicationRankSet(u32 roundIdx, u32 groupRankSize,
     145              :     std::vector<std::vector<std::pair<u32,u32>>> &partialCommRankSet)
     146              : {
     147            0 :     partialCommRankSet.clear();
     148            0 :     partialCommRankSet.resize(RANK_SET_COMPUTE_CONST + 1);
     149            0 :     u32 pairNumPerRound = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
     150            0 :     u32 pairSize = (groupRankSize < sdmaConcurrentNum_) ?
     151            0 :         (groupRankSize + RANK_SET_COMPUTE_CONST - 1) / RANK_SET_COMPUTE_CONST: pairNumPerRound;
     152            0 :     for (u32 i = roundIdx * pairNumPerRound + 1;
     153            0 :          i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
     154            0 :         u32 leftRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
     155            0 :         u32 rightRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
     156            0 :         if (leftRemoteRank == rightRemoteRank) {
     157            0 :             partialCommRankSet[2].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
     158              :         } else {
     159            0 :             partialCommRankSet[0].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
     160            0 :             partialCommRankSet[1].push_back(std::make_pair(rightRemoteRank, rightRemoteRank));
     161              :         }
     162            0 :         HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSet] round[%u] userRank [%u] i[%u]" \
     163              :             "read/write leftRemoteRank[%u] rightRemoteRank[%u]", roundIdx, userRank_, i, leftRemoteRank, rightRemoteRank);
     164              :     }
     165            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSet] round[%u] partialCommRankSet_ total size[%zu]",
     166              :         roundIdx, partialCommRankSet[0].size() + partialCommRankSet[1].size() + partialCommRankSet[2].size());
     167            0 : }
     168              : 
     169              : // 主流只需要通知当前子步骤需要收发数据的 SDMA 流,减少同步开销
     170            0 : HcclResult AlltoAllFullMeshSymmetricMemory::NotifySubStreamStart()
     171              : {
     172            0 :     for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
     173            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
     174            0 :         CHK_RET(LocalNotify::Wait(sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     175              :             INVALID_VALUE_STAGE));
     176              :     }
     177            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][NotifySubStreamStart] userRank [%u] main stream notify sdma stream [%s]",
     178              :         userRank_, GetStreamIndexString().c_str());
     179            0 :     return HCCL_SUCCESS;
     180              : }
     181              : 
     182            0 : HcclResult AlltoAllFullMeshSymmetricMemory::WaitSubStreamFinish()
     183              : {
     184            0 :     for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
     185            0 :         CHK_RET(LocalNotify::Post(sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     186              :             INVALID_VALUE_STAGE));
     187            0 :         CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     188              :             INVALID_VALUE_STAGE));
     189              :     }
     190            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][WaitSubStreamFinish] userRank [%u] main stream wait sdma stream [%s]",
     191              :         userRank_, GetStreamIndexString().c_str());
     192            0 :     return HCCL_SUCCESS;
     193              : }
     194              : 
     195            0 : HcclResult AlltoAllFullMeshSymmetricMemory::NotifyRemoteRankStart()
     196              : {
     197            0 :     u32 streamIndex = 0;
     198            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     199            0 :         for (auto& sendRecvPair : sendRecvSide) {
     200            0 :             u32 recvRank = sendRecvPair.first;
     201            0 :             u32 sendRank = sendRecvPair.second;
     202            0 :             if (sendRank == userRank_) {
     203            0 :                 continue;
     204              :             }
     205            0 :             Stream& currStream = sdmaSubStream_[streamIndex];
     206            0 :             const LINK& readTransport = links_[recvRank];
     207            0 :             const LINK& sendTransport = links_[sendRank];
     208              : 
     209            0 :             CHK_RET(sendTransport->TxAck(currStream));
     210            0 :             CHK_RET(readTransport->RxAck(currStream));
     211              : 
     212            0 :             streamIndex ++;
     213              :         }
     214              :     }
     215            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][NotifyRemoteRankStart] done");
     216            0 :     return HCCL_SUCCESS;
     217              : }
     218              : 
     219            0 : bool AlltoAllFullMeshSymmetricMemory::IsPostSyncEnable(u32 roundIdx)
     220              : {
     221            0 :     bool isPostSyncEnable = false;
     222            0 :     isPostSyncEnable = (roundIdx == lastRoundIdx_) &&
     223            0 :         algOpContext_.opRetryHandler.retryEnable;
     224            0 :     return isPostSyncEnable;
     225              : }
     226              : 
     227            0 : HcclResult AlltoAllFullMeshSymmetricMemory::SdmaMainStreamWait(u32 roundIdx)
     228              : {
     229              :     // SDMA wait
     230            0 :     u32 streamIndex = 0;
     231            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     232            0 :         for (auto& sendRecvPair : sendRecvSide) {
     233            0 :             u32 recvRank = sendRecvPair.first;
     234            0 :             u32 sendRank = sendRecvPair.second;
     235            0 :             if (sendRank == userRank_) {
     236            0 :                 continue;
     237              :             }
     238            0 :             HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
     239              :                 "sendRank[%u], sdma stream [%u], "
     240              :                 "post sync info: roundIdx[%u], lastRoundIdx_[%u] main stream wait",
     241              :                 userRank_,  recvRank, sendRank, streamIndex, roundIdx, lastRoundIdx_);
     242            0 :             CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     243              :                 INVALID_VALUE_STAGE));
     244              : 
     245            0 :             streamIndex ++;
     246              :         }
     247              :     }
     248            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamWait] done");
     249            0 :     return HCCL_SUCCESS;
     250              : }
     251              : 
     252            0 : HcclResult AlltoAllFullMeshSymmetricMemory::SdmaMainStreamPost(u32 roundIdx)
     253              : {
     254              :     // SDMA post
     255            0 :     u32 streamIndex = 0;
     256            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     257            0 :         for (auto& sendRecvPair : sendRecvSide) {
     258            0 :             u32 recvRank = sendRecvPair.first;
     259            0 :             u32 sendRank = sendRecvPair.second;
     260            0 :             if (sendRank == userRank_) {
     261            0 :                 continue;
     262              :             }
     263            0 :             HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
     264              :                 "sendRank[%u], sdma stream [%u], "
     265              :                 "post sync info: roundIdx[%u], lastRoundIdx_[%u] main stream post",
     266              :                 userRank_,  recvRank, sendRank, streamIndex, roundIdx, lastRoundIdx_);
     267            0 :             CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     268              :                 INVALID_VALUE_STAGE));
     269              : 
     270            0 :             streamIndex ++;
     271              :         }
     272              :     }
     273            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamPost] done");
     274            0 :     return HCCL_SUCCESS;
     275              : }
     276              : 
     277            0 : HcclResult AlltoAllFullMeshSymmetricMemory::SetPostSyncTasks(u32 roundIdx)
     278              : {
     279              :     // SDMA wait
     280            0 :     CHK_RET(SdmaMainStreamWait(roundIdx));
     281              :     // SDMA post
     282            0 :     CHK_RET(SdmaMainStreamPost(roundIdx));
     283            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SetPostSyncTasks] done");
     284            0 :     return HCCL_SUCCESS;
     285              : }
     286              : 
     287            0 : HcclResult AlltoAllFullMeshSymmetricMemory::SDMAwithRemoteRankAndNotifyEnd(u32 roundIdx)
     288              : {
     289            0 :     bool isPostSyncEnable = IsPostSyncEnable(roundIdx);
     290            0 :     if (isPostSyncEnable) {
     291              :         // 下发主流上的后同步wait和post
     292            0 :         CHK_RET(SetPostSyncTasks(roundIdx));
     293              :     }
     294            0 :     u32 streamIndex = 0;
     295            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     296            0 :         for (auto& sendRecvPair : sendRecvSide) {
     297            0 :             u32 recvRank = sendRecvPair.first;
     298            0 :             u32 sendRank = sendRecvPair.second;
     299            0 :             if (sendRank == userRank_) {
     300            0 :                 continue;
     301              :             }
     302            0 :             const ReadDataBlock& readInfo = subStreamReadInfo_[recvRank];
     303            0 :             Stream& currStream = sdmaSubStream_[streamIndex];
     304            0 :             const LINK& readTransport = links_[recvRank];
     305            0 :             const LINK& sendTransport = links_[sendRank];
     306              : 
     307            0 :             const LINK& intraNeighboorTransport = links_[recvRank];
     308            0 :             CHK_PTR_NULL(intraNeighboorTransport);
     309            0 :             void* remDMAMemPtr = nullptr;
     310            0 :             CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
     311            0 :             DeviceMem remoteUserInMem = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr), userInput_.size());
     312            0 :             DeviceMem srcMem = remoteUserInMem.range(readInfo.remoteOffset, readInfo.recvLen);
     313            0 :             DeviceMem dstMem = userOutput_.range(readInfo.recvOffset, readInfo.recvLen);
     314            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, currStream,
     315              :                 readTransport->GetRemoteRank(), readTransport->GetLinkType()));
     316            0 :             HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SendRecvData] userRank [%u], recvRank[%u]," \
     317              :                 "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
     318              :                 "post sync info: roundIdx[%u], lastRoundIdx_[%u]",
     319              :                 userRank_,  recvRank, streamIndex, readInfo.remoteOffset,
     320              :                 readInfo.recvLen, readInfo.recvOffset, roundIdx, lastRoundIdx_);
     321            0 :             if (isPostSyncEnable) {
     322            0 :                 HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SendRecvData] post sync begins");
     323            0 :                 CHK_RET(LocalNotify::Post(currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     324              :                     INVALID_VALUE_STAGE));
     325            0 :                 CHK_RET(LocalNotify::Wait(currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     326              :                     INVALID_VALUE_STAGE));
     327              :             }
     328            0 :             CHK_RET(readTransport->TxDataSignal(currStream));
     329            0 :             CHK_RET(sendTransport->RxDataSignal(currStream));
     330              : 
     331            0 :             streamIndex ++;
     332            0 :         }
     333              :     }
     334            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SDMAwithRemoteRankAndNotifyEnd] done");
     335            0 :     return HCCL_SUCCESS;
     336              : }
     337              : 
     338            0 : HcclResult AlltoAllFullMeshSymmetricMemory::SendRecvData(u32 roundIdx)
     339              : {
     340            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SendRecvData] userRank [%u] sdma stream [%s] wait main stream",
     341              :         userRank_, GetStreamIndexString().c_str());
     342            0 :     CHK_RET(NotifyRemoteRankStart());
     343            0 :     CHK_RET(WaitSubStreamFinish());
     344            0 :     CHK_RET(NotifySubStreamStart());
     345            0 :     CHK_RET(SDMAwithRemoteRankAndNotifyEnd(roundIdx));
     346              : 
     347            0 :     return HCCL_SUCCESS;
     348              : }
     349              : 
     350            0 : HcclResult AlltoAllFullMeshSymmetricMemory::LocalCopy()
     351              : {
     352            0 :     const ZCopySendRecvInfo& sendRecvInfo = *sendRecvInfoPtr_;
     353            0 :     DeviceMem src = userInput_.range(sendRecvInfo.remoteSendOffset[userRank_],
     354            0 :         sendRecvInfo.localRecvLength[userRank_]);
     355            0 :     DeviceMem dst = userOutput_.range(sendRecvInfo.localRecvOffset[userRank_],
     356            0 :         sendRecvInfo.localRecvLength[userRank_]);
     357            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][LocalCopy]userRank [%u] copy from userInput [%llu]" \
     358              :         "to userOutput [%llu] dstLen[%llu]", userRank_, sendRecvInfo.remoteSendOffset[userRank_],
     359              :         sendRecvInfo.localRecvOffset, sendRecvInfo.localRecvLength[userRank_]);
     360            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
     361            0 :     return HCCL_SUCCESS;
     362            0 : }
     363              : 
     364            0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunGroupFullMeshAlltoall(u32 roundIdx)
     365              : {
     366            0 :     subStreamReadInfo_.clear();
     367            0 :     UpdateSendRecvInfo(roundIdx, subStreamReadInfo_, partialCommRankSet_);
     368            0 :     CHK_RET(NotifySubStreamStart());
     369            0 :     CHK_RET(SendRecvData(roundIdx));
     370            0 :     if (!islocalCpyDone_) {
     371            0 :         CHK_RET(LocalCopy());
     372            0 :         islocalCpyDone_ = true;
     373              :     }
     374            0 :     CHK_RET(WaitSubStreamFinish());
     375            0 :     return HCCL_SUCCESS;
     376              : }
     377              : 
     378            0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunSDMATasks(u32 roundIdx, u32 groupRankSize, u32 leftRankSize)
     379              : {
     380            0 :     UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
     381            0 :     CHK_RET(RunGroupFullMeshAlltoall(roundIdx));
     382            0 :     return HCCL_SUCCESS;
     383              : }
     384              : 
     385            0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunSDMA(HcclOpMetaInfoDef &opMeta)
     386              : {
     387              :     // 计算每个rank分组fullmesh后需要通信的轮次,向上取整
     388            0 :     commRounds_ = (devNumInlocalPod_ + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_;
     389            0 :     u32 leftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
     390            0 :     lastRoundIdx_ = std::min((leftRankSize + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_, static_cast<u32>(commRounds_)) - 1;
     391            0 :     HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][RunSDMA] userRank [%u] communication rounds[%llu]"
     392              :         "post sync info: lastRoundIdx_[%u] devNumInlocalPod_[%u] sdmaConcurrentNum_[%u]",
     393              :         userRank_, commRounds_,
     394              :         lastRoundIdx_, devNumInlocalPod_, sdmaConcurrentNum_);
     395              : 
     396            0 :     u32 currentLeftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
     397            0 :     for (u32 roundIdx = 0; roundIdx < commRounds_ && currentLeftRankSize > 0; roundIdx++) {
     398            0 :         CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
     399            0 :         u32 groupRankSize = (currentLeftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : currentLeftRankSize;
     400            0 :         CHK_RET(RunSDMATasks(roundIdx, groupRankSize, currentLeftRankSize));
     401            0 :         currentLeftRankSize -= groupRankSize;
     402            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
     403              :     }
     404              : 
     405            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunSDMA] finished.");
     406            0 :     return HCCL_SUCCESS;
     407              : }
     408              : 
     409            0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunAsync()
     410              : {   
     411            0 :     HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, userInput_.size(), true);
     412            0 :     CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
     413              : 
     414            0 :     if (userRankSize_ == 1) {
     415            0 :         HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunAsync] do localcopy with 1 rank");
     416            0 :         CHK_RET(LocalCopy());
     417            0 :         return HCCL_SUCCESS;
     418              :     }
     419              : 
     420            0 :     if (devNumInlocalPod_ > 1) {
     421            0 :         CHK_RET(RunSDMA(opMeta));
     422              :     }
     423              : 
     424            0 :     HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunAsync] finished.");
     425            0 :     return HCCL_SUCCESS;
     426              : }
     427              : 
     428              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_FULL_MESH_SYMMETRIC_MEMORY, AlltoAllFullMeshSymmetricMemory);
     429              : } // namespace hccl
        

Generated by: LCOV version 2.0-1