LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_broadcast - broadcast_star.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 74 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 8 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 "broadcast_star.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : // Gather的入口函数
      16            0 : BroadcastStar::BroadcastStar(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
      17              : 
      18            0 : BroadcastStar::~BroadcastStar() {}
      19              : 
      20            0 : HcclResult BroadcastStar::Prepare(
      21              :     DeviceMem& inputMem, DeviceMem& outputMem, DeviceMem& scratchMem, const u64 count, const HcclDataType dataType,
      22              :     const Stream& stream, const HcclReduceOp reductionOp, const u32 root, const std::vector<Slice>& slices,
      23              :     const u64 baseOffset, std::vector<u32> nicRankList, u32 userRank)
      24              : {
      25            0 :     userRank_ = userRank;
      26            0 :     return AlgTemplateBase::Prepare(
      27            0 :         inputMem, outputMem, scratchMem, count, dataType, stream, reductionOp, root, slices, baseOffset, nicRankList);
      28              : }
      29              : 
      30              : HcclResult
      31            0 : BroadcastStar::RunAsync(const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
      32              : {
      33              :     // task下发接口
      34            0 :     CHK_SMART_PTR_NULL(dispatcher_);
      35              :     // ==1的处理
      36            0 :     if (rankSize == 1) {
      37            0 :         if (inputMem_ != outputMem_) {
      38            0 :             HcclResult ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
      39            0 :             CHK_PRT_RET(
      40              :                 ret != HCCL_SUCCESS,
      41              :                 HCCL_ERROR(
      42              :                     "[BroadcastStar][RunAsync]rank[%u] copy input[%p] to output[%p] failed", rank, inputMem_.ptr(),
      43              :                     outputMem_.ptr()),
      44              :                 ret);
      45              :         }
      46            0 :         return HCCL_SUCCESS;
      47              :     }
      48              :     // links本rank_id与通信域内其它rank的通信连接,rankSize本executor所在通信域的rank个数
      49            0 :     if (links.size() < rankSize) {
      50            0 :         HCCL_ERROR(
      51              :             "[BroadcastStar][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank, links.size(), rankSize);
      52            0 :         return HCCL_E_INTERNAL;
      53              :     }
      54              : 
      55            0 :     Slice sendSlice;
      56            0 :     sendSlice.offset = 0;
      57            0 :     sendSlice.size = dataBytes_;
      58            0 :     Slice recvSlice;
      59            0 :     recvSlice.offset = 0;
      60            0 :     recvSlice.size = dataBytes_;
      61            0 :     if (rank == root_) {
      62              :         // root 给其他rank发
      63              :         HcclResult ret;
      64            0 :         for (u32 dstRank = 0; dstRank < rankSize; dstRank++) {
      65            0 :             ret = RunSendBroadcast(dstRank, sendSlice, links);
      66            0 :             CHK_PRT_RET(
      67              :                 ret != HCCL_SUCCESS,
      68              :                 HCCL_ERROR(
      69              :                     "[BroadcastStar][RunAsync] root [%u] send broadcast to"
      70              :                     "other rank[%u] run failed!",
      71              :                     root_, dstRank),
      72              :                 ret);
      73              :         }
      74              :     } else {
      75              :         // 非root 接收来自root的数据
      76            0 :         HcclResult ret = RunRecvBroadcast(root_, rank, recvSlice, links);
      77            0 :         CHK_PRT_RET(
      78              :             ret == HCCL_E_AGAIN, HCCL_WARNING("[BroadcastStar][RunAsync]group has been destroyed. Break!"), ret);
      79            0 :         CHK_PRT_RET(
      80              :             ret != HCCL_SUCCESS,
      81              :             HCCL_ERROR(
      82              :                 "[BroadcastStar][RunAsync] dstrank [%u] recv broadcast from"
      83              :                 "root [%u] run failed!",
      84              :                 rank, root_),
      85              :             ret);
      86              :     }
      87            0 :     HCCL_INFO("BroadBastStar finished: rank[%u]", rank);
      88            0 :     return HCCL_SUCCESS;
      89              : }
      90              : 
      91            0 : HcclResult BroadcastStar::RunRecvBroadcast(
      92              :     const u32 srcRank, const u32 dstRank, const Slice& slice, const std::vector<LINK>& links)
      93              : {
      94              :     // 非root 接受数据
      95            0 :     DeviceMem dst;
      96            0 :     if (slice.size > 0) {
      97            0 :         if (srcRank >= links.size()) {
      98            0 :             HCCL_ERROR("[RunRecvBroadcast] root [%u] is out of range, linksize[%llu]", srcRank, links.size());
      99            0 :             return HCCL_E_INTERNAL;
     100              :         }
     101            0 :         dst = outputMem_.range(slice.offset, slice.size);
     102            0 :         HCCL_DEBUG(
     103              :             "rank [%u] will recv with output's offset[%llu], size[%llu], dstmem[%p]", dstRank, slice.offset, slice.size,
     104              :             dst.ptr());
     105              : 
     106            0 :         if (links[srcRank]->IsTransportRoce()) {
     107            0 :             CHK_RET(links[srcRank]->RxEnv(stream_));
     108              :         } else {
     109            0 :             CHK_RET(links[srcRank]->TxAck(stream_));
     110              :         }
     111              : 
     112            0 :         HcclResult ret = links[srcRank]->RxAsync(UserMemType::OUTPUT_MEM, slice.offset, dst.ptr(), slice.size, stream_);
     113            0 :         CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[RunRecvBroadcast]group has been destroyed. Break!"), ret);
     114            0 :         CHK_PRT_RET(
     115              :             ret != HCCL_SUCCESS,
     116              :             HCCL_ERROR(
     117              :                 "[RunRecvBroadcast]root rank[%u] rx async to dstrank[%u] run "
     118              :                 "failed",
     119              :                 srcRank, dstRank),
     120              :             ret);
     121              : 
     122            0 :         if (!links[srcRank]->IsTransportRoce()) {
     123            0 :             ret = ExecuteBarrier(links[srcRank], stream_); // 多server走rdma可以不用
     124            0 :             CHK_PRT_RET(
     125              :                 ret != HCCL_SUCCESS,
     126              :                 HCCL_ERROR("[RunRecvBroadcast]dstRank[%u] Broadcast star run tempAlg barrier failed", dstRank), ret);
     127            0 :             ret = links[srcRank]->RxWaitDone(stream_); // 多server走rdma可以不用
     128            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[RunRecvBroadcast]RxWaitDone failed"), ret);
     129              :         }
     130              :     }
     131            0 :     return HCCL_SUCCESS;
     132            0 : }
     133              : 
     134            0 : HcclResult BroadcastStar::RunSendBroadcast(const u32 dstRank, const Slice& slice, const std::vector<LINK>& links)
     135              : {
     136            0 :     DeviceMem src;
     137              :     // root发送数据
     138            0 :     if (slice.size > 0 && dstRank != root_) {
     139            0 :         src = inputMem_.range(slice.offset, slice.size);
     140              : 
     141            0 :         if (links[dstRank]->IsTransportRoce()) {
     142            0 :             CHK_RET(links[dstRank]->TxEnv(src.ptr(), slice.size, stream_));
     143              :         } else {
     144            0 :             CHK_RET(links[dstRank]->RxAck(stream_));
     145              :         }
     146              : 
     147            0 :         HcclResult ret = links[dstRank]->TxAsync(UserMemType::OUTPUT_MEM, slice.offset, src.ptr(), slice.size, stream_);
     148            0 :         CHK_PRT_RET(
     149              :             ret != HCCL_SUCCESS,
     150              :             HCCL_ERROR("[RunSendBroadcast]rank[%u] tx async with output's offset[%llu] failed", dstRank, slice.offset),
     151              :             ret);
     152              : 
     153            0 :         if (!links[dstRank]->IsTransportRoce()) {
     154            0 :             ret = ExecuteBarrierSrcRank(links[dstRank], stream_);
     155            0 :             CHK_PRT_RET(
     156              :                 ret != HCCL_SUCCESS,
     157              :                 HCCL_ERROR("[RunSendBroadcast] srcRank[%u] broadcast star run tempAlg barrier failed", root_), ret);
     158              : 
     159            0 :             ret = links[dstRank]->TxWaitDone(stream_);
     160            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[RunSendBroadcast]TxWaitDone failed"), ret);
     161              :         }
     162              :     }
     163            0 :     return HCCL_SUCCESS;
     164            0 : }
     165              : 
     166            0 : HcclResult BroadcastStar::ExecuteBarrierSrcRank(std::shared_ptr<Transport> link, Stream& stream) const
     167              : {
     168            0 :     CHK_RET(link->RxAck(stream));
     169              : 
     170            0 :     CHK_RET(link->TxAck(stream));
     171              : 
     172            0 :     CHK_RET(link->RxDataSignal(stream));
     173              : 
     174            0 :     CHK_RET(link->TxDataSignal(stream));
     175              : 
     176            0 :     return HCCL_SUCCESS;
     177              : }
     178              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_BROADCAST_STAR, BroadcastStar);
     179              : } // namespace hccl
        

Generated by: LCOV version 2.0-1