LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/component - reducer.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 20.7 % 169 35
Test Date: 2026-08-04 10:52:23 Functions: 45.5 % 11 5

            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 "reducer.h"
      12              : 
      13              : namespace hccl {
      14           11 : Reducer::Reducer(const HcclDataType dataType, const HcclReduceOp reductionOp, const u64 reduceAttribute)
      15           11 :     : dataType_(dataType), reductionOp_(reductionOp), reduceAttribute_(reduceAttribute)
      16              : {
      17           11 :     SetPreSyncFunc([](){ return HCCL_SUCCESS; });
      18           11 :     SetPostSyncFunc([](){ return HCCL_SUCCESS; });
      19           11 : }
      20              : 
      21           11 : Reducer::~Reducer()
      22              : {
      23           11 : }
      24              : 
      25           11 : void Reducer::SetPreSyncFunc(std::function<HcclResult()> lambda)
      26              : {
      27           11 :     preSync_ = std::move(lambda);
      28           11 : }
      29              : 
      30           11 : void Reducer::SetPostSyncFunc(std::function<HcclResult()> lambda)
      31              : {
      32           11 :     postSync_ = std::move(lambda);
      33           11 : }
      34              : 
      35            9 : HcclResult Reducer::run(const HcclDispatcher dispatcher, const std::shared_ptr<Transport> &link,
      36              :     const u64 remoteMemOffset, DeviceMem &localSrc, DeviceMem &localDst, DeviceMem &remoteRcvTemp, Stream &stream,
      37              :     DstMemType resultMem, const UserMemType srcMemType) const
      38              : {
      39            9 :     CHK_PTR_NULL(localSrc.ptr());
      40            9 :     CHK_PTR_NULL(localDst.ptr());
      41            9 :     CHK_PTR_NULL(remoteRcvTemp.ptr());
      42            9 :     CHK_PTR_NULL(stream.ptr());
      43              : 
      44            9 :     HcclResult ret = HCCL_SUCCESS;
      45              : 
      46            9 :     u64 dataBytes = remoteRcvTemp.size();
      47            9 :     HCCL_DEBUG("localSrc[%p] localDst[%p] remoteRcvtmep[%p] offset[%llu]", localSrc.ptr(), localDst.ptr(),
      48              :         remoteRcvTemp.ptr(), remoteMemOffset);
      49              : 
      50              :     // server 内 reduce 并且 reduceAttribute_ 也支持,走该分支
      51            9 :     bool isSpInlineReduce = link->IsSpInlineReduce();
      52            9 :     if (link->IsSupportTransportWithReduce() && (RDMA_REDUCE_BITMASK & reduceAttribute_)) {
      53              :         // 数据接收端执行接收动作
      54              :         // RDMA的RxAsync不需要接收端内存信息
      55            0 :         CHK_RET(link->RxAsync(UserMemType::INPUT_MEM, remoteMemOffset, localDst.ptr(), localDst.size(), stream));
      56            0 :         if (link->GetSupportDataReceivedAck()) {
      57            0 :             CHK_RET(link->DataReceivedAck(stream));
      58              :         }
      59            0 :         if (resultMem == DstMemType::RESULT_OUTPUT_MEM) {
      60            0 :             ret = HcclD2DMemcpyAsync(dispatcher, localDst, localSrc, stream);
      61            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
      62              :                 HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", localSrc.ptr(),
      63              :                 localDst.ptr()),
      64              :                 ret);
      65              :         }
      66            9 :     } else if (link->IsSupportTransportWithReduce() && link->GetLinkType() == LinkType::LINK_STANDARD_ROCE) {
      67            0 :         u64 dataCount = localDst.size() / SIZE_TABLE[dataType_];
      68            0 :         DeviceMem &reduceSrc = (localSrc == localDst) ? remoteRcvTemp : localSrc;
      69            0 :         CHK_RET(link->RxWithReduce(srcMemType, remoteMemOffset, remoteRcvTemp.ptr(), dataBytes,
      70              :             reduceSrc.ptr(), localDst.ptr(), dataCount, dataType_, reductionOp_, stream, reduceAttribute_));
      71            9 :     } else if (isSpInlineReduce && (INLINE_REDUCE_BITMASK & reduceAttribute_)) {
      72              :         //  runtime 的inline reduce 接口参数为数据的字节长度
      73            0 :         CHK_RET(link->RxDataSignal(stream));
      74            0 :         void *remoteMem = nullptr;
      75            0 :         CHK_RET(link->GetRemoteMem(srcMemType, &remoteMem));
      76            0 :         CHK_RET(HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + remoteMemOffset,
      77              :             dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, localSrc.ptr(), link->GetRemoteRank(),
      78              :             link->GetLinkType(), INLINE_REDUCE_BIT));
      79              : 
      80            0 :         if (localSrc != localDst) {
      81            0 :             ret = HcclD2DMemcpyAsync(dispatcher, localDst, localSrc, stream);
      82            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
      83              :                 HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", localSrc.ptr(),
      84              :                 localDst.ptr()),
      85              :                 ret);
      86              :         }
      87            0 :         if (link -> GetSupportDataReceivedAck()) {
      88            0 :             CHK_RET(link->TxAck(stream));
      89            0 :             CHK_RET(link->RxAck(stream));
      90            0 :             CHK_RET(link->TxDataSignal(stream));
      91            0 :             CHK_RET(link->RxDataSignal(stream));
      92              :         }
      93            0 :     } else {
      94              :         // 从上一个节点接收数据
      95            9 :         ret = link->RxAsync(UserMemType::INPUT_MEM, remoteMemOffset, remoteRcvTemp.ptr(), dataBytes, stream);
      96            9 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
      97              :             HCCL_ERROR("[Reducer][Run]rx_sync remoteRcvTemp[%p] offset[%llu] size[%llu] "
      98              :             "failed",
      99              :             remoteRcvTemp.ptr(), remoteMemOffset, dataBytes),
     100              :             ret);
     101              : 
     102            9 :         if (link->GetSupportDataReceivedAck()) {
     103            0 :             ret = link->DataReceivedAck(stream);
     104            0 :             CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Reducer][Run]rx_sync data received ack failed"), ret);
     105              :         }
     106              : 
     107            9 :         u64 dataCount = localDst.size() / SIZE_TABLE[dataType_];
     108              : 
     109              :         // 根据目的内存执行reduce
     110            9 :         DeviceMem reduceSrc = (localSrc == localDst) ? remoteRcvTemp : localSrc;
     111            9 :         ret = HcclReduceAsync(dispatcher, reduceSrc.ptr(), dataCount, dataType_, reductionOp_, stream, localDst.ptr(),
     112            9 :             INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttribute_);
     113              : 
     114            9 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     115              :             HCCL_ERROR("[Reducer][Run]reduce_async remoteRcvTemp[%p] localSrc[%p] "
     116              :             "localDst[%p] failed",
     117              :             remoteRcvTemp.ptr(), localSrc.ptr(), localDst.ptr()),
     118              :             ret);
     119            9 :     }
     120              : 
     121            9 :     return ret;
     122              : }
     123              : 
     124            0 : HcclResult Reducer::PrepareRxMems(const std::vector<ReducerMemoryInfo> &reducerMems,
     125              :     std::vector<RxMemoryInfo> &rxMems) const
     126              : {
     127            0 :     rxMems.reserve(reducerMems.size());
     128            0 :     for (const ReducerMemoryInfo &reduceMem : reducerMems) {
     129            0 :         rxMems.emplace_back(RxMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     130            0 :             reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() });
     131              :     }
     132            0 :     return HCCL_SUCCESS;
     133              : }
     134              : 
     135            0 : HcclResult Reducer::PrepareRxWithReduceMems(const std::vector<ReducerMemoryInfo> &reducerMems,
     136              :     std::vector<RxWithReduceMemoryInfo> &rxWithReduceMems) const
     137              : {
     138            0 :     rxWithReduceMems.reserve(reducerMems.size());
     139            0 :     for (const ReducerMemoryInfo &reduceMem : reducerMems) {
     140            0 :         u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
     141            0 :         DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
     142              : 
     143            0 :         rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
     144            0 :             reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size(), reduceSrc.ptr(), reduceMem.localdst.ptr(),
     145              :             dataCount});
     146            0 :     }
     147            0 :     return HCCL_SUCCESS;
     148              : }
     149              : 
     150            0 : HcclResult Reducer::run(const HcclDispatcher dispatcher, const std::shared_ptr<Transport> &link,
     151              :     const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream, DstMemType resultMem) const
     152              : {
     153            0 :     CHK_PTR_NULL(stream.ptr());
     154              : 
     155            0 :     LinkType linkType = link->GetLinkType();
     156            0 :     bool isSpInlineReduce = link->IsSpInlineReduce();
     157            0 :     bool isSpRdmaReduce = RDMA_REDUCE_BITMASK & reduceAttribute_;
     158            0 :     bool isSpTransportWithReduce = link->IsSupportTransportWithReduce();
     159            0 :     HcclResult ret = HCCL_SUCCESS;
     160              : 
     161            0 :     if (isSpTransportWithReduce && isSpRdmaReduce) {
     162              :         // 数据接收端执行接收动作
     163              :         // RDMA的RxAsync不需要接收端内存信息
     164            0 :         std::vector<RxMemoryInfo> rxMems;
     165            0 :         CHK_RET(PrepareRxMems(reducerMems, rxMems));
     166            0 :         CHK_RET(link->RxAsync(rxMems, stream));
     167            0 :         if (link->GetSupportDataReceivedAck()) {
     168            0 :             CHK_RET(link->DataReceivedAck(stream));
     169              :         }
     170            0 :         CHK_RET(preSync_());
     171            0 :         if (resultMem == DstMemType::RESULT_OUTPUT_MEM) {
     172            0 :             for (ReducerMemoryInfo reduceMem : reducerMems) {
     173            0 :                 ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     174            0 :                 CHK_PRT_RET(ret != HCCL_SUCCESS,
     175              :                     HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     176              :                     reduceMem.localdst.ptr()),
     177              :                     ret);
     178            0 :             }
     179              :         }
     180            0 :         CHK_RET(postSync_());
     181            0 :     } else if (isSpTransportWithReduce && (linkType == LinkType::LINK_STANDARD_ROCE)) {
     182            0 :         std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     183            0 :         CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
     184            0 :         CHK_RET(preSync_());
     185            0 :         CHK_RET(link->RxWithReduce(rxWithReduceMems, dataType_, reductionOp_, stream, reduceAttribute_));
     186            0 :         CHK_RET(postSync_());
     187            0 :     } else if (isSpInlineReduce && (INLINE_REDUCE_BITMASK & reduceAttribute_)) {
     188            0 :         CHK_RET(link->RxDataSignal(stream));
     189            0 :         void *remoteMem = nullptr;
     190            0 :         CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     191            0 :         CHK_RET(preSync_());
     192            0 :         for (ReducerMemoryInfo reduceMem : reducerMems) {
     193            0 :             const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     194            0 :             CHK_RET(
     195              :                 HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
     196              :                 dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
     197              :                 link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
     198              : 
     199            0 :             if (reduceMem.localsrc != reduceMem.localdst) {
     200            0 :                 ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     201            0 :                 CHK_PRT_RET(ret != HCCL_SUCCESS,
     202              :                     HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     203              :                     reduceMem.localdst.ptr()),
     204              :                     ret);
     205              :             }
     206            0 :             HCCL_DEBUG("[Reducer][Run]memcpy_async localSrc is [%p]", reduceMem.localsrc.ptr());
     207            0 :         }
     208            0 :         CHK_RET(postSync_());
     209            0 :     } else {
     210            0 :         std::vector<RxMemoryInfo> rxMems;
     211            0 :         CHK_RET(PrepareRxMems(reducerMems, rxMems));
     212            0 :         std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     213            0 :         CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
     214            0 :         CHK_RET(preSync_());
     215            0 :         CHK_RET(link->RxAsync(rxMems, stream));
     216            0 :         CHK_RET(postSync_());
     217            0 :         if (link->GetSupportDataReceivedAck()) {
     218            0 :             CHK_RET(link->DataReceivedAck(stream));
     219              :         }
     220            0 :         for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
     221            0 :             CHK_RET(HcclReduceAsync(dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_,
     222              :                 reductionOp_, stream, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttribute_));
     223              :         }
     224            0 :     }
     225              : 
     226            0 :     return HCCL_SUCCESS;
     227              : }
     228              : 
     229            0 : HcclResult Reducer::run(const HcclDispatcher dispatcher, const std::shared_ptr<Transport> &link,
     230              :     const std::vector<ReducerMemoryInfo> &reducerMems, u32 notifyIdx, Stream &stream, DstMemType resultMem) const
     231              : {
     232              :     (void) resultMem;
     233            0 :     CHK_PTR_NULL(stream.ptr());
     234            0 :     CHK_SMART_PTR_NULL(link);
     235              : 
     236            0 :     bool isSpInlineReduce = link->IsSpInlineReduce();
     237            0 :     HcclResult ret = HCCL_SUCCESS;
     238              :     
     239            0 :     if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttribute_))) {
     240            0 :         CHK_RET(link->Wait(notifyIdx, stream));
     241            0 :         void *remoteMem = nullptr;
     242            0 :         CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
     243            0 :         CHK_RET(preSync_());
     244            0 :         for (ReducerMemoryInfo reduceMem : reducerMems) {
     245            0 :             const u64 dataBytes = reduceMem.remoteRcvTemp.size();
     246            0 :             CHK_RET(
     247              :                 HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
     248              :                 dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
     249              :                 link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
     250            0 :             HCCL_DEBUG("[Reducer][Run]memcpy_async localSrc[%p]", reduceMem.localsrc.ptr());
     251            0 :             if (reduceMem.localsrc != reduceMem.localdst) {
     252            0 :                 ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
     253            0 :                 CHK_PRT_RET(ret != HCCL_SUCCESS,
     254              :                     HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
     255              :                     reduceMem.localdst.ptr()),
     256              :                     ret);
     257              :             }
     258            0 :         }
     259            0 :         CHK_RET(postSync_());
     260            0 :     }
     261              :     else {
     262            0 :         std::vector<RxMemoryInfo> rxMems;
     263            0 :         CHK_RET(PrepareRxMems(reducerMems, rxMems));
     264              : 
     265            0 :         std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
     266            0 :         CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
     267            0 :         CHK_RET(preSync_());
     268              :         
     269            0 :         CHK_RET(link->Wait(notifyIdx, stream));
     270            0 :         for(auto& mem : rxMems){
     271            0 :             CHK_PTR_NULL(mem.dst);
     272            0 :             void *srcMemPtr = nullptr;
     273            0 :             CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
     274            0 :             DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + mem.srcOffset, mem.len);
     275            0 :             DeviceMem dstDevMem(static_cast<s8 *>(mem.dst),mem.len);
     276            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher, dstDevMem, srcDevMem, stream, link->GetRemoteRank(), 
     277              :                 link->GetLinkType()));
     278            0 :         }
     279            0 :         CHK_RET(postSync_());
     280            0 :         if (link->GetSupportDataReceivedAck()) {
     281            0 :             CHK_RET(link->DataReceivedAck(stream));
     282              :         }
     283            0 :         for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
     284            0 :             CHK_RET(HcclReduceAsync(dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_,
     285              :                 reductionOp_, stream, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
     286              :                 reduceAttribute_));
     287              :         }
     288            0 :     }
     289            0 :     return HCCL_SUCCESS;
     290              : }
     291              : } // namespace hccl
        

Generated by: LCOV version 2.0-1