LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_scatter - scatter_ring.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 271 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 14 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 "scatter_ring.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15            0 : ScatterRing::ScatterRing(const HcclDispatcher dispatcher)
      16              :     : AlgTemplateBase(dispatcher),
      17            0 :       interRank_(0),
      18            0 :       interRankSize_(0)
      19            0 : {}
      20              : 
      21            0 : ScatterRing::~ScatterRing() {}
      22              : 
      23            0 : HcclResult ScatterRing::RunScatterOnRootRank()
      24              : {
      25            0 :     DeviceMem src;
      26            0 :     DeviceMem dst;
      27              :     // rank存放scatter 结果的偏移
      28            0 :     u64 scatterOffset = slices_[interRank_].offset;
      29            0 :     u64 scatterResult = slices_[interRank_].size;
      30              : 
      31            0 :     HcclResult ret = HCCL_SUCCESS;
      32              :     // 需要判断input不等于outputmem,scatter 输入只有一个input时不用拷贝
      33            0 :     if (inputMem_ != outputMem_) {
      34            0 :         src = inputMem_.range(scatterOffset, scatterResult);
      35            0 :         dst = outputMem_.range(scatterOffset, scatterResult);
      36              : 
      37            0 :         HCCL_DEBUG(
      38              :             "rootrank[%u] copy input[%p] to output[%p] scatter_offset[%llu] copysize[%llu]", interRank_, src.ptr(),
      39              :             dst.ptr(), scatterOffset, scatterResult);
      40            0 :         ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
      41            0 :         CHK_PRT_RET(
      42              :             ret != HCCL_SUCCESS,
      43              :             HCCL_ERROR(
      44              :                 "[Run][ScatterOnRootRank]root rank[%u] memcpy async from input[%p] "
      45              :                 "failed to output[%p]",
      46              :                 interRank_, inputMem_.ptr(), outputMem_.ptr()),
      47              :             ret);
      48              :     }
      49              : 
      50              :     // 数据向下一个rank发送,依次发送后继所有rank的数据
      51            0 :     for (u32 i = 1; i < interRankSize_; i++) {
      52            0 :         u32 preRank = (interRank_ - i + interRankSize_) % interRankSize_;
      53            0 :         scatterOffset = slices_[preRank].offset;
      54            0 :         scatterResult = slices_[preRank].size;
      55              : 
      56            0 :         src = inputMem_.range(scatterOffset, scatterResult);
      57              :         // 等待后一节点同步信号,进行下一轮操作
      58            0 :         CHK_RET(linkRight_->RxAck(stream_));
      59              : 
      60              :         // 向root rank的后一rank发送
      61            0 :         HCCL_DEBUG(
      62              :             " root rank[%u] sendto dstrank[%u] from srcmem offset[%llu] size[%llu]", interRank_, preRank, scatterOffset,
      63              :             scatterResult);
      64            0 :         CHK_RET(linkRight_->TxAsync(
      65              :             UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, src.ptr(), scatterResult, stream_));
      66              : 
      67            0 :         HCCL_DEBUG("root rank[%u] will rx_ack", interRank_);
      68            0 :         ret = linkRight_->TxWaitDone(stream_);
      69            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnRootRank]TxWaitDone failed"), ret);
      70              :     }
      71            0 :     return HCCL_SUCCESS;
      72            0 : }
      73            0 : HcclResult ScatterRing::RunScatterOnEndRank()
      74              : {
      75            0 :     DeviceMem src;
      76            0 :     DeviceMem dst;
      77            0 :     u64 scatterOffset = slices_[interRank_].offset;
      78            0 :     u64 scatterResult = slices_[interRank_].size;
      79              :     // 给前一节点发送同步,以便前一rank进行下一轮的操作
      80            0 :     CHK_RET(linkLeft_->TxAck(stream_));
      81              : 
      82            0 :     dst = outputMem_.range(scatterOffset, scatterResult);
      83            0 :     HCCL_DEBUG("last rank[%u] rx data ouputoffset[%llu] size[%llu]", interRank_, scatterOffset, scatterResult);
      84              :     HcclResult ret
      85            0 :         = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(), scatterResult, stream_);
      86            0 :     CHK_PRT_RET(
      87              :         ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnEndRank]last rank[%u] rx sync failed", interRank_), ret);
      88            0 :     ret = linkLeft_->RxWaitDone(stream_);
      89            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnRootRank]RxWaitDone failed"), ret);
      90            0 :     return HCCL_SUCCESS;
      91            0 : }
      92            0 : HcclResult ScatterRing::RunScatterOnMidRank()
      93              : {
      94            0 :     DeviceMem src;
      95            0 :     DeviceMem dst;
      96            0 :     DeviceMem dstLast;
      97              :     // 与root的rank号之差 + 接收的轮数 = rank_size,  每个rank 接收的次数为 root_+ranksize-rank%interRankSize_
      98            0 :     u32 round = (root_ + interRankSize_ - interRank_) % interRankSize_;
      99            0 :     HCCL_DEBUG("rank:[%u] will receive %u rounds data", interRank_, round);
     100              : 
     101            0 :     UserMemType memType
     102            0 :         = (interRank_ == ((root_ + 1) % interRankSize_)) ? UserMemType::INPUT_MEM : UserMemType::OUTPUT_MEM;
     103              : 
     104            0 :     HcclResult ret = HCCL_SUCCESS;
     105              :     // 需要接收的和发送的轮数,包含接收自己的数据
     106            0 :     for (u32 i = 1; i <= round; i++) {
     107            0 :         u32 dataRank = (interRank_ + round - i) % interRankSize_; // 收到的数据应当是哪个rank的
     108            0 :         u64 scatterOffset = slices_[dataRank].offset;
     109            0 :         u64 scatterResult = slices_[dataRank].size;
     110              : 
     111            0 :         u32 lastDataRank = (interRank_ + round - i + 1) % interRankSize_; // 加1得到发送的数据应当是哪个rank的
     112            0 :         u64 scatterLastOffset = slices_[lastDataRank].offset;
     113            0 :         u64 scatterLastResult = slices_[lastDataRank].size;
     114              : 
     115            0 :         dst = outputMem_.range(scatterOffset, scatterResult);
     116            0 :         dstLast = outputMem_.range(scatterLastOffset, scatterLastResult);
     117              : 
     118            0 :         if (i != 1) {
     119              :             // 给前一节点发送同步,以便前一rank进行下一轮的操作
     120            0 :             ret = linkLeft_->TxAck(stream_);
     121            0 :             CHK_PRT_RET(
     122              :                 ret != HCCL_SUCCESS,
     123              :                 HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx ack failed", interRank_, i), ret);
     124              :             // 从后一rank接收同步信号
     125            0 :             ret = linkRight_->RxAck(stream_);
     126            0 :             CHK_PRT_RET(
     127              :                 ret != HCCL_SUCCESS,
     128              :                 HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u]round[%u] rx ack failed", interRank_, i), ret);
     129              :             // 向后一rank发送数据
     130            0 :             HCCL_DEBUG(
     131              :                 "rank[%u] round[%u] tx async offset[%llu] size[%llu]", interRank_, i, scatterLastOffset,
     132              :                 scatterLastResult);
     133            0 :             ret = linkRight_->TxAsync(
     134            0 :                 UserMemType::OUTPUT_MEM, scatterLastOffset + baseOffset_, dstLast.ptr(), scatterLastResult, stream_);
     135            0 :             CHK_PRT_RET(
     136              :                 ret != HCCL_SUCCESS,
     137              :                 HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx async failed", interRank_, i), ret);
     138              :         } else { // 最后一轮接收数据,拷贝到自己的outputmem
     139              :             // 给前一节点发送同步,以便前一rank进行下一轮的操作
     140            0 :             ret = linkLeft_->TxAck(stream_);
     141            0 :             CHK_PRT_RET(
     142              :                 ret != HCCL_SUCCESS,
     143              :                 HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx ack failed", interRank_, i), ret);
     144              :         }
     145            0 :         HCCL_DEBUG(
     146              :             "rank[%u] round[%u] rcv with rank[%u]'s offset[%llu] size[%llu]", interRank_, i, dataRank, scatterOffset,
     147              :             scatterResult);
     148            0 :         CHK_RET(linkLeft_->RxAsync(memType, scatterOffset + baseOffset_, dst.ptr(), scatterResult, stream_));
     149            0 :         ret = linkRight_->TxWaitDone(stream_);
     150            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]TxWaitDone failed"), ret);
     151            0 :         ret = linkLeft_->RxWaitDone(stream_);
     152            0 :         CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]RxWaitDone failed"), ret);
     153              :     }
     154            0 :     return HCCL_SUCCESS;
     155            0 : }
     156              : 
     157            0 : void ScatterRing::PrepareSlicesData(const u32 unitSize, const u64 totalCount, const u32 rankSize) const
     158              : {
     159            0 :     slices_.resize(rankSize);
     160            0 :     u64 sliceSize = (totalCount / rankSize) * unitSize;
     161              : 
     162            0 :     for (u32 i = 0; i < rankSize; i++) {
     163            0 :         slices_[i].offset = i * sliceSize;
     164            0 :         slices_[i].size = sliceSize;
     165            0 :         HCCL_DEBUG("rank[%u] default slice[%u]: offset: [%llu] size[%llu]", interRank_, i, i * sliceSize, sliceSize);
     166              :     }
     167            0 : }
     168              : 
     169              : // scatter的入口函数
     170              : HcclResult
     171            0 : ScatterRing::RunAsync(const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
     172              : {
     173            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     174            0 :     CHK_PTR_NULL(stream_.ptr());
     175            0 :     if (!outputMem_ || !inputMem_) {
     176            0 :         HCCL_ERROR("[ScatterRing][RunAsync]run_async inputmem or outputmem is null");
     177            0 :         return HCCL_E_PTR;
     178              :     }
     179              : 
     180            0 :     interRank_ = rank;
     181            0 :     interRankSize_ = rankSize;
     182              : 
     183            0 :     HCCL_INFO(
     184              :         "ScatterRing run: rank[%u] totalrank[%u] count[%llu] input[%p] output[%p]", interRank_, interRankSize_, count_,
     185              :         inputMem_.ptr(), outputMem_.ptr());
     186              : 
     187              :     // ranksize为1时,只有当input!=output 时候进行拷贝
     188            0 :     if (interRankSize_ == 1) {
     189            0 :         if (inputMem_ != outputMem_) {
     190            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
     191              :         }
     192            0 :         return HCCL_SUCCESS;
     193              :     }
     194              : 
     195            0 :     u32 unitSize = DataUnitSize(dataType_);
     196            0 :     CHK_PRT_RET(
     197              :         unitSize == 0, HCCL_ERROR("[ScatterRing][RunAsync]rank[%u] unit data size is zero", rank), HCCL_E_INTERNAL);
     198              : 
     199              :     // 带入vecotr为空,计算每个rank的结果偏移和大小
     200            0 :     if (slices_.size() == 0) {
     201            0 :         PrepareSlicesData(unitSize, count_, interRankSize_);
     202              :     }
     203              : 
     204              :     // 获取link的收、发缓存, 计算chunk_size
     205            0 :     u32 ringPrevRank = (rank + rankSize - 1) % rankSize;
     206            0 :     u32 ringNextRank = (rank + 1) % rankSize;
     207              : 
     208            0 :     if (links.size() < rankSize) {
     209            0 :         HCCL_ERROR("[ScatterRing][RunAsync]rank[%u] link size[%llu] is less than rank size", rank, links.size());
     210            0 :         return HCCL_E_INTERNAL;
     211              :     }
     212              : 
     213            0 :     linkLeft_ = links[ringPrevRank];
     214            0 :     CHK_SMART_PTR_NULL(linkLeft_);
     215              : 
     216            0 :     linkRight_ = links[ringNextRank];
     217            0 :     CHK_SMART_PTR_NULL(linkRight_);
     218              : 
     219            0 :     CHK_RET(ScatterSlicesPrep(rankSize, nicRankList_.size()));
     220              : 
     221              :     // 单环场景下 nicRankList_ 长度默认为 8。
     222              :     // 多环场景下 nicRankList_ 长度为网口数量。此时若 rankSize != nicRankList_ 则为网口裁剪场景
     223            0 :     if (rankSize != HCCL_NIC_MAX_NUM || nicRankList_.size() == HCCL_NIC_MAX_NUM) {
     224              :         // 非网口裁剪场景:
     225              :         // root rank向其他rank发送数据,
     226            0 :         if (interRank_ == root_) {
     227            0 :             CHK_RET(RunScatterOnRootRank());
     228            0 :         } else if (ringNextRank == root_) { // 最后一个节点只负责接收数据,拷贝至outputmem
     229            0 :             CHK_RET(RunScatterOnEndRank());
     230              :         } else {
     231            0 :             CHK_RET(RunScatterOnMidRank());
     232              :         }
     233              :     } else {
     234              :         // 网口裁剪场景:当前仅在 910A 8P_RING (4环),且网口不满配情况下使用
     235            0 :         CHK_RET(RunScatterChunk(rank, rankSize, slices_));
     236              :     }
     237              : 
     238            0 :     if (barrierSwitchOn_) {
     239              :         // 执行barrier,保证数据收发完成
     240            0 :         CHK_RET(ExecuteBarrier(linkLeft_, linkRight_));
     241              :     }
     242            0 :     HCCL_INFO("ScatterRing finished: rank:[%u] end", interRank_);
     243              : 
     244            0 :     return HCCL_SUCCESS;
     245              : }
     246              : 
     247            0 : HcclResult ScatterRing::RunScatterChunk(const u32 rank, const u32 rankSize, const std::vector<Slice>& outputSlices)
     248              : {
     249              :     HcclResult ret;
     250            0 :     DeviceMem dst;
     251            0 :     u32 sendSliceLen = rankSliceLists_[rank].size();
     252            0 :     u32 chunkSize = HCCL_NIC_MAX_NUM / nicRankList_.size();
     253            0 :     if (sendSliceLen >= chunkSize) {
     254            0 :         CHK_RET(HeadScatterChunk(rank, rankSize, outputSlices));
     255            0 :         for (u32 midRankIdx = 1; midRankIdx < sendSliceLen - 1; midRankIdx++) {
     256            0 :             ret = MidScatterChunk(rank, rankSize, midRankIdx, outputSlices);
     257            0 :             CHK_PRT_RET(
     258              :                 ret != HCCL_SUCCESS,
     259              :                 HCCL_ERROR("[Run][ScatterChunk]rank[%u] run mid[%u] ReduceScatter chunk failed", rank, midRankIdx),
     260              :                 HCCL_E_INTERNAL);
     261              :         }
     262            0 :         CHK_RET(TailScatterChunk(rank, rankSize, sendSliceLen - 1, outputSlices));
     263            0 :     } else if (rankSliceLists_[(rank + rankSize - 1) % rankSize].size() != 0) {
     264            0 :         for (u32 sliceIdx = 0; sliceIdx < chunkSize; sliceIdx++) {
     265            0 :             std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), rank);
     266            0 :             u32 nicIdx = distance(nicRankList_.begin(), iterNic);
     267            0 :             u32 chunkStart = nicIdx * chunkSize;
     268            0 :             u32 rxSliceIndex = chunkStart + sliceIdx;
     269            0 :             u64 rxScatterOffset = slices_[rxSliceIndex].offset;
     270            0 :             u64 rxScatterResult = slices_[rxSliceIndex].size;
     271            0 :             dst = outputMem_.range(rxScatterOffset, rxScatterResult);
     272            0 :             CHK_RET(linkLeft_->TxAck(stream_));
     273              : 
     274            0 :             ret = linkLeft_->RxAsync(
     275            0 :                 UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(), rxScatterResult, stream_);
     276            0 :             CHK_PRT_RET(
     277              :                 ret != HCCL_SUCCESS,
     278              :                 HCCL_ERROR("[Run][ScatterChunk]rank[%u] Left Link rx outputSlices[%u] Failed", rank, rxSliceIndex),
     279              :                 ret);
     280              :         }
     281              :     }
     282            0 :     return HCCL_SUCCESS;
     283            0 : }
     284              : 
     285            0 : HcclResult ScatterRing::HeadScatterChunk(u32 rank, u32 rankSize, const std::vector<Slice>& outputSlices)
     286              : {
     287              :     HcclResult ret;
     288            0 :     DeviceMem dst;
     289            0 :     u32 rxSliceIndex = rankSliceLists_[rank][0];
     290            0 :     u32 txSliceIndex = rxSliceIndex;
     291            0 :     u64 scatterOffset = slices_[rxSliceIndex].offset;
     292            0 :     u64 scatterResult = slices_[rxSliceIndex].size;
     293            0 :     dst = outputMem_.range(scatterOffset, scatterResult);
     294            0 :     std::vector<u32> preRankSlices(rankSliceLists_[(rank - 1 + rankSize) % rankSize]);
     295            0 :     std::vector<u32>::iterator iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rxSliceIndex);
     296            0 :     if (iterSlice != preRankSlices.end()) {
     297            0 :         CHK_RET(linkLeft_->TxAck(stream_));
     298              : 
     299            0 :         ret = linkLeft_->RxAsync(
     300            0 :             UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(), scatterResult, stream_);
     301            0 :         CHK_PRT_RET(
     302              :             ret != HCCL_SUCCESS,
     303              :             HCCL_ERROR(
     304              :                 "[ScatterRing][HeadScatterChunk]rank[%u] Left Link rx "
     305              :                 "outputSlices[%u] Failed",
     306              :                 rank, rxSliceIndex),
     307              :             ret);
     308              :     }
     309            0 :     iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rankSliceLists_[rank][1]);
     310            0 :     if (iterSlice != preRankSlices.end()) {
     311            0 :         CHK_RET(MidScatterChunk(rank, rankSize, 0, outputSlices));
     312              :     } else {
     313            0 :         CHK_RET(linkRight_->RxAck(stream_));
     314              : 
     315            0 :         ret = linkRight_->TxAsync(
     316            0 :             UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(), scatterResult, stream_);
     317            0 :         CHK_PRT_RET(
     318              :             ret != HCCL_SUCCESS,
     319              :             HCCL_ERROR(
     320              :                 "[ScatterRing][HeadScatterChunk]rank[%u] Right Link tx "
     321              :                 "outputSlices[%u] Failed",
     322              :                 rank, txSliceIndex),
     323              :             ret);
     324              :     }
     325            0 :     return HCCL_SUCCESS;
     326            0 : }
     327              : 
     328            0 : HcclResult ScatterRing::MidScatterChunk(u32 rank, u32 rankSize, u32 sliceIdx, const std::vector<Slice>& outputSlices)
     329              : {
     330              :     (void)outputSlices;
     331              :     HcclResult ret;
     332            0 :     DeviceMem dst;
     333            0 :     u32 rxSliceIndex = rankSliceLists_[rank][sliceIdx + 1];
     334            0 :     u32 txSliceIndex = rankSliceLists_[rank][sliceIdx];
     335            0 :     u64 rxScatterOffset = slices_[rxSliceIndex].offset;
     336            0 :     u64 rxScatterResult = slices_[rxSliceIndex].size;
     337            0 :     u64 txScatterOffset = slices_[txSliceIndex].offset;
     338            0 :     u64 txScatterResult = slices_[txSliceIndex].size;
     339              : 
     340            0 :     std::vector<u32> preRankSlices(rankSliceLists_[(rank - 1 + rankSize) % rankSize]);
     341            0 :     std::vector<u32>::iterator iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rxSliceIndex);
     342            0 :     if (iterSlice != preRankSlices.end()) {
     343            0 :         CHK_RET(linkLeft_->TxAck(stream_));
     344              : 
     345            0 :         dst = outputMem_.range(txScatterOffset, txScatterResult);
     346            0 :         CHK_RET(linkRight_->RxAck(stream_));
     347              : 
     348            0 :         ret = linkRight_->TxAsync(
     349            0 :             UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(), txScatterResult, stream_);
     350            0 :         CHK_PRT_RET(
     351              :             ret != HCCL_SUCCESS,
     352              :             HCCL_ERROR(
     353              :                 "[ScatterRing][MidScatterChunk]rank[%u] Right Link tx "
     354              :                 "outputSlices[%u] Failed",
     355              :                 rank, txSliceIndex),
     356              :             ret);
     357            0 :         dst = outputMem_.range(rxScatterOffset, rxScatterResult);
     358            0 :         ret = linkLeft_->RxAsync(
     359            0 :             UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(), rxScatterResult, stream_);
     360            0 :         CHK_PRT_RET(
     361              :             ret != HCCL_SUCCESS,
     362              :             HCCL_ERROR(
     363              :                 "[ScatterRing][MidScatterChunk]rank[%u] Left Link rx "
     364              :                 "outputSlices[%u] Failed",
     365              :                 rank, rxSliceIndex),
     366              :             ret);
     367              :     } else {
     368            0 :         dst = outputMem_.range(txScatterOffset, txScatterResult);
     369            0 :         CHK_RET(linkRight_->RxAck(stream_));
     370              : 
     371            0 :         ret = linkRight_->TxAsync(
     372            0 :             UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(), txScatterResult, stream_);
     373            0 :         CHK_PRT_RET(
     374              :             ret != HCCL_SUCCESS,
     375              :             HCCL_ERROR(
     376              :                 "[ScatterRing][MidScatterChunk]rank[%u] Right Link tx "
     377              :                 "outputSlices[%u] Failed",
     378              :                 rank, txSliceIndex),
     379              :             ret);
     380              :     }
     381            0 :     return HCCL_SUCCESS;
     382            0 : }
     383              : 
     384            0 : HcclResult ScatterRing::TailScatterChunk(u32 rank, u32 rankSize, u32 sliceIdx, const std::vector<Slice>& outputSlices)
     385              : {
     386              :     (void)rankSize;
     387              :     (void)outputSlices;
     388              :     HcclResult ret;
     389            0 :     DeviceMem dst;
     390            0 :     u32 chunkSize = HCCL_NIC_MAX_NUM / nicRankList_.size();
     391            0 :     u32 txSliceIndex = rankSliceLists_[rank][sliceIdx];
     392            0 :     u64 txScatterOffset = slices_[txSliceIndex].offset;
     393            0 :     u64 txScatterResult = slices_[txSliceIndex].size;
     394            0 :     std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), rank);
     395            0 :     if (iterNic != nicRankList_.end() && rank != root_) {
     396            0 :         u32 nicIdx = distance(nicRankList_.begin(), iterNic);
     397            0 :         u32 chunkStart = nicIdx * chunkSize;
     398            0 :         u32 rxSliceIndex = chunkStart;
     399            0 :         u64 rxScatterOffset = slices_[rxSliceIndex].offset;
     400            0 :         u64 rxScatterResult = slices_[rxSliceIndex].size;
     401            0 :         CHK_RET(linkLeft_->TxAck(stream_));
     402              : 
     403            0 :         dst = outputMem_.range(txScatterOffset, txScatterResult);
     404            0 :         CHK_RET(linkRight_->RxAck(stream_));
     405              : 
     406            0 :         ret = linkRight_->TxAsync(
     407            0 :             UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(), txScatterResult, stream_);
     408            0 :         CHK_PRT_RET(
     409              :             ret != HCCL_SUCCESS,
     410              :             HCCL_ERROR(
     411              :                 "[ScatterRing][TailScatterChunk]rank[%u] Right Link tx "
     412              :                 "outputSlices[%u] Failed",
     413              :                 rank, txSliceIndex),
     414              :             ret);
     415            0 :         dst = outputMem_.range(rxScatterOffset, rxScatterResult);
     416            0 :         ret = linkLeft_->RxAsync(
     417            0 :             UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(), rxScatterResult, stream_);
     418            0 :         CHK_PRT_RET(
     419              :             ret != HCCL_SUCCESS,
     420              :             HCCL_ERROR(
     421              :                 "[ScatterRing][TailScatterChunk]rank[%u] Left Link rx "
     422              :                 "outputSlices[%u] Failed",
     423              :                 rank, rxSliceIndex),
     424              :             ret);
     425              : 
     426            0 :         for (u32 sliceIdx = 1; sliceIdx < chunkSize; sliceIdx++) {
     427            0 :             rxSliceIndex = chunkStart + sliceIdx;
     428            0 :             u64 rxScatterOffset = slices_[rxSliceIndex].offset;
     429            0 :             u64 rxScatterResult = slices_[rxSliceIndex].size;
     430            0 :             dst = outputMem_.range(rxScatterOffset, rxScatterResult);
     431            0 :             CHK_RET(linkLeft_->TxAck(stream_));
     432              : 
     433            0 :             ret = linkLeft_->RxAsync(
     434            0 :                 UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(), rxScatterResult, stream_);
     435            0 :             CHK_PRT_RET(
     436              :                 ret != HCCL_SUCCESS,
     437              :                 HCCL_ERROR(
     438              :                     "[ScatterRing][TailScatterChunk]rank[%u] Left Link rx "
     439              :                     "outputSlices[%u] Failed",
     440              :                     rank, rxSliceIndex),
     441              :                 ret);
     442              :         }
     443              :     } else {
     444            0 :         dst = outputMem_.range(txScatterOffset, txScatterResult);
     445            0 :         CHK_RET(linkRight_->RxAck(stream_));
     446              : 
     447            0 :         ret = linkRight_->TxAsync(
     448            0 :             UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(), txScatterResult, stream_);
     449            0 :         CHK_PRT_RET(
     450              :             ret != HCCL_SUCCESS,
     451              :             HCCL_ERROR(
     452              :                 "[ScatterRing][TailScatterChunk]rank[%u] Right Link tx "
     453              :                 "outputSlices[%u] Failed",
     454              :                 rank, txSliceIndex),
     455              :             ret);
     456              :     }
     457            0 :     return HCCL_SUCCESS;
     458            0 : }
     459              : 
     460            0 : HcclResult ScatterRing::ScatterSlicesPrep(u32 rankSize, u32 nicSize)
     461              : {
     462            0 :     u32 chunkSize = HCCL_NIC_MAX_NUM / nicSize;
     463            0 :     for (u32 rankIdx = 0; rankIdx < rankSize; rankIdx++) {
     464            0 :         std::vector<u32> sliceList;                          // 单个rank上的发送slice编号
     465            0 :         for (u32 nicDis = 1; nicDis <= rankSize; nicDis++) { // 递减从root遍历至当前rank的位置
     466            0 :             u32 nicIdx = (root_ + rankSize - nicDis) % rankSize;
     467            0 :             if (rankIdx == nicIdx) {
     468            0 :                 break;
     469              :             }
     470            0 :             std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), nicIdx);
     471            0 :             if (iterNic != nicRankList_.end()) { // 当前rank为网口所在位置,将网口对应的chunksize份silce放入sliceList
     472            0 :                 u32 nicListIdx = distance(nicRankList_.begin(), iterNic);
     473            0 :                 for (u32 chunkIdx = 0; chunkIdx < chunkSize; chunkIdx++) {
     474            0 :                     sliceList.push_back(chunkSize * nicListIdx + chunkIdx);
     475              :                 }
     476              :             }
     477              :         }
     478            0 :         rankSliceLists_.push_back(sliceList);
     479            0 :     }
     480            0 :     return HCCL_SUCCESS;
     481              : }
     482              : HcclResult
     483            0 : ScatterRing::GetNslbAdjInfo(const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
     484              : {
     485            0 :     if (rankSize == 1) {
     486            0 :         return HCCL_E_NOT_SUPPORT;
     487              :     }
     488            0 :     u32 ringNextRank = (rank + 1) % rankSize;
     489            0 :     LINK nslbNext = links[ringNextRank];
     490              : 
     491            0 :     NslbDpAdjInfo adjInfoStep = {};
     492            0 :     nslbAdjInfo.dstRankNum = 1;
     493            0 :     adjInfoStep.dstLocalRankId = nslbNext->GetRemoteRank();
     494            0 :     adjInfoStep.phaseId = 1;
     495            0 :     adjInfoStep.rev = 0;
     496            0 :     nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
     497              : 
     498            0 :     return HCCL_SUCCESS;
     499            0 : }
     500              : 
     501              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_SCATTER_RING, ScatterRing);
     502              : } // namespace hccl
        

Generated by: LCOV version 2.0-1