LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_reduce - all_reduce_hd_optim.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 163 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 13 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 <cmath>
      12              : #include "alg_template_register.h"
      13              : #include "all_reduce_hd_optim_pub.h"
      14              : namespace hccl {
      15            0 : AllReduceHDOptim::AllReduceHDOptim(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
      16              : {
      17            0 : }
      18              : 
      19            0 : AllReduceHDOptim::~AllReduceHDOptim()
      20              : {
      21            0 : }
      22              : 
      23            0 : HcclResult AllReduceHDOptim::Prepare(u64 reduceAttrBitMap, std::vector<Stream> &meshStreams,
      24              :     std::vector<std::shared_ptr<LocalNotify>> &meshSignal, std::vector<std::shared_ptr<LocalNotify>> &meshSignalAux,
      25              :     u32 userRank, HcomCollOpInfo *opInfo, bool aicpu)
      26              : {
      27            0 :     reduceAttr_ = reduceAttrBitMap;
      28            0 :     userRank_ = userRank;
      29            0 :     meshStreams_ = meshStreams;
      30            0 :     meshSignal_ = &meshSignal;
      31            0 :     meshSignalAux_ = &meshSignalAux;
      32            0 :     opInfo_ = opInfo;
      33            0 :     aicpu_ = aicpu;
      34            0 :     return HCCL_SUCCESS;
      35              : }
      36              : 
      37            0 : HcclResult AllReduceHDOptim::MainRecordSub(u32 streamNum)
      38              : {
      39            0 :     if(aicpu_) {
      40            0 :         return HCCL_SUCCESS;
      41              :     }
      42            0 :     for (u32 signalIndex = 0; signalIndex < streamNum; signalIndex++) {
      43            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
      44              :     }
      45            0 :     return HCCL_SUCCESS;
      46              : }
      47              : 
      48            0 : HcclResult AllReduceHDOptim::SubWaitMain(u32 streamNum)
      49              : {
      50            0 :     if(aicpu_) {
      51            0 :         return HCCL_SUCCESS;
      52              :     }
      53            0 :     for (u32 streamIndex = 0; streamIndex < streamNum; streamIndex++) {
      54            0 :         CHK_RET(LocalNotify::Wait(
      55              :             meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex], profilerInput_.stage));
      56              :     }
      57            0 :     return HCCL_SUCCESS;
      58              : }
      59              : 
      60            0 : HcclResult AllReduceHDOptim::MainWaitSub(u32 streamNum)
      61              : {
      62            0 :     if(aicpu_) {
      63            0 :         return HCCL_SUCCESS;
      64              :     }
      65            0 :     for (u32 signalIndex = 0; signalIndex < streamNum; signalIndex++) {
      66            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
      67              :     }
      68            0 :     return HCCL_SUCCESS;
      69              : }
      70              : 
      71            0 : HcclResult AllReduceHDOptim::SubRecordMain(u32 streamNum)
      72              : {
      73            0 :     if(aicpu_) {
      74            0 :         return HCCL_SUCCESS;
      75              :     }
      76            0 :     for (u32 streamIndex = 0; streamIndex < streamNum; streamIndex++) {
      77            0 :         CHK_RET(
      78              :             LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex], profilerInput_.stage));
      79              :     }
      80            0 :     return HCCL_SUCCESS;
      81              : }
      82              : 
      83              : // allreduce算法的函数入口
      84            0 : HcclResult AllReduceHDOptim::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
      85              : {
      86            0 :     HcclResult ret = HCCL_SUCCESS;
      87            0 :     CHK_SMART_PTR_NULL(dispatcher_);
      88            0 :     CHK_PTR_NULL(stream_.ptr());
      89            0 :     HCCL_INFO("AllReduceHDOptim run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
      90              :         rank, rankSize, outputMem_.ptr(), outputMem_.ptr(), count_);
      91              : 
      92            0 :     if (links.size() < rankSize) {
      93            0 :         HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]",
      94              :             rank, links.size(), rankSize);
      95            0 :         return HCCL_E_INTERNAL;
      96              :     }
      97              : 
      98            0 :     if (meshStreams_.size() < base) {
      99            0 :         HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] meshStreams_[%llu] is less than need[%u]",
     100              :             rank, meshStreams_.size(), base);
     101            0 :         return HCCL_E_INTERNAL;
     102              :     }
     103            0 :     u32 totalSize = SIZE_TABLE[dataType_] * count_;
     104            0 :     userMemIn = DeviceMem::create(opInfo_->inputAddr, totalSize);
     105            0 :     userMemOut = DeviceMem::create(opInfo_->outputAddr, totalSize);
     106            0 :     emptyMem_ = outputMem_.range(0, 0);
     107            0 :     nSteps = static_cast<u32>(log2(rankSize));
     108            0 :     stepPow = static_cast<u32>(pow(base, nSteps));
     109              : 
     110            0 :     ret = RunPreCopy(rank, rankSize, links);
     111            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     112              :         HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunPreCopy step" ,
     113              :             rank, count_), ret);
     114              : 
     115            0 :     if (rank < stepPow) {
     116            0 :         ret = RunAllReduceHDOptim(rank, rankSize, links);
     117            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     118              :             HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunAllReduceHDOptim step",
     119              :                 rank, count_), ret);
     120              :     }
     121              :     
     122            0 :     if (stepPow != rankSize) {
     123            0 :         ret = RunFinalStep(rank, rankSize, links);
     124            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     125              :             HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunFinalStep step",
     126              :                 rank, count_), ret);
     127              :     }
     128              : 
     129            0 :     HCCL_INFO("AllReduceHDOptim finished: rank[%u] ranksize[%u]", rank, rankSize);
     130            0 :     return HCCL_SUCCESS;
     131              : }
     132              : 
     133            0 : HcclResult AllReduceHDOptim::RunPreCopy(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     134              : {
     135            0 :     u32 totalSize = SIZE_TABLE[dataType_] * count_;
     136              : 
     137            0 :     DeviceMem src = userMemIn.range(0, totalSize);
     138            0 :     DeviceMem dst = outputMem_.range(0, totalSize);
     139            0 :     DeviceMem nextDst = outputMem_.range(totalSize, totalSize);
     140              : 
     141            0 :     if (stepPow == rankSize) {
     142            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, src, stream_));
     143            0 :         return HCCL_SUCCESS;
     144              :     }
     145              : 
     146            0 :     u32 neighCur = rank ^ (1 << nSteps);
     147            0 :     if (neighCur < rankSize) {
     148              :         // reduce写
     149            0 :         if (rank >= pow(base, nSteps)) {
     150            0 :             CHK_PTR_NULL(links[neighCur]);
     151            0 :             CHK_RET(links[neighCur]->RxAck(stream_));
     152            0 :             void *remMemPtr = nullptr;
     153            0 :             CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     154            0 :             dst = DeviceMem::create(static_cast<u8 *>(remMemPtr), totalSize);
     155            0 :             CHK_RET(HcclReduceAsync(
     156              :                 dispatcher_,
     157              :                 static_cast<void *>(src.ptr()),
     158              :                 count_,
     159              :                 dataType_,
     160              :                 reductionOp_,
     161              :                 stream_,
     162              :                 static_cast<void *>(dst.ptr()),
     163              :                 links[neighCur]->GetRemoteRank(),
     164              :                 links[neighCur]->GetLinkType(),
     165              :                 INLINE_REDUCE_BIT));
     166            0 :             CHK_RET(links[neighCur]->TxDataSignal(stream_));
     167              :         } else {
     168            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     169            0 :             CHK_RET(links[neighCur]->TxAck(stream_));
     170            0 :             CHK_RET(links[neighCur]->RxDataSignal(stream_));
     171            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, dst, stream_));
     172              :         }
     173              :     } else {
     174            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     175            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, dst, stream_));
     176              :     }
     177              : 
     178            0 :     return HCCL_SUCCESS;
     179            0 : }
     180              : 
     181            0 : HcclResult AllReduceHDOptim::RunBetweenStep(
     182              :     u32 rank, u32 step, u32 neighBefore, u32 neighNext, u32 rankSize, const std::vector<LINK> &links)
     183              : {
     184              :     (void) rank;
     185            0 :     u32 totalSize = SIZE_TABLE[dataType_] * count_;
     186              : 
     187            0 :     DeviceMem src;
     188            0 :     DeviceMem dst;
     189              : 
     190            0 :     if ((step == 1) && (stepPow == rankSize)) {
     191              :         // 二次幂整第一步写 串行同步
     192            0 :         CHK_RET(links[neighBefore]->TxDataSignal(stream_));
     193            0 :         CHK_RET(links[neighBefore]->RxDataSignal(stream_));
     194            0 :         CHK_RET(MainRecordSub(base));
     195            0 :         CHK_RET(SubWaitMain(base));
     196            0 :     } else {
     197            0 :         CHK_RET(MainRecordSub(base));
     198            0 :         CHK_RET(SubWaitMain(base));
     199            0 :         CHK_RET(links[neighBefore]->TxDataSignal(stream_));
     200            0 :         CHK_RET(links[neighBefore]->RxDataSignal(stream_));
     201              :     }
     202              : 
     203            0 :     src = outputMem_.range(step * totalSize, totalSize);
     204            0 :     if ((step == (nSteps - 1)) && (static_cast<u32>(pow(base, nSteps)) == rankSize)) {
     205            0 :         dst = userMemOut.range(0, totalSize);
     206              :     } else {
     207            0 :         dst = outputMem_.range((step + 1) * totalSize, totalSize);
     208              :     }
     209            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, aicpu_?stream_:meshStreams_[0]));
     210              : 
     211            0 :     CHK_RET(links[neighNext]->TxAck(aicpu_?stream_:meshStreams_[1]));
     212            0 :     CHK_RET(links[neighNext]->RxAck(aicpu_?stream_:meshStreams_[1]));
     213              : 
     214            0 :     CHK_RET(SubRecordMain(base));
     215            0 :     CHK_RET(MainWaitSub(base));
     216            0 :     return HCCL_SUCCESS;
     217            0 : }
     218              : 
     219            0 : HcclResult AllReduceHDOptim::RunAllReduceHDOptim(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     220              : {
     221            0 :     u32 unitSize = SIZE_TABLE[dataType_];
     222            0 :     u32 totalSize = unitSize * count_;
     223              : 
     224            0 :     DeviceMem src;
     225            0 :     DeviceMem dst;
     226              : 
     227            0 :     u32 neighCur = rank ^ (1 << 0);
     228            0 :     CHK_RET(links[neighCur]->TxAck(stream_));
     229            0 :     CHK_RET(links[neighCur]->RxAck(stream_));
     230              : 
     231            0 :     u32 neighNext = 0;
     232            0 :     void *remMemPtr = nullptr;
     233            0 :     for (u32 step = 1; step <= nSteps; step++) {
     234            0 :         if ((step != nSteps) || (stepPow != rankSize)) {
     235            0 :             dst = outputMem_.range(step * totalSize, totalSize);
     236              :         } else {
     237            0 :             dst = userMemOut.range(0, totalSize);
     238              :         }
     239            0 :         CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     240            0 :         src = DeviceMem::create(static_cast<u8 *>(remMemPtr) + (step - 1) * totalSize, totalSize);
     241            0 :         if ((step == 1) && (stepPow == rankSize)) {
     242              :             // 二次幂整第一步写
     243            0 :             src = userMemIn.range(0, totalSize);
     244            0 :             dst = DeviceMem::create(static_cast<u8 *>(remMemPtr) + totalSize, totalSize);
     245              :         }
     246              : 
     247            0 :         CHK_RET(HcclReduceAsync(dispatcher_,
     248              :             static_cast<void *>(src.ptr()),
     249              :             count_,
     250              :             dataType_,
     251              :             reductionOp_,
     252              :             stream_,
     253              :             static_cast<void *>(dst.ptr()),
     254              :             links[neighCur]->GetRemoteRank(),
     255              :             links[neighCur]->GetLinkType(),
     256              :             INLINE_REDUCE_BIT));
     257              : 
     258            0 :         if (step != nSteps) {
     259            0 :             neighNext = rank ^ (1 << (step));
     260            0 :             CHK_RET(RunBetweenStep(rank, step, neighCur, neighNext, rankSize, links));
     261            0 :             neighCur = neighNext;
     262              :         }
     263              :     }
     264              : 
     265            0 :     CHK_RET(links[neighCur]->TxDataSignal(stream_));
     266            0 :     CHK_RET(links[neighCur]->RxDataSignal(stream_));
     267              : 
     268            0 :     return HCCL_SUCCESS;
     269            0 : }
     270              : 
     271            0 : HcclResult AllReduceHDOptim::RunFinalStep(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     272              : {
     273            0 :     u32 unitSize = SIZE_TABLE[dataType_];
     274            0 :     u32 totalSize = unitSize * count_;
     275              : 
     276            0 :     DeviceMem src = outputMem_.range(nSteps * totalSize, totalSize);
     277            0 :     DeviceMem dst = userMemOut.range(0, totalSize);
     278              : 
     279            0 :     u32 neighCur = rank ^ (1 << nSteps);
     280            0 :     if (neighCur >= rankSize) {
     281            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     282            0 :     } else if (rank < pow(base, nSteps)) {
     283            0 :         CHK_RET(links[neighCur]->TxAck(stream_));
     284            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     285            0 :         CHK_RET(links[neighCur]->RxDataSignal(stream_));
     286              :     } else {
     287            0 :         CHK_RET(links[neighCur]->RxAck(stream_));
     288            0 :         void *remMemPtr = nullptr;
     289            0 :         CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     290            0 :         src = DeviceMem::create(static_cast<u8 *>(remMemPtr) + nSteps * totalSize, totalSize);
     291            0 :         CHK_RET(HcclD2DMemcpyAsync(
     292              :             dispatcher_, dst, src, stream_, links[neighCur]->GetRemoteRank(), links[neighCur]->GetLinkType()));
     293            0 :         CHK_RET(links[neighCur]->TxDataSignal(stream_));
     294              :     }
     295            0 :     return HCCL_SUCCESS;
     296            0 : }
     297              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_HD_OPTIM, AllReduceHDOptim);
     298              : }  // namespace hccl
        

Generated by: LCOV version 2.0-1