LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_reduce_scatter - reduce_scatter_v_pipeline.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 133 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 "reduce_scatter_v_pipeline.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : constexpr u32 STEP_OFFSET_TWO = 2;
      15              : 
      16              : namespace hccl {
      17            0 : ReduceScatterVPipeline::ReduceScatterVPipeline(const HcclDispatcher dispatcher) : ReduceScatterPipeline(dispatcher) {}
      18              : 
      19            0 : ReduceScatterVPipeline::~ReduceScatterVPipeline() {}
      20              : 
      21            0 : HcclResult ReduceScatterVPipeline::RunIntraServer(u32 step, u64 remoteOffset)
      22              : {
      23            0 :     for (u32 i = 1; i < intraRankSize_; i++) {
      24            0 :         u32 remIntraRankId = (intraRankId_ + i) % intraRankSize_;
      25            0 :         CHK_RET(intraLinks_[remIntraRankId]->TxAck(subStream_[i]));
      26            0 :         CHK_RET(intraLinks_[remIntraRankId]->RxAck(subStream_[i]));
      27            0 :         void* remoteMemPtr = nullptr;
      28            0 :         CHK_RET(intraLinks_[remIntraRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMemPtr));
      29              : 
      30              :         // 本次机内待发送数据的index坐标
      31            0 :         u32 index = (((interRankId_ + 1 + step) % interRankSize_) * intraRankSize_ + remIntraRankId);
      32            0 :         Slice userSlice = slices_[index];
      33            0 :         u64 srcOffset = userSlice.offset;
      34              : 
      35            0 :         u64 offset = srcOffset % HCCL_MIN_SLICE_ALIGN_910B;
      36            0 :         DeviceMem src = DeviceMem::create(static_cast<u8*>(usrInMem_) + srcOffset, userSlice.size);
      37            0 :         DeviceMem dst = DeviceMem::create(static_cast<u8*>(remoteMemPtr) + remoteOffset + offset, userSlice.size);
      38              : 
      39            0 :         CHK_RET(HcclReduceAsync(
      40              :             dispatcher_, src.ptr(), userSlice.size / unitSize_, dataType_, reductionOp_, subStream_[i], dst.ptr(),
      41              :             intraLinks_[remIntraRankId]->GetRemoteRank(), intraLinks_[remIntraRankId]->GetLinkType(),
      42              :             INLINE_REDUCE_BIT));
      43              : 
      44            0 :         CHK_RET(intraLinks_[remIntraRankId]->TxDataSignal(subStream_[i]));
      45            0 :         CHK_RET(intraLinks_[remIntraRankId]->RxDataSignal(subStream_[i]));
      46            0 :     }
      47            0 :     return HCCL_SUCCESS;
      48              : }
      49              : 
      50            0 : HcclResult ReduceScatterVPipeline::RunInterServer(u32 step, const LINK& prevInterLink, const LINK& nextInterLink)
      51              : {
      52            0 :     u32 dmaMemSliceNum = dmaMem_.size();
      53            0 :     u32 rxDMAMemSliceId = (step + 1) % dmaMemSliceNum;
      54            0 :     u32 txDMAMemSliceId = step % dmaMemSliceNum;
      55              : 
      56            0 :     u32 txindex = (((interRankId_ + 1 + step) % interRankSize_) * intraRankSize_ + intraRankId_);
      57            0 :     Slice txUserSlice = slices_[txindex];
      58            0 :     u64 txSliceOffset = txUserSlice.offset;
      59            0 :     u64 offset = txSliceOffset % HCCL_MIN_SLICE_ALIGN_910B;
      60            0 :     u64 rxInterOffset = rxDMAMemSliceId * blockSize_ + offset;
      61            0 :     void* txLocalAddr = static_cast<u8*>(dmaMem_[txDMAMemSliceId].ptr()) + offset;
      62              : 
      63            0 :     DeviceMem srcMem = DeviceMem::create(txLocalAddr, txUserSlice.size);
      64            0 :     CHK_RET(senderInfo_->run(nextInterLink, rxInterOffset, srcMem, subStream_[0])); // 发 srcMem -> rxInterOffset(rem)
      65            0 :     HCCL_DEBUG(
      66              :         "[ReduceScatterVPipeline][RunInterServer] local rank[%u] localOffset[%llu]tx with slice[%llu]", rankId_,
      67              :         rxInterOffset, txUserSlice.size);
      68              : 
      69            0 :     u32 rxindex = (((interRankId_ + 2 + step) % interRankSize_) * intraRankSize_ + intraRankId_);
      70            0 :     Slice rxUserSlice = slices_[rxindex];
      71            0 :     u64 rxSliceOffset = rxUserSlice.offset;
      72            0 :     u64 rxOffset = rxSliceOffset % HCCL_MIN_SLICE_ALIGN_910B;
      73            0 :     u64 rxMemOffset = txDMAMemSliceId * blockSize_ + rxOffset;
      74            0 :     void* rxLocalAddr = static_cast<u8*>(dmaMem_[rxDMAMemSliceId].ptr()) + rxOffset;
      75              : 
      76            0 :     DeviceMem rxLocalMem = DeviceMem::create(rxLocalAddr, rxUserSlice.size);
      77            0 :     CHK_RET(reducerInfo_->run(
      78              :         dispatcher_, prevInterLink, rxMemOffset, rxLocalMem, rxLocalMem,
      79              :         rxLocalMem, // 收 rxMemOffset(rem) -> rxLocalMem
      80              :         subStream_[0]));
      81            0 :     return HCCL_SUCCESS;
      82            0 : }
      83              : 
      84            0 : HcclResult ReduceScatterVPipeline::CopyToScratchBuffer(u32 step)
      85              : {
      86            0 :     u32 dmaMemSliceNum = dmaMem_.size();
      87            0 :     u32 dmaMemSliceId = step % dmaMemSliceNum;
      88              : 
      89            0 :     u32 index = (((interRankId_ + 1 + step) % interRankSize_) * intraRankSize_ + intraRankId_);
      90            0 :     Slice userslice = slices_[index];
      91            0 :     u64 srcOffset = userslice.offset;
      92            0 :     u64 offset = srcOffset % HCCL_MIN_SLICE_ALIGN_910B;
      93              : 
      94            0 :     void* srcAddr = static_cast<u8*>(usrInMem_) + srcOffset;
      95            0 :     DeviceMem locSrc = DeviceMem::create(srcAddr, userslice.size);
      96            0 :     DeviceMem locDst = DeviceMem::create(static_cast<u8*>(dmaMem_[dmaMemSliceId].ptr()) + offset, userslice.size);
      97            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDst, locSrc, stream_));
      98            0 :     return HCCL_SUCCESS;
      99            0 : }
     100              : 
     101            0 : HcclResult ReduceScatterVPipeline::RunAsync()
     102              : {
     103              :     // inter ring algo
     104            0 :     u32 prevInterRankId = (interRankId_ + 1) % interRankSize_;
     105            0 :     u32 nextInterRankId = (interRankId_ - 1 + interRankSize_) % interRankSize_;
     106            0 :     LINK prevInterLink = interLinks_[prevInterRankId];
     107            0 :     LINK nextInterLink = interLinks_[nextInterRankId];
     108              :     // 当前使用3块DMAMem buffer
     109            0 :     u32 dmaMemSliceNum = dmaMem_.size();
     110              : 
     111            0 :     for (u32 step = 0; step < interRankSize_; step++) {
     112            0 :         u32 begin = 0;
     113            0 :         if (step == 0) {
     114            0 :             begin = 1;
     115            0 :             CHK_RET(CopyToScratchBuffer(step));
     116            0 :             CHK_RET(MainRecordSub(begin));
     117            0 :             CHK_RET(SubWaitMain(begin));
     118              :         }
     119              :         // // server内做SDMA的reduce
     120            0 :         u64 remoteOffset = (step % dmaMemSliceNum) * blockSize_;
     121            0 :         CHK_RET(RunIntraServer(step, remoteOffset));
     122            0 :         CHK_RET(SubRecordMain(begin));
     123            0 :         CHK_RET(MainWaitSub(begin));
     124            0 :         if (step < interRankSize_ - 1) {
     125              :             // 把下一块切片从userIn 做拷贝到CCLBuffer
     126            0 :             CHK_RET(CopyToScratchBuffer(step + 1));
     127              :             // // 全部流同步,确保SDMA执行完成
     128            0 :             CHK_RET(MainRecordSub(0));
     129            0 :             CHK_RET(SubWaitMain(0));
     130            0 :             CHK_RET(prevInterLink->TxAck(subStream_[0]));
     131            0 :             CHK_RET(nextInterLink->RxAck(subStream_[0]));
     132              :             // server间做RDMA的reduce,可与下一个step的SDMA并发执行
     133            0 :             CHK_RET(RunInterServer(step, prevInterLink, nextInterLink));
     134            0 :             CHK_RET(prevInterLink->PostFinAck(subStream_[0]));
     135            0 :             CHK_RET(nextInterLink->WaitFinAck(subStream_[0]));
     136              :             // inter的最后一步需要barrier确保数据发完
     137            0 :             if (step == interRankSize_ - STEP_OFFSET_TWO) {
     138            0 :                 CHK_RET(ExecuteBarrier(prevInterLink, nextInterLink, subStream_[0]));
     139              :             }
     140              :         }
     141              :     }
     142              :     // 把对应的切片从CCLBuffer拷贝到userOut
     143            0 :     Slice userSlice = slices_[rankId_];
     144            0 :     u64 srcOffset = userSlice.offset % HCCL_MIN_SLICE_ALIGN_910B;
     145            0 :     void* locSrcAddr = static_cast<u8*>(dmaMem_[(interRankSize_ - 1) % dmaMemSliceNum].ptr()) + srcOffset;
     146              : 
     147            0 :     DeviceMem locSrc = DeviceMem::create(locSrcAddr, userSlice.size);
     148            0 :     DeviceMem locDst = DeviceMem::create(static_cast<u8*>(usrOutMem_), userSlice.size);
     149            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDst, locSrc, stream_));
     150            0 :     HCCL_INFO("[ReduceScatterVPipeline][RunAsync]ReduceScatterVPipeline finished rankId[%u] ", rankId_);
     151            0 :     return HCCL_SUCCESS;
     152            0 : }
     153              : 
     154              : // 适配新CollExecutor接口
     155            0 : HcclResult ReduceScatterVPipeline::Prepare(
     156              :     HcomCollOpInfo* opInfo, DeviceMem& cclBuffer, const u64 bufferSize, const std::vector<Slice>& slices,
     157              :     const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo, Stream& mainStream,
     158              :     std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
     159              :     std::vector<std::shared_ptr<LocalNotify>>& notifySub, u64 reduceAttrBitMap)
     160              : {
     161            0 :     reduceAttr_ = reduceAttrBitMap;
     162            0 :     opInfo_ = opInfo;
     163              : 
     164            0 :     unitSize_ = SIZE_TABLE[opInfo_->dataType];
     165            0 :     usrInMem_ = opInfo_->inputAddr;
     166            0 :     usrOutMem_ = opInfo_->outputAddr;
     167            0 :     reductionOp_ = opInfo_->reduceOp;
     168            0 :     dataType_ = opInfo_->dataType;
     169            0 :     slices_ = slices;
     170              : 
     171            0 :     stream_ = mainStream;
     172            0 :     subStream_ = subStream;
     173              : 
     174            0 :     intraRankSize_ = level0CommInfo.localRankSize;
     175            0 :     interRankSize_ = level1CommInfo.localRankSize;
     176            0 :     intraRankId_ = level0CommInfo.localRank;
     177            0 :     interRankId_ = level1CommInfo.localRank;
     178            0 :     rankId_ = intraRankId_ + interRankId_ * intraRankSize_;
     179              : 
     180            0 :     streamNotifyMain_ = notifyMain;
     181            0 :     if (streamNotifyMain_.size() < intraRankSize_) {
     182            0 :         HCCL_ERROR(
     183              :             "[ReduceScatterVPipeline][Prepare]rank[%u] streamNotifyMain_ size [%u] error, is smaller than,"
     184              :             "intraRankSize_[%u]",
     185              :             rankId_, streamNotifyMain_.size(), intraRankSize_);
     186            0 :         return HCCL_E_INTERNAL;
     187              :     }
     188            0 :     streamNotifySub_ = notifySub;
     189            0 :     if (streamNotifySub_.size() < intraRankSize_) {
     190            0 :         HCCL_ERROR(
     191              :             "[ReduceScatterVPipeline][Prepare]rank[%u] streamNotifySub_ size [%u] error, is smaller than,"
     192              :             "intraRankSize_[%u]",
     193              :             rankId_, streamNotifySub_.size(), intraRankSize_);
     194            0 :         return HCCL_E_INTERNAL;
     195              :     }
     196              : 
     197            0 :     intraLinks_ = level0CommInfo.links;
     198            0 :     interLinks_ = level1CommInfo.links;
     199              : 
     200              :     // 3级流水,使用3块DMAMem
     201            0 :     cclBuffer_ = cclBuffer;
     202            0 :     bufferSize_ = bufferSize;
     203              : 
     204            0 :     blockSize_ = (bufferSize_ / (HCCL_MIN_SLICE_ALIGN_910B * PIPELINE_DEPTH)) * HCCL_MIN_SLICE_ALIGN_910B;
     205              : 
     206            0 :     for (u32 i = 0; i < pipDepth_; i++) {
     207            0 :         DeviceMem mem = DeviceMem::create(static_cast<u8*>(cclBuffer_.ptr()) + blockSize_ * i, blockSize_);
     208            0 :         dmaMem_.push_back(mem);
     209            0 :     }
     210              : 
     211            0 :     HCCL_INFO(
     212              :         "[ReduceScatterVPipeline][Prepare]streamNum[%u], streamNotifyMainNum[%u], streamNotifySubNum[%u]",
     213              :         subStream_.size(), streamNotifyMain_.size(), streamNotifySub_.size());
     214            0 :     HCCL_INFO(
     215              :         "[ReduceScatterVPipeline][Prepare]interLinksNum[%u], intraLinksNum[%u]", interLinks_.size(),
     216              :         intraLinks_.size());
     217            0 :     senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
     218            0 :     CHK_SMART_PTR_NULL(senderInfo_);
     219            0 :     reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
     220            0 :     CHK_SMART_PTR_NULL(reducerInfo_);
     221            0 :     return HCCL_SUCCESS;
     222              : }
     223              : 
     224              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_V_PIPELINE, ReduceScatterVPipeline);
     225              : } // namespace hccl
        

Generated by: LCOV version 2.0-1