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

Generated by: LCOV version 2.0-1