LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - ins_temp_all_to_all_mesh_2D.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 157 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 7 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 "log.h"
      12              : #include "alg_data_trans_wrapper.h"
      13              : #include "executor_utils.h"
      14              : #include "ins_temp_all_to_all_mesh_2D.h"
      15              : 
      16              : namespace Hccl {
      17            0 : InsTempAlltoAllMesh2D::InsTempAlltoAllMesh2D(
      18              :     const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
      19            0 :     const std::map<RankId, u32>& tempVirtRankMap)
      20            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      21            0 : {}
      22              : 
      23            0 : InsTempAlltoAllMesh2D::~InsTempAlltoAllMesh2D() {}
      24              : 
      25            0 : HcclResult InsTempAlltoAllMesh2D::CalcRes(AlgTempResReq& tempResReq)
      26              : {
      27            0 :     if (tempVTopo_.size() >= TEMPVTOPOSIZE) {
      28            0 :         rankId_ = myRank_;
      29            0 :         rankSize_ = tempRankSize_;
      30            0 :         xRankSize_ = tempVTopo_[0].size();
      31            0 :         yRankSize_ = tempVTopo_[1].size();
      32            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[0], xRankId_));
      33            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[1], yRankId_));
      34            0 :         CHK_PRT_RET((xRankSize_ == 0), HCCL_ERROR("xRankSize_ equals to zero."), HcclResult::HCCL_E_PARA);
      35            0 :         CHK_PRT_RET((yRankSize_ == 0), HCCL_ERROR("yRankSize_ equals to zero."), HcclResult::HCCL_E_PARA);
      36              :     } else {
      37            0 :         HCCL_ERROR("tempVTopo_.size() is [%zu]", tempVTopo_.size());
      38            0 :         return HcclResult::HCCL_E_INTERNAL;
      39              :     }
      40              : 
      41            0 :     HCCL_DEBUG(
      42              :         "rankId_ is [%u], rankSize_ is [%u], xRankSize_ is [%u], yRankSize_ is [%u]", rankId_, rankSize_, xRankSize_,
      43              :         yRankSize_);
      44              : 
      45            0 :     tempResReq.queNum = tempVTopo_[0].size() + tempVTopo_[1].size();
      46            0 :     tempResReq.streamNum = tempResReq.queNum;
      47            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      48              : 
      49            0 :     QId centerQ = 0;
      50            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      51            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      52              : 
      53              :     uint32_t myAlgRank;
      54            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      55            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      56            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      57            0 :             u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
      58            0 :             RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
      59            0 :             HCCL_INFO(
      60              :                 "InsTempAlltoAllMesh2D::CalcRes Rank[%d], Dim[%u], NeighborRank[%d].", myRank_, dim, neighborRank);
      61              :             // LinkNum
      62            0 :             tempResReq.links[neighborRank] = 1;
      63              :         }
      64              :     }
      65            0 :     HCCL_INFO(
      66              :         "[InsTempAlltoAllMesh2D] Calculate resource, stream number is[%u], queNotifys size is[%u]",
      67              :         tempResReq.streamNum, tempResReq.queNotifys.size());
      68            0 :     return HcclResult::HCCL_SUCCESS;
      69              : }
      70              : 
      71            0 : HcclResult InsTempAlltoAllMesh2D::RunMeshX(
      72              :     std::vector<u64>& xDataInAddr, std::vector<u64>& xDataOutAddr, u64 xSize, BufferType srcBufferType,
      73              :     BufferType dstBufferType, DmaMode dmaMode, std::vector<InsQuePtr>& xInsQues, const ResLinks& tempLinks) const
      74              : {
      75            0 :     HCCL_DEBUG("RunMeshX begin, xSize is [%u]", xSize);
      76            0 :     if (xSize == 0) {
      77            0 :         HCCL_INFO("[InsTempAlltoAllMesh2D] RunMeshX, xSize is 0");
      78            0 :         return HcclResult::HCCL_SUCCESS;
      79              :     }
      80            0 :     std::vector<DataSlice> txSrcSlices, txDstSlices, rxSrcSlices, rxDstSlices;
      81            0 :     for (u32 i = 0; i < rankSize_; i++) {
      82              :         // 计算send
      83            0 :         u64 xOffset = xRankId_;
      84            0 :         u64 yOffset = i / xRankSize_;
      85            0 :         u64 dstOffset = yOffset * xRankSize_ + xOffset;
      86              : 
      87            0 :         DataSlice txSrcSlice = DataSlice(srcBufferType, xDataInAddr[i], xSize);
      88            0 :         DataSlice txDstSlice = DataSlice(dstBufferType, xDataOutAddr[dstOffset], xSize);
      89            0 :         txSrcSlices.push_back(txSrcSlice);
      90            0 :         txDstSlices.push_back(txDstSlice);
      91              : 
      92              :         // 计算recv,recv侧的 xOffset,yOffset,dstOffset的计算方式和send侧一样
      93            0 :         DataSlice rxSrcSlice = DataSlice(srcBufferType, xDataInAddr[dstOffset], xSize);
      94            0 :         DataSlice rxDstSlice = DataSlice(dstBufferType, xDataOutAddr[i], xSize);
      95            0 :         rxSrcSlices.push_back(rxSrcSlice);
      96            0 :         rxDstSlices.push_back(rxDstSlice);
      97              :     }
      98              : 
      99              :     // 同一列的用一个队列
     100            0 :     std::vector<DataSlice> txLocalSrcSlices, txLocalDstSlices;
     101            0 :     for (u32 i = 0; i < rankSize_; i++) {
     102            0 :         if (i % xRankSize_ == xRankId_) {
     103            0 :             txLocalSrcSlices.push_back(txSrcSlices[i]);
     104            0 :             txLocalDstSlices.push_back(txDstSlices[i]);
     105              :         }
     106              :     }
     107              :     // 本地拷贝
     108            0 :     CHK_RET(LocalCopySlices(xInsQues[xRankId_], txLocalSrcSlices, txLocalDstSlices));
     109              : 
     110              :     // 拷贝到其他卡
     111            0 :     for (u32 queIdx = 0; queIdx < xRankSize_; queIdx++) {
     112            0 :         if (queIdx == xRankId_) {
     113            0 :             continue;
     114              :         }
     115              : 
     116            0 :         std::vector<DataSlice> txRmtSrcSlices, txRmtDstSlices, rxRmtSrcSlices, rxRmtDstSlices;
     117            0 :         for (u32 i = 0; i < rankSize_; i++) {
     118            0 :             if (i % xRankSize_ == queIdx) {
     119            0 :                 txRmtSrcSlices.push_back(txSrcSlices[i]);
     120            0 :                 txRmtDstSlices.push_back(txDstSlices[i]);
     121            0 :                 rxRmtSrcSlices.push_back(rxSrcSlices[i]);
     122            0 :                 rxRmtDstSlices.push_back(rxDstSlices[i]);
     123              :             }
     124              :         }
     125            0 :         TxRxSlicesList sendRecvSlicesList({txRmtSrcSlices, txRmtDstSlices}, {rxRmtSrcSlices, rxRmtDstSlices});
     126              : 
     127            0 :         RankId rankSendRecv = queIdx + yRankId_ * xRankSize_;
     128            0 :         const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankSendRecv);
     129            0 :         TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
     130              : 
     131            0 :         SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     132            0 :         CHK_PRT_RET(
     133              :             SendRecv(sendRecvInfo, xInsQues[queIdx], 0, true, dmaMode),
     134              :             HCCL_ERROR("[InsTempAlltoAllMesh2D] RunMeshX SendRecv failed"), HcclResult::HCCL_E_INTERNAL);
     135            0 :     }
     136              : 
     137            0 :     return HcclResult::HCCL_SUCCESS;
     138            0 : }
     139              : 
     140            0 : HcclResult InsTempAlltoAllMesh2D::RunMeshY(
     141              :     std::vector<u64>& yDataInAddr, std::vector<u64>& yDataOutAddr, u64 ySize, BufferType srcBufferType,
     142              :     BufferType dstBufferType, DmaMode dmaMode, std::vector<InsQuePtr>& yInsQues, const ResLinks& tempLinks) const
     143              : {
     144            0 :     HCCL_DEBUG("RunMeshX begin, ySize is [%u]", ySize);
     145            0 :     if (ySize == 0) {
     146            0 :         HCCL_INFO("[InsTempAlltoAllMesh2D] RunMeshY, ySize is 0");
     147            0 :         return HcclResult::HCCL_SUCCESS;
     148              :     }
     149              : 
     150            0 :     std::vector<DataSlice> txSrcSlices, txDstSlices, rxSrcSlices, rxDstSlices;
     151            0 :     for (u32 i = 0; i < rankSize_; i++) {
     152              :         // 计算send
     153            0 :         u64 xOffset = i % xRankSize_;
     154            0 :         u64 yOffset = yRankId_;
     155            0 :         u64 dstOffset = yOffset * xRankSize_ + xOffset;
     156              : 
     157            0 :         DataSlice txSrcSlice = DataSlice(srcBufferType, yDataInAddr[i], ySize);
     158            0 :         DataSlice txDstSlice = DataSlice(dstBufferType, yDataOutAddr[dstOffset], ySize);
     159            0 :         txSrcSlices.push_back(txSrcSlice);
     160            0 :         txDstSlices.push_back(txDstSlice);
     161              : 
     162              :         // 计算recv,recv侧的 xOffset,yOffset,dstOffset的计算方式和send侧一样
     163            0 :         DataSlice rxSrcSlice = DataSlice(srcBufferType, yDataInAddr[dstOffset], ySize);
     164            0 :         DataSlice rxDstSlice = DataSlice(dstBufferType, yDataOutAddr[i], ySize);
     165            0 :         rxSrcSlices.push_back(rxSrcSlice);
     166            0 :         rxDstSlices.push_back(rxDstSlice);
     167              :     }
     168              : 
     169              :     // 同一行的用一个队列
     170            0 :     std::vector<DataSlice> txLocalSrcSlices, txLocalDstSlices;
     171            0 :     for (u32 i = 0; i < rankSize_; i++) {
     172            0 :         if (i / xRankSize_ == yRankId_) {
     173            0 :             txLocalSrcSlices.push_back(txSrcSlices[i]);
     174            0 :             txLocalDstSlices.push_back(txDstSlices[i]);
     175              :         }
     176              :     }
     177              :     // 本地拷贝
     178            0 :     CHK_RET(LocalCopySlices(yInsQues[yRankId_], txLocalSrcSlices, txLocalDstSlices));
     179              : 
     180              :     // 拷贝到其他卡
     181            0 :     for (u32 queIdx = 0; queIdx < yRankSize_; queIdx++) {
     182            0 :         if (queIdx == yRankId_) {
     183            0 :             continue;
     184              :         }
     185              : 
     186            0 :         std::vector<DataSlice> txRmtSrcSlices, txRmtDstSlices, rxRmtSrcSlices, rxRmtDstSlices;
     187            0 :         for (u32 i = 0; i < rankSize_; i++) {
     188            0 :             if (i / xRankSize_ == queIdx) {
     189            0 :                 txRmtSrcSlices.push_back(txSrcSlices[i]);
     190            0 :                 txRmtDstSlices.push_back(txDstSlices[i]);
     191            0 :                 rxRmtSrcSlices.push_back(rxSrcSlices[i]);
     192            0 :                 rxRmtDstSlices.push_back(rxDstSlices[i]);
     193              :             }
     194              :         }
     195            0 :         TxRxSlicesList sendRecvSlicesList({txRmtSrcSlices, txRmtDstSlices}, {rxRmtSrcSlices, rxRmtDstSlices});
     196              : 
     197            0 :         RankId rankSendRecv = queIdx * xRankSize_ + xRankId_;
     198            0 :         const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankSendRecv);
     199            0 :         TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
     200              : 
     201            0 :         SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     202            0 :         CHK_PRT_RET(
     203              :             SendRecv(sendRecvInfo, yInsQues[queIdx], 0, true, dmaMode),
     204              :             HCCL_ERROR("[InsTempAlltoAllMesh2D] RunMeshY SendRecv failed"), HcclResult::HCCL_E_INTERNAL);
     205            0 :     }
     206              : 
     207            0 :     return HcclResult::HCCL_SUCCESS;
     208            0 : }
     209              : 
     210            0 : HcclResult InsTempAlltoAllMesh2D::GenExtIns(
     211              :     const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
     212              :     std::vector<InsQuePtr>& tempInsQues) const
     213              : {
     214              :     (void)tempFuncs;
     215            0 :     HCCL_INFO("[InsTempAlltoAllMesh2D] Run algorithm start: rank[%d]", myRank_);
     216              : 
     217            0 :     u64 xSize = tempAlgParams.sliceSize / 2;
     218            0 :     u64 ySize = tempAlgParams.sliceSize - xSize;
     219              : 
     220              :     // queue arrangement
     221            0 :     std::vector<InsQuePtr> xInsQues, yInsQues;
     222            0 :     for (u32 queIdx = 0; queIdx < xRankSize_ + yRankSize_; queIdx++) {
     223            0 :         if (queIdx < xRankSize_) {
     224            0 :             xInsQues.push_back(tempInsQues[queIdx]);
     225              :         } else {
     226            0 :             yInsQues.push_back(tempInsQues[queIdx]);
     227              :         }
     228              :     }
     229              : 
     230              :     // stage1
     231            0 :     if (rankSize_ > 1) {
     232            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
     233              :     }
     234              : 
     235            0 :     std::vector<u64> stage1XDataInAddr, stage1YDataInAddr, stage1XDataOutAddr, stage1YDataOutAddr;
     236            0 :     for (u32 i = 0; i < rankSize_; i++) {
     237            0 :         stage1XDataInAddr.push_back(tempAlgParams.inputSliceStride * i + tempAlgParams.buffInfo.inBuffBaseOff);
     238            0 :         stage1YDataInAddr.push_back(tempAlgParams.inputSliceStride * i + tempAlgParams.buffInfo.inBuffBaseOff + xSize);
     239            0 :         stage1XDataOutAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i);
     240            0 :         stage1YDataOutAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i + xSize);
     241              :     }
     242              : 
     243            0 :     CHK_RET(RunMeshX(
     244              :         stage1XDataInAddr, stage1XDataOutAddr, xSize, BufferType::INPUT, BufferType::SCRATCH, DmaMode::PUT, xInsQues,
     245              :         tempLinks));
     246            0 :     CHK_RET(RunMeshY(
     247              :         stage1YDataInAddr, stage1YDataOutAddr, ySize, BufferType::INPUT, BufferType::SCRATCH, DmaMode::PUT, yInsQues,
     248              :         tempLinks));
     249              : 
     250            0 :     if (rankSize_ > 1) {
     251            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     252              :     }
     253              : 
     254              :     // stage2
     255            0 :     if (rankSize_ > 1) {
     256            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
     257              :     }
     258              : 
     259            0 :     std::vector<u64> stage2XDataInAddr, stage2YDataInAddr, stage2XDataOutAddr, stage2YDataOutAddr;
     260            0 :     for (u32 i = 0; i < rankSize_; i++) {
     261            0 :         stage2XDataInAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i);
     262            0 :         stage2YDataInAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i + xSize);
     263            0 :         stage2XDataOutAddr.push_back(tempAlgParams.outputSliceStride * i + tempAlgParams.buffInfo.outBuffBaseOff);
     264            0 :         stage2YDataOutAddr.push_back(
     265            0 :             tempAlgParams.outputSliceStride * i + tempAlgParams.buffInfo.outBuffBaseOff + xSize);
     266              :     }
     267              : 
     268            0 :     BufferType outType = !tempFuncs.isBottom ? BufferType::SCRATCH : BufferType::OUTPUT;
     269            0 :     CHK_RET(RunMeshY(
     270              :         stage2XDataInAddr, stage2XDataOutAddr, xSize, BufferType::SCRATCH, outType, DmaMode::GET, yInsQues, tempLinks));
     271            0 :     CHK_RET(RunMeshX(
     272              :         stage2YDataInAddr, stage2YDataOutAddr, ySize, BufferType::SCRATCH, outType, DmaMode::GET, xInsQues, tempLinks));
     273              : 
     274            0 :     if (rankSize_ > 1) {
     275            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     276              :     }
     277              : 
     278            0 :     HCCL_INFO("[InsTempAlltoAllMesh2D] Run algorithm end: rank[%d]", myRank_);
     279              : 
     280            0 :     return HcclResult::HCCL_SUCCESS;
     281            0 : }
     282              : 
     283              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1