LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template - template_utils.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 15.7 % 280 44
Test Date: 2026-08-18 17:47:01 Functions: 14.3 % 21 3

            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 "template_utils.h"
      12              : #include "log.h"
      13              : #include "buffer.h"
      14              : 
      15              : namespace Hccl {
      16            0 : HcclResult GetUnitAllignSize(const AllignInfo& allignInfo, u64& unitAllignSize)
      17              : {
      18            0 :     u32 dataSizePerVolume = DataTypeSizeGet(allignInfo.dataType);
      19              : 
      20            0 :     if (allignInfo.enableAllign) {
      21            0 :         CHK_PRT_RET(
      22              :             allignInfo.allignSize < dataSizePerVolume,
      23              :             HCCL_ERROR("[CollAlgFactory] Invalid input allignSize [%u].", allignInfo.allignSize),
      24              :             HcclResult::HCCL_E_PARA);
      25            0 :         unitAllignSize = (allignInfo.allignSize % dataSizePerVolume == 0) ? allignInfo.allignSize :
      26            0 :                                                                             allignInfo.allignSize * dataSizePerVolume;
      27              :     } else {
      28            0 :         unitAllignSize = dataSizePerVolume;
      29              :     }
      30            0 :     return HcclResult::HCCL_SUCCESS;
      31              : }
      32              : 
      33            6 : HcclResult GetAlgRank(const RankId virtRank, const std::vector<RankId>& tempVTopo, u32& algRank)
      34              : {
      35            6 :     std::vector<RankId>::const_iterator topoVecIter = std::find(tempVTopo.begin(), tempVTopo.end(), virtRank);
      36            6 :     CHK_PRT_RET(
      37              :         topoVecIter == tempVTopo.end(), HCCL_ERROR("[CollAlgFactory] Invalid virtual Rank!"), HcclResult::HCCL_E_PARA);
      38            6 :     algRank = distance(tempVTopo.begin(), topoVecIter);
      39              : 
      40            6 :     return HcclResult::HCCL_SUCCESS;
      41              : }
      42              : 
      43            0 : HcclResult CalcRsAgSliceInfoConcurrMesh(
      44              :     const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const AllignInfo& allignInfo,
      45              :     const u64 dataSize, RankSliceInfo& sliceInfoVec)
      46              : {
      47              :     // multi-dimensional mesh
      48              :     u64 unitAllignSize;
      49            0 :     CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
      50              : 
      51            0 :     u32 dimSize0 = tempVTopo[0].size();
      52            0 :     u32 dimSize1 = tempVTopo[1].size();
      53              :     u64 sliceSize0
      54            0 :         = min(dataSize, RoundUp(dataSize, ((dimSize0 + dimSize1) * unitAllignSize)) * dimSize0 * unitAllignSize);
      55            0 :     u64 sliceSize1 = dataSize - sliceSize0;
      56            0 :     u64 accumOff = 0;
      57            0 :     u32 tempRankSize = dimSize0 * dimSize1;
      58            0 :     for (u32 rankIdx = 0; rankIdx < tempRankSize; rankIdx++) {
      59            0 :         SliceInfo slice0 = {accumOff, sliceSize0};
      60            0 :         sliceInfoVec[rankIdx][0] = slice0;
      61            0 :         accumOff += sliceSize0;
      62              : 
      63            0 :         SliceInfo slice1 = {accumOff, sliceSize1};
      64            0 :         sliceInfoVec[rankIdx][1] = slice1;
      65            0 :         accumOff += sliceSize1;
      66              :     }
      67              : 
      68            0 :     CHK_PRT_RET(
      69              :         (sliceInfoVec[tempRankSize - 1][1].offset + sliceInfoVec[tempRankSize - 1][1].size != dataSize * tempRankSize),
      70              :         HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
      71              : 
      72            0 :     return HcclResult::HCCL_SUCCESS;
      73              : }
      74              : 
      75            0 : HcclResult CalcRsAgSliceInfoMesh(
      76              :     const RankId myRank, const u32 tempRankSize, const AllignInfo& allignInfo, const u64 dataSize,
      77              :     RankSliceInfo& sliceInfoVec)
      78              : {
      79              :     (void)allignInfo;
      80            0 :     u64 accumOff = 0;
      81            0 :     for (u32 rankIdx = 0; rankIdx < sliceInfoVec.size(); rankIdx++) {
      82            0 :         SliceInfo slice = {accumOff, dataSize};
      83            0 :         sliceInfoVec[rankIdx][0] = slice;
      84            0 :         accumOff += dataSize;
      85              :     }
      86            0 :     CHK_PRT_RET(
      87              :         (sliceInfoVec[tempRankSize - 1][0].offset + sliceInfoVec[tempRankSize - 1][0].size != dataSize * tempRankSize),
      88              :         HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
      89              : 
      90            0 :     return HcclResult::HCCL_SUCCESS;
      91              : }
      92              : 
      93            0 : HcclResult CalcRsAgSliceInfoRing(
      94              :     const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const AllignInfo& allignInfo,
      95              :     const u64 dataSize, RankSliceInfo& sliceInfoVec)
      96              : {
      97            0 :     u32 queNum = tempVTopo.size();
      98            0 :     u32 tempRankSize = tempVTopo[0].size();
      99              :     u64 unitAllignSize;
     100            0 :     CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
     101              : 
     102            0 :     u64 queSliceSize = RoundUp(dataSize, (queNum * unitAllignSize)) * unitAllignSize;
     103              : 
     104            0 :     u64 resChunkSize = dataSize;
     105            0 :     std::vector<u64> queSlice;
     106            0 :     for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
     107              :         // split data on queues
     108            0 :         u64 currQueSliceSize = (resChunkSize > queSliceSize) ? queSliceSize : resChunkSize;
     109            0 :         queSlice.push_back(currQueSliceSize);
     110            0 :         resChunkSize -= currQueSliceSize;
     111              :     }
     112            0 :     CHK_PRT_RET(
     113              :         resChunkSize != 0, HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank),
     114              :         HcclResult::HCCL_E_INTERNAL);
     115              : 
     116            0 :     u64 accumOff = 0;
     117            0 :     for (u32 rankIdx = 0; rankIdx < tempRankSize; rankIdx++) {
     118            0 :         for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
     119            0 :             u64 currSliceSize = queSlice[queIdx];
     120            0 :             SliceInfo currSlice = {accumOff, currSliceSize};
     121            0 :             accumOff += currSliceSize;
     122            0 :             sliceInfoVec[rankIdx][queIdx] = currSlice;
     123              :         }
     124              : 
     125            0 :         CHK_PRT_RET(
     126              :             (accumOff != dataSize * (rankIdx + 1)),
     127              :             HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank),
     128              :             HcclResult::HCCL_E_INTERNAL);
     129              :     }
     130              : 
     131            0 :     CHK_PRT_RET(
     132              :         (sliceInfoVec[tempRankSize - 1][queNum - 1].offset + sliceInfoVec[tempRankSize - 1][queNum - 1].size
     133              :          != dataSize * tempRankSize),
     134              :         HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
     135              : 
     136            0 :     return HcclResult::HCCL_SUCCESS;
     137            0 : }
     138              : 
     139            0 : HcclResult CalcRsAgSliceInfoNHR(
     140              :     const RankId myRank, const u32 tempRankSize, const AllignInfo& allignInfo, const u64 dataSize,
     141              :     RankSliceInfo& sliceInfoVec)
     142              : {
     143              :     (void)allignInfo;
     144            0 :     u64 accumOff = 0;
     145            0 :     for (u32 rankIdx = 0; rankIdx < sliceInfoVec.size(); rankIdx++) {
     146            0 :         SliceInfo slice = {accumOff, dataSize};
     147            0 :         sliceInfoVec[rankIdx][0] = slice;
     148            0 :         accumOff += dataSize;
     149              :     }
     150              : 
     151            0 :     CHK_PRT_RET(
     152              :         (sliceInfoVec[tempRankSize - 1][0].offset + sliceInfoVec[tempRankSize - 1][0].size != dataSize * tempRankSize),
     153              :         HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
     154              : 
     155            0 :     return HcclResult::HCCL_SUCCESS;
     156              : }
     157              : 
     158            0 : HcclResult CalcResLinksMesh(
     159              :     const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
     160              :     const u32 linkNumBtwPeers, AlgTempResReq& tempResReq)
     161              : {
     162              :     u32 myAlgRank;
     163            0 :     CHK_RET(GetAlgRank(myRank, tempVTopo[0], myAlgRank));
     164              : 
     165            0 :     for (u32 queIdx = 0; queIdx < tempVTopo[0].size() - 1; queIdx++) {
     166              :         // find neighbors : virtualRank
     167            0 :         RankId neighborRank = tempVTopo[0][(myAlgRank + 1 + queIdx) % tempRankSize];
     168              : 
     169              :         // LinkNum
     170            0 :         tempResReq.links[neighborRank] = linkNumBtwPeers;
     171              :     }
     172              : 
     173            0 :     return HcclResult::HCCL_SUCCESS;
     174              : }
     175              : 
     176            0 : HcclResult CalcResLinksMesh2D(
     177              :     const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const u32 linkNumBtwPeers,
     178              :     AlgTempResReq& tempResReq)
     179              : {
     180              :     u32 myAlgRank;
     181            0 :     for (u32 dim = 0; dim < tempVTopo.size(); dim++) {
     182            0 :         CHK_RET(GetAlgRank(myRank, tempVTopo[dim], myAlgRank));
     183            0 :         for (u32 queIdx = 0; queIdx < tempVTopo[dim].size() - 1; queIdx++) {
     184            0 :             u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo[dim].size());
     185            0 :             CHK_PRT_RET(neighborAlgRank > (tempVTopo[dim].size() - 1),
     186              :                         HCCL_ERROR(
     187              :                             "[CalcResLinksMesh2D] neighborAlgRank[%u] is invalid,"
     188              :                             "the Max rank[%zu].",
     189              :                             neighborAlgRank, tempVTopo[dim].size() - 1);
     190              :                         , HcclResult::HCCL_E_INTERNAL);
     191            0 :             RankId neighborRank = tempVTopo[dim][neighborAlgRank];
     192            0 :             tempResReq.links[neighborRank] = linkNumBtwPeers;
     193              :         }
     194              :     }
     195              : 
     196            0 :     return HcclResult::HCCL_SUCCESS;
     197              : }
     198              : 
     199            0 : HcclResult GetDetourSendRecvLinksIn4P(
     200              :     const RankId myRank, const RankId neighborRank, const ResLinks& tempLinks,
     201              :     std::vector<std::vector<LinkDataIterator>>& sendRecvLinks)
     202              : {
     203            0 :     HCCL_DEBUG("[CollAlgFactory] [GetDetourSendRecvLinksIn4P] Rank [%d], NeighborRank [%d].", myRank, neighborRank);
     204            0 :     LinkDataIterator neighborLinkDataIter = tempLinks.at(neighborRank).begin();
     205            0 :     while (neighborLinkDataIter != tempLinks.at(neighborRank).end()) {
     206            0 :         if ((*neighborLinkDataIter).GetDirection() == LinkDirection::BOTH) {
     207            0 :             sendRecvLinks[0][0] = (neighborLinkDataIter);
     208            0 :             sendRecvLinks[0][1] = (neighborLinkDataIter);
     209            0 :         } else if ((*neighborLinkDataIter).GetDirection() == LinkDirection::RECV_ONLY) {
     210              :             // 当前算法是根据Linkdata属性来判断哪条绕路链路负责收数据,哪条负责发数据,后续方案会改进
     211            0 :             sendRecvLinks[1][1] = (neighborLinkDataIter);
     212              :         } else {
     213            0 :             sendRecvLinks[1][0] = (neighborLinkDataIter);
     214              :         }
     215            0 :         neighborLinkDataIter++;
     216              :     }
     217            0 :     return HcclResult::HCCL_SUCCESS;
     218              : }
     219              : 
     220            0 : HcclResult CalcResLinksRing(
     221              :     const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
     222              :     AlgTempResReq& tempResReq)
     223              : {
     224            0 :     std::vector<std::vector<RankId>>::const_iterator tempVTopoIter;
     225            0 :     for (tempVTopoIter = tempVTopo.begin(); tempVTopoIter != tempVTopo.end(); tempVTopoIter++) {
     226              :         // locate myRank in tempVTopo -> algRank
     227              :         u32 myAlgRank;
     228            0 :         CHK_RET(GetAlgRank(myRank, (*tempVTopoIter), myAlgRank));
     229              : 
     230              :         // find neighbors -> virtualRank
     231            0 :         RankId sendToRank = tempVTopoIter->at((myAlgRank + 1) % tempRankSize);
     232            0 :         RankId recvFromRank = tempVTopoIter->at((myAlgRank - 1 + tempRankSize) % tempRankSize); // virtualRank
     233              : 
     234              :         // LinkNum
     235            0 :         tempResReq.links[sendToRank] = 1;
     236            0 :         tempResReq.links[recvFromRank] = 1;
     237              :     }
     238            0 :     return HcclResult::HCCL_SUCCESS;
     239              : }
     240              : 
     241            0 : u32 GetLinkNum(const RankGraph* rankGraph, RankId srcRank, RankId dstRank)
     242              : {
     243            0 :     std::set<u32> levelSet = rankGraph->GetLevels(srcRank);
     244            0 :     u32 linkNum = 0;
     245            0 :     for (u32 levelIdx : levelSet) {
     246            0 :         std::vector<NetInstance::Path> paths = rankGraph->GetPaths(levelIdx, srcRank, dstRank);
     247            0 :         linkNum += paths.size();
     248            0 :     }
     249            0 :     return linkNum;
     250            0 : }
     251              : 
     252              : // NHR的算法步数 = Ceil(log2(N))
     253            0 : u32 GetNHRStepNum(u32 rankSize)
     254              : {
     255            0 :     u32 nSteps = 0;
     256            0 :     for (u32 tmp = rankSize - 1; tmp != 0; tmp >>= 1, nSteps++) {
     257              :     }
     258            0 :     HCCL_DEBUG("[NHRBase][GetStepNumInterServer] rankSize[%u] nSteps[%u]", rankSize, nSteps);
     259              : 
     260            0 :     return nSteps;
     261              : }
     262              : 
     263            0 : HcclResult CalcResLinksNHR(
     264              :     const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
     265              :     AlgTempResReq& tempResReq)
     266              : {
     267            0 :     CHK_PRT_RET(
     268              :         tempVTopo.size() != 1,
     269              :         HCCL_ERROR("[CollAlgFactory][CalcResLinksNHR] invalid tempVTopo size[%zu]", tempVTopo.size()),
     270              :         HcclResult::HCCL_E_PARA);
     271            0 :     const std::vector<RankId>& tree = tempVTopo[0];
     272            0 :     CHK_PRT_RET(
     273              :         tree.size() != tempRankSize,
     274              :         HCCL_ERROR("[CollAlgFactory][CalcResLinksNHR] tempRankSize[%u] != tree.size[%zu]", tempRankSize, tree.size()),
     275              :         HcclResult::HCCL_E_PARA);
     276            0 :     u32 nSteps = GetNHRStepNum(tempRankSize);
     277              : 
     278              :     RankId sendToRank;
     279              :     RankId recvFromRank;
     280              :     // locate myRank in tempVTopo -> algRank
     281              :     u32 myAlgRank;
     282            0 :     CHK_RET(GetAlgRank(myRank, tree, myAlgRank));
     283              : 
     284            0 :     for (u32 currentStep = 0; currentStep < nSteps; currentStep++) {
     285            0 :         u32 deltaRank = nSteps - 1 - currentStep;
     286              :         // send info
     287            0 :         sendToRank = tree[(myAlgRank + (1 << deltaRank)) % tempRankSize];
     288              :         // receive Info
     289            0 :         recvFromRank = tree[(myAlgRank + tempRankSize - (1 << deltaRank)) % tempRankSize];
     290            0 :         tempResReq.links[sendToRank] = 1;
     291            0 :         tempResReq.links[recvFromRank] = 1;
     292              :     }
     293            0 :     return HcclResult::HCCL_SUCCESS;
     294              : }
     295              : 
     296            1 : HcclResult GetLocalSendRecvInfoforAlltoall(
     297              :     const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
     298              : {
     299            1 :     u64 curSendDispls = 0;
     300            1 :     u64 curSendOffset = 0;
     301            1 :     u64 curRecvDispls = 0;
     302            1 :     u64 curRecvOffset = 0;
     303            5 :     for (u32 j = 0; j < userRankSize; j++) {
     304            4 :         u64 curSendCounts = opParam.all2AllDataDes.sendCount;
     305            4 :         u64 curSendLength = curSendCounts * DataTypeSizeGet(opParam.all2AllDataDes.sendType);
     306            4 :         localSendRecvInfo.sendCounts[j] = curSendCounts;
     307            4 :         localSendRecvInfo.sendDispls[j] = curSendDispls;
     308            4 :         localSendRecvInfo.sendLength[j] = curSendLength;
     309            4 :         localSendRecvInfo.sendOffset[j] = curSendOffset;
     310            4 :         curSendDispls += curSendCounts;
     311            4 :         curSendOffset += curSendLength;
     312              : 
     313            4 :         u64 curRecvCounts = opParam.all2AllDataDes.sendCount;
     314            4 :         u64 curRecvLength = curRecvCounts * DataTypeSizeGet(opParam.all2AllDataDes.recvType);
     315            4 :         localSendRecvInfo.recvCounts[j] = curRecvCounts;
     316            4 :         localSendRecvInfo.recvDispls[j] = curRecvDispls;
     317            4 :         localSendRecvInfo.recvLength[j] = curRecvLength;
     318            4 :         localSendRecvInfo.recvOffset[j] = curRecvOffset;
     319            4 :         curRecvDispls += curRecvCounts;
     320            4 :         curRecvOffset += curRecvLength;
     321           12 :         HCCL_DEBUG(
     322              :             "[GetLocalSendRecvInfoforAlltoall] rank[%u], sendCounts[%llu], sendDispls[%llu] "
     323              :             "recvCounts[%llu], recvDispls[%llu], sendLength[%llu], recvLength[%llu]",
     324              :             userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
     325              :             localSendRecvInfo.recvDispls[j], localSendRecvInfo.sendLength[j], localSendRecvInfo.recvLength[j]);
     326              :     }
     327            1 :     return HcclResult::HCCL_SUCCESS;
     328              : }
     329              : 
     330            0 : HcclResult GetLocalSendRecvInfoforAlltoallV(
     331              :     const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
     332              : {
     333            0 :     CHK_PTR_NULL(opParam.all2AllVDataDes.sendCounts);
     334            0 :     CHK_PTR_NULL(opParam.all2AllVDataDes.sdispls);
     335            0 :     CHK_PTR_NULL(opParam.all2AllVDataDes.recvCounts);
     336            0 :     CHK_PTR_NULL(opParam.all2AllVDataDes.rdispls);
     337            0 :     for (u32 j = 0; j < userRankSize; j++) {
     338            0 :         u64 curSendCounts = *(static_cast<const u64*>(opParam.all2AllVDataDes.sendCounts) + j);
     339            0 :         u64 curSendDispls = *(static_cast<const u64*>(opParam.all2AllVDataDes.sdispls) + j);
     340            0 :         localSendRecvInfo.sendCounts[j] = curSendCounts;
     341            0 :         localSendRecvInfo.sendDispls[j] = curSendDispls;
     342            0 :         localSendRecvInfo.sendLength[j] = curSendCounts * DataTypeSizeGet(opParam.all2AllVDataDes.sendType);
     343            0 :         localSendRecvInfo.sendOffset[j] = curSendDispls * DataTypeSizeGet(opParam.all2AllVDataDes.sendType);
     344              : 
     345            0 :         u64 curRecvCounts = *(static_cast<const u64*>(opParam.all2AllVDataDes.recvCounts) + j);
     346            0 :         u64 curRecvDispls = *(static_cast<const u64*>(opParam.all2AllVDataDes.rdispls) + j);
     347            0 :         localSendRecvInfo.recvCounts[j] = curRecvCounts;
     348            0 :         localSendRecvInfo.recvDispls[j] = curRecvDispls;
     349            0 :         localSendRecvInfo.recvLength[j] = curRecvCounts * DataTypeSizeGet(opParam.all2AllVDataDes.recvType);
     350            0 :         localSendRecvInfo.recvOffset[j] = curRecvDispls * DataTypeSizeGet(opParam.all2AllVDataDes.recvType);
     351              : 
     352            0 :         HCCL_DEBUG(
     353              :             "[GetLocalSendRecvInfoforAlltoallV] rank[%u], sendCounts[%llu], sendDispls[%llu] "
     354              :             "recvCounts[%llu], recvDispls[%llu], sendLength[%llu], recvLength[%llu]",
     355              :             userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
     356              :             localSendRecvInfo.recvDispls[j], localSendRecvInfo.sendLength[j], localSendRecvInfo.recvLength[j]);
     357              :     }
     358            0 :     return HcclResult::HCCL_SUCCESS;
     359              : }
     360              : 
     361            0 : HcclResult GetLocalSendRecvInfoforAlltoallVC(
     362              :     const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
     363              : {
     364            0 :     u64 curSendDispls = 0;
     365            0 :     u64 curSendOffset = 0;
     366            0 :     u64 curRecvDispls = 0;
     367            0 :     u64 curRecvOffset = 0;
     368            0 :     for (u32 j = 0; j < userRankSize; j++) {
     369            0 :         u64 curSendCounts
     370            0 :             = *(static_cast<const u64*>(opParam.all2AllVCDataDes.sendCountMatrix) + userRank * userRankSize + j);
     371            0 :         u64 curSendLength = curSendCounts * DataTypeSizeGet(opParam.all2AllVCDataDes.sendType);
     372            0 :         localSendRecvInfo.sendCounts[j] = curSendCounts;
     373            0 :         localSendRecvInfo.sendDispls[j] = curSendDispls;
     374            0 :         localSendRecvInfo.sendLength[j] = curSendLength;
     375            0 :         localSendRecvInfo.sendOffset[j] = curSendOffset;
     376            0 :         curSendDispls += curSendCounts;
     377            0 :         curSendOffset += curSendLength;
     378              : 
     379            0 :         u64 curRecvCounts
     380            0 :             = *(static_cast<const u64*>(opParam.all2AllVCDataDes.sendCountMatrix) + userRank + userRankSize * j);
     381            0 :         u64 curRecvLength = curRecvCounts * DataTypeSizeGet(opParam.all2AllVCDataDes.recvType);
     382            0 :         localSendRecvInfo.recvCounts[j] = curRecvCounts;
     383            0 :         localSendRecvInfo.recvDispls[j] = curRecvDispls;
     384            0 :         localSendRecvInfo.recvLength[j] = curRecvLength;
     385            0 :         localSendRecvInfo.recvOffset[j] = curRecvOffset;
     386            0 :         curRecvDispls += curRecvCounts;
     387            0 :         curRecvOffset += curRecvLength;
     388            0 :         HCCL_DEBUG(
     389              :             "[GetLocalSendRecvInfoforAlltoallVC] rank[%u], sendCounts[%llu], sendDispls[%llu] "
     390              :             "recvCounts[%llu], recvDispls[%llu]",
     391              :             userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
     392              :             localSendRecvInfo.recvDispls[j]);
     393              :     }
     394            0 :     return HcclResult::HCCL_SUCCESS;
     395              : }
     396              : 
     397            1 : HcclResult GetAlltoAllLocalSendRecvInfo(
     398              :     const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
     399              : {
     400            3 :     HCCL_DEBUG("[GetAlltoAllLocalSendRecvInfo] rank[%u], userRankSize[%u]", userRank, userRankSize);
     401            1 :     localSendRecvInfo.sendCounts.resize(userRankSize, 0);
     402            1 :     localSendRecvInfo.sendDispls.resize(userRankSize, 0);
     403            1 :     localSendRecvInfo.sendLength.resize(userRankSize, 0);
     404            1 :     localSendRecvInfo.sendOffset.resize(userRankSize, 0);
     405              : 
     406            1 :     localSendRecvInfo.recvCounts.resize(userRankSize, 0);
     407            1 :     localSendRecvInfo.recvDispls.resize(userRankSize, 0);
     408            1 :     localSendRecvInfo.recvLength.resize(userRankSize, 0);
     409            1 :     localSendRecvInfo.recvOffset.resize(userRankSize, 0);
     410            1 :     if (opParam.opType == OpType::ALLTOALLV) {
     411            0 :         CHK_RET(GetLocalSendRecvInfoforAlltoallV(opParam, userRank, userRankSize, localSendRecvInfo));
     412            1 :     } else if (opParam.opType == OpType::ALLTOALL) {
     413            1 :         CHK_RET(GetLocalSendRecvInfoforAlltoall(opParam, userRank, userRankSize, localSendRecvInfo));
     414            0 :     } else if (opParam.opType == OpType::ALLTOALLVC) {
     415            0 :         CHK_RET(GetLocalSendRecvInfoforAlltoallVC(opParam, userRank, userRankSize, localSendRecvInfo));
     416            0 :     } else if (opParam.opType != OpType::HALFALLTOALLV) {
     417            0 :         HCCL_ERROR("Only support optype alltoall , alltoallv, halfalltoallv and alltoallvc !");
     418              :     }
     419            3 :     HCCL_DEBUG("[GetAlltoAllLocalSendRecvInfo] GetAlltoAllLocalSendRecvInfo success");
     420            1 :     return HcclResult::HCCL_SUCCESS;
     421              : }
     422              : 
     423              : /*
     424              :  * 一个基本的 Allreduce 数据切分函数,用于ReduceScatter + Allgather组合成的 Allreduce 算。
     425              :  * 输入的 dataSize 是一张卡上完整的数据量
     426              :  * 函数会将 dataSize 切分成 rankSize 份,最后一份尾块可能会比其他的切分出来的子块大。
     427              :  */
     428            0 : HcclResult CalcSliceInfoAllReduce(
     429              :     const AllignInfo& allignInfo, const u32 rankSize, const u64 dataSize, RankSliceInfo& sliceInfoVec)
     430              : {
     431            0 :     sliceInfoVec.clear();
     432            0 :     sliceInfoVec.resize(rankSize);
     433              : 
     434            0 :     u32 dataSizePerVolume = DataTypeSizeGet(allignInfo.dataType);
     435              :     u64 unitAllignSize;
     436            0 :     CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
     437            0 :     u64 unitPerSlice = dataSize / unitAllignSize / rankSize;
     438            0 :     HCCL_DEBUG("unitAllignSize[%llu] unitPerSlice[%llu]", unitAllignSize, unitPerSlice);
     439              : 
     440            0 :     u64 accumOff = 0;
     441              :     SliceInfo currSlice;
     442            0 :     for (u32 rankIdx = 0; rankIdx < rankSize; rankIdx++) {
     443            0 :         if (rankIdx == rankSize - 1) {
     444            0 :             currSlice.offset = accumOff;
     445            0 :             currSlice.size = dataSize - accumOff;
     446              :         } else {
     447            0 :             currSlice.offset = accumOff;
     448            0 :             currSlice.size = unitPerSlice * unitAllignSize;
     449              :         }
     450            0 :         CHK_PRT_RET(
     451              :             currSlice.size % dataSizePerVolume != 0,
     452              :             HCCL_ERROR(
     453              :                 "[Calc][SliceInfo]rank[%u] slice size[%llu] is invalid, dataSizePerVolume[%llu]", rankIdx,
     454              :                 currSlice.size, dataSizePerVolume),
     455              :             HcclResult::HCCL_E_INTERNAL);
     456            0 :         sliceInfoVec[rankIdx].push_back(currSlice);
     457            0 :         accumOff += currSlice.size;
     458              :     }
     459              : 
     460            0 :     CHK_PRT_RET(
     461              :         (sliceInfoVec[rankSize - 1][0].offset + sliceInfoVec[rankSize - 1][0].size != dataSize),
     462              :         HCCL_ERROR(
     463              :             "[CalcSliceInfoAllReduce] SliceInfo calculation error! DataSize[%llu], "
     464              :             "lastoffset[%llu], lastsize[%llu]",
     465              :             dataSize, sliceInfoVec[rankSize - 1][0].offset, sliceInfoVec[rankSize - 1][0].size),
     466              :         HcclResult::HCCL_E_INTERNAL);
     467              : 
     468            0 :     return HcclResult::HCCL_SUCCESS;
     469              : }
     470              : 
     471            0 : HcclResult BufferTypeToAddr(const BufferType& bufferType, CollAlgOperator& op, uint64_t& addr)
     472              : {
     473            0 :     Buffer* buffer = op.GetBuffer(bufferType);
     474            0 :     CHK_PTR_NULL(buffer);
     475            0 :     addr = buffer->GetAddr();
     476            0 :     return HcclResult::HCCL_SUCCESS;
     477              : }
     478              : 
     479            0 : HcclResult CalcDataSplitRateForLinks(const std::vector<LinkData>& links, std::vector<float>& dataSplitRate)
     480              : {
     481              :     // 取到第一个对端的link数量来作为数据切分的依据
     482            0 :     std::vector<u8> linkPortGroupSizes;
     483            0 :     linkPortGroupSizes.resize(links.size());
     484            0 :     for (u32 linkIdx = 0; linkIdx < links.size(); linkIdx++) {
     485            0 :         const LinkData& linkData = links[linkIdx];
     486            0 :         linkPortGroupSizes[linkIdx] = linkData.GetPortGroupSize();
     487              :     }
     488            0 :     u32 totalPortNum = accumulate(linkPortGroupSizes.begin(), linkPortGroupSizes.end(), 0);
     489            0 :     if (totalPortNum == 0) {
     490            0 :         HCCL_ERROR("totalPortNum is zero");
     491            0 :         return HcclResult::HCCL_E_INTERNAL;
     492              :     }
     493            0 :     for (u32 linkIdx = 0; linkIdx < linkPortGroupSizes.size(); linkIdx++) {
     494            0 :         dataSplitRate[linkIdx] = static_cast<float>(linkPortGroupSizes[linkIdx]) / totalPortNum;
     495              :     }
     496            0 :     return HcclResult::HCCL_SUCCESS;
     497            0 : }
     498              : 
     499            0 : DataSlice CalcDataSliceForLinks(
     500              :     const DataSlice& recvSrcSliceAllLinks, std::vector<float> dataSplitRate, u32 j, DataType dataType_)
     501              : {
     502            0 :     BufferType type = recvSrcSliceAllLinks.GetType();
     503            0 :     u64 offset = recvSrcSliceAllLinks.GetOffset();
     504            0 :     u64 size = recvSrcSliceAllLinks.GetSize();
     505            0 :     u64 AccSize = 0;
     506            0 :     u64 typeSize = DataTypeSizeGet(dataType_);
     507            0 :     u64 dataCnt = size / typeSize;
     508            0 :     u64 linkNum = dataSplitRate.size();
     509            0 :     std::vector<DataSlice> dataSliceForLinks(linkNum);
     510            0 :     HCCL_INFO("[InsTempAllGatherNHR] Slice data for links");
     511            0 :     for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     512            0 :         if (linkIdx != linkNum - 1) {
     513            0 :             dataSliceForLinks[linkIdx].SetSize(
     514            0 :                 static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize);
     515              :         } else {
     516            0 :             dataSliceForLinks[linkIdx].SetSize(size - AccSize);
     517              :         }
     518            0 :         dataSliceForLinks[linkIdx].SetOffset(offset + AccSize);
     519            0 :         AccSize += dataSliceForLinks[linkIdx].GetSize();
     520            0 :         dataSliceForLinks[linkIdx].SetBufferType(type);
     521              :     }
     522            0 :     return dataSliceForLinks[j];
     523            0 : }
     524              : 
     525              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1