LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_gather - all_gather_mesh_atomic.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 36 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 4 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 "all_gather_mesh_atomic.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : AllGatherMeshAtomic::AllGatherMeshAtomic(const HcclDispatcher dispatcher) : AllGatherMesh(dispatcher) {}
      16              : 
      17            0 : AllGatherMeshAtomic::~AllGatherMeshAtomic() {}
      18              : 
      19            0 : HcclResult AllGatherMeshAtomic::RunAllGather(
      20              :     const std::vector<LINK>& links, const std::vector<Slice>& outputSlices, const std::vector<Slice>& inputSlices)
      21              : {
      22            0 :     for (u32 round = 1; round < interRankSize_; round++) {
      23            0 :         u32 dstRank = BackwardRank(interRank_, interRankSize_, round);
      24            0 :         Stream& subStream = (round == interRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
      25            0 :         CHK_RET(links[dstRank]->TxAck(subStream));
      26            0 :         CHK_RET(links[dstRank]->RxAck(subStream));
      27              :     }
      28              : 
      29            0 :     for (u32 round = 1; round < interRankSize_; round++) {
      30            0 :         u32 dstRank = BackwardRank(interRank_, interRankSize_, round);
      31            0 :         Stream& subStream = (round == interRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
      32            0 :         profilerInput_.streamID = subStream.id();
      33            0 :         profilerInput_.planeID = round - 1;
      34            0 :         profilerInput_.step = HCCL_EXEC_STEP_NOT_SET;
      35              : 
      36            0 :         if (round == interRankSize_ - 1) {
      37            0 :             for (u32 signalIndex = 0; signalIndex < interRankSize_ - 2; signalIndex++) { // rankSize-2: stream num
      38            0 :                 CHK_RET(LocalNotify::Wait(subStream, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
      39              :             }
      40              :             // 为子图增加一个从stream到主stream的附着点
      41            0 :             DeviceMem src = DeviceMem::create(inputMem_.ptr(), 0);
      42            0 :             DeviceMem dst = DeviceMem::create(outputMem_.ptr(), 0);
      43            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
      44            0 :             for (u32 signalIndex = 0; signalIndex < interRankSize_ - 2; signalIndex++) { // rankSize-2: stream num
      45            0 :                 CHK_RET(
      46              :                     LocalNotify::Post(subStream, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
      47              :             }
      48            0 :         } else {
      49            0 :             u32 signalIndex = round - 1;
      50            0 :             CHK_RET(LocalNotify::Post(subStream, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
      51            0 :             CHK_RET(LocalNotify::Wait(subStream, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
      52              :         }
      53              :         // 本rank要收数据
      54            0 :         void* srcMemPtr = nullptr;
      55              :         // 从对端的input内存拿数据,input==output也没有关系
      56            0 :         CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
      57              :         DeviceMem srcDevMem(
      58            0 :             static_cast<s8*>(srcMemPtr) + baseOffset_ + inputSlices[dstRank].offset, inputSlices[dstRank].size);
      59            0 :         DeviceMem dstDevMem = outputMem_.range(outputSlices[dstRank].offset, outputSlices[dstRank].size);
      60            0 :         CHK_RET(HcclD2DMemcpyAsync(
      61              :             dispatcher_, dstDevMem, srcDevMem, subStream, links[dstRank]->GetRemoteRank(),
      62              :             links[dstRank]->GetLinkType()));
      63            0 :         CHK_RET(links[dstRank]->TxDataSignal(subStream));
      64            0 :         CHK_RET(links[dstRank]->RxDataSignal(subStream));
      65            0 :     }
      66              : 
      67            0 :     CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
      68            0 :     return HCCL_SUCCESS;
      69              : }
      70              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_MESH_ATOMIC, AllGatherMeshAtomic);
      71              : } // namespace hccl
        

Generated by: LCOV version 2.0-1