LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_gather - all_gather_hccs_sio.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 113 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 9 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 "all_gather_hccs_sio.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : AllGatherHccsSio::AllGatherHccsSio(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
      16              : 
      17            0 : AllGatherHccsSio::~AllGatherHccsSio() {}
      18              : 
      19            0 : HcclResult AllGatherHccsSio::Prepare(
      20              :     SubCommInfo& outerCommInfoHccs, SubCommInfo& outerCommInfoSio, DeviceMem& usrInMem, DeviceMem& usrOutMem, u64 count,
      21              :     const HcclDataType dataType, const Stream& mainStream, std::vector<Stream>& meshStreams,
      22              :     std::vector<std::shared_ptr<LocalNotify>>& meshSignal, std::vector<std::shared_ptr<LocalNotify>>& meshSignalAux,
      23              :     u32 userRank, HcomCollOpInfo* opInfo)
      24              : {
      25            0 :     inputMem_ = usrInMem;
      26            0 :     outputMem_ = usrOutMem;
      27            0 :     stream_ = mainStream;
      28            0 :     meshStreams_ = meshStreams;
      29            0 :     meshSignal_ = meshSignal;
      30            0 :     meshSignalAux_ = meshSignalAux;
      31            0 :     userRank_ = userRank;
      32            0 :     dataType_ = dataType;
      33            0 :     dataBytes_ = count * SIZE_TABLE[dataType];
      34            0 :     count_ = count;
      35            0 :     outerCommInfoHccs_ = outerCommInfoHccs;
      36            0 :     outerCommInfoSio_ = outerCommInfoSio;
      37            0 :     opInfo_ = opInfo;
      38            0 :     totalDataBytes_ = opInfo->count * SIZE_TABLE[dataType_];
      39            0 :     return HCCL_SUCCESS;
      40              : }
      41              : 
      42              : // 主流所有从流
      43            0 : HcclResult AllGatherHccsSio::NotifySubStreamStart()
      44              : {
      45            0 :     for (u32 streamIndex = 0; streamIndex < meshStreams_.size(); streamIndex++) {
      46            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, meshSignalAux_[streamIndex], INVALID_VALUE_STAGE));
      47            0 :         CHK_RET(LocalNotify::Wait(
      48              :             meshStreams_[streamIndex], dispatcher_, meshSignalAux_[streamIndex], INVALID_VALUE_STAGE));
      49              :     }
      50            0 :     return HCCL_SUCCESS;
      51              : }
      52              : 
      53            0 : HcclResult AllGatherHccsSio::WaitSubStreamFinish()
      54              : {
      55            0 :     for (u32 streamIndex = 0; streamIndex < meshStreams_.size(); streamIndex++) {
      56            0 :         CHK_RET(
      57              :             LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, meshSignal_[streamIndex], INVALID_VALUE_STAGE));
      58            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, meshSignal_[streamIndex], INVALID_VALUE_STAGE));
      59              :     }
      60            0 :     return HCCL_SUCCESS;
      61              : }
      62              : 
      63              : HcclResult
      64            0 : AllGatherHccsSio::RunInterDie(const u32 dieRankId, const std::vector<LINK>& links, const u32 srcDMAMemSliceId)
      65              : {
      66              :     // 检查链接是否为空
      67            0 :     CHK_SMART_PTR_NULL(links[dieRankId]);
      68              : 
      69              :     // 获取远程内存指针
      70            0 :     void* remDMAMemPtr = nullptr;
      71            0 :     CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
      72              : 
      73              :     // 确定需要传输的数据部分(上半部分或下半部分)
      74            0 :     u64 dataPartOffset = dieRankId * dataBytes_;
      75            0 :     u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
      76              : 
      77            0 :     DeviceMem locDieDst;
      78            0 :     DeviceMem srcDieMem;
      79              : 
      80              :     // 定义本地目标内存和远程源内存
      81            0 :     if (srcDMAMemSliceId == 0) {
      82            0 :         locDieDst = dmaMem_[1].range(dataPartOffset, dataPartSize);
      83            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), dataPartSize);
      84              :     } else {
      85            0 :         locDieDst = dmaMem_[1].range(dataPartOffset + dataPartSize, dataBytes_ - dataPartSize);
      86            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
      87              :     }
      88              : 
      89            0 :     HCCL_INFO(
      90              :         "RunInterDie: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remDMAMemPtr[%p]", dieRankId,
      91              :         locDieDst.ptr(), locDieDst.size(), remDMAMemPtr);
      92              :     // 执行异步内存复制
      93            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
      94              : 
      95            0 :     return HCCL_SUCCESS;
      96            0 : }
      97              : 
      98              : HcclResult
      99            0 : AllGatherHccsSio::RunInterDieOpBase(const u32 dieRankId, const std::vector<LINK>& links, const u32 srcDMAMemSliceId)
     100              : {
     101              :     // 检查链接是否为空
     102            0 :     CHK_SMART_PTR_NULL(links[dieRankId]);
     103              : 
     104              :     // 获取远程CCLin内存指针
     105            0 :     void* remCCLMemPtr = nullptr;
     106            0 :     CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remCCLMemPtr));
     107              : 
     108              :     // 确定需要传输的数据部分(上半部分或下半部分)
     109            0 :     u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
     110              : 
     111            0 :     DeviceMem locDieDst;
     112            0 :     DeviceMem srcDieMem;
     113              :     DeviceMem usroutMem
     114            0 :         = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * dieRankId, dataBytes_);
     115              :     ;
     116              : 
     117              :     // 定义本地目标内存和远程源内存
     118            0 :     if (srcDMAMemSliceId == 0) {
     119            0 :         locDieDst = usroutMem.range(0, dataPartSize);
     120            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remCCLMemPtr), dataPartSize);
     121              :     } else {
     122            0 :         locDieDst = usroutMem.range(dataPartSize, dataBytes_ - dataPartSize);
     123            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remCCLMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
     124              :     }
     125            0 :     u32 linkType = static_cast<u32>(links[dieRankId]->GetLinkType());
     126            0 :     HCCL_DEBUG("[AllGatherHccsSio][RunInterDieOpbase] dstRankId[%u], linkType[%u]", dieRankId, linkType);
     127            0 :     HCCL_INFO(
     128              :         "RunInterDieOpbase: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remCCLMemPtr[%p]", dieRankId,
     129              :         locDieDst.ptr(), locDieDst.size(), remCCLMemPtr);
     130              :     // 执行异步内存复制
     131            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
     132              : 
     133            0 :     return HCCL_SUCCESS;
     134            0 : }
     135              : 
     136              : // allgather的入口函数
     137              : HcclResult
     138            0 : AllGatherHccsSio::RunAsync(const u32 rank, const u32 rankSize, [[maybe_unused]] const std::vector<LINK>& links)
     139              : {
     140              :     /*rank0:
     141              :     从userin拷贝到userout上半部分
     142              :     userout下半部分的上半部分通过sio读取rank1的userin上半部分
     143              :     userout下半部分的下半部分通过hccs读取rank1的userin下半部分
     144              :     */
     145              : 
     146              :     /*rank1:
     147              :     从userin拷贝到userout下半部分
     148              :     userout上半部分的上半部分通过sio读取rank0的userin上半部分
     149              :     userout上半部分的下半部分通过hccs读取rank0的userin下半部分
     150              :     */
     151            0 :     intraRankSize_ = rankSize;
     152            0 :     u32 dieRankId = (rank + 1) % rankSize;
     153              :     // 数据切分为2
     154              :     static u32 HCCL_ALLGATHER_SPLIT_FACTOR = 2;
     155              : 
     156              :     // dmaMem0部分userin,dmaMem1部分userout
     157            0 :     DeviceMem dmaMem0 = DeviceMem::create(inputMem_.ptr(), dataBytes_);
     158            0 :     DeviceMem dmaMem1 = DeviceMem::create(outputMem_.ptr(), dataBytes_ * intraRankSize_);
     159            0 :     DeviceMem locDieDst = dmaMem1.range(dataBytes_ * rank, dataBytes_);
     160              : 
     161            0 :     HCCL_INFO(
     162              :         "RunAsync: dmaMem0 ptr[%p], dmaMem0 size[%ld]; dmaMem1 ptr[%p], dmaMem1 size[%ld]; locDieDst ptr[%p], "
     163              :         "locDieDst size[%ld]",
     164              :         inputMem_.ptr(), dataBytes_, outputMem_.ptr(), dataBytes_ * intraRankSize_, locDieDst.ptr(), dataBytes_);
     165              : 
     166            0 :     dmaMem_.push_back(dmaMem0); // userin
     167            0 :     dmaMem_.push_back(dmaMem1); // userout
     168              : 
     169            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     170              :         // usrin 到 cclin
     171            0 :         DeviceMem locDieUsrin = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr), dataBytes_);
     172            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dmaMem0, locDieUsrin, stream_));
     173            0 :     } else {
     174              :         // step 0操作 : 所有卡本地数据从userIn-->userout
     175            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, dmaMem0, stream_));
     176              :     }
     177              : 
     178              :     // 主流启动从流
     179            0 :     CHK_RET(NotifySubStreamStart());
     180              : 
     181              :     // step 1 : die间 && device间并行收发
     182              : 
     183              :     // 数据搬运及后同步
     184            0 :     u32 srcDMAMemSliceId = 0;
     185              : 
     186            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId]));    // AckRecord
     187            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId]));    // AckWait
     188            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckRecord
     189            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckWait
     190              : 
     191            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     192              :         // 本地userout 读取die间 cclin by sio
     193            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     194            0 :         notifyIdx_++;
     195              : 
     196              :         // 本地userout 读取die间 userin by hccs
     197              :         // srcDMAMemSliceId++;
     198            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     199              : 
     200              :         // 本地usrout读取本地usrin
     201              :         DeviceMem locDieSrc = DeviceMem::create(
     202            0 :             static_cast<u8*>(opInfo_->inputAddr), count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     203            0 :         locDieDst = DeviceMem::create(
     204            0 :             static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank,
     205            0 :             count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     206            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 2]));
     207              : 
     208            0 :         locDieSrc = DeviceMem::create(
     209            0 :             static_cast<u8*>(opInfo_->inputAddr) + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     210            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     211            0 :         locDieDst = DeviceMem::create(
     212            0 :             static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank
     213            0 :                 + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     214            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     215            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 3]));
     216            0 :     } else {
     217              :         // 本地userout 读取die间 userin by sio
     218            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     219            0 :         notifyIdx_++;
     220              : 
     221              :         // 本地userout 读取die间 userin by hccs
     222              :         // srcDMAMemSliceId++;
     223            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     224              :     }
     225              : 
     226            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId]));    // DataRecord
     227            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId]));    // Datawait
     228            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // DataRecord
     229            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // Datawait
     230              : 
     231            0 :     CHK_RET(WaitSubStreamFinish());
     232              : 
     233            0 :     HCCL_INFO("[AllGatherHccsSio][RunAsync]AllGatherHccsSio finished groupRankId[%u] ", userRank_);
     234            0 :     return HCCL_SUCCESS;
     235            0 : }
     236              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_HCCS_SIO, AllGatherHccsSio);
     237              : } // namespace hccl
        

Generated by: LCOV version 2.0-1