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-17 10:19:35 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            0 : HcclResult AllGatherHccsSio::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
     138              : {
     139              :     /*rank0:
     140              :     从userin拷贝到userout上半部分
     141              :     userout下半部分的上半部分通过sio读取rank1的userin上半部分
     142              :     userout下半部分的下半部分通过hccs读取rank1的userin下半部分
     143              :     */
     144              : 
     145              :     /*rank1:
     146              :     从userin拷贝到userout下半部分
     147              :     userout上半部分的上半部分通过sio读取rank0的userin上半部分
     148              :     userout上半部分的下半部分通过hccs读取rank0的userin下半部分
     149              :     */
     150            0 :     intraRankSize_ = rankSize;
     151            0 :     u32 dieRankId = (rank + 1) % rankSize;
     152              :     // 数据切分为2
     153              :     static u32 HCCL_ALLGATHER_SPLIT_FACTOR = 2;
     154              : 
     155              :     // dmaMem0部分userin,dmaMem1部分userout
     156            0 :     DeviceMem dmaMem0 = DeviceMem::create(inputMem_.ptr(), dataBytes_);
     157            0 :     DeviceMem dmaMem1 = DeviceMem::create(outputMem_.ptr(), dataBytes_ * intraRankSize_);
     158            0 :     DeviceMem locDieDst = dmaMem1.range(dataBytes_ * rank, dataBytes_);
     159              : 
     160            0 :     HCCL_INFO(
     161              :         "RunAsync: dmaMem0 ptr[%p], dmaMem0 size[%ld]; dmaMem1 ptr[%p], dmaMem1 size[%ld]; locDieDst ptr[%p], "
     162              :         "locDieDst size[%ld]",
     163              :         inputMem_.ptr(), dataBytes_, outputMem_.ptr(), dataBytes_ * intraRankSize_, locDieDst.ptr(), dataBytes_);
     164              : 
     165            0 :     dmaMem_.push_back(dmaMem0); // userin
     166            0 :     dmaMem_.push_back(dmaMem1); // userout
     167              : 
     168            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     169              :         // usrin 到 cclin
     170            0 :         DeviceMem locDieUsrin = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr), dataBytes_);
     171            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dmaMem0, locDieUsrin, stream_));
     172            0 :     } else {
     173              :         // step 0操作 : 所有卡本地数据从userIn-->userout
     174            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, dmaMem0, stream_));
     175              :     }
     176              : 
     177              :     // 主流启动从流
     178            0 :     CHK_RET(NotifySubStreamStart());
     179              : 
     180              :     // step 1 : die间 && device间并行收发
     181              : 
     182              :     // 数据搬运及后同步
     183            0 :     u32 srcDMAMemSliceId = 0;
     184              : 
     185            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId]));    // AckRecord
     186            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId]));    // AckWait
     187            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckRecord
     188            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId + 1])); // AckWait
     189              : 
     190            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     191              :         // 本地userout 读取die间 cclin by sio
     192            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     193            0 :         notifyIdx_++;
     194              : 
     195              :         // 本地userout 读取die间 userin by hccs
     196              :         // srcDMAMemSliceId++;
     197            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     198              : 
     199              :         // 本地usrout读取本地usrin
     200              :         DeviceMem locDieSrc = DeviceMem::create(
     201            0 :             static_cast<u8*>(opInfo_->inputAddr), count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     202            0 :         locDieDst = DeviceMem::create(
     203            0 :             static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank,
     204            0 :             count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     205            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 2]));
     206              : 
     207            0 :         locDieSrc = DeviceMem::create(
     208            0 :             static_cast<u8*>(opInfo_->inputAddr) + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     209            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     210            0 :         locDieDst = DeviceMem::create(
     211            0 :             static_cast<u8*>(opInfo_->outputAddr) + totalDataBytes_ * rank
     212            0 :                 + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     213            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     214            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 3]));
     215            0 :     } else {
     216              :         // 本地userout 读取die间 userin by sio
     217            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     218            0 :         notifyIdx_++;
     219              : 
     220              :         // 本地userout 读取die间 userin by hccs
     221              :         // srcDMAMemSliceId++;
     222            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     223              :     }
     224              : 
     225            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId]));    // DataRecord
     226            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId]));    // Datawait
     227            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // DataRecord
     228            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId + 1])); // Datawait
     229              : 
     230            0 :     CHK_RET(WaitSubStreamFinish());
     231              : 
     232            0 :     HCCL_INFO("[AllGatherHccsSio][RunAsync]AllGatherHccsSio finished groupRankId[%u] ", userRank_);
     233            0 :     return HCCL_SUCCESS;
     234            0 : }
     235              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_HCCS_SIO, AllGatherHccsSio);
     236              : } // namespace hccl
        

Generated by: LCOV version 2.0-1