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 % 111 0
Test Date: 2026-08-04 10:52:23 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            0 : }
      17              :  
      18            0 : AllGatherHccsSio::~AllGatherHccsSio() {}
      19              : 
      20            0 : HcclResult AllGatherHccsSio::Prepare(SubCommInfo &outerCommInfoHccs, SubCommInfo &outerCommInfoSio, DeviceMem &usrInMem,
      21              :     DeviceMem &usrOutMem, u64 count, const HcclDataType dataType, const Stream &mainStream,
      22              :     std::vector<Stream> &meshStreams, std::vector<std::shared_ptr<LocalNotify>> &meshSignal,
      23              :     std::vector<std::shared_ptr<LocalNotify>> &meshSignalAux, 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(meshStreams_[streamIndex], dispatcher_, meshSignalAux_[streamIndex],
      48              :             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(LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, meshSignal_[streamIndex],
      57              :             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            0 : HcclResult AllGatherHccsSio::RunInterDie(const u32 dieRankId, const std::vector<LINK> &links, const u32 srcDMAMemSliceId)
      64              : {
      65              :     // 检查链接是否为空
      66            0 :     CHK_SMART_PTR_NULL(links[dieRankId]);
      67              : 
      68              :     // 获取远程内存指针
      69            0 :     void* remDMAMemPtr = nullptr;
      70            0 :     CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM,  &remDMAMemPtr));
      71              : 
      72              :     // 确定需要传输的数据部分(上半部分或下半部分)
      73            0 :     u64 dataPartOffset = dieRankId * dataBytes_;
      74            0 :     u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
      75              : 
      76            0 :     DeviceMem locDieDst;
      77            0 :     DeviceMem srcDieMem;
      78              : 
      79              :     // 定义本地目标内存和远程源内存
      80            0 :     if (srcDMAMemSliceId == 0) {
      81            0 :         locDieDst = dmaMem_[1].range(dataPartOffset, dataPartSize);
      82            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), dataPartSize);
      83              :     } else {
      84            0 :         locDieDst = dmaMem_[1].range(dataPartOffset + dataPartSize, dataBytes_ - dataPartSize);
      85            0 :         srcDieMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
      86              :     }
      87              : 
      88            0 :     HCCL_INFO("RunInterDie: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remDMAMemPtr[%p]",
      89              :         dieRankId,
      90              :         locDieDst.ptr(),
      91              :         locDieDst.size(),
      92              :         remDMAMemPtr);
      93              :     // 执行异步内存复制
      94            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
      95              : 
      96            0 :     return HCCL_SUCCESS;
      97            0 : }
      98              : 
      99            0 : HcclResult AllGatherHccsSio::RunInterDieOpBase(
     100              :     const u32 dieRankId, const std::vector<LINK> &links, const u32 srcDMAMemSliceId)
     101              : {
     102              :     // 检查链接是否为空
     103            0 :     CHK_SMART_PTR_NULL(links[dieRankId]);
     104              : 
     105              :     // 获取远程CCLin内存指针
     106            0 :     void *remCCLMemPtr = nullptr;
     107            0 :     CHK_RET(links[dieRankId]->GetRemoteMem(UserMemType::INPUT_MEM, &remCCLMemPtr));
     108              : 
     109              :     // 确定需要传输的数据部分(上半部分或下半部分)
     110            0 :     u64 dataPartSize = count_ / 2 * SIZE_TABLE[dataType_];
     111              : 
     112            0 :     DeviceMem locDieDst;
     113            0 :     DeviceMem srcDieMem;
     114            0 :     DeviceMem usroutMem = DeviceMem::create(static_cast<u8*>(opInfo_->outputAddr) +  totalDataBytes_ * dieRankId, dataBytes_);;
     115              : 
     116              :     // 定义本地目标内存和远程源内存
     117            0 :     if (srcDMAMemSliceId == 0) {
     118            0 :         locDieDst = usroutMem.range(0, dataPartSize);
     119            0 :         srcDieMem = DeviceMem::create(static_cast<u8 *>(remCCLMemPtr), dataPartSize);
     120              :     } else {
     121            0 :         locDieDst = usroutMem.range(dataPartSize, dataBytes_ - dataPartSize);
     122            0 :         srcDieMem = DeviceMem::create(static_cast<u8 *>(remCCLMemPtr) + dataPartSize, dataBytes_ - dataPartSize);
     123              :     }
     124            0 :     u32 linkType = static_cast<u32>(links[dieRankId]->GetLinkType());
     125            0 :     HCCL_DEBUG("[AllGatherHccsSio][RunInterDieOpbase] dstRankId[%u], linkType[%u]", dieRankId, linkType);
     126            0 :     HCCL_INFO("RunInterDieOpbase: dieRankId[%d], locDieDst ptr[%p], locDieDst size[%ld], remCCLMemPtr[%p]",
     127              :         dieRankId,
     128              :         locDieDst.ptr(),
     129              :         locDieDst.size(),
     130              :         remCCLMemPtr);
     131              :     // 执行异步内存复制
     132            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, srcDieMem, meshStreams_[srcDMAMemSliceId]));
     133              : 
     134            0 :     return HCCL_SUCCESS;
     135            0 : }
     136              : 
     137              : // allgather的入口函数
     138            0 : HcclResult AllGatherHccsSio::RunAsync(const u32 rank, const u32 rankSize, 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("RunAsync: dmaMem0 ptr[%p], dmaMem0 size[%ld]; dmaMem1 ptr[%p], dmaMem1 size[%ld]; locDieDst ptr[%p], "
     162              :               "locDieDst size[%ld]",
     163              :         inputMem_.ptr(),
     164              :         dataBytes_,
     165              :         outputMem_.ptr(),
     166              :         dataBytes_ * intraRankSize_,
     167              :         locDieDst.ptr(),
     168              :         dataBytes_);
     169              :         
     170            0 :     dmaMem_.push_back(dmaMem0);//userin
     171            0 :     dmaMem_.push_back(dmaMem1);//userout
     172              : 
     173            0 :     if(GetWorkflowMode() ==  HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     174              :         // usrin 到 cclin
     175            0 :         DeviceMem locDieUsrin = DeviceMem::create(static_cast<u8*>(opInfo_->inputAddr), dataBytes_);
     176            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dmaMem0, locDieUsrin, stream_));
     177            0 :     } else {
     178              :         // step 0操作 : 所有卡本地数据从userIn-->userout
     179            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, dmaMem0, stream_));
     180              :     }
     181              : 
     182              :     // 主流启动从流
     183            0 :     CHK_RET(NotifySubStreamStart());
     184              :  
     185              :     // step 1 : die间 && device间并行收发
     186              :  
     187              :     // 数据搬运及后同步
     188            0 :     u32 srcDMAMemSliceId = 0;
     189              : 
     190            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId]));     // AckRecord
     191            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId]));     // AckWait
     192            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxAck(meshStreams_[srcDMAMemSliceId + 1]));  // AckRecord
     193            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxAck(meshStreams_[srcDMAMemSliceId + 1]));  // AckWait
     194              : 
     195            0 :     if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     196              :         // 本地userout 读取die间 cclin by sio
     197            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     198            0 :         notifyIdx_++;
     199              : 
     200              :         // 本地userout 读取die间 userin by hccs
     201              :         // srcDMAMemSliceId++;
     202            0 :         CHK_RET(RunInterDieOpBase(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     203              : 
     204              :         // 本地usrout读取本地usrin
     205              :         DeviceMem locDieSrc =
     206            0 :             DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr), count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     207            0 :         locDieDst = DeviceMem::create(
     208            0 :             static_cast<u8 *>(opInfo_->outputAddr) + totalDataBytes_ * rank, count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     209            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 2]));
     210              : 
     211            0 :         locDieSrc = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     212            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     213            0 :         locDieDst = DeviceMem::create(
     214            0 :             static_cast<u8 *>(opInfo_->outputAddr) + totalDataBytes_ * rank + count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_],
     215            0 :             dataBytes_ - count_ / HCCL_ALLGATHER_SPLIT_FACTOR * SIZE_TABLE[dataType_]);
     216            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, locDieDst, locDieSrc, meshStreams_[srcDMAMemSliceId + 3]));
     217            0 :     } else {
     218              :         // 本地userout 读取die间 userin by sio
     219            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoHccs_.links, srcDMAMemSliceId));
     220            0 :         notifyIdx_++;
     221              : 
     222              :         // 本地userout 读取die间 userin by hccs
     223              :         // srcDMAMemSliceId++;
     224            0 :         CHK_RET(RunInterDie(dieRankId, outerCommInfoSio_.links, srcDMAMemSliceId + 1));
     225              :     }
     226              : 
     227            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId]));     // DataRecord
     228            0 :     CHK_RET(outerCommInfoHccs_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId]));     // Datawait
     229            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->TxDataSignal(meshStreams_[srcDMAMemSliceId + 1]));  // DataRecord
     230            0 :     CHK_RET(outerCommInfoSio_.links[dieRankId]->RxDataSignal(meshStreams_[srcDMAMemSliceId + 1]));  // Datawait
     231              :  
     232            0 :     CHK_RET(WaitSubStreamFinish());
     233              :  
     234            0 :     HCCL_INFO("[AllGatherHccsSio][RunAsync]AllGatherHccsSio finished groupRankId[%u] ", userRank_);
     235            0 :     return HCCL_SUCCESS;
     236            0 : }
     237              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_HCCS_SIO, AllGatherHccsSio);
     238              : }
        

Generated by: LCOV version 2.0-1