LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - ins_temp_all_to_all_mesh.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 332 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 20 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 "log.h"
      12              : 
      13              : #include "ins_temp_all_to_all_mesh.h"
      14              : 
      15              : namespace Hccl {
      16            0 : InsTempAlltoAllMesh::InsTempAlltoAllMesh(
      17              :     const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
      18            0 :     const std::map<RankId, u32>& tempVirtRankMap)
      19            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20              : {
      21            0 :     if (tempRankSize_ == 0) {
      22            0 :         THROW<InvalidParamsException>(StringFormat("[InsTempAlltoAllMesh] Invalid tempRankSize[%u].", tempRankSize_));
      23              :     }
      24            0 : }
      25              : 
      26            0 : InsTempAlltoAllMesh::~InsTempAlltoAllMesh() {}
      27              : 
      28            0 : HcclResult InsTempAlltoAllMesh::CalcRes(AlgTempResReq& tempResReq)
      29              : {
      30            0 :     CHK_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq));
      31              : 
      32            0 :     auto& linkReq = tempResReq.links;
      33            0 :     for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
      34            0 :         auto remoteRank = resReqIter->first;
      35            0 :         if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_.at(remoteRank) == 0) {
      36            0 :             HCCL_ERROR("No path to remoteRank[%d]", remoteRank);
      37            0 :             return HcclResult::HCCL_E_INTERNAL;
      38              :         }
      39            0 :         resReqIter->second = rank2PathNumMap_.at(remoteRank);
      40            0 :         maxPathNum = std::max(maxPathNum, rank2PathNumMap_.at(remoteRank)); // 每个rank的远端搬运最多需要pathNum条流
      41              :     }
      42            0 :     tempResReq.queNum = 1 + maxPathNum * std::min(ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE, tempRankSize_);
      43            0 :     tempResReq.streamNum = tempResReq.queNum;
      44            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      45              : 
      46            0 :     QId centerQ = 0;
      47            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      48            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      49            0 :     HCCL_DEBUG(
      50              :         "[InsTempAlltoAllMesh] Rank[%d], VtopoSize[%zu], requiredQue Num [%u].", myRank_, tempVTopo_[0].size(),
      51              :         tempResReq.queNum);
      52            0 :     return HcclResult::HCCL_SUCCESS;
      53              : }
      54              : 
      55            0 : void InsTempAlltoAllMesh::SetA2ASendRecvInfo(const A2ASendRecvInfo& sendRecvInfo) { localSendRecvInfo_ = sendRecvInfo; }
      56              : 
      57            0 : HcclResult InsTempAlltoAllMesh::GetScratchBufferInfo(const u64& scratchBufferSize, DataType dataType)
      58              : {
      59              :     // 需要变更为CCU/AICPU均只加载scratchSize
      60            0 :     u32 concurrentSendRecvNum = (tempRankSize_ > ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE) ?
      61              :                                     ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE :
      62            0 :                                     tempRankSize_;
      63              :     // scratch 不需要分两份,userIn直接一步到对端 scratch buffer
      64            0 :     u64 maxTmpMemSize = scratchBufferSize;
      65            0 :     u32 typeSize = DataTypeSizeGet(dataType);
      66            0 :     CHK_PRT_RET(
      67              :         typeSize == 0,
      68              :         HCCL_ERROR("[InsTempAlltoAllMesh] Rank [%d], Invalid dataSizePerVolume [%u].", myRank_, typeSize),
      69              :         HcclResult::HCCL_E_INTERNAL);
      70            0 :     u64 scratchInputMemSize = maxTmpMemSize / (concurrentSendRecvNum * typeSize) * typeSize;
      71            0 :     CHK_RET(SetBuffBlockSize(scratchInputMemSize));
      72            0 :     CHK_RET(SetConcurrentSendRecvNum(concurrentSendRecvNum));
      73              : 
      74            0 :     HCCL_INFO(
      75              :         "[InsTempAlltoAllMesh][GetScratchBufferInfo] concurrentSendRecvNum[%u], scratchInputMemSize[%llu]",
      76              :         concurrentSendRecvNum, scratchInputMemSize);
      77            0 :     return HcclResult::HCCL_SUCCESS;
      78              : }
      79              : 
      80            0 : HcclResult InsTempAlltoAllMesh::SetBuffBlockSize(const u64 buffBlockSize)
      81              : {
      82            0 :     CHK_PRT_RET(
      83              :         buffBlockSize == 0, HCCL_ERROR("[InsTempAlltoAllMesh][SetBuffBlockSize] buffBlockSize should not be zero"),
      84              :         HcclResult::HCCL_E_PARA);
      85            0 :     buffBlockSize_ = buffBlockSize;
      86            0 :     return HcclResult::HCCL_SUCCESS;
      87              : }
      88              : 
      89            0 : HcclResult InsTempAlltoAllMesh::SetConcurrentSendRecvNum(const u32 concurrentSendRecvNum)
      90              : {
      91            0 :     CHK_PRT_RET(
      92              :         concurrentSendRecvNum == 0,
      93              :         HCCL_ERROR("[InsTempAlltoAllMesh][SetConcurrentSendRecvNum] concurrentSendRecvNum should not be zero"),
      94              :         HcclResult::HCCL_E_PARA);
      95            0 :     concurrentSendRecvNum_ = concurrentSendRecvNum;
      96            0 :     return HcclResult::HCCL_SUCCESS;
      97              : }
      98              : 
      99            0 : u32 InsTempAlltoAllMesh::CalcStepNum()
     100              : {
     101            0 :     u32 numSubStep = 0;
     102              : 
     103            0 :     for (u32 destRank = 0; destRank < tempRankSize_; destRank++) {
     104            0 :         if (destRank == static_cast<u32>(myRank_)) {
     105            0 :             continue;
     106              :         }
     107            0 :         u32 currRankSendSubStep = ((localSendRecvInfo_.sendLength[destRank] + buffBlockSize_ - 1) / buffBlockSize_);
     108            0 :         u32 currRankRecvSubStep = ((localSendRecvInfo_.recvLength[destRank] + buffBlockSize_ - 1) / buffBlockSize_);
     109            0 :         numSubStep = std::max(numSubStep, std::max(currRankSendSubStep, currRankRecvSubStep));
     110              :     }
     111            0 :     return numSubStep;
     112              : }
     113              : 
     114            0 : HcclResult InsTempAlltoAllMesh::CalcSendSliceInfo(u32 remoteRank, UsrData& sendSliceInfo)
     115              : {
     116              :     // 判断在打平的视角下,当前的remoteRank 是在本rank的左边还是右边
     117            0 :     u32 pairNum = concurrentSendRecvNum_ / 2; // 每轮次,会与concurrentSendRecvNum个对端通信,左右各pairNum个
     118            0 :     u32 gapRight = (tempRankSize_ + remoteRank - myRank_) % tempRankSize_;
     119            0 :     u32 gapLeft = (tempRankSize_ + myRank_ - remoteRank) % tempRankSize_;
     120            0 :     u32 sendSrcBuffIdx = 0;
     121            0 :     u32 sendDstBuffIdx = 0;
     122            0 :     if (gapLeft < gapRight) {
     123              :         // 离左边更近
     124            0 :         u32 gap = gapLeft;
     125            0 :         sendSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
     126            0 :         sendDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
     127            0 :     } else if (gapLeft > gapRight) {
     128              :         // 离右边更近
     129            0 :         u32 gap = gapRight;
     130            0 :         sendSrcBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
     131            0 :         sendDstBuffIdx = pairNum + ((gap - 1) % pairNum);
     132              :     } else {
     133            0 :         sendSrcBuffIdx = 0;
     134            0 :         sendDstBuffIdx = 0;
     135              :     }
     136            0 :     u64 sendDataOffset = 0;
     137            0 :     u64 remainSendLen = localSendRecvInfo_.sendLength[remoteRank];
     138            0 :     u64 sendSrcBuffOffset = sendSrcBuffIdx * buffBlockSize_;
     139            0 :     u64 sendDstBuffOffset = sendDstBuffIdx * buffBlockSize_;
     140            0 :     while (remainSendLen > 0) {
     141            0 :         u64 currDataRemainLen = localSendRecvInfo_.sendLength[remoteRank] - sendDataOffset;
     142            0 :         u64 sendLen = std::min(buffBlockSize_, currDataRemainLen);
     143            0 :         u64 userInOffset = localSendRecvInfo_.sendOffset[remoteRank] + sendDataOffset;
     144            0 :         u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_] + sendDataOffset;
     145              : 
     146            0 :         DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendLen);
     147            0 :         DataSlice sendScratchSlice = DataSlice(buffInfo_.inBuffType, sendSrcBuffOffset, sendLen);
     148            0 :         DataSlice sendDstSlice = DataSlice(buffInfo_.outBuffType, sendDstBuffOffset, sendLen);
     149            0 :         DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, sendLen);
     150            0 :         sendSliceInfo.usrInSlices.emplace_back(userInSlice);
     151            0 :         sendSliceInfo.scratchInSlices.emplace_back(sendScratchSlice);
     152            0 :         sendSliceInfo.scratchOutSlices.emplace_back(sendDstSlice);
     153            0 :         sendSliceInfo.usrOutSlices.emplace_back(userOutSlice);
     154            0 :         sendDataOffset += sendLen;
     155            0 :         remainSendLen -= sendLen;
     156              :     }
     157            0 :     return HcclResult::HCCL_SUCCESS;
     158              : }
     159              : 
     160            0 : HcclResult InsTempAlltoAllMesh::CalcRecvSliceInfo(u32 remoteRank, UsrData& readSliceInfo)
     161              : {
     162            0 :     u32 pairNum = concurrentSendRecvNum_ / 2;
     163            0 :     u32 gapRight = (remoteRank - myRank_ + tempRankSize_) % tempRankSize_;
     164            0 :     u32 gapLeft = (myRank_ - remoteRank + tempRankSize_) % tempRankSize_;
     165            0 :     u32 recvSrcBuffIdx = 0;
     166            0 :     u32 recvDstBuffIdx = 0;
     167            0 :     if (gapLeft < gapRight) {
     168            0 :         u32 gap = gapLeft;
     169            0 :         recvSrcBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
     170            0 :         recvDstBuffIdx = pairNum + ((gap - 1) % pairNum);
     171            0 :     } else if (gapLeft > gapRight) {
     172            0 :         u32 gap = gapRight;
     173            0 :         recvSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
     174            0 :         recvDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
     175              :     } else {
     176            0 :         recvSrcBuffIdx = 0;
     177            0 :         recvDstBuffIdx = 0;
     178              :     }
     179            0 :     u64 recvDataOffset = 0;
     180            0 :     u64 remainRecvLen = localSendRecvInfo_.recvLength[remoteRank];
     181            0 :     u64 recvSrcBuffOffset = recvSrcBuffIdx * buffBlockSize_;
     182            0 :     u64 recvDstBuffOffset = recvDstBuffIdx * buffBlockSize_;
     183            0 :     while (remainRecvLen > 0) {
     184            0 :         u64 currDataRemainLen = localSendRecvInfo_.recvLength[remoteRank] - recvDataOffset;
     185            0 :         u64 recvLen = std::min(buffBlockSize_, currDataRemainLen);
     186            0 :         u64 userOutOffset = localSendRecvInfo_.recvOffset[remoteRank] + recvDataOffset;
     187            0 :         u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_] + recvDataOffset;
     188              : 
     189            0 :         DataSlice recvInSlice = DataSlice(BufferType::INPUT, userInOffset, recvLen);
     190            0 :         DataSlice recvScratchSlice = DataSlice(buffInfo_.inBuffType, recvSrcBuffOffset, recvLen);
     191            0 :         DataSlice recvDstSlice = DataSlice(buffInfo_.outBuffType, recvDstBuffOffset, recvLen);
     192            0 :         DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvLen);
     193            0 :         readSliceInfo.usrInSlices.emplace_back(recvInSlice);
     194            0 :         readSliceInfo.scratchInSlices.emplace_back(recvScratchSlice);
     195            0 :         readSliceInfo.scratchOutSlices.emplace_back(recvDstSlice);
     196            0 :         readSliceInfo.usrOutSlices.emplace_back(userOutSlice);
     197            0 :         recvDataOffset += recvLen;
     198            0 :         remainRecvLen -= recvLen;
     199              :     }
     200            0 :     return HcclResult::HCCL_SUCCESS;
     201              : }
     202              : 
     203            0 : HcclResult InsTempAlltoAllMesh::CalcSendRecvAllSliceInfo(
     204              :     std::unordered_map<u32, UsrData>& sendSliceInfoMap, std::unordered_map<u32, UsrData>& recvSliceInfoMap)
     205              : {
     206            0 :     for (u32 remoteRank = 0; remoteRank < tempRankSize_; remoteRank++) {
     207            0 :         if (remoteRank == static_cast<u32>(myRank_)) {
     208            0 :             continue;
     209              :         }
     210            0 :         UsrData readSliceInfo;
     211            0 :         UsrData sendSliceInfo;
     212            0 :         CalcSendSliceInfo(remoteRank, sendSliceInfo);
     213            0 :         CalcRecvSliceInfo(remoteRank, readSliceInfo);
     214            0 :         sendSliceInfoMap[remoteRank] = sendSliceInfo;
     215            0 :         recvSliceInfoMap[remoteRank] = readSliceInfo;
     216            0 :     }
     217            0 :     return HcclResult::HCCL_SUCCESS;
     218              : }
     219              : 
     220              : HcclResult
     221            0 : InsTempAlltoAllMesh::CalcCommRankSetforOneLoop(u32 roundIdx, const u32 groupRankSize, std::vector<u32>& commRanks) const
     222              : {
     223            0 :     commRanks.clear();
     224            0 :     u32 pairNumPerRound = concurrentSendRecvNum_ / 2;
     225            0 :     u32 pairSize = (groupRankSize < concurrentSendRecvNum_) ? (groupRankSize + 1) / 2 : pairNumPerRound;
     226            0 :     for (u32 i = roundIdx * pairNumPerRound + 1; i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
     227            0 :         u32 leftRemoteRank = (myRank_ + tempRankSize_ - i) % tempRankSize_;
     228            0 :         u32 rightRemoteRank = (myRank_ + i) % tempRankSize_;
     229            0 :         if (leftRemoteRank == rightRemoteRank) {
     230            0 :             commRanks.push_back(leftRemoteRank);
     231            0 :             break;
     232              :         } else {
     233            0 :             commRanks.push_back(leftRemoteRank);
     234            0 :             commRanks.push_back(rightRemoteRank);
     235              :         }
     236              :     }
     237            0 :     return HcclResult::HCCL_SUCCESS;
     238              : }
     239              : 
     240            0 : HcclResult InsTempAlltoAllMesh::CopySendDataToScratch(
     241              :     u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfo,
     242              :     const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
     243              : {
     244            0 :     HCCL_INFO("[InsTempAlltoAllMesh] CopySendDataToScratch");
     245            0 :     u32 queueId = 0;
     246            0 :     for (u32 i = 0; i < commRanks.size(); i++) {
     247            0 :         u32 remoteRank = commRanks[i];
     248            0 :         HCCL_INFO("remoteRank = %u", remoteRank);
     249            0 :         u32 linkNum = rank2PathNumMap_.at(remoteRank);
     250            0 :         std::vector<LinkData> links = tempLinks.at(remoteRank);
     251            0 :         std::vector<float> dataSplitRate(linkNum);
     252            0 :         CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
     253            0 :         UsrData& currSendSliceInfo = sendSliceInfo[remoteRank];
     254            0 :         for (u32 j = 0; j < linkNum; j++) {
     255            0 :             InsQuePtr queue = queues[queueId];
     256            0 :             queueId++;
     257            0 :             HCCL_INFO("queueId=%u", queueId);
     258            0 :             if (step < currSendSliceInfo.usrInSlices.size()) {
     259            0 :                 DataSlice& userInSliceAllLinks = currSendSliceInfo.usrInSlices[step];
     260            0 :                 DataSlice& sendScratchSliceAllLinks = currSendSliceInfo.scratchInSlices[step];
     261            0 :                 DataSlice userInSlice = CalcDataSliceForLinks(userInSliceAllLinks, dataSplitRate, j, dataType_);
     262              :                 DataSlice sendScratchSlice
     263            0 :                     = CalcDataSliceForLinks(sendScratchSliceAllLinks, dataSplitRate, j, dataType_);
     264              : 
     265            0 :                 InsQuePtr queue = queues[i];
     266            0 :                 CHK_RET(LocalCopy(queue, userInSlice, sendScratchSlice));
     267            0 :             }
     268            0 :         }
     269            0 :     }
     270            0 :     return HcclResult::HCCL_SUCCESS;
     271              : }
     272              : 
     273            0 : HcclResult InsTempAlltoAllMesh::SendRecvData(
     274              :     u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfo,
     275              :     std::unordered_map<u32, UsrData>& readSliceInfo, const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
     276              : {
     277            0 :     CHK_PRT_RET(
     278              :         queues.empty(), HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] empty queues"), HcclResult::HCCL_E_INTERNAL);
     279            0 :     CHK_PTR_NULL(queues[0]);
     280              : 
     281            0 :     if (commRanks.size() * maxPathNum > queues.size()) {
     282            0 :         HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] commRanks.size() * maxPathNum > queues.size() is wrong");
     283            0 :         return HcclResult::HCCL_E_INTERNAL;
     284              :     }
     285            0 :     uint64_t queuesId = 0;
     286            0 :     for (u32 i = 0; i < commRanks.size(); i++) {
     287            0 :         s32 remoteRank = static_cast<s32>(commRanks[i]);
     288            0 :         u32 linkNum = rank2PathNumMap_.at(remoteRank);
     289            0 :         UsrData& currSendSliceInfo = sendSliceInfo[remoteRank];
     290            0 :         UsrData& currReadSliceInfo = readSliceInfo[remoteRank];
     291            0 :         std::vector<LinkData> links = tempLinks.at(remoteRank);
     292            0 :         std::vector<float> dataSplitRate(linkNum);
     293            0 :         CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
     294              :         std::vector<DataSlice>& currSendSrcSlices
     295            0 :             = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.scratchInSlices : currSendSliceInfo.usrInSlices;
     296              :         std::vector<DataSlice>& currSendDstSlices
     297            0 :             = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.usrOutSlices : currSendSliceInfo.scratchOutSlices;
     298              :         std::vector<DataSlice>& currReadSrcSlices
     299            0 :             = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.scratchInSlices : currReadSliceInfo.usrInSlices;
     300              :         std::vector<DataSlice>& currReadDstSlices
     301            0 :             = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.usrOutSlices : currReadSliceInfo.scratchOutSlices;
     302            0 :         for (u32 j = 0; j < linkNum; j++) {
     303            0 :             InsQuePtr queue = queues[queuesId];
     304            0 :             queuesId++;
     305              : 
     306            0 :             LinkData link = tempLinks.at(remoteRank)[j];
     307            0 :             if (step < currSendSrcSlices.size() && step < currReadSrcSlices.size()) {
     308            0 :                 DataSlice& sendSrcSliceAllLinks = currSendSrcSlices[step];
     309            0 :                 DataSlice& sendDstSliceAllLinks = currSendDstSlices[step];
     310            0 :                 DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
     311            0 :                 DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
     312            0 :                 std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
     313            0 :                 std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
     314            0 :                 SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
     315            0 :                 DataSlice& recvSrcSliceAllLinks = currReadSrcSlices[step];
     316            0 :                 DataSlice& recvDstSliceAllLinks = currReadDstSlices[step];
     317            0 :                 DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
     318            0 :                 DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
     319            0 :                 std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
     320            0 :                 std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
     321            0 :                 SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
     322            0 :                 TxRxSlicesList sendRecvSlice(sendDataSlice, recvDataSlice);
     323              : 
     324            0 :                 TxRxLinks sendRecvLinks(link, link);
     325            0 :                 SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlice);
     326            0 :                 CHK_RET(SendRecv(sendRecvInfo, queue, 0, true, dmaMode_));
     327            0 :                 HCCL_DEBUG(
     328              :                     "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send and recv.",
     329              :                     step, i, remoteRank);
     330            0 :             } else if (step < currSendSrcSlices.size()) {
     331            0 :                 DataSlice& sendSrcSliceAllLinks = currSendSrcSlices[step];
     332            0 :                 DataSlice& sendDstSliceAllLinks = currSendDstSlices[step];
     333            0 :                 DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
     334            0 :                 DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
     335            0 :                 std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
     336            0 :                 std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
     337            0 :                 SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
     338            0 :                 DataInfo sendDataInfo(link, sendDataSlice);
     339            0 :                 CHK_RET(Send(sendDataInfo, queue));
     340            0 :                 HCCL_DEBUG(
     341              :                     "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send.", step, i,
     342              :                     remoteRank);
     343            0 :             } else if (step < currReadSrcSlices.size()) {
     344            0 :                 DataSlice& recvSrcSliceAllLinks = currReadSrcSlices[step];
     345            0 :                 DataSlice& recvDstSliceAllLinks = currReadDstSlices[step];
     346            0 :                 DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
     347            0 :                 DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
     348            0 :                 std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
     349            0 :                 std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
     350            0 :                 SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
     351            0 :                 DataInfo recvDataInfo(link, recvDataSlice);
     352            0 :                 CHK_RET(Recv(recvDataInfo, queue));
     353            0 :                 HCCL_DEBUG(
     354              :                     "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run recv.", step, i,
     355              :                     remoteRank);
     356            0 :             }
     357            0 :         }
     358            0 :     }
     359            0 :     return HcclResult::HCCL_SUCCESS;
     360              : }
     361              : 
     362            0 : HcclResult InsTempAlltoAllMesh::CopyRecvDataFromScratch(
     363              :     u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& readSliceInfo,
     364              :     const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
     365              : {
     366            0 :     HCCL_INFO("[InsTempAlltoAllMesh] CopyRecvDataFromScratch");
     367            0 :     u32 queueId = 0;
     368            0 :     for (u32 i = 0; i < commRanks.size(); i++) {
     369            0 :         u32 remoteRank = commRanks[i];
     370            0 :         HCCL_INFO("remoteRank = %u", remoteRank);
     371            0 :         u32 linkNum = rank2PathNumMap_.at(remoteRank);
     372            0 :         std::vector<LinkData> links = tempLinks.at(remoteRank);
     373            0 :         std::vector<float> dataSplitRate(linkNum);
     374            0 :         CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
     375            0 :         UsrData& currReadSliceInfo = readSliceInfo[remoteRank];
     376            0 :         for (u32 j = 0; j < linkNum; j++) {
     377            0 :             InsQuePtr queue = queues[queueId];
     378            0 :             queueId++;
     379            0 :             HCCL_INFO("queueId=%u", queueId);
     380            0 :             if (step < currReadSliceInfo.scratchOutSlices.size()) {
     381            0 :                 DataSlice& recvScratchSliceAllLinks = currReadSliceInfo.scratchOutSlices[step];
     382            0 :                 DataSlice& userOutSliceAllLinks = currReadSliceInfo.usrOutSlices[step];
     383              :                 DataSlice recvScratchSlice
     384            0 :                     = CalcDataSliceForLinks(recvScratchSliceAllLinks, dataSplitRate, j, dataType_);
     385            0 :                 DataSlice userOutSlice = CalcDataSliceForLinks(userOutSliceAllLinks, dataSplitRate, j, dataType_);
     386            0 :                 CHK_RET(LocalCopy(queue, recvScratchSlice, userOutSlice));
     387              :             }
     388            0 :         }
     389            0 :     }
     390            0 :     return HcclResult::HCCL_SUCCESS;
     391              : }
     392              : 
     393            0 : HcclResult InsTempAlltoAllMesh::RunSendRecvBufferLoop(
     394              :     u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfoMap,
     395              :     std::unordered_map<u32, UsrData>& recvSliceInfoMap, const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
     396              : {
     397            0 :     if (commRanks.size() > 1) {
     398            0 :         CHK_RET(PreSyncInterQueues(queues));
     399              :     }
     400              : 
     401            0 :     if (dmaMode_ == DmaMode::PUT) {
     402            0 :         CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
     403            0 :         CHK_RET(CopyRecvDataFromScratch(step, commRanks, recvSliceInfoMap, tempLinks, queues));
     404            0 :     } else if (dmaMode_ == DmaMode::GET) {
     405            0 :         CHK_RET(CopySendDataToScratch(step, commRanks, sendSliceInfoMap, tempLinks, queues));
     406            0 :         CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
     407              :     }
     408              : 
     409            0 :     if (commRanks.size() > 1) {
     410            0 :         CHK_RET(PostSyncInterQueues(queues));
     411              :     }
     412            0 :     return HcclResult::HCCL_SUCCESS;
     413              : }
     414              : 
     415            0 : HcclResult InsTempAlltoAllMesh::RunSendRecvForAllRanks(
     416              :     u32 step, std::unordered_map<u32, UsrData>& sendSliceInfoMap, std::unordered_map<u32, UsrData>& recvSliceInfoMap,
     417              :     const ResLinks& tempLinks, std::vector<InsQuePtr>& queues)
     418              : {
     419            0 :     u64 commLoops = (tempRankSize_ + concurrentSendRecvNum_ - 1) / concurrentSendRecvNum_;
     420            0 :     u32 leftRankSize = tempRankSize_ - 1; // leftRankSize中去掉本卡
     421            0 :     std::vector<u32> commRanks;
     422            0 :     for (u32 roundIdx = 0; roundIdx < commLoops && leftRankSize > 0; roundIdx++) {
     423            0 :         u32 groupRankSize = (leftRankSize > concurrentSendRecvNum_) ? concurrentSendRecvNum_ : leftRankSize;
     424            0 :         CHK_RET(CalcCommRankSetforOneLoop(roundIdx, groupRankSize, commRanks));
     425            0 :         CHK_RET(RunSendRecvBufferLoop(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
     426            0 :         leftRankSize -= groupRankSize;
     427            0 :         HCCL_DEBUG(
     428              :             "[InsTempAlltoAllMesh][RunSendRecvForAllRanks] commRanksSize[%zu], roundIdx[%u] finish.", commRanks.size(),
     429              :             roundIdx);
     430              :     }
     431            0 :     return HcclResult::HCCL_SUCCESS;
     432            0 : }
     433              : 
     434            0 : HcclResult InsTempAlltoAllMesh::LocalDataCopy(InsQuePtr tempInsQue)
     435              : {
     436            0 :     u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_];
     437            0 :     u64 sendSize = localSendRecvInfo_.sendLength[myRank_];
     438            0 :     u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_];
     439            0 :     u64 recvSize = localSendRecvInfo_.recvLength[myRank_];
     440            0 :     DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendSize);
     441            0 :     DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvSize);
     442            0 :     CHK_RET(LocalCopy(tempInsQue, userInSlice, userOutSlice));
     443            0 :     return HcclResult::HCCL_SUCCESS;
     444              : }
     445              : 
     446            0 : HcclResult InsTempAlltoAllMesh::Run(
     447              :     const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
     448              :     std::vector<InsQuePtr>& tempInsQues)
     449              : {
     450              :     (void)tempFuncs;
     451              :     (void)sliceInfoVec;
     452            0 :     if (tempInsQues.size() == 0) {
     453            0 :         HCCL_ERROR("[CcuTempAlltoAllMesh1D] tempInsQues.size() is zero.");
     454            0 :         return HcclResult::HCCL_E_PARA;
     455              :     }
     456            0 :     dmaMode_ = DmaMode::PUT;
     457            0 :     if (IsPcieLink(tempLinks)) {
     458            0 :         dmaMode_ = DmaMode::GET;
     459              :     }
     460            0 :     buffInfo_ = buffInfo;
     461            0 :     u32 totalStep = CalcStepNum();
     462            0 :     HCCL_INFO(
     463              :         "[InsTempAlltoAllMesh] AllToAll full mesh total step is [%u], tempInsQuesSize[%zu].", totalStep,
     464              :         tempInsQues.size());
     465            0 :     std::unordered_map<u32, UsrData> sendSliceInfoMap;
     466            0 :     std::unordered_map<u32, UsrData> recvSliceInfoMap;
     467            0 :     CHK_RET(CalcSendRecvAllSliceInfo(sendSliceInfoMap, recvSliceInfoMap));
     468            0 :     std::vector<InsQuePtr> localCopyQues;
     469            0 :     if (tempInsQues.size() > 1) {
     470            0 :         localCopyQues.push_back(tempInsQues[0]);
     471            0 :         localCopyQues.push_back(tempInsQues.back());
     472              :     }
     473              : 
     474            0 :     if (localCopyQues.size() > 1) {
     475            0 :         CHK_RET(PreSyncInterQueues(localCopyQues));
     476              :     }
     477              :     // 最后一条流做localCopy
     478            0 :     CHK_RET(LocalDataCopy(tempInsQues.back()));
     479              :     // 其余的流去做sendRecv
     480            0 :     std::vector<InsQuePtr> sendRecvQues(tempInsQues.begin(), tempInsQues.begin() + (tempInsQues.size() - 1));
     481            0 :     for (u32 step = 0; step < totalStep; step++) {
     482            0 :         CHK_RET(RunSendRecvForAllRanks(step, sendSliceInfoMap, recvSliceInfoMap, tempLinks, sendRecvQues));
     483            0 :         HCCL_DEBUG("[InsTempAlltoAllMesh] AllToAll full mesh step[%u] execute success", step);
     484              :     }
     485            0 :     if (localCopyQues.size() > 1) {
     486            0 :         CHK_RET(PostSyncInterQueues(localCopyQues));
     487              :     }
     488            0 :     HCCL_INFO("[InsTempAlltoAllMesh] AllToAll full mesh rank[%d] finish.", myRank_);
     489            0 :     return HcclResult::HCCL_SUCCESS;
     490            0 : }
     491              : 
     492              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1