LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_alltoallv - alltoallv_direct_fullmesh.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 864 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 50 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 "alltoallv_direct_fullmesh.h"
      12              : #include "dispatcher_pub.h"
      13              : 
      14              : namespace hccl {
      15            0 : AlltoAllVDirectFullMesh::AlltoAllVDirectFullMesh(const HcclDispatcher dispatcher)
      16            0 :     : AlgTemplateBase(dispatcher)
      17              : {
      18            0 : }
      19              : 
      20            0 : AlltoAllVDirectFullMesh::~AlltoAllVDirectFullMesh() {}
      21              : 
      22            0 : HcclResult AlltoAllVDirectFullMesh::GenerateSubStreamInfo(const std::vector<Stream> &subStreams,
      23              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalMainToSub,
      24              :     const std::vector<std::shared_ptr<LocalNotify>> &meshSignalSubToMain)
      25              : {
      26            0 :     u32 totalSubstreamSize = (totalRdmaRankNum_ > 0) ?
      27            0 :         (sdmaConcurrentNum_ + rdmaConcurrentNum_ + 1) : (sdmaConcurrentNum_);
      28            0 :     if (subStreams.size() < totalSubstreamSize || meshSignalMainToSub.size() < totalSubstreamSize ||
      29            0 :         meshSignalSubToMain.size() < totalSubstreamSize) {
      30            0 :         HCCL_ERROR("[AlltoAllVDirectFullMesh][GenerateSubStreamInfo]subStreamsSize[%zu], meshSignalMainToSubSize[%zu]"\
      31              :             "meshSignalSubToMainSize[%zu] is smaller than totalSubstreamSize[%u]",subStreams.size(),
      32              :             meshSignalMainToSub.size(), meshSignalSubToMain.size(), totalSubstreamSize);
      33            0 :         return HCCL_E_PARA;
      34              :     }
      35            0 :     CHK_PRT_RET(links_.size() < userRankSize_, HCCL_ERROR("[AlltoAllVDirectFullMesh][GenerateSubStreamInfo]"\
      36              :         "links_.size()[%zu] is smaller than userRankSize_[%u].", links_.size(), userRankSize_),
      37              :         HCCL_E_PARA);
      38            0 :     HCCL_DEBUG("subStreams.size[%zu], meshSignalMainToSub.size[%zu], links_.size[%zu]",
      39              :         subStreams.size(), meshSignalMainToSub.size(), links_.size());
      40            0 :     u32 index = 0;
      41            0 :     for (u32 sdmaIndex = 0; sdmaIndex < sdmaConcurrentNum_; sdmaIndex++) {
      42            0 :         sdmaSubStream_.push_back(subStreams[index]);
      43            0 :         sdmaMeshSignalMainToSub_.push_back(meshSignalMainToSub[index]);
      44            0 :         sdmaMeshSignalSubToMain_.push_back(meshSignalSubToMain[index]);
      45            0 :         index++;
      46              :     }
      47            0 :     for (u32 localIndex = 0; localIndex < sdmaConcurrentNum_; localIndex++) {
      48            0 :         localSubStream_.push_back(subStreams[index]);
      49            0 :         localSignalMainToSub_.push_back(meshSignalMainToSub[index]);
      50            0 :         localSignalSubToMain_.push_back(meshSignalSubToMain[index]);
      51            0 :         index++;
      52              :     }
      53            0 :     if (totalRdmaRankNum_ > 0) {
      54            0 :         rdmaSubStreams_.push_back(subStreams[index]);
      55            0 :         main2RdmaControlStreamNotify_ = meshSignalMainToSub[index];
      56            0 :         rdmaControl2MainStreamNotify_ = meshSignalSubToMain[index];
      57            0 :         index++;
      58            0 :         for (u32 rdmaIndex = 0; rdmaIndex < rdmaConcurrentNum_; rdmaIndex++) {
      59            0 :             rdmaSubStreams_.push_back(subStreams[index]);
      60            0 :             rdmaControl2SubNotifies_.push_back(meshSignalMainToSub[index]);
      61            0 :             rdmaSub2ControlNotifies_.push_back(meshSignalSubToMain[index]);
      62            0 :             index++;
      63              :         }
      64              :     }
      65            0 :     return HCCL_SUCCESS;
      66              : }
      67              : 
      68            0 : HcclResult AlltoAllVDirectFullMesh::Prepare(PrepareData &param)
      69              : {
      70            0 :     needAlltoallvCache_ = param.needAlltoallvCache;
      71            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][Prepare] set needAlltoallvCache_[%u] for alltoallv aicpu cache", needAlltoallvCache_);
      72              : 
      73            0 :     mainStream_ = param.stream;
      74            0 :     userRank_ = param.userRank;
      75            0 :     userRankSize_ = param.userRankSize;
      76            0 :     links_ = *param.linksPtr;
      77            0 :     localSendRecvInfoPtr_ = param.localSendRecvInfoPtr;
      78            0 :     devNumInlocalPod_ = param.devNumInlocalPod;
      79            0 :     rankIdxInPod_ = param.rankIdxInPod;
      80            0 :     opType_ = param.opType;
      81            0 :     algOpContext_ = param.algOpContext;
      82              : 
      83            0 :     podStartRank_ = userRank_ - rankIdxInPod_;
      84            0 :     podEndRank_ = podStartRank_ + devNumInlocalPod_ - 1;
      85            0 :     sdmaConcurrentNum_ = (devNumInlocalPod_ > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
      86            0 :         (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) : (devNumInlocalPod_);
      87              : 
      88            0 :     totalRdmaRankNum_ = userRankSize_ - devNumInlocalPod_;
      89            0 :     rdmaConcurrentNum_ = (totalRdmaRankNum_ > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
      90            0 :         (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) : (totalRdmaRankNum_);
      91            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh]devNumInlocalPod_[%u], userRankSize_[%u] podStartRank_[%u]" \
      92              :         "podEndRank_[%u], totalRdmaRankNum_[%u], sdmaConcurrentNum_[%u], rdmaConcurrentNum_[%u]",
      93              :         devNumInlocalPod_, userRankSize_, podStartRank_, podEndRank_, totalRdmaRankNum_,
      94              :         sdmaConcurrentNum_, rdmaConcurrentNum_);
      95              : 
      96            0 :     CHK_PRT_RET(userRankSize_ == 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][Prepare]userRankSize_ is zero."),
      97              :         HCCL_E_PARA);
      98              : 
      99            0 :     userInput_ = param.inputMem;
     100            0 :     userOutput_ = param.outputMem;
     101            0 :     cclInMem_ = param.cclInMem;
     102            0 :     cclOutMem_ = param.cclOutMem;
     103            0 :     workMode_ = param.workMode;
     104            0 :     isSuPodAsym_ = param.isSuPodAsym;
     105              : 
     106              :     // 注意: 如果isBigCount的计算逻辑发生变化, 需要同步修改IsBigCountForAlltoallv()中的代码
     107            0 :     u64 maxSendLen = CalcMaxSendLen();
     108            0 :     isBigCount_ = (maxSendLen > ALLTOALLV_DIRECT_FULLMESH_BIG_SIZE) ? true : false;
     109            0 :     CHK_RET(GenerateSubStreamInfo(*param.subStreamsPtr, *param.signalPtr, *param.signalAuxPtr));
     110              : 
     111            0 :     if (algOpContext_.mc2Handler.stepSize > 0) {
     112            0 :         sdmaConcurrentNum_ = (devNumInlocalPod_ > 1) ? 1 : (devNumInlocalPod_);
     113              :         //MC2细粒度不需要本地并发处理
     114            0 :         isBigCount_ = false;
     115              :     }
     116              : 
     117              :     /* 考虑当group0 的rank 跟 group 1的所有rank通信时,每次都要收发,所以取sdmaConcurrentNum_块;
     118              :     跟group 0内的rank通信有一块儿浪费 */
     119              :     // 注意: 如果sdmaDataBlockSize_的计算逻辑发生变化, 需要同步修改framework下CalcMetadataForFirstAlltoallv()函数中的part 1
     120            0 :     u32 blockGroup = (isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC) ? 2 : 1;
     121            0 :     sdmaDataBlockSize_= (cclInMem_.size() / std::max(1u, sdmaConcurrentNum_ * blockGroup));
     122              :     // 向下对齐到16k Byte
     123            0 :     if (sdmaDataBlockSize_> HCCL_MIN_SLICE_ALIGN_910B) {
     124            0 :         sdmaDataBlockSize_= (sdmaDataBlockSize_/ HCCL_MIN_SLICE_ALIGN_910B) * HCCL_MIN_SLICE_ALIGN_910B;
     125              :     }
     126            0 :     CHK_PRT_RET(sdmaDataBlockSize_== 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][Prepare]sdmaDataBlockSize_is zero."),
     127              :         HCCL_E_INTERNAL);
     128            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][Prepare] userRank [%u] total cclsize[%llu]," \
     129              :         "sdmaDataBlockSize_[%llu], BigCountFlag[%d], stepSize[%u]", userRank_, cclInMem_.size(), sdmaDataBlockSize_, isBigCount_,
     130              :         algOpContext_.mc2Handler.stepSize);
     131              : 
     132              :     // 一半的CCLOut用来发送RDMA数据,另一半用来接收RDMA数据,因此需要除以2
     133            0 :     rdmaDataBlockSize_ = cclOutMem_.size() / std::max(1u, rdmaConcurrentNum_) / 2;
     134              : 
     135            0 :     return HCCL_SUCCESS;
     136              : }
     137              : 
     138            0 : std::string AlltoAllVDirectFullMesh::GetStreamIndexString()
     139              : {
     140            0 :     std::string res = "";
     141            0 :     for (auto& info : subStreamReadInfo_) {
     142            0 :         u32 destRank = info.first;
     143            0 :         u32 streamIndex = destRank % sdmaConcurrentNum_;
     144            0 :         res += std::to_string(streamIndex) + ", ";
     145              :     }
     146            0 :     return res;
     147            0 : }
     148              : 
     149            0 : u64 AlltoAllVDirectFullMesh::CalcMaxSendLen()
     150              : {
     151            0 :     u64 maxSendLen = 0;
     152            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
     153              : 
     154            0 :     for (u32 dstRank = 0; dstRank < localSendRecvInfo.sendLength.size(); dstRank++) {
     155            0 :         maxSendLen = std::max(maxSendLen, localSendRecvInfo.sendLength[dstRank]);
     156              :     }
     157              : 
     158            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][CalcMaxSendLen] maxSendLen[%llu]", maxSendLen);
     159            0 :     return maxSendLen;
     160              : }
     161              : 
     162            0 : HcclResult AlltoAllVDirectFullMesh::UpdateCurrRankRecvInfo(u32 step, u32 roundIdx, u32 side, u32 destRank,
     163              :     std::vector<ReadDataBlock>& readInfo, std::unordered_map<u32, ReadDataBlock>& subStreamZcopyReadInfo, u32 maxRecvStep)
     164              : {
     165            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
     166            0 :     u64 remainRecvLen = localSendRecvInfo.recvLength[destRank];
     167            0 :     u64 scratchOffset = 0;
     168            0 :     u32 bufferIdx = 0;
     169            0 :     u32 pairNum = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
     170            0 :     if (sdmaConcurrentNum_ == 1) { // 保证和当前rank距离一样时,send/recv用的是同一块buff
     171            0 :         bufferIdx = 0;
     172            0 :     } else if (side == 0) { // 在curRank左边
     173            0 :         u32 gap = (userRank_ - destRank + devNumInlocalPod_) % devNumInlocalPod_;
     174            0 :         bufferIdx = pairNum - (gap - roundIdx * pairNum);
     175            0 :     } else if (side == 1) { // 在curRank右边
     176            0 :         u32 gap = (destRank - userRank_ + devNumInlocalPod_) % devNumInlocalPod_;
     177            0 :         bufferIdx = pairNum - 1 + (gap - roundIdx * pairNum);
     178              :     } else { // 最后一个中间位置的rank
     179            0 :         bufferIdx = 0;
     180              :     }
     181              : 
     182            0 :     if ((isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC) &&
     183            0 :         (roundIdx % RANK_SET_COMPUTE_CONST != 0)) { // 奇数轮,用下半Buffer
     184            0 :         bufferIdx += sdmaConcurrentNum_;
     185              :     }
     186              : 
     187            0 :     scratchOffset = bufferIdx * sdmaDataBlockSize_;
     188              : 
     189            0 :     u32 recvStepIdx = 0;
     190            0 :     u64 dataOffset = 0;
     191            0 :     HCCL_DEBUG("step[%u] round[%u] usrRank[%u] total recv localSendRecvInfo.recvLength[%llu] from dstRank[%u] bufferIdx[%u]",
     192              :         step, roundIdx, userRank_, remainRecvLen, destRank, bufferIdx);
     193              : 
     194              :     // alltoallv类算子的零长拷贝, 需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发 (cache不使能时为空函数调用)
     195            0 :     if (needAlltoallvCache_ && remainRecvLen == 0) {
     196              :         // 获取local user output offset
     197            0 :         const u64 recvLen = 0;
     198            0 :         u64 userOutOffset = localSendRecvInfo.recvOffset[destRank];
     199            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
     200              :                 "recvStepIdx[%u] recvLen[%lu] userOutOffset[%llu] scratchOffset[%llu]",
     201              :                 userRank_, destRank, recvStepIdx, recvLen, userOutOffset, scratchOffset);
     202              :         
     203              :         // 更新零长拷贝的read info
     204            0 :         ReadDataBlock readBlock = {recvLen, scratchOffset, userOutOffset};
     205            0 :         subStreamZcopyReadInfo[destRank] = readBlock;
     206              : 
     207              :         // sendCount为0, step和readInfo.size一定为0
     208            0 :         CHK_PRT_RET(maxRecvStep > 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] maxRecvStep[%u] != 0 for remainRecvLen[%llu]", maxRecvStep, remainRecvLen), HCCL_E_INTERNAL);
     209            0 :         CHK_PRT_RET(readInfo.size() != 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] invalid readInfo.size[%u]", readInfo.size()), HCCL_E_INTERNAL);
     210            0 :     } else {
     211            0 :         while(recvStepIdx < maxRecvStep && remainRecvLen > 0) {
     212            0 :             u64 currDataRemainLen = localSendRecvInfo.recvLength[destRank] - dataOffset;
     213            0 :             u64 recvLen = std::min(sdmaDataBlockSize_, currDataRemainLen);
     214            0 :             u64 userOutOffset = localSendRecvInfo.recvOffset[destRank] + dataOffset;
     215            0 :             HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
     216              :                 "recvStepIdx[%u] recvLen[%lu] userOutOffset[%llu] scratchOffset[%llu]",
     217              :                 userRank_, destRank, recvStepIdx, recvLen, userOutOffset, scratchOffset);
     218            0 :             readInfo.push_back({recvLen, scratchOffset, userOutOffset});
     219            0 :             dataOffset += recvLen;
     220            0 :             recvStepIdx++;
     221            0 :             remainRecvLen -= recvLen;
     222              :         }
     223              :     }
     224              : 
     225            0 :     return HCCL_SUCCESS;
     226              : }
     227              : 
     228            0 : HcclResult AlltoAllVDirectFullMesh::UpdateCurrRankSendInfo(u32 step, u32 roundIdx, u32 side, u32 destRank,
     229              :     std::vector<SendDataBlock>& sendInfo, std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo, u32 maxSendStep)
     230              : {
     231            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
     232            0 :     u64 remainSendLen = localSendRecvInfo.sendLength[destRank];
     233              : 
     234            0 :     u64 scratchOffset = 0;
     235            0 :     u32 bufferIdx = 0;
     236            0 :     u32 pairNum = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
     237            0 :     if (sdmaConcurrentNum_ == 1) { // 保证和当前rank距离一样时,send/recv用的是同一块buff
     238            0 :         bufferIdx = 0;
     239            0 :     } else if (side == 0) { // 在curRank左边
     240            0 :         u32 gap = (userRank_ - destRank + devNumInlocalPod_) % devNumInlocalPod_;
     241            0 :         bufferIdx = pairNum - 1 + (gap - roundIdx * pairNum);
     242            0 :     } else if (side == 1) { // 在curRank右边
     243            0 :         u32 gap = (destRank - userRank_ + devNumInlocalPod_) % devNumInlocalPod_;
     244            0 :         bufferIdx = pairNum - (gap - roundIdx * pairNum);
     245              :     } else { // 最后一个中间位置的rank
     246            0 :         bufferIdx = 0;
     247              :     }
     248              : 
     249            0 :     if ((isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC) &&
     250            0 :         (roundIdx % RANK_SET_COMPUTE_CONST != 0)) { // 奇数轮,用下半Buffer
     251            0 :         bufferIdx += sdmaConcurrentNum_;
     252              :     }
     253            0 :     scratchOffset = bufferIdx * sdmaDataBlockSize_;
     254              : 
     255              :     // 更新hcclOffset到dstRank的映射, 用于alltoallv算子aicpu展开的SQE缓存
     256            0 :     if (needAlltoallvCache_) {
     257              :         // alltoallv cache只针对小数据量, 至多只有1个step
     258            0 :         CHK_PRT_RET(step != 0,
     259              :             HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] needAlltoallvCache_[%u] step[%u]",
     260              :                 needAlltoallvCache_, step),
     261              :             HCCL_E_INTERNAL);
     262              : 
     263            0 :         std::unordered_map<uint64_t, std::vector<uint32_t>>::iterator mapIter = hcclOffsetDstRanksMap_.find(scratchOffset);
     264            0 :         if (mapIter == hcclOffsetDstRanksMap_.end()) {
     265            0 :             constexpr uint32_t singleRankVecSize = 1;
     266            0 :             std::pair<std::unordered_map<uint64_t, std::vector<uint32_t>>::iterator, bool> emplaceResult = hcclOffsetDstRanksMap_.emplace(scratchOffset, std::vector<uint32_t>(singleRankVecSize, destRank));
     267            0 :             CHK_PRT_RET(!emplaceResult.second, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] fail to insert hcclOffset[%llu]-dstRank[%u] pair", scratchOffset, destRank), HCCL_E_INTERNAL);
     268            0 :             mapIter = emplaceResult.first;
     269              :         } else {
     270              :             // 虽然同一个dstRank不需要重复计算sendInfo, 但不同dstRanks在multi-round case下可能对应相同的hcclOffset
     271            0 :             CHK_PRT_RET(mapIter->second.size() == 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] empty dstRanks for hcclOffset[%llu] before add destRank[%u]", mapIter->second, destRank, scratchOffset), HCCL_E_INTERNAL);
     272            0 :             mapIter->second.push_back(destRank);
     273              :         }
     274            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] mapIter->first[%llu] mapIter->second.size[%u] destRank[%u]", mapIter->first, mapIter->second.size(), destRank);
     275              :     }
     276              : 
     277            0 :     u32 sendStepIdx = 0;
     278            0 :     u64 dataOffset = 0;
     279            0 :     HCCL_DEBUG("step[%u] round[%u] usrRank[%u] total send localSendRecvInfo.sendLength[%llu] to dstRank[%u] bufferIdx[%u]",
     280              :         step, roundIdx, userRank_, remainSendLen, destRank, bufferIdx);
     281              : 
     282            0 :     if (needAlltoallvCache_ && remainSendLen == 0) { // alltoallv类算子的零长拷贝, 需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发 (cache不使能时为空函数调用)
     283              :         // 获取local user input offset
     284            0 :         const u64 sendLen = 0;
     285            0 :         u64 userInOffset = localSendRecvInfo.sendOffset[destRank];
     286            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] usrRank[%u] send to destRank [%u]"
     287              :             " sendStepIdx[%u] sendLen[%lu] userInOffset[%llu] scratchOffset[%llu]",
     288              :             userRank_, destRank, sendStepIdx, sendLen, userInOffset, scratchOffset);
     289              :         
     290              :         // 更新零长拷贝的send info
     291            0 :         SendDataBlock sendBlock = {sendLen, userInOffset, scratchOffset};
     292            0 :         subStreamZcopySendInfo[destRank] = sendBlock;
     293              : 
     294              :         // sendCount为0, step和sendInfo.size一定为0
     295            0 :         CHK_PRT_RET(maxSendStep > 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] maxSendStep[%u] != 0 for remainSendLen[%llu]", maxSendStep, remainSendLen), HCCL_E_INTERNAL);
     296            0 :         CHK_PRT_RET(sendInfo.size() != 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] invalid sendInfo.size[%u]", sendInfo.size()), HCCL_E_INTERNAL);
     297            0 :     } else {
     298            0 :         while (sendStepIdx < maxSendStep && remainSendLen > 0) {
     299            0 :             u64 currDataRemainLen = localSendRecvInfo.sendLength[destRank] - dataOffset;
     300            0 :             u64 sendLen = std::min(sdmaDataBlockSize_, currDataRemainLen);
     301            0 :             u64 userInOffset = localSendRecvInfo.sendOffset[destRank] + dataOffset;
     302            0 :             HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] usrRank[%u] send to destRank [%u]"
     303              :                 " sendStepIdx[%u] sendLen[%lu] userInOffset[%llu] scratchOffset[%llu]",
     304              :                 userRank_, destRank, sendStepIdx, sendLen, userInOffset, scratchOffset);
     305            0 :             sendInfo.push_back({sendLen, userInOffset, scratchOffset});
     306            0 :             dataOffset += sendLen;
     307            0 :             sendStepIdx++;
     308            0 :             remainSendLen -= sendLen;
     309              :         }
     310              :     }
     311              : 
     312            0 :     return HCCL_SUCCESS;
     313              : }
     314              : 
     315            0 : void AlltoAllVDirectFullMesh::UpdateSendRecvInfo(u32 step, u32 roundIdx,
     316              :     std::unordered_map<u32, std::vector<ReadDataBlock>> &subStreamReadInfo,
     317              :     std::unordered_map<u32, std::vector<SendDataBlock>> &subStreamSendInfo,
     318              :     std::unordered_map<u32, ReadDataBlock>& subStreamZcopyReadInfo,
     319              :     std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo,
     320              :     const std::vector<std::vector<std::pair<u32,u32>>> &partialCommRankSet)
     321              : {
     322            0 :     for (u32 side = 0; side < partialCommRankSet.size(); side++) {
     323            0 :         for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
     324            0 :             u32 readRemoteRank = partialCommRankSet[side][j].first;
     325            0 :             if (readRemoteRank == userRank_) {
     326            0 :                 continue;
     327              :             }
     328            0 :             u32 currDestRecvStep = recvNumSubStep_[readRemoteRank];
     329            0 :             std::vector<ReadDataBlock> readInfo;
     330            0 :             UpdateCurrRankRecvInfo(step, roundIdx, side, readRemoteRank, readInfo, subStreamZcopyReadInfo, currDestRecvStep);
     331              : 
     332            0 :             subStreamReadInfo[readRemoteRank] = readInfo;
     333            0 :         }
     334              :     }
     335              : 
     336            0 :     for (u32 side = 0; side < partialCommRankSet.size(); side++) {
     337            0 :         for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
     338            0 :             u32 sendRemoteRank = partialCommRankSet[side][j].second;
     339            0 :             if (sendRemoteRank == userRank_) {
     340            0 :                 continue;
     341              :             }
     342            0 :             u32 currDestSendStep = sendNumSubStep_[sendRemoteRank];
     343            0 :             std::vector<SendDataBlock> sendInfo;
     344            0 :             UpdateCurrRankSendInfo(step, roundIdx, side, sendRemoteRank, sendInfo, subStreamZcopySendInfo, currDestSendStep);
     345              : 
     346            0 :             subStreamSendInfo[sendRemoteRank] = sendInfo;
     347            0 :         }
     348              :     }
     349            0 : }
     350              : 
     351            0 : void AlltoAllVDirectFullMesh::UpdateOpBaseSubStreamInfo(u32 step, u32 roundIdx)
     352              : {
     353            0 :     if (roundIdx == 0 || !isBigCount_) {
     354            0 :         subStreamReadInfo_.clear();
     355            0 :         subStreamSendInfo_.clear();
     356            0 :         if (needAlltoallvCache_) {
     357            0 :             subStreamZcopyReadInfo_.clear();
     358            0 :             subStreamZcopySendInfo_.clear();
     359              :         }
     360            0 :         UpdateSendRecvInfo(step, roundIdx, subStreamReadInfo_, subStreamSendInfo_, subStreamZcopyReadInfo_, subStreamZcopySendInfo_, partialCommRankSet_);
     361              :     }
     362            0 :     if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
     363            0 :         nextSubStreamReadInfo_.clear();
     364            0 :         nextSubStreamSendInfo_.clear();
     365            0 :         if (needAlltoallvCache_) {
     366            0 :             nextSubStreamZcopyReadInfo_.clear();
     367            0 :             nextSubStreamZcopySendInfo_.clear();
     368              :         }
     369            0 :         UpdateSendRecvInfo(step, roundIdx + 1, nextSubStreamReadInfo_, nextSubStreamSendInfo_, nextSubStreamZcopyReadInfo_, nextSubStreamZcopySendInfo_, nextPartialCommRankSet_);
     370              :     }
     371            0 : }
     372              : 
     373            0 : HcclResult AlltoAllVDirectFullMesh::PrepareIntraData(u32 step,
     374              :     std::unordered_map<u32,std::vector<SendDataBlock>> &subStreamSendInfo,
     375              :     std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo)
     376              : {
     377            0 :     u32 sendDataIndex = 0;
     378            0 :     for (auto& sdmaInfo : subStreamSendInfo) {
     379            0 :         const std::vector<SendDataBlock>& sendInfo = sdmaInfo.second;
     380              : 
     381              :         // 对于alltoallv类算子, 零长拷贝需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发
     382              :         // 注意: alltoallv aicpu cache只考虑小数据量 (即max step为1), 所以只需要在step 0时下发一个placeholder SQE即可
     383            0 :         if (needAlltoallvCache_) {
     384              :             // alltoallv cache只针对小数据量, 至多只有1个step
     385            0 :             CHK_PRT_RET(step != 0,
     386              :                 HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] needAlltoallvCache_[%u] step[%u]",
     387              :                     needAlltoallvCache_, step),
     388              :                 HCCL_E_INTERNAL);
     389              : 
     390            0 :             const u32 sendRank = sdmaInfo.first;
     391            0 :             std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo.find(sendRank);
     392            0 :             if (mapIter != subStreamZcopySendInfo.end()) { // sendRank的sendCount为0
     393              :                 // 零长拷贝下, sendRank对应的step和sendInfo.size一定为0
     394            0 :                 CHK_PRT_RET(sendNumSubStep_[sdmaInfo.first] > 0, HCCL_ERROR("invalid sendNumSubStep_[%u][%u] != 0", sdmaInfo.first, sendNumSubStep_[sdmaInfo.first]), HCCL_E_INTERNAL);
     395            0 :                 CHK_PRT_RET(sendInfo.size() > 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][PrepareIntraData] invalid sendInfo.size[%u] != 0", sendInfo.size()), HCCL_E_INTERNAL);
     396              : 
     397              :                 // 获取零长拷贝的发送偏移
     398            0 :                 const SendDataBlock& sendBlock = mapIter->second;
     399            0 :                 CHK_PRT_RET(sendBlock.sendLen != 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][PrepareIntraData] invalid sendBlock.sendLen[%llu] != 0", sendBlock.sendLen), HCCL_E_INTERNAL);
     400              : 
     401              :                 // 强制调用HcclD2DMemcpyAsync下发cache-memcpy placeholder (aicpu cache使能时才会生效, 未使能时会直接返回)
     402            0 :                 DeviceMem src = userInput_.range(sendBlock.userInOffset, sendBlock.sendLen);
     403            0 :                 DeviceMem dst = cclInMem_.range(sendBlock.scratchOffset, sendBlock.sendLen);
     404            0 :                 HCCL_DEBUG("[AlltoAllVDirectFullMesh][PrepareIntraData]userRank [%u] copy from userInOffset[%llu]"
     405              :                     "len[%u] to scratchOffset [%llu]", userRank_, sendBlock.userInOffset, sendBlock.sendLen,
     406              :                     sendBlock.scratchOffset);
     407            0 :                 reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     408            0 :                 HCCL_INFO("[AlltoAllVDirectFullMesh][PrepareIntraData] generate cache-memcpy placeholder for sendRank[%u]", sendRank);
     409            0 :                 if (isBigCount_) {
     410            0 :                     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, localSubStream_[sendDataIndex]));
     411              :                 } else {
     412            0 :                     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
     413              :                 }
     414            0 :                 reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     415            0 :             }
     416              :         }
     417              :         
     418            0 :         if (step < sendNumSubStep_[sdmaInfo.first]) {
     419            0 :             DeviceMem src = userInput_.range(sendInfo[step].userInOffset, sendInfo[step].sendLen);
     420            0 :             DeviceMem dst = cclInMem_.range(sendInfo[step].scratchOffset, sendInfo[step].sendLen);
     421            0 :             HCCL_DEBUG("[AlltoAllVDirectFullMesh][PrepareIntraData]userRank [%u] copy from userInOffset[%llu]"
     422              :                 "len[%u] to scratchOffset [%llu]", userRank_, sendInfo[step].userInOffset, sendInfo[step].sendLen,
     423              :                 sendInfo[step].scratchOffset);
     424            0 :             if (isBigCount_) {
     425            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, localSubStream_[sendDataIndex]));
     426              :             } else {
     427            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
     428              :             }
     429            0 :         }
     430            0 :         sendDataIndex++;
     431              :     }
     432            0 :     return HCCL_SUCCESS;
     433              : }
     434              : 
     435            0 : void AlltoAllVDirectFullMesh::UpdateRemoteRankSet(u32 roundIdx, u32 groupRankSize)
     436              : {
     437            0 :     if (sdmaConcurrentNum_ == 1) {
     438            0 :         UpdatePartialCommunicationRankSetPairWise(roundIdx, groupRankSize);
     439              :     } else {
     440            0 :         UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
     441              :     }
     442            0 : }
     443              : 
     444            0 : void AlltoAllVDirectFullMesh::UpdatePartialCommunicationRankSetPairWise(u32 roundIdx, u32 groupRankSize)
     445              : {
     446            0 :     partialCommRankSet_.clear();
     447            0 :     partialCommRankSet_.resize(1);
     448            0 :     for (u32 i = roundIdx * sdmaConcurrentNum_; i < (roundIdx * sdmaConcurrentNum_ + groupRankSize); i++) {
     449            0 :         u32 readRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
     450            0 :         u32 sendRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
     451            0 :         partialCommRankSet_[0].push_back(std::make_pair(readRemoteRank, sendRemoteRank));
     452            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSetPairWise] userRank [%u] i[%u]" \
     453              :             "readRemoteRank[%u] writeRemoteRank[%u]", userRank_, i, readRemoteRank, sendRemoteRank);
     454              :     }
     455            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSetPairWise] partialCommRankSet_ size[%zu]",
     456              :         partialCommRankSet_[0].size());
     457            0 : }
     458              : 
     459            0 : void AlltoAllVDirectFullMesh::UpdatePartialCommunicationRankSet(u32 roundIdx, u32 groupRankSize,
     460              :     std::vector<std::vector<std::pair<u32,u32>>> &partialCommRankSet)
     461              : {
     462            0 :     partialCommRankSet.clear();
     463            0 :     partialCommRankSet.resize(RANK_SET_COMPUTE_CONST + 1);
     464            0 :     u32 pairNumPerRound = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
     465            0 :     u32 pairSize = (groupRankSize < sdmaConcurrentNum_) ?
     466            0 :         (groupRankSize + RANK_SET_COMPUTE_CONST - 1) / RANK_SET_COMPUTE_CONST: pairNumPerRound;
     467            0 :     for (u32 i = roundIdx * pairNumPerRound + 1;
     468            0 :          i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
     469            0 :         u32 leftRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
     470            0 :         u32 rightRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
     471            0 :         if (leftRemoteRank == rightRemoteRank) {
     472            0 :             partialCommRankSet[2].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
     473              :         } else {
     474            0 :             partialCommRankSet[0].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
     475            0 :             partialCommRankSet[1].push_back(std::make_pair(rightRemoteRank, rightRemoteRank));
     476              :         }
     477            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSet] round[%u] userRank [%u] i[%u]" \
     478              :             "read/write leftRemoteRank[%u] rightRemoteRank[%u]", roundIdx, userRank_, i, leftRemoteRank, rightRemoteRank);
     479              :     }
     480            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSet] round[%u] partialCommRankSet_ total size[%zu]",
     481              :         roundIdx, partialCommRankSet[0].size() + partialCommRankSet[1].size() + partialCommRankSet[2].size());
     482            0 : }
     483              : 
     484              : // 主流只需要通知当前子步骤需要收发数据的 SDMA 流,减少同步开销
     485            0 : HcclResult AlltoAllVDirectFullMesh::NotifySubStreamStart()
     486              : {
     487            0 :     for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
     488            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
     489            0 :         CHK_RET(LocalNotify::Wait(sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     490              :             INVALID_VALUE_STAGE));
     491              :     }
     492            0 :     for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
     493            0 :         CHK_RET(ExecEmptyTask(userInput_, userOutput_, sdmaSubStream_[streamIndex], dispatcher_));
     494              :     }
     495            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][NotifySubStreamStart] userRank [%u] main stream notify sdma stream [%s]",
     496              :         userRank_, GetStreamIndexString().c_str());
     497            0 :     return HCCL_SUCCESS;
     498              : }
     499              : 
     500            0 : HcclResult AlltoAllVDirectFullMesh::WaitSubStreamFinish()
     501              : {
     502            0 :     for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
     503            0 :         CHK_RET(LocalNotify::Post(sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     504              :             INVALID_VALUE_STAGE));
     505            0 :         CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     506              :             INVALID_VALUE_STAGE));
     507              :     }
     508            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][WaitSubStreamFinish] userRank [%u] main stream wait sdma stream [%s]",
     509              :         userRank_, GetStreamIndexString().c_str());
     510            0 :     return HCCL_SUCCESS;
     511              : }
     512              : 
     513            0 : HcclResult AlltoAllVDirectFullMesh::NotifyLocalSubStreamStart()
     514              : {
     515            0 :     for (u32 streamIndex = 0; streamIndex < subStreamSendInfo_.size(); streamIndex++) {
     516            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, localSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
     517            0 :         CHK_RET(LocalNotify::Wait(localSubStream_[streamIndex], dispatcher_, localSignalSubToMain_[streamIndex],
     518              :             INVALID_VALUE_STAGE));
     519              :     }
     520            0 :     return HCCL_SUCCESS;
     521              : }
     522              : 
     523            0 : HcclResult AlltoAllVDirectFullMesh::WaitLocalSubStreamFinish()
     524              : {
     525            0 :     for (u32 streamIndex = 0; streamIndex < subStreamSendInfo_.size(); streamIndex++) {
     526            0 :         CHK_RET(LocalNotify::Post(localSubStream_[streamIndex], dispatcher_, localSignalMainToSub_[streamIndex],
     527              :             INVALID_VALUE_STAGE));
     528            0 :         CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, localSignalMainToSub_[streamIndex],
     529              :             INVALID_VALUE_STAGE));
     530              :     }
     531            0 :     return HCCL_SUCCESS;
     532              : }
     533              : 
     534            0 : u32 AlltoAllVDirectFullMesh::CalcNumSubStep()
     535              : {
     536            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
     537              : 
     538            0 :     sendNumSubStep_.clear();
     539            0 :     recvNumSubStep_.clear();
     540            0 :     u32 numSubStep = 0;
     541              : 
     542            0 :     for (u32 destRank = podStartRank_; destRank < podStartRank_ + devNumInlocalPod_; destRank++) {
     543            0 :         if (destRank == userRank_) {
     544            0 :             continue;
     545              :         }
     546              : 
     547            0 :         u32 currRankSendSubStep = ((localSendRecvInfo.sendLength[destRank] + sdmaDataBlockSize_- 1) / sdmaDataBlockSize_);
     548            0 :         sendNumSubStep_[destRank] = currRankSendSubStep;
     549              : 
     550            0 :         u32 currRankRecvSubStep = ((localSendRecvInfo.recvLength[destRank] + sdmaDataBlockSize_- 1) / sdmaDataBlockSize_);
     551            0 :         recvNumSubStep_[destRank] = currRankRecvSubStep;
     552            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][CalcNumSubStep] userRank [%u] currRankSendSubStep[%u]" \
     553              :         "currRankRecvSubStep[%u]", userRank_, currRankSendSubStep, currRankRecvSubStep);
     554            0 :         numSubStep = std::max(numSubStep, std::max(currRankSendSubStep, currRankRecvSubStep));
     555              :     }
     556            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][CalcNumSubStep] userRank [%u] max communication step[%u]",
     557              :         userRank_, numSubStep);
     558            0 :     return numSubStep;
     559              : }
     560              : 
     561            0 : HcclResult AlltoAllVDirectFullMesh::NotifyRemoteRankStart(u32 step)
     562              : {
     563            0 :     u32 streamIndex = 0;
     564            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     565            0 :         for (auto& sendRecvPair : sendRecvSide) {
     566            0 :             u32 recvRank = sendRecvPair.first;
     567            0 :             u32 sendRank = sendRecvPair.second;
     568            0 :             if (sendRank == userRank_) {
     569            0 :                 continue;
     570              :             }
     571            0 :             const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
     572            0 :             const std::vector<SendDataBlock>& sendInfo = subStreamSendInfo_[sendRank];
     573            0 :             Stream& currStream = sdmaSubStream_[streamIndex];
     574            0 :             const LINK& readTransport = links_[recvRank];
     575            0 :             const LINK& sendTransport = links_[sendRank];
     576              :             
     577            0 :             if (needAlltoallvCache_) {
     578              :                 // alltoallv cache只针对小数据量, 至多只有1个step
     579            0 :                 CHK_PRT_RET(step != 0,
     580              :                     HCCL_ERROR("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] needAlltoallvCache_[%u] step[%u]",
     581              :                         needAlltoallvCache_, step),
     582              :                     HCCL_E_INTERNAL);
     583              : 
     584            0 :                 std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo_.find(sendRank);
     585            0 :                 if (mapIter != subStreamZcopySendInfo_.end()) { // sendRank的sendCount为0
     586              :                     // 零长拷贝下, sendRank对应的sendInfo.size一定为0
     587            0 :                     CHK_PRT_RET(sendInfo.size() > 0,
     588              :                         HCCL_ERROR("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] invalid sendInfo.size[%u] != 0",
     589              :                             sendInfo.size()),
     590              :                         HCCL_E_INTERNAL);
     591              : 
     592              :                     // 生成cache-write placeholder
     593            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     594            0 :                     HCCL_INFO("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] generate cache-write placeholder for sendRank[%u]", sendRank);
     595            0 :                     CHK_RET(sendTransport->TxAck(currStream));
     596            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     597              :                 }
     598              :             }
     599            0 :             if (step < sendInfo.size()) {
     600            0 :                 CHK_RET(sendTransport->TxAck(currStream));
     601              :             }
     602              : 
     603            0 :             if (needAlltoallvCache_) {
     604            0 :                 std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
     605            0 :                 if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
     606              :                     // 零长拷贝下, recvRank对应的readInfo.size一定为0
     607            0 :                     CHK_PRT_RET(readInfo.size() > 0,
     608              :                         HCCL_ERROR("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] invalid readInfo.size[%u] != 0",
     609              :                             readInfo.size()),
     610              :                         HCCL_E_INTERNAL);
     611              : 
     612              :                     // 生成cache-write placeholder
     613            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     614            0 :                     HCCL_INFO("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] generate cache-notify placeholder for recvRank[%u]", recvRank);
     615            0 :                     CHK_RET(readTransport->RxAck(currStream));
     616            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     617              :                 }
     618              :             }
     619            0 :             if (step < readInfo.size()) {
     620            0 :                 CHK_RET(readTransport->RxAck(currStream));
     621              :             }
     622            0 :             streamIndex ++;
     623              :         }
     624              :     }
     625            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] done");
     626            0 :     return HCCL_SUCCESS;
     627              : }
     628              : 
     629            0 : bool AlltoAllVDirectFullMesh::IsPostSyncEnable(u32 step, u32 roundIdx)
     630              : {
     631            0 :     bool isPostSyncEnable = false;
     632            0 :     isPostSyncEnable = (step == lastStep_) && (roundIdx == lastRoundIdx_) &&
     633            0 :         algOpContext_.opRetryHandler.retryEnable;
     634            0 :     return isPostSyncEnable;
     635              : }
     636              : 
     637            0 : HcclResult AlltoAllVDirectFullMesh::SdmaMainStreamWait(u32 step, u32 roundIdx)
     638              : {
     639              :     // SDMA wait
     640            0 :     u32 streamIndex = 0;
     641            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     642            0 :         for (auto& sendRecvPair : sendRecvSide) {
     643            0 :             u32 recvRank = sendRecvPair.first;
     644            0 :             u32 sendRank = sendRecvPair.second;
     645            0 :             if (sendRank == userRank_) {
     646            0 :                 continue;
     647              :             }
     648            0 :             const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
     649              : 
     650            0 :             if (needAlltoallvCache_) {
     651              :                 // alltoallv cache只针对小数据量, 至多只有1个step
     652            0 :                 CHK_PRT_RET(step != 0,
     653              :                     HCCL_ERROR("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] needAlltoallvCache_[%u] step[%u]",
     654              :                         needAlltoallvCache_, step),
     655              :                     HCCL_E_INTERNAL);
     656              : 
     657            0 :                 std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
     658            0 :                 if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
     659              :                     // 零长拷贝下, recvRank对应的readInfo.size一定为0
     660            0 :                     CHK_PRT_RET(readInfo.size() > 0,
     661              :                         HCCL_ERROR("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] invalid readInfo.size[%u] != 0",
     662              :                             readInfo.size()),
     663              :                         HCCL_E_INTERNAL);
     664              :                     
     665              :                     // 正常下NotifyWait SQE (本地主从流同步, 由于从流不存在跨卡数据搬运, 主流wait后会立刻wake up)
     666            0 :                     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
     667              :                         "sendRank[%u], sdma stream [%u], "
     668              :                         "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream wait",
     669              :                         userRank_,  recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
     670            0 :                     CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     671              :                         INVALID_VALUE_STAGE));
     672              :                 }
     673              :             }
     674              : 
     675            0 :             if (step < readInfo.size()) {
     676            0 :                 HCCL_DEBUG("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
     677              :                     "sendRank[%u], sdma stream [%u], "
     678              :                     "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream wait",
     679              :                     userRank_,  recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
     680            0 :                 CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     681              :                     INVALID_VALUE_STAGE));
     682              :             }
     683            0 :             streamIndex ++;
     684              :         }
     685              :     }
     686            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] done");
     687            0 :     return HCCL_SUCCESS;
     688              : }
     689              : 
     690            0 : HcclResult AlltoAllVDirectFullMesh::SdmaMainStreamPost(u32 step, u32 roundIdx)
     691              : {
     692              :     // SDMA post
     693            0 :     u32 streamIndex = 0;
     694            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     695            0 :         for (auto& sendRecvPair : sendRecvSide) {
     696            0 :             u32 recvRank = sendRecvPair.first;
     697            0 :             u32 sendRank = sendRecvPair.second;
     698            0 :             if (sendRank == userRank_) {
     699            0 :                 continue;
     700              :             }
     701            0 :             const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
     702              : 
     703            0 :             if (needAlltoallvCache_) {
     704              :                 // alltoallv cache只针对小数据量, 至多只有1个step
     705            0 :                 CHK_PRT_RET(step != 0,
     706              :                     HCCL_ERROR("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] needAlltoallvCache_[%u] step[%u]",
     707              :                         needAlltoallvCache_, step),
     708              :                     HCCL_E_INTERNAL);
     709              : 
     710            0 :                 std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
     711            0 :                 if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
     712              :                     // 零长拷贝下, recvRank对应的readInfo.size一定为0
     713            0 :                     CHK_PRT_RET(readInfo.size() > 0,
     714              :                         HCCL_ERROR("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] invalid readInfo.size[%u] != 0",
     715              :                             readInfo.size()),
     716              :                         HCCL_E_INTERNAL);
     717              :                     
     718              :                     // 正常下NotifyRecord SQE
     719            0 :                     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
     720              :                         "sendRank[%u], sdma stream [%u], "
     721              :                         "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream post",
     722              :                         userRank_,  recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
     723            0 :                     CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     724              :                         INVALID_VALUE_STAGE));
     725              :                 }
     726              :             }
     727              : 
     728            0 :             if (step < readInfo.size()) {
     729            0 :                 HCCL_DEBUG("[AlltoAllVDirectFullMesh][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
     730              :                     "sendRank[%u], sdma stream [%u], "
     731              :                     "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream post",
     732              :                     userRank_,  recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
     733            0 :                 CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     734              :                     INVALID_VALUE_STAGE));
     735              :             }
     736            0 :             streamIndex ++;
     737              :         }
     738              :     }
     739            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][SdmaMainStreamPost] done");
     740            0 :     return HCCL_SUCCESS;
     741              : }
     742              : 
     743            0 : HcclResult AlltoAllVDirectFullMesh::SetPostSyncTasks(u32 step, u32 roundIdx)
     744              : {
     745              :     // SDMA wait
     746            0 :     CHK_RET(SdmaMainStreamWait(step, roundIdx));
     747            0 :     if (rdmaConcurrentNum_ > 0) {
     748              :         // RDMA wait
     749            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] rdma post sync info: main stream wait");
     750            0 :         CHK_RET(RdmaControlNotifyMainFinish());
     751              :     }
     752              :     // SDMA post
     753            0 :     CHK_RET(SdmaMainStreamPost(step, roundIdx));
     754            0 :     if (rdmaConcurrentNum_ > 0) {
     755              :         // RDMA post
     756            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] rdma post sync info: main stream post");
     757            0 :         CHK_RET(MainNotifyRdmaControlStart());
     758              :     }
     759            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] done");
     760            0 :     return HCCL_SUCCESS;
     761              : }
     762              : 
     763            0 : HcclResult AlltoAllVDirectFullMesh::SDMAwithRemoteRankAndNotifyEnd(u32 step, u32 roundIdx)
     764              : {
     765            0 :     bool isPostSyncEnable = IsPostSyncEnable(step, roundIdx);
     766            0 :     if (isPostSyncEnable) {
     767              :         // 下发主流上的后同步wait和post
     768            0 :         CHK_RET(SetPostSyncTasks(step, roundIdx));
     769              :     }
     770            0 :     u32 streamIndex = 0;
     771            0 :     for (auto& sendRecvSide : partialCommRankSet_) {
     772            0 :         for (auto& sendRecvPair : sendRecvSide) {
     773            0 :             u32 recvRank = sendRecvPair.first;
     774            0 :             u32 sendRank = sendRecvPair.second;
     775            0 :             if (sendRank == userRank_) {
     776            0 :                 continue;
     777              :             }
     778            0 :             const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
     779            0 :             const std::vector<SendDataBlock>& sendInfo = subStreamSendInfo_[sendRank];
     780            0 :             Stream& currStream = sdmaSubStream_[streamIndex];
     781            0 :             const LINK& readTransport = links_[recvRank];
     782            0 :             const LINK& sendTransport = links_[sendRank];
     783              : 
     784              :             // 对于alltoallv类算子, 零长拷贝需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发
     785              :             // 注意: alltoallv aicpu cache只考虑小数据量 (即max step为1), 所以只需要在step 0时下发一个placeholder SQE即可
     786            0 :             if (needAlltoallvCache_) {
     787              :                 // alltoallv cache只针对小数据量, 至多只有1个step
     788            0 :                 CHK_PRT_RET(step != 0,
     789              :                     HCCL_ERROR("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] needAlltoallvCache_[%u] step[%u]",
     790              :                         needAlltoallvCache_, step),
     791              :                     HCCL_E_INTERNAL);
     792              : 
     793            0 :                 std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
     794            0 :                 if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
     795              :                     // 零长拷贝下, recvRank对应的readInfo.size一定为0
     796            0 :                     CHK_PRT_RET(readInfo.size() > 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid readInfo.size[%u] != 0", readInfo.size()), HCCL_E_INTERNAL);
     797              : 
     798              :                     // 获取零长拷贝的接收偏移
     799            0 :                     const ReadDataBlock& readBlock = mapIter->second;
     800            0 :                     CHK_PRT_RET(readBlock.recvLen != 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid readBlock.recvLen[%llu] != 0", readBlock.recvLen), HCCL_E_INTERNAL);
     801              : 
     802              :                     // 强制调用HcclD2DMemcpyAsync下发cache-memcpy/write placeholder (aicpu cache使能时才会下发, 未使能时会直接返回)
     803            0 :                     const LINK& intraNeighboorTransport = links_[recvRank];
     804            0 :                     CHK_PTR_NULL(intraNeighboorTransport);
     805            0 :                     void* remDMAMemPtr = nullptr;
     806            0 :                     CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
     807            0 :                     DeviceMem remoteCCLInMem = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr), cclInMem_.size());
     808            0 :                     DeviceMem srcMem = remoteCCLInMem.range(readBlock.remoteOffset, readBlock.recvLen);
     809            0 :                     DeviceMem dstMem = userOutput_.range(readBlock.recvOffset, readBlock.recvLen);
     810            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     811            0 :                     HCCL_INFO("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-memcpy placeholder for recvRank[%u]", recvRank);
     812            0 :                     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, currStream,
     813              :                         readTransport->GetRemoteRank(), readTransport->GetLinkType()));
     814            0 :                     HCCL_INFO("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-write placeholder for recvRank[%u]", recvRank);
     815            0 :                     CHK_RET(readTransport->TxDataSignal(currStream));
     816            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     817              : 
     818              :                     // 正常下NotifyRecord/Wait SQE (本地主从流同步, 从流不存在跨卡数据拷贝, 下发placeholder后会立刻post主流并进入wait)
     819            0 :                     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] userRank [%u], recvRank[%u], sendRank[%u]," \
     820              :                         "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
     821              :                         "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u]",
     822              :                         userRank_,  recvRank, sendRank, streamIndex, readBlock.remoteOffset,
     823              :                         readBlock.recvLen, readBlock.recvOffset, step, roundIdx, lastStep_, lastRoundIdx_);
     824            0 :                     if (isPostSyncEnable) {
     825            0 :                         HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] post sync begins");
     826            0 :                         CHK_RET(LocalNotify::Post(currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     827              :                             INVALID_VALUE_STAGE));
     828            0 :                         CHK_RET(LocalNotify::Wait(currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     829              :                             INVALID_VALUE_STAGE));
     830              :                     }
     831            0 :                 }
     832              :             }
     833              :             
     834            0 :             if (step < readInfo.size()) {
     835            0 :                 const LINK& intraNeighboorTransport = links_[recvRank];
     836            0 :                 CHK_PTR_NULL(intraNeighboorTransport);
     837            0 :                 void* remDMAMemPtr = nullptr;
     838            0 :                 CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
     839            0 :                 DeviceMem remoteCCLInMem = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr), cclInMem_.size());
     840            0 :                 DeviceMem srcMem = remoteCCLInMem.range(readInfo[step].remoteOffset, readInfo[step].recvLen);
     841            0 :                 DeviceMem dstMem = userOutput_.range(readInfo[step].recvOffset, readInfo[step].recvLen);
     842            0 :                 CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, currStream,
     843              :                     readTransport->GetRemoteRank(), readTransport->GetLinkType()));
     844            0 :                 HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] userRank [%u], recvRank[%u], sendRank[%u]," \
     845              :                     "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
     846              :                     "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u]",
     847              :                     userRank_,  recvRank, sendRank, streamIndex, readInfo[step].remoteOffset,
     848              :                     readInfo[step].recvLen, readInfo[step].recvOffset, step, roundIdx, lastStep_, lastRoundIdx_);
     849            0 :                 if (isPostSyncEnable) {
     850            0 :                     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] post sync begins");
     851            0 :                     CHK_RET(LocalNotify::Post(currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex],
     852              :                         INVALID_VALUE_STAGE));
     853            0 :                     CHK_RET(LocalNotify::Wait(currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex],
     854              :                         INVALID_VALUE_STAGE));
     855              :                 }
     856            0 :                 CHK_RET(readTransport->TxDataSignal(currStream));
     857            0 :             }
     858              : 
     859            0 :             if (needAlltoallvCache_) {
     860            0 :                 std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo_.find(sendRank);
     861            0 :                 if (mapIter != subStreamZcopySendInfo_.end()) { // sendRank的sendCount为0
     862              :                     // 零长拷贝下, sendRank对应的sendInfo.size一定为0
     863            0 :                     CHK_PRT_RET(sendInfo.size() > 0,
     864              :                         HCCL_ERROR("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid sendInfo.size[%u] != 0",
     865              :                             sendInfo.size()),
     866              :                         HCCL_E_INTERNAL);
     867              : 
     868              :                     // 生成cache-notify placeholder
     869            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     870            0 :                     HCCL_INFO("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-notify placeholder for sendRank[%u]", sendRank);
     871            0 :                     CHK_RET(sendTransport->RxDataSignal(currStream));
     872            0 :                     reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     873              :                 }
     874              :             }
     875              : 
     876            0 :             if (step < sendInfo.size()) {
     877            0 :                 CHK_RET(sendTransport->RxDataSignal(currStream));
     878              :             }
     879            0 :             streamIndex ++;
     880              :         }
     881              :     }
     882            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] done");
     883            0 :     return HCCL_SUCCESS;
     884              : }
     885              : 
     886            0 : HcclResult AlltoAllVDirectFullMesh::SendRecvData(u32 step, u32 roundIdx)
     887              : {
     888            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvData] userRank [%u] sdma stream [%s] wait main stream",
     889              :         userRank_, GetStreamIndexString().c_str());
     890            0 :     CHK_RET(NotifyRemoteRankStart(step));
     891            0 :     CHK_RET(WaitSubStreamFinish());
     892            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     893            0 :     CHK_RET(NotifySubStreamStart());
     894            0 :     if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
     895            0 :         CHK_RET(NotifyLocalSubStreamStart());
     896            0 :         CHK_RET(PrepareIntraData(step, nextSubStreamSendInfo_, nextSubStreamZcopySendInfo_));
     897              :     }
     898            0 :     CHK_RET(SDMAwithRemoteRankAndNotifyEnd(step, roundIdx));
     899              : 
     900            0 :     return HCCL_SUCCESS;
     901              : }
     902              : 
     903            0 : HcclResult AlltoAllVDirectFullMesh::LocalCopy()
     904              : {
     905            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
     906            0 :     DeviceMem src = userInput_.range(localSendRecvInfo.sendOffset[userRank_],
     907            0 :         localSendRecvInfo.sendLength[userRank_]);
     908            0 :     DeviceMem dst = userOutput_.range(localSendRecvInfo.recvOffset[userRank_],
     909            0 :         localSendRecvInfo.recvLength[userRank_]);
     910            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][LocalCopy]userRank [%u] copy from userInput [%llu] len [%llu]" \
     911              :         "to userOutput [%llu] dstLen[%llu]", userRank_, localSendRecvInfo.sendOffset[userRank_],
     912              :         localSendRecvInfo.sendLength[userRank_],
     913              :         localSendRecvInfo.recvOffset[userRank_],
     914              :         localSendRecvInfo.recvLength[userRank_]);
     915            0 :     if (needAlltoallvCache_ && localSendRecvInfo.sendLength[userRank_] == 0) {
     916            0 :         reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
     917              :     }
     918            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
     919            0 :     if (needAlltoallvCache_ && localSendRecvInfo.sendLength[userRank_] == 0) {
     920            0 :         reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
     921              :     }
     922              : 
     923            0 :     return HCCL_SUCCESS;
     924            0 : }
     925              : 
     926            0 : HcclResult AlltoAllVDirectFullMesh::RunGroupFullMeshAlltoall(u32 roundIdx, u32 step)
     927              : {
     928            0 :     UpdateOpBaseSubStreamInfo(step, roundIdx);
     929            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     930            0 :     if (isBigCount_ && (roundIdx == 0) ) {
     931            0 :         CHK_RET(NotifyLocalSubStreamStart());
     932            0 :         CHK_RET(PrepareIntraData(step, subStreamSendInfo_, subStreamZcopySendInfo_));
     933            0 :         CHK_RET(WaitLocalSubStreamFinish());
     934            0 :         CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     935            0 :     } else if (!isBigCount_) {
     936            0 :         CHK_RET(PrepareIntraData(step, subStreamSendInfo_, subStreamZcopySendInfo_));
     937              :     }
     938            0 :     CHK_RET(NotifySubStreamStart());
     939            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     940            0 :     CHK_RET(SendRecvData(step, roundIdx));
     941            0 :     if (step == 0 && !islocalCpyDone_) {
     942            0 :         CHK_RET(LocalCopy());
     943            0 :         islocalCpyDone_ = true;
     944              :     }
     945            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     946            0 :     CHK_RET(WaitSubStreamFinish());
     947            0 :     if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
     948            0 :         CHK_RET(WaitLocalSubStreamFinish());
     949              :     }
     950            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
     951            0 :     return HCCL_SUCCESS;
     952              : }
     953              : 
     954              : // 主流通知RDMA控制流启动
     955            0 : HcclResult AlltoAllVDirectFullMesh::MainNotifyRdmaControlStart()
     956              : {
     957            0 :     CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, rdmaControl2MainStreamNotify_, INVALID_VALUE_STAGE));
     958            0 :     CHK_RET(LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2MainStreamNotify_, INVALID_VALUE_STAGE));
     959            0 :     return HCCL_SUCCESS;
     960              : }
     961              : 
     962              : // RDMA控制流通知主流任务完成
     963            0 : HcclResult AlltoAllVDirectFullMesh::RdmaControlNotifyMainFinish()
     964              : {
     965            0 :     CHK_RET(LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, main2RdmaControlStreamNotify_, INVALID_VALUE_STAGE));
     966            0 :     CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, main2RdmaControlStreamNotify_, INVALID_VALUE_STAGE));
     967            0 :     return HCCL_SUCCESS;
     968              : }
     969              : 
     970              : // RDMA控制流通知从流启动任务
     971            0 : HcclResult AlltoAllVDirectFullMesh::RdmaControlNotifySubStart()
     972              : {
     973            0 :     for (u32 i = 1; i < rdmaSubStreams_.size(); i++) {
     974            0 :         CHK_RET(LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, rdmaSub2ControlNotifies_[i-1], INVALID_VALUE_STAGE));
     975            0 :         CHK_RET(LocalNotify::Wait(rdmaSubStreams_[i], dispatcher_, rdmaSub2ControlNotifies_[i-1], INVALID_VALUE_STAGE));
     976              :     }
     977              : 
     978            0 :     return HCCL_SUCCESS;
     979              : }
     980              : 
     981              : // 从流通知RDMA控制流任务结束
     982            0 : HcclResult AlltoAllVDirectFullMesh::SubNotifyRdmaControlFinish()
     983              : {
     984            0 :     for (u32 i = 1; i < rdmaSubStreams_.size(); i++) {
     985            0 :         CHK_RET(LocalNotify::Post(rdmaSubStreams_[i], dispatcher_, rdmaControl2SubNotifies_[i-1], INVALID_VALUE_STAGE));
     986            0 :         CHK_RET(LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2SubNotifies_[i-1], INVALID_VALUE_STAGE));
     987              :     }
     988              : 
     989            0 :     return HCCL_SUCCESS;
     990              : }
     991              : 
     992            0 : u32 AlltoAllVDirectFullMesh::GetNextDstRank(u32& curDstRank)
     993              : {
     994            0 :     if (curDstRank >= userRankSize_) {
     995            0 :         curDstRank = curDstRank % userRankSize_;
     996              :     }
     997            0 :     if (curDstRank == podStartRank_) {
     998            0 :         curDstRank += devNumInlocalPod_;
     999              :     }
    1000            0 :     curDstRank = curDstRank % userRankSize_;
    1001            0 :     return curDstRank++;
    1002              : }
    1003              : 
    1004            0 : u32 AlltoAllVDirectFullMesh::GetPreSrcRank(u32& curDstRank)
    1005              : {
    1006            0 :     if (curDstRank == podStartRank_ + devNumInlocalPod_ - 1) {
    1007            0 :         curDstRank = (curDstRank + userRankSize_ - devNumInlocalPod_) % userRankSize_;
    1008              :     }
    1009              : 
    1010            0 :     if (curDstRank == 0) {
    1011            0 :         curDstRank = userRankSize_ - 1;
    1012            0 :         return 0;
    1013              :     }
    1014            0 :     return curDstRank--;
    1015              : }
    1016              : 
    1017            0 : void AlltoAllVDirectFullMesh::GenRdmaSendInfo(u32 dstRank, std::vector<SendDataBlock>& sendInfo)
    1018              : {
    1019            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
    1020            0 :     u64 sendOffset = localSendRecvInfo.sendOffset[dstRank];
    1021            0 :     u64 sendLength = localSendRecvInfo.sendLength[dstRank];
    1022            0 :     while (sendLength > 0) {
    1023            0 :         u64 curSendLength = std::min(sendLength, rdmaDataBlockSize_);
    1024              :         SendDataBlock sendData;
    1025            0 :         sendData.userInOffset = sendOffset;
    1026            0 :         sendData.sendLen = curSendLength;
    1027            0 :         u32 index = dstRank % rdmaConcurrentNum_;
    1028            0 :         sendData.scratchOffset = rdmaDataBlockSize_ * index;
    1029            0 :         sendInfo.push_back(sendData);
    1030            0 :         sendOffset += curSendLength;
    1031            0 :         sendLength -= curSendLength;
    1032            0 :         HCCL_DEBUG("[GenRdmaSendInfo] userRank[%u], dstRank[%u], sendData.userInOffset[%llu]," \
    1033              :             "sendData.sendLen[%llu], sendData.scratchOffset[%llu]", userRank_, dstRank,
    1034              :             sendData.userInOffset, sendData.sendLen, sendData.scratchOffset);
    1035              :     }
    1036            0 :     return;
    1037              : }
    1038              : 
    1039            0 : void AlltoAllVDirectFullMesh::GenRdmaRecvInfo(u32 srcRank, std::vector<RecvDataBlock>& recvInfo)
    1040              : {
    1041            0 :     const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
    1042            0 :     u64 recvOffset = localSendRecvInfo.recvOffset[srcRank];
    1043            0 :     u64 recvLength = localSendRecvInfo.recvLength[srcRank];
    1044            0 :     while (recvLength > 0) {
    1045            0 :         u64 curRecvLength = std::min(recvLength, rdmaDataBlockSize_);
    1046              :         RecvDataBlock recvData;
    1047            0 :         recvData.recvOffset = recvOffset;
    1048            0 :         recvData.recvLen = curRecvLength;
    1049            0 :         u32 index = srcRank % rdmaConcurrentNum_;
    1050            0 :         recvData.scratchOffset = rdmaDataBlockSize_ * index + rdmaDataBlockSize_ * rdmaConcurrentNum_;
    1051            0 :         recvInfo.push_back(recvData);
    1052            0 :         recvOffset += curRecvLength;
    1053            0 :         recvLength -= curRecvLength;
    1054            0 :         HCCL_DEBUG("[GenRdmaRecvInfo] userRank[%llu], srcRank[%u], recvData.recvOffset[%llu]," \
    1055              :             "recvData.recvLen[%llu], recvData.scratchOffset[%llu]", userRank_, srcRank,
    1056              :             recvData.recvOffset, recvData.recvLen, recvData.scratchOffset);
    1057              :     }
    1058            0 :     return;
    1059              : }
    1060              : 
    1061              : // 将数据从userIn拷贝到CCL out
    1062            0 : HcclResult AlltoAllVDirectFullMesh::CopyDataForSend(u32 dstRank, std::vector<SendDataBlock>& sendInfo, u32 curStep, Stream stream)
    1063              : {
    1064            0 :     if (curStep >= sendInfo.size()) {
    1065            0 :         return HCCL_SUCCESS;
    1066              :     }
    1067            0 :     DeviceMem src = userInput_.range(sendInfo[curStep].userInOffset, sendInfo[curStep].sendLen);
    1068            0 :     DeviceMem dst = cclOutMem_.range(sendInfo[curStep].scratchOffset, sendInfo[curStep].sendLen);
    1069            0 :     HCCL_DEBUG("[CopyDataForSend] userRank[%u], dstRank[%u], userInOffset[%llu], sendLen[%llu], scratchOffset[%llu]",
    1070              :         userRank_, dstRank, sendInfo[curStep].userInOffset, sendInfo[curStep].sendLen, sendInfo[curStep].scratchOffset);
    1071            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream));
    1072            0 :     return HCCL_SUCCESS;
    1073            0 : }
    1074              : 
    1075            0 : HcclResult AlltoAllVDirectFullMesh::RdmaPostSync(Stream& stream)
    1076              : {
    1077            0 :     CHK_RET(LocalNotify::Post(stream, dispatcher_, rdmaControl2SubNotifies_[0], INVALID_VALUE_STAGE));
    1078            0 :     CHK_RET(LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2SubNotifies_[0], INVALID_VALUE_STAGE));
    1079              : 
    1080            0 :     CHK_RET(LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, rdmaSub2ControlNotifies_[0], INVALID_VALUE_STAGE));
    1081            0 :     CHK_RET(LocalNotify::Wait(stream, dispatcher_, rdmaSub2ControlNotifies_[0], INVALID_VALUE_STAGE));
    1082            0 :     return HCCL_SUCCESS;
    1083              : }
    1084              : 
    1085              : // 从流完成RDMA数据的收发
    1086            0 : HcclResult AlltoAllVDirectFullMesh::SendRecvRdmaData(u32 dstRank, u32 srcRank, std::vector<SendDataBlock>& sendInfo,
    1087              :     std::vector<RecvDataBlock>& recvInfo, u32 round, u32 index, u32 curStep, Stream stream)
    1088              : {
    1089            0 :     const LINK& sendTransport = links_[dstRank];
    1090            0 :     const LINK& recvTransport = links_[srcRank];
    1091            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] userRank[%u], dstRank[%u], srcRank[%u]",
    1092              :         userRank_, dstRank, srcRank);
    1093            0 :     u32 minStep = std::min(sendInfo.size(), recvInfo.size());
    1094            0 :     CHK_PTR_NULL(sendTransport);
    1095            0 :     CHK_PTR_NULL(recvTransport);
    1096            0 :     if (curStep < minStep) {
    1097            0 :         CHK_RET(recvTransport->TxAck(stream));
    1098            0 :         CHK_RET(sendTransport->RxAck(stream));
    1099            0 :         u64 sendSrcOffset = (dstRank % rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1100            0 :         void* srcPtr = static_cast<u8 *>(cclOutMem_.ptr()) + sendSrcOffset;
    1101            0 :         u32 dstIndex = userRank_ % rdmaConcurrentNum_;
    1102            0 :         u64 sendDstOffset = (dstIndex + rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1103            0 :         CHK_RET(sendTransport->TxAsync(UserMemType::OUTPUT_MEM, sendDstOffset, srcPtr,
    1104              :             sendInfo[curStep].sendLen, stream));
    1105              : 
    1106            0 :         u64 recvDstOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1107            0 :         void* dstPtr = static_cast<u8 *>(cclOutMem_.ptr()) + recvDstOffset;
    1108            0 :         u64 recvSrcOffset = (userRank_ % rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1109            0 :         CHK_RET(recvTransport->RxAsync(UserMemType::OUTPUT_MEM, recvSrcOffset, dstPtr,
    1110              :             recvInfo[curStep].recvLen, stream));
    1111            0 :         if ((round == lastRdmaRoundIdx_) && (index == lastRdmaDstRanksIdx_) && (curStep == lastRdmaStep_) &&
    1112            0 :             (sdmaConcurrentNum_ > 1) && algOpContext_.opRetryHandler.retryEnable) {
    1113            0 :             HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] post sync begins");
    1114            0 :             CHK_RET(RdmaPostSync(stream));
    1115              :         }
    1116            0 :         CHK_RET(recvTransport->PostFinAck(stream));
    1117            0 :         CHK_RET(sendTransport->WaitFinAck(stream));
    1118            0 :         HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] sendSrcOffset[%llu], sendDstOffset[%llu]," \
    1119              :             "recvDstOffset[%llu], recvSrcOffset[%llu], srcPtr[%p], dstPtr[%p]",sendSrcOffset,
    1120              :             sendDstOffset, recvDstOffset, recvSrcOffset, srcPtr, dstPtr);
    1121            0 :     } else if (curStep < sendInfo.size()) {
    1122            0 :         CHK_RET(sendTransport->RxAck(stream));
    1123            0 :         u64 sendSrcOffset = (dstRank % rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1124            0 :         void* srcPtr = static_cast<u8 *>(cclOutMem_.ptr()) + sendSrcOffset;
    1125            0 :         u32 dstIndex = userRank_ % rdmaConcurrentNum_;
    1126            0 :         u64 sendDstOffset = (dstIndex + rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1127            0 :         CHK_RET(sendTransport->TxAsync(UserMemType::OUTPUT_MEM, sendDstOffset, srcPtr,
    1128              :             sendInfo[curStep].sendLen, stream));
    1129            0 :         CHK_RET(sendTransport->WaitFinAck(stream));
    1130              :     } else {
    1131            0 :         CHK_RET(recvTransport->TxAck(stream));
    1132            0 :         u64 recvDstOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1133            0 :         void* dstPtr = static_cast<u8 *>(cclOutMem_.ptr()) + recvDstOffset;
    1134            0 :         u64 recvSrcOffset = (userRank_ % rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1135            0 :         CHK_RET(recvTransport->RxAsync(UserMemType::OUTPUT_MEM, recvSrcOffset, dstPtr,
    1136              :             recvInfo[curStep].recvLen, stream));
    1137            0 :         if ((round == lastRdmaRoundIdx_) && (index == lastRdmaDstRanksIdx_) && (curStep == lastRdmaStep_) &&
    1138            0 :             (sdmaConcurrentNum_ > 1) && algOpContext_.opRetryHandler.retryEnable) {
    1139            0 :             HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] post sync begins");
    1140            0 :             CHK_RET(RdmaPostSync(stream));
    1141              :         }
    1142            0 :         CHK_RET(recvTransport->PostFinAck(stream));
    1143              :     }
    1144            0 :     return HCCL_SUCCESS;
    1145              : }
    1146              : 
    1147              : // 从流将接收到的数据拷贝到输出
    1148            0 : HcclResult AlltoAllVDirectFullMesh::CopyRecvDataToOutput(u32 srcRank, std::vector<RecvDataBlock>& recvInfo,
    1149              :     u32 curStep, Stream stream)
    1150              : {
    1151            0 :     if (curStep >= recvInfo.size()) {
    1152            0 :         return HCCL_SUCCESS;
    1153              :     }
    1154            0 :     u64 srcOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
    1155            0 :     DeviceMem src = cclOutMem_.range(srcOffset, recvInfo[curStep].recvLen);
    1156            0 :     DeviceMem dst = userOutput_.range(recvInfo[curStep].recvOffset, recvInfo[curStep].recvLen);
    1157            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][CopyRecvDataToOutput] userRank[%u], srcRank[%u], srcOffset[%llu]," \
    1158              :         "recvInfo[curStep].recvOffset[%llu], recvLen[%llu]", userRank_, srcRank, srcOffset,
    1159              :         recvInfo[curStep].recvOffset, recvInfo[curStep].recvLen);
    1160            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream));
    1161            0 :     return HCCL_SUCCESS;
    1162            0 : }
    1163              : 
    1164            0 : HcclResult AlltoAllVDirectFullMesh::ProcessSingleGroupRdmaData(std::vector<u32>& dstRanks, std::vector<u32>& srcRanks, u32 round)
    1165              : {
    1166            0 :     lastRdmaDstRanksIdx_ = dstRanks.size() - 1;
    1167            0 :     for (u32 index = 0; index < dstRanks.size(); index++) {
    1168            0 :         u32 dstRank = dstRanks[index];
    1169            0 :         u32 srcRank = srcRanks[index];
    1170            0 :         Stream stream = rdmaSubStreams_[index + 1];
    1171              : 
    1172            0 :         std::vector<SendDataBlock> sendInfo;
    1173            0 :         std::vector<RecvDataBlock> recvInfo;
    1174            0 :         GenRdmaSendInfo(dstRank, sendInfo);
    1175            0 :         GenRdmaRecvInfo(srcRank, recvInfo);
    1176            0 :         u32 totalStep = std::max(sendInfo.size(), recvInfo.size());
    1177            0 :         lastRdmaStep_ = totalStep - 1;
    1178            0 :         for (u32 curStep = 0; curStep < totalStep; curStep++) {
    1179            0 :             CHK_RET(CopyDataForSend(dstRank, sendInfo, curStep, stream));
    1180            0 :             CHK_RET(SendRecvRdmaData(dstRank, srcRank, sendInfo, recvInfo, round, index, curStep, stream));
    1181            0 :             CHK_RET(CopyRecvDataToOutput(srcRank, recvInfo, curStep, stream));
    1182              :         }
    1183            0 :     }
    1184              : 
    1185            0 :     return HCCL_SUCCESS;
    1186              : }
    1187              : 
    1188            0 : HcclResult AlltoAllVDirectFullMesh::ProcessRdmaData()
    1189              : {
    1190              :     // RDMA通信轮次
    1191            0 :     u32 rdmaRoundNum = (totalRdmaRankNum_ + rdmaConcurrentNum_ - 1) / rdmaConcurrentNum_;
    1192            0 :     lastRdmaRoundIdx_ = rdmaRoundNum - 1;
    1193              : 
    1194            0 :     u32 leftRankNum = totalRdmaRankNum_;
    1195            0 :     u32 curSrcRank = INVALID_VALUE_RANKID;
    1196            0 :     u32 curDstRank = INVALID_VALUE_RANKID;
    1197            0 :     if (isSuPodAsym_) {
    1198            0 :         for (u32 i = 0; i < userRankSize_; i++) {
    1199            0 :             if (i < podStartRank_ || i > podEndRank_) {
    1200            0 :                 curSrcRank = i;
    1201            0 :                 curDstRank = i;
    1202            0 :                 break;
    1203              :             }
    1204              :         }
    1205              :     } else {
    1206            0 :         curDstRank = (userRank_ + devNumInlocalPod_) % userRankSize_;
    1207            0 :         curSrcRank = (userRank_ + userRankSize_ - devNumInlocalPod_) % userRankSize_;
    1208              :     }
    1209              : 
    1210            0 :     for (u32 round = 0; round < rdmaRoundNum; round++) {
    1211            0 :         u32 curProcessRankNum = leftRankNum >= rdmaConcurrentNum_ ? rdmaConcurrentNum_ : leftRankNum;
    1212            0 :         leftRankNum -= curProcessRankNum;
    1213              : 
    1214            0 :         std::vector<u32> dstRanks;
    1215            0 :         std::vector<u32> srcRanks;
    1216            0 :         for (u32 i = 0; i < curProcessRankNum; i++) {
    1217            0 :             dstRanks.push_back(GetNextDstRank(curDstRank));
    1218            0 :             if (isSuPodAsym_) {
    1219            0 :                 srcRanks.push_back(dstRanks.back());
    1220              :             } else {
    1221            0 :                 srcRanks.push_back(GetPreSrcRank(curSrcRank));
    1222              :             }
    1223              :         }
    1224            0 :         CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
    1225            0 :         CHK_RET(RdmaControlNotifySubStart());
    1226            0 :         CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
    1227            0 :         CHK_RET(ProcessSingleGroupRdmaData(dstRanks, srcRanks, round));
    1228            0 :         CHK_RET(SubNotifyRdmaControlFinish());
    1229            0 :         CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
    1230            0 :     }
    1231            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][ProcessRdmaData] done");
    1232            0 :     return HCCL_SUCCESS;
    1233              : }
    1234              : 
    1235            0 : HcclResult AlltoAllVDirectFullMesh::RunRDMA()
    1236              : {
    1237              :     // 先启动RDMA通信
    1238            0 :     CHK_RET(MainNotifyRdmaControlStart());
    1239            0 :     CHK_RET(ProcessRdmaData());
    1240            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
    1241            0 :     CHK_RET(LocalCopy());
    1242            0 :     islocalCpyDone_ = true;
    1243            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][RunRDMA] finished.");
    1244            0 :     return HCCL_SUCCESS;
    1245              : }
    1246              : 
    1247            0 : HcclResult AlltoAllVDirectFullMesh::RunSDMATasks(u32 roundIdx, u32 step, u32 groupRankSize, u32 leftRankSize)
    1248              : {
    1249            0 :     if (isBigCount_) {
    1250            0 :         if (roundIdx == 0) {
    1251            0 :             UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
    1252              :         }
    1253            0 :         if (roundIdx < commRounds_ - 1) {
    1254            0 :             u32 nextgroupRankSize = (leftRankSize - groupRankSize > sdmaConcurrentNum_) ?
    1255              :                 sdmaConcurrentNum_ : leftRankSize - groupRankSize;
    1256            0 :             UpdatePartialCommunicationRankSet(roundIdx + 1, nextgroupRankSize, nextPartialCommRankSet_);
    1257              :         }
    1258            0 :         CHK_RET(RunGroupFullMeshAlltoall(roundIdx, step));
    1259              : 
    1260            0 :         if (roundIdx < commRounds_ - 1) {
    1261            0 :             partialCommRankSet_ = nextPartialCommRankSet_;
    1262            0 :             subStreamSendInfo_ = nextSubStreamSendInfo_;
    1263            0 :             subStreamReadInfo_ = nextSubStreamReadInfo_;
    1264            0 :             if (needAlltoallvCache_) {
    1265            0 :                 subStreamZcopySendInfo_ = nextSubStreamZcopySendInfo_;
    1266            0 :                 subStreamZcopyReadInfo_ = nextSubStreamZcopyReadInfo_;
    1267              :             }
    1268              :         }
    1269            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, localSubStream_));
    1270              :     } else {
    1271            0 :         UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
    1272            0 :         CHK_RET(RunGroupFullMeshAlltoall(roundIdx, step));
    1273              :     }
    1274            0 :     return HCCL_SUCCESS;
    1275              : }
    1276              : 
    1277            0 : HcclResult AlltoAllVDirectFullMesh::RunSDMAFineGrained(u32 totalStep, HcclOpMetaInfoDef &opMeta){
    1278            0 :     if (totalStep > 1){
    1279              :         //细粒度场景不支持切分
    1280            0 :         HCCL_ERROR("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] AlltoAllV is not supported when totalStep[%u] > 1, "\
    1281              :             "HCCL buffer is insufficient. stepSize : %u ", totalStep, algOpContext_.mc2Handler.stepSize);
    1282            0 :         return HCCL_E_NOT_SUPPORT;
    1283            0 :     } else if (totalStep == 0){
    1284              :         //totalStep不需要通信,但是需要适配高阶API wait/write
    1285            0 :         for (u32 roundIdx = 0; roundIdx < commRounds_; roundIdx++) {
    1286            0 :             CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), roundIdx));
    1287            0 :             CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
    1288            0 :             HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] step is 0 finished.");
    1289              :         }
    1290              :     } else {
    1291              :         // totalStep == 1 细粒度修改
    1292            0 :         u32 leftRankSize = devNumInlocalPod_; // leftRankSize中去掉本卡
    1293            0 :         for (u32 roundIdx = 0; roundIdx < commRounds_ && leftRankSize > 0; roundIdx++) {
    1294            0 :             CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), roundIdx));
    1295            0 :             CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1296            0 :             u32 groupRankSize = (leftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : leftRankSize;
    1297            0 :             UpdateRemoteRankSet(roundIdx, groupRankSize);
    1298            0 :             CHK_RET(RunGroupFullMeshAlltoall(roundIdx, 0));
    1299            0 :             leftRankSize -= groupRankSize;
    1300            0 :             CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
    1301            0 :             CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
    1302              :         }
    1303            0 :         HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] fine-grained finished.");
    1304            0 :         return HCCL_SUCCESS;
    1305              :     }
    1306              : 
    1307            0 :     if (totalStep == 0 && !islocalCpyDone_) {
    1308            0 :         CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1309            0 :         CHK_RET(LocalCopy());
    1310            0 :         islocalCpyDone_ = true;
    1311            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
    1312            0 :         return HCCL_SUCCESS;
    1313              :     }
    1314              : 
    1315            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] finished.");
    1316            0 :     return HCCL_SUCCESS;
    1317              : }
    1318              : 
    1319            0 : HcclResult AlltoAllVDirectFullMesh::RunSDMA(HcclOpMetaInfoDef &opMeta)
    1320              : {
    1321            0 :     u32 totalStep = CalcNumSubStep();
    1322            0 :     lastStep_ = totalStep - 1;
    1323              :     // 计算每个rank分组fullmesh后需要通信的轮次,向上取整
    1324            0 :     commRounds_ = (devNumInlocalPod_ + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_;
    1325            0 :     u32 leftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
    1326            0 :     lastRoundIdx_ = std::min((leftRankSize + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_, static_cast<u32>(commRounds_)) - 1;
    1327            0 :     HCCL_DEBUG("[AlltoAllVDirectFullMesh][RunSDMA] userRank [%u] communication rounds[%llu] totalStep [%u] "
    1328              :         "stepSize [%u], post sync info: lastStep_[%u] lastRoundIdx_[%u] devNumInlocalPod_[%u] sdmaConcurrentNum_[%u]",
    1329              :         userRank_, commRounds_, totalStep, algOpContext_.mc2Handler.stepSize,
    1330              :         lastStep_, lastRoundIdx_, devNumInlocalPod_, sdmaConcurrentNum_);
    1331              : 
    1332            0 :     if (UNLIKELY(algOpContext_.mc2Handler.stepSize > 0)){
    1333            0 :         CHK_RET(RunSDMAFineGrained(totalStep, opMeta));
    1334              :     } else {
    1335            0 :         if (totalStep == 0 && !islocalCpyDone_) {
    1336            0 :             CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1337            0 :             CHK_RET(LocalCopy());
    1338            0 :             islocalCpyDone_ = true;
    1339            0 :             CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
    1340            0 :             return HCCL_SUCCESS;
    1341              :         }
    1342              : 
    1343            0 :         for (u32 step = 0; step < totalStep; step++) {
    1344            0 :             u32 currentLeftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
    1345            0 :             for (u32 roundIdx = 0; roundIdx < commRounds_ && currentLeftRankSize > 0; roundIdx++) {
    1346            0 :                 CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1347            0 :                 u32 groupRankSize = (currentLeftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : currentLeftRankSize;
    1348            0 :                 CHK_RET(RunSDMATasks(roundIdx, step, groupRankSize, currentLeftRankSize));
    1349            0 :                 currentLeftRankSize -= groupRankSize;
    1350            0 :                 CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
    1351              :             }
    1352              :         }
    1353              :     }
    1354              :     
    1355            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMA] finished.");
    1356            0 :     return HCCL_SUCCESS;
    1357              : }
    1358              : 
    1359            0 : HcclResult AlltoAllVDirectFullMesh::RunAsync()
    1360              : {   
    1361            0 :     HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, cclInMem_.size(), true);
    1362            0 :     CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1363              : 
    1364            0 :     if (algOpContext_.mc2Handler.stepSize > 0){
    1365            0 :         if(algOpContext_.mc2Handler.stepSize > userRankSize_ || userRankSize_ % algOpContext_.mc2Handler.stepSize != 0){
    1366            0 :             HCCL_ERROR("[AlltoAllVDirectFullMesh][RunAsync] Step size should be less than or equal to the rank size, "\
    1367              :                         "and the rank size should be a multiple of the step size, but the step size is [%u] and the rank size is [%u].",
    1368              :                         algOpContext_.mc2Handler.stepSize, userRankSize_);
    1369            0 :             return HCCL_E_PARA;
    1370              :         }
    1371            0 :         if(userRankSize_ == 1){
    1372            0 :             HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] AlltoAllV do localcopy with 1 rank");
    1373            0 :             CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), 0));
    1374            0 :             CHK_RET(LocalCopy());
    1375            0 :             CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
    1376            0 :             return HCCL_SUCCESS;
    1377              :         }
    1378              :     }
    1379              : 
    1380            0 :     if (userRankSize_ == 1) {
    1381            0 :         HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] do localcopy with 1 rank");
    1382            0 :         CHK_RET(LocalCopy());
    1383            0 :         return HCCL_SUCCESS;
    1384              :     }
    1385              : 
    1386            0 :     CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
    1387            0 :     if (totalRdmaRankNum_ > 0) {
    1388            0 :         CHK_RET(RunRDMA());
    1389              :     }
    1390              : 
    1391            0 :     CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, rdmaSubStreams_));
    1392              : 
    1393            0 :     if (devNumInlocalPod_ > 1) {
    1394            0 :         CHK_RET(RunSDMA(opMeta));
    1395              :     }
    1396              : 
    1397            0 :     if (totalRdmaRankNum_ > 0) {
    1398              :         // 等待RDMA通信结束
    1399            0 :         CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
    1400            0 :         CHK_RET(RdmaControlNotifyMainFinish());
    1401            0 :         CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, rdmaSubStreams_));
    1402              :     }
    1403              : 
    1404            0 :     HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] finished.");
    1405            0 :     return HCCL_SUCCESS;
    1406              : }
    1407              : 
    1408            0 : HcclResult AlltoAllVDirectFullMesh::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
    1409              :                                                    const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
    1410              : {
    1411              :     (void) links;
    1412            0 :     if (rankSize == 1) {
    1413            0 :         return HCCL_SUCCESS;
    1414              :     }
    1415              : 
    1416            0 :     u32 devNumInlocalPod = nslbAdjInfo.dstRankNum;
    1417            0 :     u32 totalRdmaRankNum = rankSize - devNumInlocalPod;
    1418              : 
    1419            0 :     u32 rdmaConcurrentNum = (totalRdmaRankNum > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
    1420            0 :         (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) : (totalRdmaRankNum);
    1421            0 :     if (rdmaConcurrentNum == 0) {
    1422            0 :         return HCCL_SUCCESS;
    1423              :     }
    1424              :     // RDMA通信轮次
    1425            0 :     u32 rdmaRoundNum = (totalRdmaRankNum + rdmaConcurrentNum - 1) / rdmaConcurrentNum;
    1426            0 :     if (rdmaRoundNum == 0) {
    1427            0 :         return HCCL_SUCCESS;
    1428              :     }
    1429            0 :     u32 currStage = rank / devNumInlocalPod;
    1430              : 
    1431            0 :     for (u32 step = 0; step < rdmaRoundNum; step++) {
    1432            0 :         u32 sendTo =(rank + devNumInlocalPod + step) % rankSize;
    1433            0 :         u32 sendToStag = sendTo / devNumInlocalPod;
    1434            0 :         if(currStage == sendToStag) {
    1435              :             //此时认为时同一个超节点内通讯
    1436            0 :             sendTo = (sendTo + devNumInlocalPod) % rankSize;
    1437              :         }
    1438            0 :         NslbDpAdjInfo adjInfoStep = {0};
    1439            0 :         adjInfoStep.dstLocalRankId = sendTo;
    1440            0 :         adjInfoStep.phaseId = step + 1;
    1441            0 :         adjInfoStep.rev = 0;
    1442            0 :         nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
    1443              :     }
    1444            0 :     nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
    1445            0 :     return HCCL_SUCCESS;
    1446              : }
    1447              : 
    1448            0 : HcclResult AlltoAllVDirectFullMesh::GetHcclOffsetDstRanksMap(std::unordered_map<uint64_t, std::vector<uint32_t>>& hcclOffsetDstRanksMap) const {
    1449            0 :     hcclOffsetDstRanksMap.clear();
    1450            0 :     hcclOffsetDstRanksMap = hcclOffsetDstRanksMap_; // Deep copy
    1451              : 
    1452            0 :     return HCCL_SUCCESS;
    1453              : }
    1454              : 
    1455              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, AlltoAllVDirectFullMesh);
    1456              : } // namespace hccl
        

Generated by: LCOV version 2.0-1