LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_alltoall - allltoall_pipeline_mesh_pairwise_ccl_enough.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 183 0
Test Date: 2026-07-28 12:11:00 Functions: 0.0 % 16 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 "allltoall_pipeline_mesh_pairwise_ccl_enough.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : 
      16              : static const u32 INTRA_STREAM_INFO_SENDLEN_INDEX = 0; // intraStreamInfo 中 sendLen 的下标
      17              : static const u32 INTRA_STREAM_INFO_RECVLEN_INDEX = 1; // intraStreamInfo 中 recvLen 的下标
      18              : static const u32 INTRA_STREAM_INFO_RECV_REMOTE_OFFSET_INDEX = 2; // intraStreamInfo 中 recvRemoteOffset 的下标
      19              : static const u32 INTRA_STREAM_INFO_RECV_LOCAL_OFFSET_INDEX = 3; // intraStreamInfo 中 recvLocalOffset 的下标
      20              : 
      21            0 : AlltoallPipelineMeshPairwiseCCLEnough::AlltoallPipelineMeshPairwiseCCLEnough(
      22            0 :     const HcclDispatcher dispatcher): AlltoallPipelineBase(dispatcher) {}
      23              : 
      24            0 : AlltoallPipelineMeshPairwiseCCLEnough::~AlltoallPipelineMeshPairwiseCCLEnough() {}
      25              : 
      26            0 : u32 AlltoallPipelineMeshPairwiseCCLEnough::CalcInterNumSteps()
      27              : {
      28            0 :     return interRankSize_ - 1;
      29              : }
      30              : 
      31              : // 适配新CollExecutor接口
      32            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::Prepare(u32 userRank, A2aPipelineMemory A2aPipelineMemory,
      33              :     const SubCommInfo &level0CommInfo, const SubCommInfo &level1CommInfo,
      34              :     Stream &mainStream, std::vector<Stream> &subStream,
      35              :     std::vector<std::shared_ptr<LocalNotify>> &notifyMain, std::vector<std::shared_ptr<LocalNotify>> &notifySub,
      36              :     std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, HcclWorkflowMode workMode)
      37              : {
      38            0 :     AlltoallPipelineBase::Prepare(userRank, A2aPipelineMemory, level0CommInfo, level1CommInfo, mainStream,
      39              :         subStream, notifyMain, notifySub, allMeshAggregationSendRecvInfo, workMode);
      40            0 :     GetIntraScratchOffset();
      41            0 :     CHK_RET(DeviceMemMapping());
      42            0 :     return HCCL_SUCCESS;
      43              : }
      44              : 
      45              : // 统一计算每步 mesh 内收发时从各卡 scratch 读取的 offset 和 length
      46            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::GetIntraScratchOffset()
      47              : {
      48            0 :     for (u32 i = 0; i < intraRankSize_; i++) {
      49            0 :         intraScratchOffsetMap_[i] = std::vector<u64>();
      50            0 :         intraScratchLengMap_[i] = std::vector<u64>();
      51            0 :         u64 startOffset = 0;
      52            0 :         for (u32 remoteRank = i; remoteRank < groupRankSize_; remoteRank += intraRankSize_) {
      53            0 :             if (remoteRank == userRank_) {
      54            0 :                 localScratchOffset_ = startOffset;
      55              :             }
      56            0 :             const std::vector<u64>& remoteSendOffset = (*allMeshAggregationSendRecvInfo_)[remoteRank].sendOffset;
      57            0 :             const std::vector<u64>& remoteSendLength = (*allMeshAggregationSendRecvInfo_)[remoteRank].sendLength;
      58            0 :             intraScratchOffsetMap_[i].push_back(startOffset + (remoteSendOffset[userRank_] -
      59            0 :                 remoteSendOffset[meshRankStart_]));
      60            0 :             startOffset += (remoteSendOffset[meshRankEnd_] + remoteSendLength[meshRankEnd_] -
      61            0 :                 remoteSendOffset[meshRankStart_]);
      62            0 :             intraScratchLengMap_[i].push_back(remoteSendLength[userRank_]);
      63              :         }
      64              :     }
      65            0 :     return HCCL_SUCCESS;
      66              : }
      67              : 
      68            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::DeviceMemMapping()
      69              : {
      70            0 :     if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      71            0 :         interTransportSend_ = cclIn_;
      72            0 :         interTransportRecv_ = cclOut_;
      73            0 :         intraTransportSend_ = cclOut_;
      74              :     } else {
      75            0 :         interTransportSend_ = inputMem_;
      76            0 :         interTransportRecv_ = scratchMem_;
      77            0 :         intraTransportSend_ = scratchMem_;
      78              :     }
      79            0 :     for (u32 intraRank = 0; intraRank < intraRankSize_; intraRank++) {
      80            0 :         if (intraRank == intraRankId_) {
      81            0 :             continue;
      82              :         }
      83            0 :         LINK& intraNeighboorTransport = intraLinks_[intraRank];
      84            0 :         void* remDMAMemPtr = nullptr;
      85            0 :         CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
      86              :         DeviceMem remoteAlltoallScratch = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr),
      87            0 :             workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE ? cclIn_.size() : scratchMem_.size());
      88            0 :         intraNeighBoorMemory_[intraRank] = {remoteAlltoallScratch};
      89            0 :     }
      90            0 :     return HCCL_SUCCESS;
      91            0 : }
      92              : 
      93            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PrepareInterData(u32 step)
      94              : {
      95              :     // 准备 mesh 间发送信息
      96            0 :     nextInterSendData_.clear();
      97            0 :     u32 interSendRankStart = ((interRankId_ + 1 + step) % interRankSize_) * intraRankSize_;
      98            0 :     u32 interSendRankEnd = interSendRankStart + intraRankSize_ - 1;
      99            0 :     u64 startMemOffset = localSendRecvInfo_.sendOffset[interSendRankStart];
     100            0 :     u64 meshSendLength = localSendRecvInfo_.sendOffset[interSendRankEnd] +
     101            0 :         localSendRecvInfo_.sendLength[interSendRankEnd] - startMemOffset;
     102            0 :     u64 sendDestOffset = 0;
     103            0 :     for (u32 relatedRank = intraRankId_; relatedRank < userRank_; relatedRank += intraRankSize_) {
     104            0 :         const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
     105            0 :         sendDestOffset += (info.sendOffset[interSendRankEnd] + info.sendLength[interSendRankEnd] -
     106            0 :             info.sendOffset[interSendRankStart]);
     107              :     }
     108            0 :     DeviceMem srcMem = inputMem_.range(startMemOffset, meshSendLength);
     109            0 :     DeviceMem dstMem = interTransportSend_.range(startMemOffset, meshSendLength);
     110            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][PrepareInterSendData] userRank %u, interRank %u, "
     111              :         "intraRank %u move from userInput offset %llu length %llu to interTransportSend, send to remote offset "
     112              :         "%llu", userRank_, interRankId_, intraRankId_, startMemOffset, meshSendLength, sendDestOffset);
     113              : 
     114            0 :     HCCL_DEBUG("user size %u, inter size %u, intra size %u",
     115              :         groupRankSize_, interRankSize_, intraRankSize_);
     116            0 :     if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     117            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
     118              :     }
     119            0 :     nextInterSendData_.emplace_back(TxMemoryInfo{UserMemType::OUTPUT_MEM, sendDestOffset, dstMem.ptr(),
     120              :         meshSendLength});
     121              : 
     122              :     // 准备 mesh 间接收信息
     123            0 :     nextInterRecvData_.clear();
     124            0 :     u32 recvBlockStart = (((interRankId_ + interRankSize_ - 1u - step) % interRankSize_) * intraRankSize_);
     125            0 :     u64 recvLocalOffset = 0;
     126            0 :     for (u32 relatedRank = intraRankId_; relatedRank < recvBlockStart; relatedRank += intraRankSize_) {
     127            0 :         const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
     128            0 :         recvLocalOffset += (info.sendOffset[meshRankEnd_] + info.sendLength[meshRankEnd_] -
     129            0 :             info.sendOffset[meshRankStart_]);
     130              :     }
     131            0 :     const SendRecvInfo& recvRankInfo = (*allMeshAggregationSendRecvInfo_)[recvBlockStart + intraRankId_];
     132            0 :     u64 recvRemoteOffset = recvRankInfo.sendOffset[meshRankStart_];
     133            0 :     u64 recvLength = (recvRankInfo.sendOffset[meshRankEnd_] + recvRankInfo.sendLength[meshRankEnd_] -
     134            0 :         recvRankInfo.sendOffset[meshRankStart_]);
     135            0 :     nextInterRecvData_.emplace_back(RxMemoryInfo{UserMemType::INPUT_MEM, recvRemoteOffset,
     136            0 :         interTransportRecv_.range(recvLocalOffset, recvLength).ptr(), recvLength});
     137            0 :     return HCCL_SUCCESS;
     138            0 : }
     139              : 
     140              : // 将原先在userInput,且需要发到本mesh内其他卡的数据搬到CCL
     141            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PrepareIntraData()
     142              : {
     143            0 :     u64 startMemOffset = localSendRecvInfo_.sendOffset[meshRankStart_];
     144            0 :     u64 meshSendLength = localSendRecvInfo_.sendOffset[meshRankEnd_] +
     145            0 :         localSendRecvInfo_.sendLength[meshRankEnd_] - startMemOffset;
     146            0 :     u64 intraSendOffset = 0;
     147            0 :     for (u32 relatedRank = intraRankId_; relatedRank < meshRankStart_; relatedRank += intraRankSize_) {
     148            0 :         const SendRecvInfo& info = (*allMeshAggregationSendRecvInfo_)[relatedRank];
     149            0 :         intraSendOffset += (info.sendOffset[meshRankEnd_] + info.sendLength[meshRankEnd_] -
     150            0 :             info.sendOffset[meshRankStart_]);
     151              :     }
     152            0 :     DeviceMem srcMem = inputMem_.range(startMemOffset, meshSendLength);
     153            0 :     DeviceMem dstMem = intraTransportSend_.range(intraSendOffset, meshSendLength);
     154            0 :     if (workMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     155            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
     156              :     }
     157            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][PrepareIntraData] userRank %u, interRank %u, intraRank %u "
     158              :         "copy from userInput offset %llu length %llu to intraTransportSend offset %llu", userRank_, interRankId_,
     159              :         intraRankId_, startMemOffset, meshSendLength, intraSendOffset);
     160            0 :     return HCCL_SUCCESS;
     161            0 : }
     162              : 
     163            0 : void AlltoallPipelineMeshPairwiseCCLEnough::UpdateIntraStreamInfo(u32 step)
     164              : {
     165            0 :     intraStreamInfo_.clear();
     166            0 :     u32 localMeshIndex = (interRankId_ + interRankSize_ - step) % interRankSize_;
     167            0 :     u32 firstDataBlockIndex = (meshRankStart_ + groupRankSize_ - step * intraRankSize_) % groupRankSize_;
     168            0 :     const std::vector<u64>& sendLengths = (*allMeshAggregationSendRecvInfo_)[firstDataBlockIndex +
     169            0 :         intraRankId_].sendLength;
     170            0 :     const std::vector<u64>& recvLengths = localSendRecvInfo_.recvLength;
     171            0 :     const std::vector<u64>& recvOffsets = localSendRecvInfo_.recvOffset;
     172            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][UpdateIntraStreamInfo] userRank %u, "
     173              :         "interRank %u, intraRank %u, step %u", userRank_, interRankId_, intraRankId_, step);
     174            0 :     for (u32 intraRank = 0; intraRank < intraRankSize_; intraRank++) {
     175            0 :         u64 sendLen = sendLengths[meshRankStart_ + intraRank];
     176            0 :         u64 recvLen = recvLengths[firstDataBlockIndex + intraRank];
     177            0 :         u64 recvRemoteOffset = intraScratchOffsetMap_[intraRank][localMeshIndex];
     178            0 :         u64 recvLocalOffset = recvOffsets[firstDataBlockIndex + intraRank];
     179            0 :         if (intraRank != intraRankId_) {
     180            0 :             intraStreamInfo_[intraRank] = {sendLen, recvLen, recvRemoteOffset, recvLocalOffset};
     181            0 :             HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][UpdateIntraStreamInfo] userRank %u, interRank %u, "
     182              :                 "intraRank %u, sdma stream %u need send %llu and read length %llu from remote offset %llu "
     183              :                 "to local offset %llu", userRank_, interRankId_, intraRankId_, intraRank, sendLen, recvLen,
     184              :                 recvRemoteOffset, recvLocalOffset);
     185              :         }
     186              :     }
     187            0 : }
     188              : 
     189            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::SendRecvDataIntraMesh()
     190              : {
     191            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, "
     192              :         "interRank %u, intraRank %u, sdma stream %s wait main stream", userRank_, interRankId_,
     193              :         intraRankId_, GetStreamIndexString().c_str());
     194            0 :     for (auto& intraInfo : intraStreamInfo_) {
     195            0 :         u32 streamIndex = intraInfo.first;
     196            0 :         u64 recvLen = intraInfo.second[INTRA_STREAM_INFO_RECVLEN_INDEX];
     197            0 :         Stream& currStream = subStream_[streamIndex];
     198            0 :         LINK& readTransport = intraLinks_[streamIndex];
     199            0 :         CHK_RET(readTransport->TxAck(currStream));
     200            0 :         CHK_RET(readTransport->RxAck(currStream));
     201            0 :         u64 recvRemoteOffset = intraInfo.second[INTRA_STREAM_INFO_RECV_REMOTE_OFFSET_INDEX];
     202            0 :         u64 recvLocalOffset = intraInfo.second[INTRA_STREAM_INFO_RECV_LOCAL_OFFSET_INDEX];
     203            0 :         DeviceMem src = intraNeighBoorMemory_[streamIndex][0].range(recvRemoteOffset, recvLen);
     204            0 :         DeviceMem dst = outputMem_.range(recvLocalOffset, recvLen);
     205            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, currStream, readTransport->GetRemoteRank(),
     206              :             readTransport->GetLinkType()));
     207            0 :         CHK_RET(readTransport->TxDataSignal(currStream));
     208            0 :         HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, interRank %u, "
     209              :             "intraRank %u, sdma stream %llu read data from remote offset %llu len %llu to local %llu",
     210              :             userRank_, interRankId_, intraRankId_, streamIndex, recvRemoteOffset, recvLen, recvLocalOffset);
     211            0 :         CHK_RET(readTransport->RxDataSignal(currStream));
     212            0 :     }
     213            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][SendRecvDataIntraMesh] userRank %u, interRank %u, "
     214              :         "intraRank %u, sdma stream %s notify main stream", userRank_, interRankId_, intraRankId_,
     215              :         GetStreamIndexString().c_str());
     216            0 :     return HCCL_SUCCESS;
     217              : }
     218              : 
     219            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::SendRecvDataInterMesh(u32 step)
     220              : {
     221            0 :     Stream& interStream = subStream_[intraRankId_];
     222            0 :     LINK& interRecvTransport = interLinks_[(interRankId_ + interRankSize_ - 1 - step) % interRankSize_];
     223            0 :     LINK& interSendTransport = interLinks_[(interRankId_ + 1 + step) % interRankSize_];
     224            0 :     CHK_RET(interRecvTransport->TxAck(interStream));
     225            0 :     CHK_RET(interSendTransport->RxAck(interStream));
     226            0 :     CHK_RET(interSendTransport->TxAsync(nextInterSendData_, interStream));
     227            0 :     CHK_RET(interRecvTransport->RxAsync(nextInterRecvData_, interStream));
     228            0 :     CHK_RET(interRecvTransport->PostFinAck(interStream));
     229            0 :     CHK_RET(interSendTransport->WaitFinAck(interStream));
     230            0 :     CHK_RET(ExecuteBarrier(interRecvTransport, interSendTransport, interStream));
     231            0 :     return HCCL_SUCCESS;
     232              : }
     233              : 
     234            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::LocalCopyDataRecvFromInter(u32 interRankDistance)
     235              : {
     236            0 :     u64 scratchOffset = intraScratchOffsetMap_[intraRankId_][(interRankId_ +
     237            0 :         interRankSize_ - interRankDistance) % interRankSize_];
     238            0 :     u64 recvLen = localSendRecvInfo_.recvLength[(userRank_ + groupRankSize_ -
     239            0 :         interRankDistance * intraRankSize_) % groupRankSize_];
     240            0 :     u64 userOutOffset = localSendRecvInfo_.recvOffset[(userRank_ + groupRankSize_ -
     241            0 :         interRankDistance * intraRankSize_) % groupRankSize_];
     242            0 :     DeviceMem src = interTransportRecv_.range(scratchOffset, recvLen);
     243            0 :     DeviceMem dst = outputMem_.range(userOutOffset, recvLen);
     244            0 :     HCCL_DEBUG("[AlltoallPipelineMeshPairwiseCCLEnough][LocalCopyDataRecvFromInter]local move from "
     245              :         "interTransportRecv_ offset %llu length %llu to outputMem_ %llu", scratchOffset, recvLen, userOutOffset);
     246            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
     247            0 :     return HCCL_SUCCESS;
     248            0 : }
     249              : 
     250            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PreProcess()
     251              : {
     252              :     // server 间收发时间较长,先搬 server 间收发所需数据然后马上让server间开始收发
     253            0 :     CHK_RET(PrepareInterData(0));
     254            0 :     CHK_RET(NotifyInterStreamStart());
     255              :     // 之后将 server 内所需数据准备好之后唤醒server内从流收发
     256            0 :     UpdateIntraStreamInfo(0);
     257            0 :     CHK_RET(PrepareIntraData());
     258            0 :     CHK_RET(NotifyIntraStreamStart());
     259            0 :     CHK_RET(SendRecvDataIntraMesh());
     260            0 :     return HCCL_SUCCESS;
     261              : }
     262              : 
     263            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PipelineSend(u32 step, bool isLastStep)
     264              : {
     265            0 :     CHK_RET(SendRecvDataInterMesh(step));
     266            0 :     CHK_RET(PrepareInterData(step + 1u));
     267            0 :     CHK_RET(WaitInterStreamFinish());
     268            0 :     CHK_RET(WaitIntraStreamFinish());
     269            0 :     CHK_RET(ExecEmptyTask(inputMem_, outputMem_, mainStream_, dispatcher_));
     270            0 :     if (!isLastStep) {
     271            0 :         CHK_RET(NotifyInterStreamStart());
     272              :     }
     273            0 :     UpdateIntraStreamInfo(step + 1u);
     274            0 :     CHK_RET(NotifyIntraStreamStart());
     275            0 :     CHK_RET(SendRecvDataIntraMesh());
     276            0 :     CHK_RET(LocalCopyDataRecvFromInter(step + 1u));
     277            0 :     return HCCL_SUCCESS;
     278              : }
     279              : 
     280            0 : HcclResult AlltoallPipelineMeshPairwiseCCLEnough::PostProcess()
     281              : {
     282              :     // 最后的收尾工作
     283            0 :     CHK_RET(LocalCopyDataRecvFromInter(0));
     284            0 :     CHK_RET(WaitIntraStreamFinish());
     285            0 :     return HCCL_SUCCESS;
     286              : }
     287              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_PIPELINE_MESH_PAIRWISE_CCL_ENOUGH,
     288              :                   AlltoallPipelineMeshPairwiseCCLEnough);
     289              : } // namespace hccl
        

Generated by: LCOV version 2.0-1