LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_broadcast - broadcast_nhr_oneshot.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 74.0 % 104 77
Test Date: 2026-08-04 10:52:23 Functions: 80.0 % 10 8

            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_nhr_oneshot.h"
      12              : #include <cmath>
      13              : #include "alg_template_register.h"
      14              : 
      15              : namespace hccl {
      16            3 : BroadcastNHROneshot::BroadcastNHROneshot(const HcclDispatcher dispatcher)
      17            3 :     : NHRBase(dispatcher), localBaseOffset_(0), isForAllReduce_(false)
      18              : {
      19            3 : }
      20              : 
      21            3 : BroadcastNHROneshot::~BroadcastNHROneshot()
      22              : {
      23            3 : }
      24              : 
      25            3 : HcclResult BroadcastNHROneshot::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      26              : {
      27              :     // 基本的检查
      28            3 :     CHK_RET(SimpleCheck(rank, rankSize, links));
      29            3 :     HCCL_INFO("[BroadcastNHROneshot][RunAsync] rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
      30              :         rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
      31              : 
      32            3 :     u32 unitSize = DataUnitSize(dataType_);
      33            3 :     CHK_PRT_RET(unitSize == 0, HCCL_ERROR("[BroadcastNHROneshot][RunAsync] unitSize is zero"), HCCL_E_INTERNAL);
      34              : 
      35            3 :     if (!isForAllReduce_) {
      36            0 :         localBaseOffset_ = baseOffset_; // broadcast的本地偏移量和baseOffset_一致
      37              :     }
      38              : 
      39              :     // 双buffer下, 先将input拷贝到output的合适位置
      40            3 :     if (inputMem_ != outputMem_ && rank == root_) {
      41            3 :         u64 totalSize = count_ * SIZE_TABLE[dataType_];
      42            3 :         DeviceMem src = inputMem_.range(localBaseOffset_, totalSize);
      43            3 :         DeviceMem dst = outputMem_.range(localBaseOffset_, totalSize);
      44            3 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
      45            3 :     }
      46              : 
      47              :     // 如果ranksize为1, 从input->output就结束
      48            3 :     if (rankSize == 1) {
      49            0 :         return HCCL_SUCCESS;
      50              :     }
      51              : 
      52              :     // 运行bcast, rst算法
      53            3 :     CHK_RET(RunBroadcastNHROneshot(rank, rankSize, links));
      54              : 
      55            3 :     HCCL_INFO("[BroadcastNHROneshot][RunAsync] finished: rank[%u] ranksize[%u]", rank, rankSize);
      56            3 :     return HCCL_SUCCESS;
      57              : }
      58              : 
      59            3 : HcclResult BroadcastNHROneshot::RunAsyncForAllReduce(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      60              : {
      61            3 :     isForAllReduce_ = true;
      62            3 :     return RunAsync(rank, rankSize, links);
      63              : }
      64              : 
      65            3 : HcclResult BroadcastNHROneshot::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      66              : {
      67              :     // 判断stream, dispatcher是否为空
      68            3 :     CHK_SMART_PTR_NULL(dispatcher_);
      69            3 :     CHK_PTR_NULL(stream_.ptr());
      70              : 
      71              :     // 检查memory
      72            3 :     CHK_PRT_RET(!outputMem_ || !inputMem_,
      73              :         HCCL_ERROR("[BroadcastNHROneshot][SimpleCheck] rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
      74              : 
      75              :     // 判断links数量是否正确
      76            3 :     CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[BroadcastNHROneshot][SimpleCheck] rank[%u] link size[%llu] is "
      77              :         "less than rank size[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
      78            3 :     return HCCL_SUCCESS;
      79              : }
      80              : 
      81            0 : HcclResult BroadcastNHROneshot::SdmaRx(LINK &linkLeft, LINK &linkRight, InterServerAlgoStep &stepInfo, 
      82              :     const std::vector<LINK> &links)
      83              : {
      84            0 :     u64 totalSize = count_ * SIZE_TABLE[dataType_];
      85            0 :     DeviceMem srcMem = outputMem_.range(localBaseOffset_, totalSize);
      86              : 
      87            0 :     if (linkRight != nullptr) {
      88            0 :         CHK_RET(linkRight->TxAck(stream_));
      89              :     }
      90            0 :     if (linkLeft != nullptr) {
      91            0 :         CHK_RET(linkLeft->RxAck(stream_));
      92            0 :         void *srcMemPtr = nullptr;
      93            0 :         CHK_RET(linkLeft->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
      94            0 :         DeviceMem srcMemLeft(static_cast<s8 *>(srcMemPtr) + baseOffset_, totalSize);
      95            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, srcMem, srcMemLeft, stream_, linkLeft->GetRemoteRank(), // Memecpy
      96              :                     linkLeft->GetLinkType()));
      97            0 :         CHK_RET(linkLeft->TxDataSignal(stream_));
      98            0 :     }
      99            0 :     if (linkRight != nullptr) {
     100            0 :         CHK_RET(linkRight->RxDataSignal(stream_));
     101              :     }
     102            0 :     return HCCL_SUCCESS;
     103            0 : }
     104              : 
     105            9 : HcclResult BroadcastNHROneshot::RdmaTxRx(LINK &linkLeft, LINK &linkRight, InterServerAlgoStep &stepInfo, 
     106              :     const std::vector<LINK> &links)
     107              : {
     108            9 :     u64 totalSize = count_ * SIZE_TABLE[dataType_];
     109            9 :     DeviceMem srcMem = outputMem_.range(localBaseOffset_, totalSize);
     110              : 
     111            9 :     if (linkLeft != nullptr) {
     112            0 :         CHK_RET(linkLeft->TxAck(stream_));
     113              :     }
     114              : 
     115            9 :     if (linkRight != nullptr) {
     116            9 :         CHK_RET(linkRight->RxAck(stream_));
     117            9 :         CHK_RET(linkRight->TxAsync(UserMemType::OUTPUT_MEM, baseOffset_, srcMem.ptr(), srcMem.size(), stream_));
     118            9 :         CHK_RET(linkRight->WaitFinAck(stream_));
     119              :     }
     120              : 
     121            9 :     if (linkLeft != nullptr) {
     122            0 :         CHK_RET(linkLeft->RxAsync(UserMemType::OUTPUT_MEM, baseOffset_, srcMem.ptr(), srcMem.size(), stream_));
     123            0 :         CHK_RET(linkLeft->PostFinAck(stream_));
     124              :     }
     125            9 :     return HCCL_SUCCESS;
     126            9 : }
     127              : 
     128            3 : HcclResult BroadcastNHROneshot::RunBroadcastNHROneshot(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     129              : {
     130              :     // 计算通信步数
     131            3 :     u32 nSteps = GetStepNumInterServer(rankSize);
     132            3 :     HCCL_DEBUG("[BroadcastNHROneshot][RunBroadcastNHROneshot] rank[%u] rankSize[%u] nSteps[%u]",
     133              :         rank, rankSize, nSteps);
     134              : 
     135              :     // 逐步编排任务
     136           12 :     for (u32 step = 0; step < nSteps; step++) {
     137            9 :         InterServerAlgoStep stepInfo;
     138            9 :         GetStepInfo(step, nSteps, rank, rankSize, stepInfo);
     139              : 
     140            9 :         HCCL_DEBUG("[BroadcastNHROneshot][RunBroadcastNHROneshot] recvFrom[%u] sendTo[%u] step[%u]",
     141              :             stepInfo.fromRank, stepInfo.toRank, step);
     142              : 
     143            9 :         LINK linkLeft;
     144            9 :         LINK linkRight;
     145            9 :         if (stepInfo.txSliceIdxs.size() > 0) {
     146            9 :             linkRight = links[stepInfo.toRank];
     147            9 :             CHK_SMART_PTR_NULL(linkRight);
     148              :         }
     149            9 :         if (stepInfo.rxSliceIdxs.size() > 0) {
     150            0 :             linkLeft = links[stepInfo.fromRank];
     151            0 :             CHK_SMART_PTR_NULL(linkLeft);
     152              :         }
     153              : 
     154           18 :         if ((linkRight != nullptr && linkRight->IsSpInlineReduce()) || 
     155            9 :             (linkLeft != nullptr && linkLeft->IsSpInlineReduce())) {
     156            0 :             CHK_RET(SdmaRx(linkLeft, linkRight, stepInfo, links));
     157              :         } else {
     158            9 :             CHK_RET(RdmaTxRx(linkLeft, linkRight, stepInfo, links));
     159              :         }
     160            9 :     }
     161            3 :     return HCCL_SUCCESS;
     162              : }
     163              : 
     164              : // NHR每步的算法描述原理函数
     165            9 : HcclResult BroadcastNHROneshot::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep &stepInfo)
     166              : {
     167            9 :     stepInfo.txSliceIdxs.clear();
     168            9 :     stepInfo.rxSliceIdxs.clear();
     169            9 :     stepInfo.nSlices = 1;
     170            9 :     stepInfo.toRank = rankSize;
     171            9 :     stepInfo.fromRank = rankSize;
     172            9 :     stepInfo.step = step;
     173            9 :     stepInfo.myRank = rank;
     174              : 
     175            9 :     u32 nRanks = (rankSize - 1 + (1 << (nSteps - 1 - step))) / (1 << (nSteps - step)); // 本步需要进行收/发的rank数
     176              : 
     177            9 :     u32 deltaRoot = (rank + rankSize - root_) % rankSize;
     178              : 
     179            9 :     u32 deltaRankPair = 1 << (nSteps - 1 - step);
     180            9 :     u32 deltaRankGroup = 1 << (nSteps - step);
     181              : 
     182            9 :     if (deltaRoot / deltaRankGroup < nRanks) {
     183            9 :         if (deltaRoot % deltaRankGroup == 0) {
     184            9 :             stepInfo.toRank = (rank + deltaRankPair) % rankSize;
     185            9 :             stepInfo.txSliceIdxs.push_back(0);
     186              :         }
     187              : 
     188            9 :         if ((deltaRoot + deltaRankPair) % deltaRankGroup == 0) {
     189            0 :             stepInfo.fromRank = (rank + rankSize - deltaRankPair) % rankSize;
     190            0 :             stepInfo.rxSliceIdxs.push_back(0);
     191              :         }
     192              :     }
     193            9 :     return HCCL_SUCCESS;
     194              : }
     195              : 
     196              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_BROADCAST_NHR_ONESHOT, BroadcastNHROneshot);
     197              : }  // namespace hccl
        

Generated by: LCOV version 2.0-1