LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_gather - coll_all_gather_ring_zerocopy_exchange_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 92 0
Test Date: 2026-07-28 12:11:00 Functions: 0.0 % 7 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 "coll_all_gather_ring_zerocopy_exchange_executor.h"
      12              : 
      13              : namespace hccl {
      14            0 : CollAllGatherRingZerocopyExchangeExecutor::CollAllGatherRingZerocopyExchangeExecutor(const HcclDispatcher dispatcher,
      15            0 :                                                                    std::unique_ptr<TopoMatcher> &topoMatcher)
      16            0 :     : CollAllGatherRingZerocopyExecutor(dispatcher, topoMatcher)
      17              : {
      18            0 :     DMAReduceFlag_ = true;      // 设为true,以禁用RunLoop中的本地拷贝
      19            0 :     desc_.isZeroCopy = true;
      20            0 : }
      21              : 
      22            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      23              : {
      24              :     // 调用父类编排函数建链关系计算函数
      25            0 :     CHK_RET(CollAllGatherRingZerocopyExecutor::CalcCommInfo(opTransport));
      26              :     // 额外增加数据交换的建链
      27            0 :     CHK_RET(CalcExchangeCommInfo(opTransport));
      28            0 :     return HCCL_SUCCESS;
      29              : }
      30              : 
      31            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::CalExchangeRemoteRank(u32 &remoteRankSend, u32 &remoteRankRecv)
      32              : {
      33              :     // AllGather的发端与收端与ReduceScatter相反
      34            0 :     return CalExchangeRemoteRankForReduceScatter(remoteRankRecv, remoteRankSend);
      35              : }
      36              : 
      37            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::CalcExchangeCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      38              : {
      39            0 :     std::set<u32> commTargetUserRankSet;
      40            0 :     u32 remoteRankSend = 0;
      41            0 :     u32 remoteRankRecv = 0;
      42              : 
      43            0 :     CHK_RET(CalExchangeRemoteRank(remoteRankSend, remoteRankRecv));
      44            0 :     HCCL_DEBUG("[%s] remoteRankSend:%d, remoteRankRecv:%d", __func__, remoteRankSend, remoteRankRecv);
      45            0 :     commTargetUserRankSet.insert(remoteRankSend);
      46            0 :     commTargetUserRankSet.insert(remoteRankRecv);
      47              :     CommParaInfo commParaInfo(COMM_COMBINE_ORDER, CommType::COMM_TAG_PARTIAL_MESH_COMBINED, INVALID_VALUE_RANKID,
      48            0 :         INVALID_VALUE_RANKID, false, false, commTargetUserRankSet);
      49              : 
      50            0 :     TransportMemType inputType = TransportMemType::CCL_INPUT;
      51            0 :     TransportMemType outputType = TransportMemType::CCL_OUTPUT;
      52              : 
      53            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaInfo, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
      54            0 :     LevelNSubCommTransport &subCommTransport = opTransport[COMM_COMBINE_ORDER];
      55            0 :     for (u32 subCommIndex = 0; subCommIndex < subCommTransport.size(); subCommIndex++) {
      56            0 :         for (auto &transportRequest : subCommTransport[subCommIndex].transportRequests) {
      57            0 :             transportRequest.isUsedRdma = (topoAttr_.superPodNum > 1 ||
      58            0 :                 (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) && topoAttr_.serverNum > 1));
      59              :         }
      60              :     }
      61            0 :     return HCCL_SUCCESS;
      62            0 : }
      63              : 
      64            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::KernelRunInterServerPreProcess(const OpParam &param, const ExecMem &execMem)
      65              : {
      66            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[AllGatherRingZerocopyExchangeExecutor] KernelRunInterServerPreProcess");
      67              :     // 计算需要交换数据的通信对端
      68            0 :     u32 remoteRankSend = 0;
      69            0 :     u32 remoteRankRecv = 0;
      70            0 :     CHK_RET(CalExchangeRemoteRank(remoteRankSend, remoteRankRecv));
      71              : 
      72            0 :     Stream stream = param.stream;
      73            0 :     u64 inputMemSize = execMem.inputMem.size();
      74            0 :     if (remoteRankSend != topoAttr_.userRank && remoteRankRecv != topoAttr_.userRank) {     // 需要交换数据
      75              :         // 获取通信对端的link
      76            0 :         LINK sendLink;
      77            0 :         LINK recvLink;
      78            0 :         CHK_RET(GetTransportForExchange(remoteRankSend, sendLink));
      79            0 :         CHK_RET(GetTransportForExchange(remoteRankRecv, recvLink));
      80              :         // 当通信对端恰好是同server的邻居时,复用Level0的建链,其注册的内存是UserMem,需要特殊处理
      81              :         // 否则,在CommCombineOrder上建链,其注册内存是CCL Buffer
      82            0 :         if (!IsLevel0Neighbor(remoteRankSend, level0RankSize_)) {
      83              :             // user in mem -> ccl in mem
      84            0 :             DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.inputPtr), inputMemSize);
      85            0 :             DeviceMem dstMem = execMem.inputMem.range(0, inputMemSize);
      86            0 :             HCCL_DEBUG("[%s] not neighbor, srcPtr:%p, dstPtr:%p, size:%llu", __func__, srcMem.ptr(), dstMem.ptr(), inputMemSize);
      87            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
      88            0 :         }
      89              :         // 执行通信
      90            0 :         recvLink->TxAck(stream);
      91            0 :         sendLink->RxAck(stream);
      92            0 :         u32 remoteLevel1Index = remoteRankSend % (level0RankSize_ * level1RankSize_) / level0RankSize_;
      93            0 :         u32 remoteLevel2Index = remoteRankSend / level0RankSize_ / level1RankSize_;   
      94            0 :         u64 txDstOffset = (remoteLevel1Index * level2RankSize_ + remoteLevel2Index) * inputMemSize;
      95            0 :         HCCL_DEBUG("[%s] remoteLevel1Index:%d, remoteLevel2Index:%d, txDstOffset:%llu", __func__, remoteLevel1Index, remoteLevel2Index, txDstOffset);
      96            0 :         if (IsLevel0Neighbor(remoteRankSend, level0RankSize_)) {
      97            0 :             sendLink->TxAsync(UserMemType::OUTPUT_MEM, txDstOffset, execMem.inputPtr, inputMemSize, stream);
      98            0 :             HCCL_DEBUG("[%s] neighbor, send data to userMem", __func__);
      99              :         } else {
     100            0 :             sendLink->TxAsync(UserMemType::OUTPUT_MEM, txDstOffset, execMem.inputMem.ptr(), inputMemSize, stream);
     101            0 :             HCCL_DEBUG("[%s] not neighbor, send data to ccl buffer", __func__);
     102              :         }
     103            0 :         u64 rxDstOffset = (level1Rank_ * level2RankSize_ + level2Rank_) * inputMemSize;
     104            0 :         u64 rxSrcOffset = IsLevel0Neighbor(remoteRankRecv, level0RankSize_) ? static_cast<u8 *>(execMem.inputPtr) - static_cast<u8 *>(param.inputPtr) : 0;
     105            0 :         HCCL_DEBUG("[%s] rxDstOffset:%llu, rxSrcOffset:%llu", __func__, rxDstOffset, rxSrcOffset);
     106            0 :         recvLink->RxAsync(UserMemType::INPUT_MEM, rxSrcOffset, static_cast<u8 *>(execMem.outputMem.ptr()) + rxDstOffset, inputMemSize, stream);
     107              :         // 交换数据的两端之间Barrier,确认收发完成
     108            0 :         CHK_RET(recvLink->TxAck(stream));
     109            0 :         CHK_RET(sendLink->RxAck(stream));
     110            0 :         CHK_RET(sendLink->TxDataSignal(stream));
     111            0 :         CHK_RET(recvLink->RxDataSignal(stream));
     112            0 :     } else {    // 不需要交换数据,将数据从user in拷到ccl out
     113            0 :         u64 dstMemOffset = (level1Rank_ * level2RankSize_ + level2Rank_) * inputMemSize;
     114            0 :         DeviceMem dstMem = execMem.outputMem.range(dstMemOffset, inputMemSize);
     115            0 :         DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.inputPtr), inputMemSize);
     116            0 :         HCCL_DEBUG("[%s] not exchange, just copy data from CCLOut[%p] to UserInput[%p]", __func__, dstMem.ptr(), srcMem.ptr());
     117            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     118            0 :     }
     119              : 
     120            0 :     return HCCL_SUCCESS;
     121            0 : }
     122              : 
     123            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::KernelRunInterServerPostProcess(const OpParam &param, const ExecMem &execMem)
     124              : {
     125              :     // 将通信结果从ccl output搬到user output
     126            0 :     if (level1RankSize_ > 1 || level2RankSize_ > 1) {
     127            0 :         u32 unitSize = SIZE_TABLE[param.DataDes.dataType];
     128            0 :         u64 curSize = execMem.inputMem.size();
     129            0 :         Stream stream = param.stream;
     130            0 :         for (u32 i = 0; i < level1RankSize_ * level2RankSize_; i++) {
     131            0 :             DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(execMem.outputPtr) + param.DataDes.count * unitSize * (level0Rank_ * level1RankSize_ * level2RankSize_ + i), curSize);
     132            0 :             DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.outputMem.ptr()) + i * curSize, curSize);
     133            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     134            0 :             HCCL_DEBUG("[%s] memcopy from CCLOut[%p] to UserOut[%p]", __func__, srcMem.ptr(), dstMem.ptr());
     135            0 :         }
     136            0 :     }
     137            0 :     return HCCL_SUCCESS;
     138              : }
     139              : 
     140            0 : HcclResult CollAllGatherRingZerocopyExchangeExecutor::CalcLevel0DataSlices(const OpParam &param, const ExecMem &execMem, std::vector<Slice> &dataSegsSlice)
     141              : {
     142            0 :     return CalcIntraServerDataSlicesContinuous(param, execMem, level0RankSize_, level1RankSize_, level2RankSize_, dataSegsSlice);
     143              : }
     144              : 
     145              : REGISTER_EXEC("AllGatherRingZerocopyExchangeExecutor", AllGatherRingZerocopyExchange, CollAllGatherRingZerocopyExchangeExecutor);
     146              : 
     147              : } // namespace hccl
        

Generated by: LCOV version 2.0-1