LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - alg_data_trans_wrapper_high.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 262 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 18 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 "alg_data_trans_wrapper.h"
      12              : #include "log.h"
      13              : 
      14              : namespace Hccl {
      15            0 : HcclResult Send(const DataInfo& sendInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      16              : {
      17            0 :     CHK_RET(TxReady(sendInfo.link_, queue, topicId, dmaMode));
      18            0 :     CHK_RET(TxDataWithFin(sendInfo.link_, queue, sendInfo.slices_, topicId, dmaMode));
      19            0 :     if (needNetFinAck) {
      20            0 :         CHK_RET(TxFinAck(sendInfo.link_, queue, topicId, dmaMode));
      21              :     }
      22              : 
      23            0 :     return HcclResult::HCCL_SUCCESS;
      24              : }
      25              : 
      26            0 : HcclResult Recv(const DataInfo& recvInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      27              : {
      28            0 :     CHK_RET(RxReady(recvInfo.link_, queue, topicId, dmaMode));
      29            0 :     CHK_RET(RxDataWithFin(recvInfo.link_, queue, recvInfo.slices_, topicId, dmaMode));
      30            0 :     if (needNetFinAck) {
      31            0 :         CHK_RET(RxFinAck(recvInfo.link_, queue, topicId, dmaMode));
      32              :     }
      33              : 
      34            0 :     return HcclResult::HCCL_SUCCESS;
      35              : }
      36              : 
      37            0 : HcclResult SendRecv(const SendRecvInfo& sendRecvInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      38              : {
      39            0 :     CHK_RET(TxRxReady(sendRecvInfo.sendRecvLinks_, queue, topicId, dmaMode));
      40            0 :     CHK_RET(TxRxDataWithFin(sendRecvInfo.sendRecvLinks_, queue, sendRecvInfo.sendRecvSlices_, topicId, dmaMode));
      41            0 :     if (needNetFinAck) {
      42            0 :         CHK_RET(TxRxFinAck(sendRecvInfo.sendRecvLinks_, queue, topicId, dmaMode));
      43              :     }
      44              : 
      45            0 :     return HcclResult::HCCL_SUCCESS;
      46              : }
      47              : 
      48              : HcclResult
      49            0 : SendReduce(const DataReduceInfo& sendReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      50              : {
      51            0 :     CHK_RET(TxReady(sendReduceInfo.link_, queue, topicId, dmaMode));
      52            0 :     CHK_RET(TxReduceWithFin(
      53              :         sendReduceInfo.link_, queue, {sendReduceInfo.slices_, sendReduceInfo.dataType_, sendReduceInfo.reduceOp_},
      54              :         topicId, dmaMode));
      55            0 :     if (needNetFinAck) {
      56            0 :         CHK_RET(TxFinAck(sendReduceInfo.link_, queue, topicId, dmaMode));
      57              :     }
      58              : 
      59            0 :     return HcclResult::HCCL_SUCCESS;
      60              : }
      61              : 
      62              : HcclResult
      63            0 : RecvReduce(const DataReduceInfo& recvReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      64              : {
      65            0 :     CHK_RET(RxReady(recvReduceInfo.link_, queue, topicId, dmaMode));
      66            0 :     CHK_RET(RxReduceWithFin(
      67              :         recvReduceInfo.link_, queue, {recvReduceInfo.slices_, recvReduceInfo.dataType_, recvReduceInfo.reduceOp_},
      68              :         topicId, dmaMode));
      69            0 :     if (needNetFinAck) {
      70            0 :         CHK_RET(RxFinAck(recvReduceInfo.link_, queue, topicId, dmaMode));
      71              :     }
      72              : 
      73            0 :     return HcclResult::HCCL_SUCCESS;
      74              : }
      75              : 
      76            0 : HcclResult SendRecvReduce(
      77              :     const SendRecvReduceInfo& sendRecvReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
      78              : {
      79            0 :     CHK_RET(TxRxReady(sendRecvReduceInfo.sendRecvLinks_, queue, topicId, dmaMode));
      80            0 :     CHK_RET(TxRxReduceWithFin(
      81              :         sendRecvReduceInfo.sendRecvLinks_, queue,
      82              :         {sendRecvReduceInfo.sendRecvSlices_, sendRecvReduceInfo.dataType_, sendRecvReduceInfo.reduceOp_}, topicId,
      83              :         dmaMode));
      84            0 :     if (needNetFinAck) {
      85            0 :         CHK_RET(TxRxFinAck(sendRecvReduceInfo.sendRecvLinks_, queue, topicId, dmaMode));
      86              :     }
      87              : 
      88            0 :     return HcclResult::HCCL_SUCCESS;
      89              : }
      90              : 
      91            0 : HcclResult MultiSendCounter(const MultiDataInfo& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
      92              : {
      93            0 :     if (sendInfo.links_.size() == 0) {
      94            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: link size equals 0, do nothing.");
      95            0 :         return HcclResult::HCCL_SUCCESS;
      96              :     }
      97              : 
      98            0 :     CHK_PRT_RET(
      99              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     100              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: inter-rank CounterNotify is "
     101              :                    "supported only when device supports StarsPollNetCq."),
     102              :         HcclResult::HCCL_E_INTERNAL);
     103              : 
     104            0 :     CHK_PRT_RET(
     105              :         (sendInfo.links_.size() != queues.size()) || (sendInfo.slices_.size() != queues.size()),
     106              :         HCCL_ERROR(
     107              :             "[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: invalid input with link num [%zu], "
     108              :             "slice num [%zu], queue num [%zu].",
     109              :             sendInfo.links_.size(), sendInfo.slices_.size(), queues.size()),
     110              :         HcclResult::HCCL_E_INTERNAL);
     111              : 
     112            0 :     auto linkIter = sendInfo.links_.begin();
     113            0 :     auto queIter = queues.begin();
     114              : 
     115            0 :     for (; linkIter != sendInfo.links_.end(); linkIter++, queIter++) {
     116            0 :         CHK_RET(TxReady((*linkIter), (*queIter), topicId, dmaMode));
     117              :     }
     118              : 
     119            0 :     CHK_RET(MultiTxDataWithFinCounter(sendInfo.links_, queues, sendInfo.slices_, topicId, dmaMode));
     120              : 
     121            0 :     return HcclResult::HCCL_SUCCESS;
     122              : }
     123              : 
     124            0 : HcclResult MultiRecvCounter(const MultiDataInfo& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
     125              : {
     126            0 :     if (recvInfo.links_.size() == 0) {
     127            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: link size equals 0, do nothing.");
     128            0 :         return HcclResult::HCCL_SUCCESS;
     129              :     }
     130              : 
     131            0 :     CHK_PRT_RET(
     132              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     133              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: inter-rank CounterNotify is "
     134              :                    "supported only when device supports StarsPollNetCq."),
     135              :         HcclResult::HCCL_E_INTERNAL);
     136              : 
     137            0 :     CHK_PRT_RET(
     138              :         (recvInfo.links_.size() != queues.size()) || (recvInfo.slices_.size() != queues.size()),
     139              :         HCCL_ERROR(
     140              :             "[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: invalid input with link num [%zu], "
     141              :             "slice num [%zu], queue num [%zu].",
     142              :             recvInfo.links_.size(), recvInfo.slices_.size(), queues.size()),
     143              :         HcclResult::HCCL_E_INTERNAL);
     144              : 
     145            0 :     auto linkIter = recvInfo.links_.begin();
     146            0 :     auto queIter = queues.begin();
     147              : 
     148            0 :     for (; linkIter != recvInfo.links_.end(); linkIter++, queIter++) {
     149            0 :         CHK_RET(RxReady((*linkIter), (*queIter), topicId, dmaMode));
     150              :     }
     151              : 
     152            0 :     CHK_RET(MultiRxDataWithFinCounter(recvInfo.links_, queues, recvInfo.slices_, topicId, dmaMode));
     153              : 
     154            0 :     return HcclResult::HCCL_SUCCESS;
     155              : }
     156              : 
     157            0 : HcclResult MultiSendRecvCounter(
     158              :     const MultiSendRecvInfo& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
     159              : {
     160            0 :     if (sendRecvInfo.txRxLinks_.size() == 0) {
     161            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: link size equals 0, do nothing.");
     162            0 :         return HcclResult::HCCL_SUCCESS;
     163              :     }
     164              : 
     165            0 :     CHK_PRT_RET(
     166              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     167              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: inter-rank CounterNotify is "
     168              :                    "supported only when device supports StarsPollNetCq."),
     169              :         HcclResult::HCCL_E_INTERNAL);
     170              : 
     171            0 :     CHK_PRT_RET(
     172              :         (sendRecvInfo.txRxLinks_.size() != queues.size()) || (sendRecvInfo.txRxSlices_.size() != queues.size()),
     173              :         HCCL_ERROR(
     174              :             "[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: invalid input with link num [%zu], "
     175              :             "slice num [%zu], queue num [%zu].",
     176              :             sendRecvInfo.txRxLinks_.size(), sendRecvInfo.txRxSlices_.size(), queues.size()),
     177              :         HcclResult::HCCL_E_INTERNAL);
     178              : 
     179            0 :     auto linkIter = sendRecvInfo.txRxLinks_.begin();
     180            0 :     auto queIter = queues.begin();
     181              : 
     182            0 :     for (; linkIter != sendRecvInfo.txRxLinks_.end(); linkIter++, queIter++) {
     183            0 :         CHK_RET(TxRxReady((*linkIter), (*queIter), topicId, dmaMode));
     184              :     }
     185              : 
     186            0 :     CHK_RET(MultiTxRxDataWithFinCounter(sendRecvInfo.txRxLinks_, queues, sendRecvInfo.txRxSlices_, topicId, dmaMode));
     187              : 
     188            0 :     return HcclResult::HCCL_SUCCESS;
     189              : }
     190              : 
     191            0 : HcclResult MultiSendReduceCounter(
     192              :     const MultiDataReduceInfo& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
     193              : {
     194            0 :     if (sendInfo.links_.size() == 0) {
     195            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: link size equals 0, do nothing.");
     196            0 :         return HcclResult::HCCL_SUCCESS;
     197              :     }
     198              : 
     199            0 :     CHK_PRT_RET(
     200              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     201              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: inter-rank CounterNotify is "
     202              :                    "supported only when device supports StarsPollNetCq."),
     203              :         HcclResult::HCCL_E_INTERNAL);
     204              : 
     205            0 :     CHK_PRT_RET(
     206              :         (sendInfo.links_.size() != queues.size()) || (sendInfo.slices_.size() != queues.size()),
     207              :         HCCL_ERROR(
     208              :             "[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: invalid input with link num [%zu], "
     209              :             "slice num [%zu], queue num [%zu].",
     210              :             sendInfo.links_.size(), sendInfo.slices_.size(), queues.size()),
     211              :         HcclResult::HCCL_E_INTERNAL);
     212              : 
     213            0 :     auto linkIter = sendInfo.links_.begin();
     214            0 :     auto queIter = queues.begin();
     215              : 
     216            0 :     for (; linkIter != sendInfo.links_.end(); linkIter++, queIter++) {
     217            0 :         CHK_RET(TxReady((*linkIter), (*queIter), topicId, dmaMode));
     218              :     }
     219              : 
     220            0 :     CHK_RET(MultiTxReduceWithFinCounter(sendInfo.links_, queues, sendInfo.slices_, topicId, dmaMode));
     221              : 
     222            0 :     return HcclResult::HCCL_SUCCESS;
     223              : }
     224              : 
     225            0 : HcclResult MultiRecvReduceCounter(
     226              :     const MultiDataReduceInfo& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
     227              : {
     228            0 :     if (recvInfo.links_.size() == 0) {
     229            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: link size equals 0, do nothing.");
     230            0 :         return HcclResult::HCCL_SUCCESS;
     231              :     }
     232              : 
     233            0 :     CHK_PRT_RET(
     234              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     235              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: inter-rank CounterNotify is "
     236              :                    "supported only when device supports StarsPollNetCq."),
     237              :         HcclResult::HCCL_E_INTERNAL);
     238              : 
     239            0 :     CHK_PRT_RET(
     240              :         (recvInfo.links_.size() != queues.size()) || (recvInfo.slices_.size() != queues.size()),
     241              :         HCCL_ERROR(
     242              :             "[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: invalid input with link num [%zu], "
     243              :             "slice num [%zu], queue num [%zu].",
     244              :             recvInfo.links_.size(), recvInfo.slices_.size(), queues.size()),
     245              :         HcclResult::HCCL_E_INTERNAL);
     246              : 
     247            0 :     auto linkIter = recvInfo.links_.begin();
     248            0 :     auto queIter = queues.begin();
     249              : 
     250            0 :     for (; linkIter != recvInfo.links_.end(); linkIter++, queIter++) {
     251            0 :         CHK_RET(RxReady((*linkIter), (*queIter), topicId, dmaMode));
     252              :     }
     253              : 
     254            0 :     CHK_RET(MultiRxReduceWithFinCounter(recvInfo.links_, queues, recvInfo.slices_, topicId, dmaMode));
     255              : 
     256            0 :     return HcclResult::HCCL_SUCCESS;
     257              : }
     258              : 
     259            0 : HcclResult MultiSendRecvReduceCounter(
     260              :     const MultiSendRecvReduceInfo& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
     261              : {
     262            0 :     if (sendRecvInfo.txRxLinks_.size() == 0) {
     263            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: link size equals 0, do nothing.");
     264            0 :         return HcclResult::HCCL_SUCCESS;
     265              :     }
     266              : 
     267            0 :     CHK_PRT_RET(
     268              :         !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
     269              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: inter-rank CounterNotify is "
     270              :                    "supported only when device supports StarsPollNetCq."),
     271              :         HcclResult::HCCL_E_INTERNAL);
     272              : 
     273            0 :     CHK_PRT_RET(
     274              :         (sendRecvInfo.txRxLinks_.size() != queues.size()) || (sendRecvInfo.txRxSlices_.size() != queues.size()),
     275              :         HCCL_ERROR(
     276              :             "[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: invalid input with link num [%zu], "
     277              :             "slice num [%zu], queue num [%zu].",
     278              :             sendRecvInfo.txRxLinks_.size(), sendRecvInfo.txRxSlices_.size(), queues.size()),
     279              :         HcclResult::HCCL_E_INTERNAL);
     280              : 
     281            0 :     auto linkIter = sendRecvInfo.txRxLinks_.begin();
     282            0 :     auto queIter = queues.begin();
     283              : 
     284            0 :     for (; linkIter != sendRecvInfo.txRxLinks_.end(); linkIter++, queIter++) {
     285            0 :         CHK_RET(TxRxReady((*linkIter), (*queIter), topicId, dmaMode));
     286              :     }
     287              : 
     288            0 :     CHK_RET(MultiTxRxReduceWithFinCounter(sendRecvInfo.txRxLinks_, queues, sendRecvInfo.txRxSlices_, topicId, dmaMode));
     289              : 
     290            0 :     return HcclResult::HCCL_SUCCESS;
     291              : }
     292              : 
     293            0 : HcclResult SendThruMultiLinks(
     294              :     const std::vector<DataInfo>& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
     295              :     DmaMode dmaMode)
     296              : {
     297            0 :     if (sendInfo.size() == 0) {
     298            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: sendInfo size equals 0, do nothing.");
     299            0 :         return HcclResult::HCCL_SUCCESS;
     300              :     }
     301              : 
     302            0 :     CHK_PRT_RET(
     303              :         sendInfo.size() != queues.size(),
     304              :         HCCL_ERROR(
     305              :             "[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: sendInfo size [%zu] is non-equal to queue num "
     306              :             "[%zu].",
     307              :             sendInfo.size(), queues.size()),
     308              :         HcclResult::HCCL_E_INTERNAL);
     309              : 
     310              :     // only those worker queues required to be sync: put mode in send
     311            0 :     std::vector<InsQuePtr> syncQues = {queues[0]};
     312            0 :     bool hasDiffDmaMode = false;
     313              : 
     314            0 :     CHK_RET(ProceedMultiLinks(
     315              :         sendInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::PUT, dmaMode), syncQues, hasDiffDmaMode));
     316              : 
     317            0 :     if (hasDiffDmaMode) {
     318            0 :         HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: current send links have two DmaMode.");
     319            0 :         CHK_RET(TxRxReady({sendInfo[0].link_, sendInfo[0].link_}, queues[0], topicId, dmaMode));
     320              :     } else {
     321            0 :         CHK_RET(TxReady(sendInfo[0].link_, queues[0], topicId, dmaMode));
     322              :     }
     323              : 
     324            0 :     CHK_RET(PreSyncQues(syncQues, 0));
     325              : 
     326            0 :     auto dataInfoIter = sendInfo.begin();
     327            0 :     auto queIter = queues.begin();
     328            0 :     for (; dataInfoIter != sendInfo.end(); dataInfoIter++, queIter++) {
     329            0 :         if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
     330            0 :             CHK_RET(TxData(dataInfoIter->link_, (*queIter), dataInfoIter->slices_, dmaMode));
     331              :         }
     332              :     }
     333              : 
     334            0 :     CHK_RET(PostSyncQues(syncQues, 0));
     335              : 
     336            0 :     if (hasDiffDmaMode) {
     337            0 :         CHK_RET(TxRxFin({sendInfo[0].link_, sendInfo[0].link_}, queues[0], topicId, dmaMode));
     338              :     } else {
     339            0 :         CHK_RET(TxFin(sendInfo[0].link_, queues[0], topicId, dmaMode));
     340              :     }
     341              : 
     342            0 :     if (needNetFinAck) {
     343            0 :         TxFinAck(sendInfo[0].link_, queues[0], topicId, dmaMode);
     344              :     }
     345              : 
     346            0 :     return HcclResult::HCCL_SUCCESS;
     347            0 : }
     348              : 
     349            0 : HcclResult RecvThruMultiLinks(
     350              :     const std::vector<DataInfo>& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
     351              :     DmaMode dmaMode)
     352              : {
     353            0 :     if (recvInfo.size() == 0) {
     354            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: recvInfo size equals 0, do nothing.");
     355            0 :         return HcclResult::HCCL_SUCCESS;
     356              :     }
     357              : 
     358            0 :     CHK_PRT_RET(
     359              :         recvInfo.size() != queues.size(),
     360              :         HCCL_ERROR(
     361              :             "[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: invalid input with recvInfo size [%u], "
     362              :             "queue num [%u].",
     363              :             recvInfo.size(), queues.size()),
     364              :         HcclResult::HCCL_E_INTERNAL);
     365              : 
     366              :     // only those worker queues required to be sync: put mode in send
     367            0 :     std::vector<InsQuePtr> syncQues = {queues[0]};
     368            0 :     bool hasDiffDmaMode = false;
     369              : 
     370            0 :     CHK_RET(ProceedMultiLinks(
     371              :         recvInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::GET, dmaMode), syncQues,
     372              :         hasDiffDmaMode)); // Get mode should be sync for Recv
     373              : 
     374            0 :     if (hasDiffDmaMode) {
     375            0 :         HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: current recv links have two DmaMode.");
     376            0 :         CHK_RET(TxRxReady({recvInfo[0].link_, recvInfo[0].link_}, queues[0], topicId, dmaMode));
     377              :     } else {
     378            0 :         CHK_RET(RxReady(recvInfo[0].link_, queues[0], topicId, dmaMode));
     379              :     }
     380              : 
     381            0 :     CHK_RET(PreSyncQues(syncQues, 0));
     382              : 
     383            0 :     auto dataInfoIter = recvInfo.begin();
     384            0 :     auto queIter = queues.begin();
     385            0 :     for (; dataInfoIter != recvInfo.end(); dataInfoIter++, queIter++) {
     386            0 :         if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
     387            0 :             CHK_RET(RxData(dataInfoIter->link_, (*queIter), dataInfoIter->slices_, dmaMode));
     388              :         }
     389              :     }
     390              : 
     391            0 :     CHK_RET(PostSyncQues(syncQues, 0));
     392              : 
     393            0 :     if (hasDiffDmaMode) {
     394            0 :         CHK_RET(TxRxFin({recvInfo[0].link_, recvInfo[0].link_}, queues[0], topicId, dmaMode));
     395              :     } else {
     396            0 :         CHK_RET(RxFin(recvInfo[0].link_, queues[0], topicId, dmaMode));
     397              :     }
     398              : 
     399            0 :     if (needNetFinAck) {
     400            0 :         RxFinAck(recvInfo[0].link_, queues[0], topicId, dmaMode);
     401              :     }
     402              : 
     403            0 :     return HcclResult::HCCL_SUCCESS;
     404            0 : }
     405              : 
     406            0 : HcclResult SendRecvThruMultiLinks(
     407              :     const std::vector<SendRecvInfo>& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
     408              :     DmaMode dmaMode)
     409              : {
     410            0 :     if (sendRecvInfo.size() == 0) {
     411            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: empty sendRecvInfo, do nothing.");
     412            0 :         return HcclResult::HCCL_SUCCESS;
     413              :     }
     414              : 
     415            0 :     CHK_PRT_RET(
     416              :         sendRecvInfo.size() != queues.size(),
     417              :         HCCL_ERROR(
     418              :             "[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: invalid input with recvInfo size [%zu], "
     419              :             "queue num [%zu].",
     420              :             sendRecvInfo.size(), queues.size()),
     421              :         HcclResult::HCCL_E_INTERNAL);
     422              : 
     423            0 :     auto sendRecvInfoIter = sendRecvInfo.begin();
     424            0 :     auto queIter = queues.begin();
     425            0 :     u32 netTxLinksNum = 0;
     426            0 :     u32 netRxLinksNum = 0;
     427              : 
     428            0 :     CHK_RET(TxRxReady(sendRecvInfoIter->sendRecvLinks_, (*queIter), topicId, dmaMode));
     429              : 
     430            0 :     u32 mainQueIdx = 0;
     431            0 :     CHK_RET(PreSyncQues(queues, mainQueIdx));
     432              : 
     433            0 :     for (; sendRecvInfoIter != sendRecvInfo.end(); sendRecvInfoIter++, queIter++) {
     434            0 :         if (((sendRecvInfoIter->sendRecvLinks_).txLink_).GetType() == PortDeploymentType::DEV_NET) {
     435            0 :             netTxLinksNum++;
     436              :         }
     437            0 :         if (((sendRecvInfoIter->sendRecvLinks_).rxLink_).GetType() == PortDeploymentType::DEV_NET) {
     438            0 :             netRxLinksNum++;
     439              :         }
     440            0 :         CHK_RET(TxRxData(sendRecvInfoIter->sendRecvLinks_, (*queIter), sendRecvInfoIter->sendRecvSlices_, dmaMode));
     441              :     }
     442              : 
     443            0 :     CHK_PRT_RET(
     444              :         ((netTxLinksNum > 1) || (netRxLinksNum > 1)),
     445              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: multi net links is not "
     446              :                    "supported as NET operations are async, use mid-level wrapper instead."),
     447              :         HcclResult::HCCL_E_INTERNAL);
     448              : 
     449            0 :     CHK_PRT_RET(
     450              :         (((netTxLinksNum == 1) && (sendRecvInfo[0].sendRecvLinks_.txLink_.GetType() == PortDeploymentType::P2P))
     451              :          || ((netRxLinksNum == 1) && (sendRecvInfo[0].sendRecvLinks_.rxLink_.GetType() == PortDeploymentType::P2P))),
     452              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: first link must be NET when there "
     453              :                    "exists NET links."),
     454              :         HcclResult::HCCL_E_INTERNAL);
     455              : 
     456            0 :     CHK_RET(PostSyncQues(queues, mainQueIdx));
     457              : 
     458            0 :     CHK_RET(TxRxFin(sendRecvInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
     459              : 
     460            0 :     if (needNetFinAck) {
     461            0 :         CHK_RET(TxRxFinAck(sendRecvInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
     462              :     }
     463              : 
     464            0 :     return HcclResult::HCCL_SUCCESS;
     465              : }
     466              : 
     467            0 : HcclResult SendReduceThruMultiLinks(
     468              :     const std::vector<DataReduceInfo>& sendReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
     469              :     DmaMode dmaMode)
     470              : {
     471            0 :     if (sendReduceInfo.size() == 0) {
     472            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: empty sendReduceInfo, do nothing.");
     473            0 :         return HcclResult::HCCL_SUCCESS;
     474              :     }
     475              : 
     476            0 :     CHK_PRT_RET(
     477              :         sendReduceInfo.size() != queues.size(),
     478              :         HCCL_ERROR(
     479              :             "[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: sendReduceInfo size [%zu] is "
     480              :             "non-equal to queue num [%zu].",
     481              :             sendReduceInfo.size(), queues.size()),
     482              :         HcclResult::HCCL_E_INTERNAL);
     483              : 
     484              :     // only those worker queues required to be sync: put mode in send
     485            0 :     std::vector<InsQuePtr> syncQues = {queues[0]};
     486            0 :     bool hasDiffDmaMode = false;
     487              : 
     488            0 :     CHK_RET(ProceedMultiLinks(
     489              :         sendReduceInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::PUT, dmaMode), syncQues, hasDiffDmaMode));
     490              : 
     491            0 :     if (hasDiffDmaMode) {
     492            0 :         HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: current send links have two DmaMode.");
     493            0 :         CHK_RET(TxRxReady({sendReduceInfo[0].link_, sendReduceInfo[0].link_}, queues[0], topicId, dmaMode));
     494              :     } else {
     495            0 :         CHK_RET(TxReady(sendReduceInfo[0].link_, queues[0], topicId, dmaMode));
     496              :     }
     497              : 
     498            0 :     CHK_RET(PreSyncQues(syncQues, 0));
     499              : 
     500            0 :     auto dataInfoIter = sendReduceInfo.begin();
     501            0 :     auto queIter = queues.begin();
     502            0 :     for (; dataInfoIter != sendReduceInfo.end(); dataInfoIter++, queIter++) {
     503            0 :         if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
     504            0 :             CHK_RET(TxReduce(
     505              :                 dataInfoIter->link_, (*queIter),
     506              :                 {dataInfoIter->slices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
     507              :         }
     508              :     }
     509              : 
     510            0 :     CHK_RET(PostSyncQues(syncQues, 0));
     511              : 
     512            0 :     if (hasDiffDmaMode) {
     513            0 :         CHK_RET(TxRxFin({sendReduceInfo[0].link_, sendReduceInfo[0].link_}, queues[0], topicId, dmaMode));
     514              :     } else {
     515            0 :         CHK_RET(TxFin(sendReduceInfo[0].link_, queues[0], topicId, dmaMode));
     516              :     }
     517              : 
     518            0 :     if (needNetFinAck) {
     519            0 :         TxFinAck(sendReduceInfo[0].link_, queues[0], topicId, dmaMode);
     520              :     }
     521              : 
     522            0 :     return HcclResult::HCCL_SUCCESS;
     523            0 : }
     524              : 
     525            0 : HcclResult RecvReduceThruMultiLinks(
     526              :     const std::vector<DataReduceInfo>& recvReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
     527              :     DmaMode dmaMode)
     528              : {
     529            0 :     if (recvReduceInfo.size() == 0) {
     530            0 :         HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: empty recvReduceInfo, do nothing.");
     531            0 :         return HcclResult::HCCL_SUCCESS;
     532              :     }
     533              : 
     534            0 :     CHK_PRT_RET(
     535              :         recvReduceInfo.size() != queues.size(),
     536              :         HCCL_ERROR(
     537              :             "[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: invalid input with recvReduceInfo size "
     538              :             "[%zu], "
     539              :             "queue num [%zu].",
     540              :             recvReduceInfo.size(), queues.size()),
     541              :         HcclResult::HCCL_E_INTERNAL);
     542              : 
     543              :     // only those worker queues required to be sync: put mode in send
     544            0 :     std::vector<InsQuePtr> syncQues = {queues[0]};
     545            0 :     bool hasDiffDmaMode = false;
     546              : 
     547            0 :     CHK_RET(ProceedMultiLinks(
     548              :         recvReduceInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::GET, dmaMode), syncQues,
     549              :         hasDiffDmaMode)); // Get mode should be sync for Recv
     550              : 
     551            0 :     if (hasDiffDmaMode) {
     552            0 :         HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: current recv links have two DmaMode.");
     553            0 :         CHK_RET(TxRxReady({recvReduceInfo[0].link_, recvReduceInfo[0].link_}, queues[0], topicId, dmaMode));
     554              :     } else {
     555            0 :         CHK_RET(RxReady(recvReduceInfo[0].link_, queues[0], topicId, dmaMode));
     556              :     }
     557              : 
     558            0 :     CHK_RET(PreSyncQues(syncQues, 0));
     559              : 
     560            0 :     auto dataInfoIter = recvReduceInfo.begin();
     561            0 :     auto queIter = queues.begin();
     562            0 :     for (; dataInfoIter != recvReduceInfo.end(); dataInfoIter++, queIter++) {
     563            0 :         if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
     564            0 :             CHK_RET(RxReduce(
     565              :                 dataInfoIter->link_, (*queIter),
     566              :                 {dataInfoIter->slices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
     567              :         }
     568              :     }
     569              : 
     570            0 :     CHK_RET(PostSyncQues(syncQues, 0));
     571              : 
     572            0 :     if (hasDiffDmaMode) {
     573            0 :         CHK_RET(TxRxFin({recvReduceInfo[0].link_, recvReduceInfo[0].link_}, queues[0], topicId, dmaMode));
     574              :     } else {
     575            0 :         CHK_RET(RxFin(recvReduceInfo[0].link_, queues[0], topicId, dmaMode));
     576              :     }
     577              : 
     578            0 :     if (needNetFinAck) {
     579            0 :         RxFinAck(recvReduceInfo[0].link_, queues[0], topicId, dmaMode);
     580              :     }
     581              : 
     582            0 :     return HcclResult::HCCL_SUCCESS;
     583            0 : }
     584              : 
     585            0 : HcclResult SendRecvReduceThruMultiLinks(
     586              :     const std::vector<SendRecvReduceInfo>& sendRecvReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId,
     587              :     bool needNetFinAck, DmaMode dmaMode)
     588              : {
     589            0 :     if (sendRecvReduceInfo.size() == 0) {
     590            0 :         HCCL_WARNING(
     591              :             "[InsCollAlgFactory] [AlgDataTrans] SendRecvReduceThruMultiLinks: empty sendRecvReduceInfo, do nothing.");
     592            0 :         return HcclResult::HCCL_SUCCESS;
     593              :     }
     594              : 
     595            0 :     CHK_PRT_RET(
     596              :         sendRecvReduceInfo.size() != queues.size(),
     597              :         HCCL_ERROR(
     598              :             "[InsCollAlgFactory] [AlgDataTrans] SendRecvReduceThruMultiLinks: sendRecvReduceInfo size [%u] is "
     599              :             "non-equal to queue num [%u].",
     600              :             sendRecvReduceInfo.size(), queues.size()),
     601              :         HcclResult::HCCL_E_INTERNAL);
     602              : 
     603            0 :     auto dataInfoIter = sendRecvReduceInfo.begin();
     604            0 :     auto queIter = queues.begin();
     605            0 :     u32 netTxLinksNum = 0;
     606            0 :     u32 netRxLinksNum = 0;
     607              : 
     608            0 :     CHK_RET(TxRxReady(dataInfoIter->sendRecvLinks_, (*queIter), topicId, dmaMode));
     609              : 
     610            0 :     u32 mainQueIdx = 0;
     611            0 :     CHK_RET(PreSyncQues(queues, mainQueIdx));
     612              : 
     613            0 :     for (; dataInfoIter != sendRecvReduceInfo.end(); dataInfoIter++, queIter++) {
     614            0 :         if (((dataInfoIter->sendRecvLinks_).txLink_).GetType() == PortDeploymentType::DEV_NET) {
     615            0 :             netTxLinksNum++;
     616              :         }
     617            0 :         if (((dataInfoIter->sendRecvLinks_).rxLink_).GetType() == PortDeploymentType::DEV_NET) {
     618            0 :             netRxLinksNum++;
     619              :         }
     620            0 :         CHK_RET(TxRxReduce(
     621              :             dataInfoIter->sendRecvLinks_, (*queIter),
     622              :             {dataInfoIter->sendRecvSlices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
     623              :     }
     624              : 
     625            0 :     CHK_PRT_RET(
     626              :         ((netTxLinksNum > 1) || (netRxLinksNum > 1)),
     627              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: multi net links is not "
     628              :                    "supported as NET operations are async, use mid-level wrapper instead."),
     629              :         HcclResult::HCCL_E_INTERNAL);
     630              : 
     631            0 :     CHK_PRT_RET(
     632              :         (((netTxLinksNum == 1) && (sendRecvReduceInfo[0].sendRecvLinks_.txLink_.GetType() == PortDeploymentType::P2P))
     633              :          || ((netRxLinksNum == 1)
     634              :              && (sendRecvReduceInfo[0].sendRecvLinks_.rxLink_.GetType() == PortDeploymentType::P2P))),
     635              :         HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: first link must be NET when there "
     636              :                    "exists NET links."),
     637              :         HcclResult::HCCL_E_INTERNAL);
     638              : 
     639            0 :     CHK_RET(PostSyncQues(queues, mainQueIdx));
     640              : 
     641            0 :     CHK_RET(TxRxFin(sendRecvReduceInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
     642              : 
     643            0 :     if (needNetFinAck) {
     644            0 :         CHK_RET(TxRxFinAck(sendRecvReduceInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
     645              :     }
     646              : 
     647            0 :     return HcclResult::HCCL_SUCCESS;
     648              : }
     649              : 
     650              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1