LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template/temp_all_reduce - all_reduce_chunk_mesh.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 7.1 % 198 14
Test Date: 2026-07-28 12:11:00 Functions: 30.8 % 13 4

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "alg_template_register.h"
      12              : #include "all_reduce_chunk_mesh.h"
      13              : 
      14              : namespace hccl {
      15            2 : AllReduceChunkMesh::AllReduceChunkMesh(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
      16              : {
      17            2 : }
      18              : 
      19            4 : AllReduceChunkMesh::~AllReduceChunkMesh()
      20              : {
      21            4 : }
      22              : 
      23            2 : HcclResult AllReduceChunkMesh::Prepare(u64 reduceAttrBitMap, std::vector<Stream> &meshStreams,
      24              :     std::vector<std::shared_ptr<LocalNotify>> &meshSignal, std::vector<std::shared_ptr<LocalNotify>> &meshSignalAux,
      25              :     u32 interRank, u32 interRankSize, u32 userRank, HcomCollOpInfo *opInfo)
      26              : {
      27            2 :     reduceAttr_ = reduceAttrBitMap;
      28            2 :     localRank_ = interRank;
      29            2 :     localRankSize_ = interRankSize;
      30            2 :     userRank_ = userRank;
      31            2 :     meshStreams_ = meshStreams;
      32            2 :     meshSignal_ = &meshSignal;
      33            2 :     meshSignalAux_ = &meshSignalAux;
      34            2 :     opInfo_ = opInfo;
      35            2 :     return HCCL_SUCCESS;
      36              : }
      37            0 : HcclResult AllReduceChunkMesh::MainRecordSub()
      38              : {
      39            0 :     for (u32 signalIndex = 0; signalIndex < meshSignalAux_->size(); signalIndex++) {
      40            0 :         CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex],
      41              :             profilerInput_.stage));
      42              :     }
      43            0 :     return HCCL_SUCCESS;
      44              : }
      45              : 
      46            0 : HcclResult AllReduceChunkMesh::SubWaitMain()
      47              : {
      48            0 :     for (u32 streamIndex = 0; streamIndex < meshSignalAux_->size(); streamIndex++) {
      49            0 :         CHK_RET(LocalNotify::Wait(meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex],
      50              :             profilerInput_.stage));
      51              :     }
      52            0 :     return HCCL_SUCCESS;
      53              : }
      54              : 
      55            0 : HcclResult AllReduceChunkMesh::MainWaitSub()
      56              : {
      57            0 :     for (u32 signalIndex = 0; signalIndex < meshSignal_->size(); signalIndex++) {
      58            0 :         CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
      59              :     }
      60            0 :     return HCCL_SUCCESS;
      61              : }
      62              : 
      63            0 : HcclResult AllReduceChunkMesh::SubRecordMain()
      64              : {
      65            0 :     for (u32 streamIndex = 0; streamIndex < meshSignal_->size(); streamIndex++) {
      66            0 :         CHK_RET(LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex],
      67              :             profilerInput_.stage));
      68              :     }
      69            0 :     return HCCL_SUCCESS;
      70              : }
      71              : 
      72              : // 将数据均分,最小单位是128
      73            0 : HcclResult AllReduceChunkMesh::PrepareSlice(
      74              :     u64 dataCount, u32 unitSize, u32 sliceNum, std::vector<Slice> &dataSlice)
      75              : {
      76            0 :     u64 totalSize = dataCount * unitSize;
      77            0 :     Slice temp;
      78            0 :     dataSlice.clear();
      79            0 :     dataSlice.reserve(sliceNum);
      80            0 :     if (sliceNum == 0) {
      81            0 :         HCCL_ERROR("[Prepare][SliceData]data slice prepare, sliceNum is 0");
      82            0 :         return HCCL_E_PARA;
      83              :     }
      84            0 :     u64 sizePerSlice = (totalSize + sliceNum - 1) / sliceNum; /* 1是为了向上取整 */
      85            0 :     sizePerSlice = RoundUpWithDivisor(sizePerSlice, HCCL_MIN_SLICE_ALIGN);
      86            0 :     u64 residueSize = totalSize;
      87            0 :     u32 i = 0;
      88            0 :     while (residueSize > 0) {
      89            0 :         u64 sliceSize = sizePerSlice < residueSize ? sizePerSlice : residueSize;
      90            0 :         temp.size = sliceSize;
      91            0 :         temp.offset = totalSize - residueSize;
      92            0 :         i++;
      93            0 :         if (sliceSize <= 0) {
      94            0 :             HCCL_ERROR("[Prepare][SliceData]data_slice_prepare sliceSize[%llu]", sliceSize);
      95            0 :             return HCCL_E_PARA;
      96              :         }
      97            0 :         residueSize -= sliceSize;
      98            0 :         dataSlice.push_back(temp);
      99              :     }
     100            0 :     while (i < sliceNum) {
     101            0 :         temp.size = 0;
     102            0 :         temp.offset = totalSize;
     103            0 :         i++;
     104            0 :         dataSlice.push_back(temp);
     105              :     }
     106            0 :     return HCCL_SUCCESS;
     107              : }
     108              : 
     109            0 : HcclResult AllReduceChunkMesh::PrepareAllreduceSliceData()
     110              : {
     111            0 :     u32 unitSize = SIZE_TABLE[dataType_];
     112            0 :     HcclResult ret = HCCL_SUCCESS;
     113            0 :     CHK_RET(PrepareSlice(count_, unitSize, localRankSize_, slices_));
     114            0 :     for (u32 rank = 0; rank < localRankSize_; rank++) {
     115            0 :         std::vector<Slice> dataSegsSlice;
     116            0 :         ret = PrepareSlice(slices_[rank].size / unitSize, unitSize, localRankSize_ - 1, dataSegsSlice);
     117            0 :         sliceMap[rank] = dataSegsSlice;
     118            0 :         CHK_PRT_RET(
     119              :             ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceChunkMesh][PrepareSlice]rank[%u] failed", rank), ret);
     120            0 :     }
     121            0 :     return HCCL_SUCCESS;
     122              : }
     123              : 
     124            0 : HcclResult AllReduceChunkMesh::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
     125              : {
     126            0 :     HcclResult ret = HCCL_SUCCESS;
     127            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     128            0 :     CHK_PTR_NULL(stream_.ptr());
     129            0 :     HCCL_INFO("AllReduceChunkMesh run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
     130              :         rank,
     131              :         rankSize,
     132              :         inputMem_.ptr(),
     133              :         outputMem_.ptr(),
     134              :         count_);
     135              : 
     136            0 :     if (links.size() < rankSize) {
     137            0 :         HCCL_ERROR("[AllReduceChunkMesh][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]",
     138              :             rank,
     139              :             links.size(),
     140              :             rankSize);
     141            0 :         return HCCL_E_INTERNAL;
     142              :     }
     143              : 
     144              :     // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
     145            0 :     if (rankSize == 1) {
     146            0 :         if (opInfo_->inputAddr != opInfo_->outputAddr) {
     147            0 :             DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * DataUnitSize(dataType_));
     148            0 :             DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * DataUnitSize(dataType_));
     149            0 :             ret = HcclD2DMemcpyAsync(dispatcher_, userMemOut, userMemIn, stream_);
     150            0 :             CHK_PRT_RET(
     151              :                 ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] memcpy async failed", rank), ret);
     152            0 :         }
     153            0 :         return ret;
     154              :     }
     155              : 
     156            0 :     ret = PrepareAllreduceSliceData();
     157            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     158              :         HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in PrepareSliceData step",
     159              :             rank,
     160              :             count_),
     161              :         ret);
     162              : 
     163            0 :     ret = RunReduceScatter(rank, rankSize, links);
     164            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     165              :         HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in reducescater "
     166              :                    "step",
     167              :             rank,
     168              :             count_),
     169              :         ret);
     170              : 
     171            0 :     ret = RunAllGather(rank, rankSize, links);
     172            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     173              :         HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in AllGather "
     174              :                 "step",
     175              :             rank,
     176              :             count_),
     177              :         ret);
     178              : 
     179            0 :     HCCL_INFO("AllReduceChunkMesh finished: rank[%u] ranksize[%u]", rank, rankSize);
     180            0 :     return HCCL_SUCCESS;
     181              : }
     182              : 
     183            0 : HcclResult AllReduceChunkMesh::RunReduceScatter(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     184              : {
     185            0 :     HCCL_INFO("ReduceScatterMeshAtomicOpbase run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]",
     186              :         rank,
     187              :         rankSize,
     188              :         inputMem_.ptr(),
     189              :         outputMem_.ptr(),
     190              :         count_);
     191              : 
     192              :     // 数据准备
     193            0 :     u32 unitSize = DataUnitSize(dataType_);
     194              : 
     195            0 :     DeviceMem commMemOut = outputMem_;
     196              : 
     197            0 :     DeviceMem src;
     198            0 :     DeviceMem dst;
     199              : 
     200            0 :     src = DeviceMem::create(static_cast<char *>(opInfo_->inputAddr), count_ * unitSize);
     201              : 
     202            0 :     if (commMemOut.ptr() == opInfo_-> outputAddr) {
     203              :         // 图模式
     204            0 :         src = src.range(slices_[rank].offset, slices_[rank].size);
     205            0 :         dst = commMemOut.range(slices_[rank].offset, slices_[rank].size);
     206              :     } else {
     207              :         // 单算子
     208            0 :         src = src.range(0, count_ * unitSize);
     209            0 :         dst = commMemOut.range(0, count_ * unitSize);
     210              :     }
     211            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
     212              : 
     213            0 :     DeviceMem emptySrc = commMemOut.range(0, 0);
     214            0 :     DeviceMem emptyDst = commMemOut.range(0, 0);
     215              : 
     216              :     // 主从流之前加空拷贝 防止成环
     217            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     218              : 
     219            0 :     CHK_RET(MainRecordSub());
     220            0 :     CHK_RET(SubWaitMain());
     221              : 
     222            0 :     for (u32 round = 1; round < rankSize; round++) {
     223            0 :         u32 dstRank = (round + rank) % rankSize;
     224            0 :         Stream &subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
     225            0 :         CHK_RET(links[dstRank]->TxAck(subStream));
     226            0 :         CHK_RET(links[dstRank]->RxAck(subStream));
     227              :     }
     228              : 
     229            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     230              : 
     231            0 :     for (u32 round = 1; round < rankSize; round++) {
     232              :         // 主从流同步
     233              : 
     234            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     235              : 
     236            0 :         CHK_RET(SubRecordMain());
     237            0 :         CHK_RET(MainWaitSub());
     238              : 
     239            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     240              : 
     241            0 :         CHK_RET(MainRecordSub());
     242            0 :         CHK_RET(SubWaitMain());
     243              : 
     244            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     245              : 
     246              :         // 跨片reduceinline写
     247            0 :         for (u32 peer = 1; peer < rankSize; peer++) {
     248            0 :             u32 gap = (peer + round) > rankSize ? (peer + round - 1)%(rankSize -1):(peer + round - 1);
     249            0 :             u32 dstRank = (gap + rank) % rankSize;
     250            0 :             Stream &subStream = (peer == localRankSize_ - 1) ? stream_ : meshStreams_[peer - 1];
     251            0 :             u32 dstSlice = peer - 1;
     252            0 :             void *remMemPtr = nullptr;
     253              : 
     254            0 :             if (commMemOut.ptr() == opInfo_-> outputAddr) {
     255            0 :                 CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::INPUT_MEM, &remMemPtr));
     256              :             } else {
     257            0 :                 CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     258              :             }
     259              : 
     260            0 :             src = DeviceMem::create(
     261            0 :                 static_cast<char *>(remMemPtr) + slices_[rank].offset + sliceMap[rank][dstSlice].offset,
     262            0 :                 sliceMap[rank][dstSlice].size);
     263            0 :             dst = commMemOut.range(
     264            0 :                 slices_[rank].offset + sliceMap[rank][dstSlice].offset, sliceMap[rank][dstSlice].size);
     265            0 :             CHK_RET(HcclReduceAsync(dispatcher_, static_cast<void *>(src.ptr()),
     266              :                 sliceMap[rank][dstSlice].size / unitSize,
     267              :                 dataType_,
     268              :                 reductionOp_,
     269              :                 subStream,
     270              :                 static_cast<void *>(dst.ptr()),
     271              :                 links[dstRank]->GetRemoteRank(),
     272              :                 links[dstRank]->GetLinkType(), INLINE_REDUCE_BIT));
     273              :         }
     274              :     }
     275              : 
     276            0 :     for (u32 round = 1; round < rankSize; round++) {
     277            0 :         u32 gap = (round - 1) == 0 ? (rankSize - 1) : (round - 1);
     278            0 :         u32 dstRank = (rank + gap) % rankSize;
     279            0 :         Stream &subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
     280            0 :         CHK_RET(links[dstRank]->TxDataSignal(subStream));
     281            0 :         CHK_RET(links[dstRank]->RxDataSignal(subStream));
     282              :     }
     283              : 
     284            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     285              : 
     286            0 :     CHK_RET(SubRecordMain());
     287            0 :     CHK_RET(MainWaitSub());
     288              : 
     289            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     290            0 :     return HCCL_SUCCESS;
     291            0 : }
     292              : 
     293            0 : HcclResult AllReduceChunkMesh::RunAllGather(u32 rank, u32 rankSize, const std::vector<LINK> &links)
     294              : {
     295            0 :     HCCL_INFO("AllGatherMesh run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]",
     296              :         rank,
     297              :         rankSize,
     298              :         inputMem_.ptr(),
     299              :         outputMem_.ptr(),
     300              :         count_);
     301            0 :     u32 unitSize = DataUnitSize(dataType_);
     302              : 
     303            0 :     DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * unitSize);
     304            0 :     DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
     305              : 
     306            0 :     DeviceMem emptySrc = userMemOut.range(0, 0);
     307            0 :     DeviceMem emptyDst = commMemOut.range(0, 0);
     308              : 
     309            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     310              : 
     311            0 :     CHK_RET(MainRecordSub());
     312            0 :     CHK_RET(SubWaitMain());
     313              : 
     314            0 :     for (u32 round = 1; round < rankSize; round++) {
     315            0 :         u32 dstRank = BackwardRank(rank, rankSize, round);
     316            0 :         Stream &subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
     317            0 :         CHK_RET(links[dstRank]->TxAck(subStream));
     318            0 :         CHK_RET(links[dstRank]->RxAck(subStream));
     319              :     }
     320              : 
     321            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     322              : 
     323            0 :     CHK_RET(SubRecordMain());
     324            0 :     CHK_RET(MainWaitSub());
     325              : 
     326            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     327              : 
     328            0 :     CHK_RET(MainRecordSub());
     329            0 :     CHK_RET(SubWaitMain());
     330              : 
     331            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     332              : 
     333            0 :     DeviceMem src;
     334            0 :     DeviceMem dst;
     335            0 :     if (opInfo_->outputAddr != outputMem_.ptr()) {
     336            0 :         dst = userMemOut.range(slices_[rank].offset, slices_[rank].size);
     337            0 :         src = commMemOut.range(slices_[rank].offset, slices_[rank].size);
     338            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, meshStreams_[meshStreams_.size()-1]));
     339              :     }
     340              : 
     341            0 :     for (u32 round = 1; round < rankSize; round++) {
     342            0 :         u32 dstRank = BackwardRank(rank, rankSize, round);
     343            0 :         Stream &subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
     344            0 :         void *remMemPtr = nullptr;
     345            0 :         CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
     346            0 :         src = DeviceMem::create(static_cast<char *>(remMemPtr) + slices_[dstRank].offset, slices_[dstRank].size);
     347            0 :         dst = userMemOut.range(slices_[dstRank].offset, slices_[dstRank].size);
     348            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStream,
     349              :             links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
     350              : 
     351            0 :         CHK_RET(links[dstRank]->TxDataSignal(subStream));
     352            0 :         CHK_RET(links[dstRank]->RxDataSignal(subStream));
     353            0 :         HCCL_DEBUG("[AllReduceChunkMesh]round %u success");
     354              :     }
     355              : 
     356            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     357              : 
     358            0 :     CHK_RET(SubRecordMain());
     359            0 :     CHK_RET(MainWaitSub());
     360              : 
     361            0 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
     362              : 
     363            0 :     HCCL_INFO("[AllGatherMesh] finished: rank[%u]", rank);
     364            0 :     return HCCL_SUCCESS;
     365            0 : }
     366              : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_CHUNK_MESH, AllReduceChunkMesh);
     367              : }  // namespace hccl
        

Generated by: LCOV version 2.0-1