LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_reduce - all_reduce_local_reduce.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 20.3 % 251 51
Test Date: 2026-08-29 17:38:31 Functions: 35.7 % 14 5

            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 <cmath>
      12              : #include "alg_template_register.h"
      13              : #include "all_reduce_local_reduce.h"
      14              : 
      15              : namespace hccl {
      16            2 : AllReduceLocalReduce::AllReduceLocalReduce(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
      17              : 
      18            4 : AllReduceLocalReduce::~AllReduceLocalReduce() {}
      19              : 
      20            2 : HcclResult AllReduceLocalReduce::Prepare(
      21              :     u64 reduceAttrBitMap, std::vector<Stream>& meshStreams, std::vector<std::shared_ptr<LocalNotify>>& meshSignal,
      22              :     std::vector<std::shared_ptr<LocalNotify>>& meshSignalAux, u32 interRank, u32 interRankSize, u32 userRank,
      23              :     HcomCollOpInfo* opInfo)
      24              : {
      25            2 :     reduceAttr_ = reduceAttrBitMap;
      26            2 :     localRank_ = interRank;
      27            2 :     localRankSize_ = interRankSize;
      28            2 :     userRank_ = userRank;
      29            2 :     meshStreams_ = meshStreams;
      30            2 :     meshSignal_ = &meshSignal;
      31            2 :     meshSignalAux_ = &meshSignalAux;
      32            2 :     opInfo_ = opInfo;
      33            2 :     return HCCL_SUCCESS;
      34              : }
      35            0 : HcclResult AllReduceLocalReduce::MainRecordSub()
      36              : {
      37            0 :     for (u32 signalIndex = 0; signalIndex < meshSignalAux_->size(); signalIndex++) {
      38            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
      39              :     }
      40            0 :     return HCCL_SUCCESS;
      41              : }
      42              : 
      43            0 : HcclResult AllReduceLocalReduce::SubWaitMain()
      44              : {
      45            0 :     for (u32 streamIndex = 0; streamIndex < meshSignalAux_->size(); streamIndex++) {
      46            0 :         CHK_RET(LocalNotify::Wait(
      47              :             meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex], profilerInput_.stage));
      48              :     }
      49            0 :     return HCCL_SUCCESS;
      50              : }
      51              : 
      52            0 : HcclResult AllReduceLocalReduce::MainWaitSub()
      53              : {
      54            0 :     for (u32 signalIndex = 0; signalIndex < meshSignal_->size(); signalIndex++) {
      55            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
      56              :     }
      57            0 :     return HCCL_SUCCESS;
      58              : }
      59              : 
      60            0 : HcclResult AllReduceLocalReduce::SubRecordMain()
      61              : {
      62            0 :     for (u32 streamIndex = 0; streamIndex < meshSignal_->size(); streamIndex++) {
      63            0 :         CHK_RET(LocalNotify::Post(
      64              :             meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex], profilerInput_.stage));
      65              :     }
      66            0 :     return HCCL_SUCCESS;
      67              : }
      68              : 
      69              : // 将数据均分,最小单位是128
      70            1 : HcclResult AllReduceLocalReduce::PrepareSlice(
      71              :     u64 dataCount, u32 unitSize, u32 sliceNum, std::vector<Slice>& dataSlice, std::vector<Slice>& startSlice)
      72              : {
      73            1 :     Slice temp;
      74            1 :     Slice startTemp;
      75            1 :     u64 totalSize = dataCount * unitSize;
      76            1 :     dataSlice.clear();
      77            1 :     dataSlice.reserve(sliceNum);
      78            1 :     if (sliceNum == 0) {
      79            0 :         HCCL_ERROR("[Prepare][SliceData]data slice prepare, sliceNum is 0");
      80            0 :         return HCCL_E_PARA;
      81              :     }
      82            1 :     u64 sizePerSliceOri = (totalSize + sliceNum - 1) / sliceNum; /* 1是为了向上取整 */
      83            1 :     u64 sizeLimit = 0;
      84            1 :     if (outputMem_.ptr() == opInfo_->outputAddr) {
      85            1 :         sizeLimit = outputMem_.size();
      86              :     } else {
      87            0 :         sizeLimit = totalSize;
      88              :     }
      89              : 
      90            1 :     u64 sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, HCCL_MIN_SLICE_ALIGN_910B); // 512B对齐
      91            1 :     if (sizePerSlice * (localRankSize_ - 1) > sizeLimit) {
      92            1 :         sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, HCCL_MIN_SLICE_ALIGN_ONCHIP);
      93              :     }
      94            1 :     if (sizePerSlice * (localRankSize_ - 1) > sizeLimit) {
      95            1 :         sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, unitSize);
      96              :     }
      97            1 :     u64 residueSize = totalSize;
      98            1 :     u32 i = 0;
      99            7 :     while (residueSize > 0) {
     100            6 :         u64 sliceSize = sizePerSlice < residueSize ? sizePerSlice : residueSize;
     101            6 :         temp.size = sliceSize;
     102            6 :         temp.offset = totalSize - residueSize;
     103            6 :         i++;
     104            6 :         if (sliceSize <= 0) {
     105            0 :             HCCL_ERROR("[Prepare][SliceData]data_slices_prepare sliceSize[%llu]", sliceSize);
     106            0 :             return HCCL_E_PARA;
     107              :         }
     108            6 :         residueSize -= sliceSize;
     109            6 :         dataSlice.push_back(temp);
     110            6 :         if (i != sliceNum) {
     111            6 :             startTemp.size = sizePerSlice;
     112            6 :             startTemp.offset = 0;
     113            6 :             startSlice.push_back(startTemp);
     114              :         } else {
     115            0 :             startTemp.size = sizePerSlice;
     116            0 :             startTemp.offset = sizePerSlice;
     117            0 :             startSlice.push_back(startTemp);
     118              :         }
     119              :     }
     120            3 :     while (i < sliceNum) {
     121            2 :         temp.size = 0;
     122            2 :         temp.offset = totalSize;
     123            2 :         i++;
     124            2 :         dataSlice.push_back(temp);
     125            2 :         startTemp.size = 0;
     126            2 :         startTemp.offset = 0;
     127            2 :         startSlice.push_back(startTemp);
     128              :     }
     129            1 :     return HCCL_SUCCESS;
     130              : }
     131              : 
     132            0 : HcclResult AllReduceLocalReduce::PrepareAllreduceSliceData()
     133              : {
     134            0 :     return PrepareSlice(count_, DataUnitSize(dataType_), localRankSize_, slices_, startOffset);
     135              : }
     136              : 
     137              : // ringallreduce算法的函数入口
     138            0 : HcclResult AllReduceLocalReduce::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
     139              : {
     140            0 :     HcclResult ret = HCCL_SUCCESS;
     141            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     142            0 :     CHK_PTR_NULL(stream_.ptr());
     143            0 :     HCCL_INFO(
     144              :         "AllReduceLocalReduce run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
     145              :         inputMem_.ptr(), outputMem_.ptr(), count_);
     146              : 
     147            0 :     if (links.size() < rankSize) {
     148            0 :         HCCL_ERROR(
     149              :             "[AllReduceLocalReduce][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank, links.size(),
     150              :             rankSize);
     151            0 :         return HCCL_E_INTERNAL;
     152              :     }
     153              : 
     154              :     // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
     155            0 :     if (rankSize == 1) {
     156            0 :         if (opInfo_->inputAddr != opInfo_->outputAddr) {
     157            0 :             DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * DataUnitSize(dataType_));
     158            0 :             DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * DataUnitSize(dataType_));
     159            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, userMemOut, userMemIn, stream_);
     160            0 :             CHK_PRT_RET(
     161              :                 ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] memcpy async failed", rank),
     162              :                 ret);
     163            0 :         }
     164            0 :         return ret;
     165              :     }
     166              : 
     167            0 :     ret = PrepareAllreduceSliceData();
     168            0 :     CHK_PRT_RET(
     169              :         ret != HCCL_SUCCESS,
     170              :         HCCL_ERROR(
     171              :             "[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in PrepareSliceData step", rank, count_),
     172              :         ret);
     173              : 
     174            0 :     ret = RunReduceScatter(rank, rankSize, links);
     175            0 :     CHK_PRT_RET(
     176              :         ret != HCCL_SUCCESS,
     177              :         HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in ReduceScatter step", rank, count_),
     178              :         ret);
     179              : 
     180            0 :     ret = RunAllGather(rank, rankSize, links);
     181            0 :     CHK_PRT_RET(
     182              :         ret != HCCL_SUCCESS,
     183              :         HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in AllGather step", rank, count_), ret);
     184              : 
     185            0 :     HCCL_INFO("AllReduceLocalReduce finished: rank[%u] ranksize[%u]", rank, rankSize);
     186            0 :     return HCCL_SUCCESS;
     187              : }
     188              : 
     189            0 : HcclResult AllReduceLocalReduce::RunReduceScatter(u32 rank, u32 rankSize, const std::vector<LINK>& links)
     190              : {
     191            0 :     HCCL_INFO(
     192              :         "ReduceScatterMeshLocalReduce run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank,
     193              :         rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
     194              : 
     195              :     // 数据准备
     196            0 :     u32 unitSize = DataUnitSize(dataType_);
     197            0 :     DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * unitSize);
     198            0 :     DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
     199              : 
     200            0 :     DeviceMem src;
     201            0 :     DeviceMem dst;
     202              : 
     203            0 :     src = DeviceMem::create(static_cast<char*>(opInfo_->inputAddr) + slices_[rank].offset, slices_[rank].size);
     204            0 :     dst = commMemOut.range(slices_[rank].offset, slices_[rank].size);
     205            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     206              : 
     207            0 :     DeviceMem emptySrc = userMemIn.range(0, 0);
     208            0 :     DeviceMem emptyDst = commMemOut.range(0, 0);
     209              : 
     210            0 :     CHK_RET(MainRecordSub());
     211            0 :     CHK_RET(SubWaitMain());
     212              : 
     213            0 :     for (u32 round = 1; round < rankSize; round++) {
     214            0 :         u32 dstRank = (round + rank) % rankSize;
     215            0 :         Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
     216              : 
     217            0 :         CHK_RET(links[dstRank]->TxAck(subStream));
     218            0 :         CHK_RET(links[dstRank]->RxAck(subStream));
     219              :     }
     220            0 :     HCCL_DEBUG("[ReduceScatterMeshLocalReduce] D2DMemcpy start");
     221            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     222              : 
     223            0 :     CHK_RET(SubRecordMain());
     224            0 :     CHK_RET(MainWaitSub());
     225              : 
     226            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     227              : 
     228            0 :     CHK_RET(MainRecordSub());
     229            0 :     CHK_RET(SubWaitMain());
     230              : 
     231            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     232              : 
     233            0 :     for (u32 round = 1; round < rankSize; round++) {
     234            0 :         Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
     235            0 :         void* remMemPtr = nullptr;
     236            0 :         u32 dstRank = (rank + round) % rankSize;
     237            0 :         u32 dstSlice = (dstRank + round) % (rankSize - 1);
     238            0 :         if (dstRank == (rankSize - 1)) {
     239            0 :             dstSlice = (dstSlice + rankSize - 1 - 1) % (rankSize - 1);
     240              :         }
     241            0 :         if (round == (rankSize - 1)) {
     242            0 :             CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     243              : 
     244            0 :             dstSlice = (dstRank == (rankSize - 1)) ? (dstRank - 1) : dstRank;
     245            0 :             dst = DeviceMem::create(
     246            0 :                 static_cast<char*>(remMemPtr) + startOffset[dstRank].offset + dstSlice * startOffset[dstRank].size,
     247            0 :                 slices_[dstRank].size);
     248            0 :             src = userMemIn.range(slices_[dstRank].offset, slices_[dstRank].size);
     249              : 
     250            0 :             HCCL_INFO(
     251              :                 "AllReducelocalreduce reduce dst offset1 %llu offset2 %llu size %llu, rank %u, dstrank %u",
     252              :                 startOffset[dstRank].offset, dstSlice * startOffset[dstRank].size, slices_[dstRank].size, rank,
     253              :                 dstRank);
     254              : 
     255            0 :             HCCL_INFO(
     256              :                 "AllReducelocalreduce reduce src offset %llu size %llu, rank %u, dstrank %u", slices_[dstRank].offset,
     257              :                 slices_[dstRank].size, rank, dstRank);
     258              : 
     259            0 :             CHK_RET(HcclReduceAsync(
     260              :                 dispatcher_, static_cast<void*>(src.ptr()), slices_[dstRank].size / unitSize, dataType_, reductionOp_,
     261              :                 subStream, static_cast<void*>(dst.ptr()), links[dstRank]->GetRemoteRank(),
     262              :                 links[dstRank]->GetLinkType(), INLINE_REDUCE_BIT));
     263              :         } else {
     264            0 :             CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     265              : 
     266            0 :             dst = DeviceMem::create(
     267            0 :                 static_cast<char*>(remMemPtr) + startOffset[dstRank].offset + dstSlice * startOffset[dstRank].size,
     268            0 :                 slices_[dstRank].size);
     269            0 :             src = userMemIn.range(slices_[dstRank].offset, slices_[dstRank].size);
     270              : 
     271            0 :             HCCL_INFO(
     272              :                 "AllReducelocalreduce memcpy dst offset1 %llu offset2 %llu size %llu, rank %u, dstrank %u",
     273              :                 startOffset[dstRank].offset, dstSlice * startOffset[dstRank].size, slices_[dstRank].size, rank,
     274              :                 dstRank);
     275              : 
     276            0 :             HCCL_INFO(
     277              :                 "AllReducelocalreduce memcpy src offset %llu size %llu, rank %u, dstrank %u", slices_[dstRank].offset,
     278              :                 slices_[dstRank].size, rank, dstRank);
     279              : 
     280            0 :             CHK_RET(HcclD2DMemcpyAsync(
     281              :                 dispatcher_, dst, src, subStream, links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
     282              :         }
     283              : 
     284            0 :         CHK_RET(links[dstRank]->TxDataSignal(subStream));
     285            0 :         CHK_RET(links[dstRank]->RxDataSignal(subStream));
     286              :     }
     287              : 
     288            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     289              : 
     290            0 :     CHK_RET(SubRecordMain());
     291            0 :     CHK_RET(MainWaitSub());
     292              : 
     293            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     294              : 
     295            0 :     HcclResult ret = HCCL_SUCCESS;
     296            0 :     ret = RunLocalReduce(rank, rankSize);
     297              : 
     298            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceLocalReduce]rank[%u] ReduceScatter failed", rank), ret);
     299            0 :     return HCCL_SUCCESS;
     300            0 : }
     301              : 
     302            0 : HcclResult AllReduceLocalReduce::RunLocalReduce(u32 rank, u32 rankSize)
     303              : {
     304            0 :     DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
     305            0 :     u32 power = static_cast<u32>(log2(rankSize - 1));
     306            0 :     u32 rankPower = static_cast<u32>(pow(2, power));
     307            0 :     u32 unitSize = SIZE_TABLE[dataType_];
     308            0 :     u64 align = startOffset[rank].size;
     309            0 :     u64 totalSize = slices_[rank].size;
     310            0 :     DeviceMem src;
     311            0 :     DeviceMem dst;
     312            0 :     for (u32 i = 0u; i < rankSize - rankPower - 1; ++i) {
     313            0 :         u64 size = totalSize;
     314            0 :         if (rank < rankPower) {
     315            0 :             src = commMemOut.range(startOffset[rank].offset + (rankPower + i) * align, size);
     316            0 :             dst = commMemOut.range(startOffset[rank].offset + i * align, size);
     317            0 :             HCCL_INFO(
     318              :                 "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
     319              :                 startOffset[rank].offset + (rankPower + i) * align, startOffset[rank].offset + i * align, size);
     320              :         } else {
     321            0 :             dst = commMemOut.range(startOffset[rank].offset + (rankPower + i) * align, size);
     322            0 :             src = commMemOut.range(startOffset[rank].offset + i * align, size);
     323            0 :             HCCL_INFO(
     324              :                 "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
     325              :                 startOffset[rank].offset + i * align, startOffset[rank].offset + (rankPower + i) * align, size);
     326              :         }
     327            0 :         CHK_RET(HcclReduceAsync(
     328              :             dispatcher_, static_cast<void*>(src.ptr()), size / unitSize, dataType_, reductionOp_, stream_,
     329              :             static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
     330              :     }
     331            0 :     u32 center = rank < rankPower ? rank : (rank - rankPower + 1);
     332            0 :     center = std::min(center, rankPower - 1);
     333            0 :     u64 offset = rank < rankPower ? 0 : ((rankSize - rankPower - 1) * align);
     334            0 :     offset += startOffset[rank].offset;
     335            0 :     for (u32 round = 0; round < power; round++) {
     336            0 :         u32 slices_num = static_cast<u32>(rankPower / pow(2, round + 1));
     337            0 :         u64 size = totalSize;
     338            0 :         if (center < slices_num) {
     339            0 :             for (auto i = 0u; i < slices_num; ++i) {
     340            0 :                 src = commMemOut.range(offset + (slices_num + i) * align, size);
     341            0 :                 dst = commMemOut.range(offset + i * align, size);
     342            0 :                 HCCL_INFO(
     343              :                     "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
     344              :                     offset + (slices_num + i) * align, offset + i * align, size);
     345            0 :                 CHK_RET(HcclReduceAsync(
     346              :                     dispatcher_, static_cast<void*>(src.ptr()), src.size() / unitSize, dataType_, reductionOp_, stream_,
     347              :                     static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
     348              :             }
     349              :         } else {
     350            0 :             for (auto i = 0u; i < slices_num; ++i) {
     351            0 :                 dst = commMemOut.range(offset + (slices_num + i) * align, size);
     352            0 :                 src = commMemOut.range(offset + i * align, size);
     353            0 :                 HCCL_INFO(
     354              :                     "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank, offset + i * align,
     355              :                     offset + (slices_num + i) * align, size);
     356            0 :                 CHK_RET(HcclReduceAsync(
     357              :                     dispatcher_, static_cast<void*>(src.ptr()), src.size() / unitSize, dataType_, reductionOp_, stream_,
     358              :                     static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
     359              :             }
     360            0 :             offset = offset + slices_num * align;
     361            0 :             center -= slices_num;
     362              :         }
     363              :     }
     364            0 :     return HCCL_SUCCESS;
     365            0 : }
     366              : 
     367            0 : HcclResult AllReduceLocalReduce::RunAllGather(u32 rank, u32 rankSize, const std::vector<LINK>& links)
     368              : {
     369            0 :     HCCL_INFO(
     370              :         "AllGatherMesh run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
     371              :         inputMem_.ptr(), outputMem_.ptr(), count_);
     372            0 :     u32 unitSize = DataUnitSize(dataType_);
     373            0 :     DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * unitSize);
     374            0 :     DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
     375              : 
     376            0 :     DeviceMem src;
     377            0 :     DeviceMem dst;
     378              : 
     379            0 :     DeviceMem emptySrc = commMemOut.range(0, 0);
     380            0 :     DeviceMem emptyDst = userMemOut.range(0, 0);
     381              : 
     382            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     383              : 
     384            0 :     CHK_RET(MainRecordSub());
     385            0 :     CHK_RET(SubWaitMain());
     386              : 
     387            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     388              : 
     389            0 :     for (u32 round = 1; round < rankSize; round++) {
     390            0 :         u32 dstRank = BackwardRank(rank, rankSize, round);
     391            0 :         Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
     392            0 :         CHK_RET(links[dstRank]->TxAck(subStream));
     393            0 :         CHK_RET(links[dstRank]->RxAck(subStream));
     394              :     }
     395              : 
     396            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     397              : 
     398            0 :     CHK_RET(SubRecordMain());
     399            0 :     CHK_RET(MainWaitSub());
     400              : 
     401            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     402              : 
     403            0 :     CHK_RET(MainRecordSub());
     404            0 :     CHK_RET(SubWaitMain());
     405              : 
     406            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     407              : 
     408            0 :     if (userMemOut.ptr() != commMemOut.ptr()) {
     409            0 :         src = commMemOut.range(slices_[rank].offset, slices_[rank].size);
     410            0 :         dst = userMemOut.range(slices_[rank].offset, slices_[rank].size);
     411            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, meshStreams_[meshStreams_.size() - 1]));
     412              :     }
     413              : 
     414            0 :     for (u32 round = 1; round < rankSize; round++) {
     415            0 :         u32 dstRank = BackwardRank(rank, rankSize, round);
     416            0 :         Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
     417            0 :         void* remMemPtr = nullptr;
     418            0 :         CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     419              : 
     420            0 :         src = DeviceMem::create(static_cast<char*>(remMemPtr) + dstRank * slices_[0].size, slices_[dstRank].size);
     421            0 :         dst = userMemOut.range(slices_[dstRank].offset, slices_[dstRank].size);
     422              : 
     423            0 :         CHK_RET(HcclD2DMemcpyAsync(
     424              :             dispatcher_, dst, src, subStream, links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
     425            0 :         CHK_RET(links[dstRank]->TxDataSignal(subStream));
     426            0 :         CHK_RET(links[dstRank]->RxDataSignal(subStream));
     427              :     }
     428              : 
     429            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     430              : 
     431            0 :     CHK_RET(SubRecordMain());
     432            0 :     CHK_RET(MainWaitSub());
     433              : 
     434            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     435              : 
     436            0 :     HCCL_INFO("AllGatherMesh finished: rank[%u]", rank);
     437            0 :     return HCCL_SUCCESS;
     438            0 : }
     439              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_LOCAL_REDUCE, AllReduceLocalReduce);
     440              : } // namespace hccl
        

Generated by: LCOV version 2.0-1