LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/base/alg_template - alg_template_multi_deter_pipeline.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 286 0
Test Date: 2026-08-17 10:19:35 Functions: 0.0 % 33 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              : #include "alg_template_multi_deter_pipeline.h"
      11              : namespace hccl {
      12            0 : MultiDeterPipeline::MultiDeterPipeline(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
      13              : 
      14            0 : MultiDeterPipeline::~MultiDeterPipeline() {}
      15              : 
      16            0 : HcclResult MultiDeterPipeline::RunAsync() { return HCCL_SUCCESS; }
      17              : 
      18              : // ReduceScatterDeterPipeline
      19            0 : HcclResult MultiDeterPipeline::Prepare(
      20              :     HcomCollOpInfo* opInfo, DeviceMem& buffer, const u64 count, const u64 offset, const std::vector<Slice>& slices,
      21              :     const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo, Stream& mainStream,
      22              :     std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
      23              :     std::vector<std::shared_ptr<LocalNotify>>& notifySub)
      24              : {
      25            0 :     return HCCL_SUCCESS;
      26              : }
      27              : 
      28              : // AllReduceDeterPipeline
      29            0 : HcclResult MultiDeterPipeline::Prepare(
      30              :     HcomCollOpInfo* opInfo, DeviceMem& inBuffer, DeviceMem& outBuffer, const u64 count,
      31              :     const std::vector<Slice>& slices, const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo,
      32              :     Stream& mainStream, std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
      33              :     std::vector<std::shared_ptr<LocalNotify>>& notifySub)
      34              : {
      35            0 :     return HCCL_SUCCESS;
      36              : }
      37              : 
      38            0 : HcclResult MultiDeterPipeline::MainWaitSub(u32 begin, u32 end)
      39              : {
      40            0 :     for (u32 signalIndex = begin; signalIndex < end; signalIndex++) {
      41            0 :         CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[signalIndex], INVALID_VALUE_STAGE));
      42              :     }
      43            0 :     return HCCL_SUCCESS;
      44              : }
      45              : 
      46            0 : HcclResult MultiDeterPipeline::SubRecordMain(u32 begin, u32 end)
      47              : {
      48            0 :     for (u32 streamIndex = begin; streamIndex < end; streamIndex++) {
      49            0 :         CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, streamNotifyMain_[streamIndex], -1));
      50              :     }
      51            0 :     return HCCL_SUCCESS;
      52              : }
      53              : 
      54            0 : HcclResult MultiDeterPipeline::MainRecordSub(u32 begin, u32 end)
      55              : {
      56            0 :     for (u32 signalIndex = begin; signalIndex < end; signalIndex++) {
      57            0 :         CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[signalIndex], -1));
      58              :     }
      59            0 :     return HCCL_SUCCESS;
      60              : }
      61              : 
      62              : // begin max = 7, end max = 11
      63            0 : HcclResult MultiDeterPipeline::SubWaitMain(u32 begin, u32 end)
      64              : {
      65            0 :     for (u32 streamIndex = begin; streamIndex < end; streamIndex++) {
      66            0 :         CHK_RET(LocalNotify::Wait(
      67              :             subStreams_[streamIndex], dispatcher_, streamNotifySub_[streamIndex], INVALID_VALUE_STAGE));
      68              :     }
      69            0 :     return HCCL_SUCCESS;
      70              : }
      71              : 
      72            0 : HcclResult MultiDeterPipeline::GetRemoteCclbufferDeviceMem(
      73              :     u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
      74              : {
      75            0 :     return HCCL_SUCCESS;
      76              : }
      77              : 
      78            0 : HcclResult MultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& locaMem)
      79              : {
      80            0 :     return HCCL_SUCCESS;
      81              : }
      82              : 
      83            0 : HcclResult MultiDeterPipeline::GetLocalUserOutDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
      84              : {
      85            0 :     return HCCL_SUCCESS;
      86              : }
      87              : 
      88              : HcclResult
      89            0 : MultiDeterPipeline::GetLocalInCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
      90              : {
      91            0 :     return HCCL_SUCCESS;
      92              : }
      93              : 
      94              : HcclResult
      95            0 : MultiDeterPipeline::GetLocalOutCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
      96              : {
      97            0 :     return HCCL_SUCCESS;
      98              : }
      99              : 
     100            0 : HcclResult MultiDeterPipeline::RunLocalCopy() { return HCCL_SUCCESS; }
     101              : 
     102            0 : HcclResult MultiDeterPipeline::RunIntraAlltoallPreSync(u32 step) { return HCCL_SUCCESS; }
     103              : 
     104            0 : HcclResult MultiDeterPipeline::RunIntraAlltoall(u32 step)
     105              : {
     106            0 :     u32 recvServerId = GetPreServerIdByStep(step);  // 从上一个收 2
     107            0 :     u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个 1
     108              :     // 机内alltoall full mesh收集数据,是为了收集机内第sendServerId整块的内存(包含intraRankId_块)
     109              :     // 该索引是为了计算第sendServerId整块内的内存序号块(0~intraRankId_-1)
     110            0 :     std::vector<u32> localUsrInIndex;
     111            0 :     for (u32 i = intraRankId_ + 1; i < intraRankSize_ + intraRankId_; ++i) {
     112            0 :         localUsrInIndex.push_back(i % intraRankSize_);
     113              :     }
     114            0 :     for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
     115            0 :         u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
     116            0 :         LINK sendIntraLink = intraLinks_[sendIntraRankId];
     117            0 :         DeviceMem srcMem;
     118            0 :         DeviceMem dstMem;
     119              :         // 从usrin收集发给下一个cclbufer的数据, 收集的所有数据需要发送给机间序号为sendServerId的server
     120            0 :         u32 needSendInputndex = GetRankIdx(sendServerId, localUsrInIndex[i]);
     121              :         // 发送数据到cclbufer 索引为[intraRankId_, localUsrInIndex[i]]
     122            0 :         u32 recvIntraRankIdx = alltoallRecvBlockIdxMap_[intraRankId_][localUsrInIndex[i]];
     123            0 :         u32 recvCclbufferIndex = GetRankIdx(recvServerId, recvIntraRankIdx);
     124            0 :         CHK_RET(GetLocalUserInDeviceMem(needSendInputndex, srcMem));
     125            0 :         CHK_RET(GetRemoteCclbufferDeviceMem(needSendInputndex, sendIntraLink, recvCclbufferIndex, dstMem));
     126              :         // 发送给机内rank索引 [serverId_, sendIntraRankId]
     127            0 :         u32 remoteUserRank = GetRankIdx(serverId_, sendIntraRankId);
     128              :         // SDMA copy write语义,因为只能写到cclbuffer
     129            0 :         CHK_RET(HcclD2DMemcpyAsync(
     130              :             dispatcher_, dstMem, srcMem, subStreams_[i], remoteUserRank, sendIntraLink->GetLinkType()));
     131            0 :         CHK_RET(sendIntraLink->TxDataSignal(subStreams_[i]));
     132            0 :         CHK_RET(sendIntraLink->RxDataSignal(subStreams_[i]));
     133            0 :         HCCL_DEBUG("[%s] intra-server SDMA send, intraRank: [%u] -> [%u]", __func__, intraRankId_, sendIntraRankId);
     134            0 :         HCCL_DEBUG(
     135              :             "[%s] intra-server SDMA send, mem: inputMem[%u, %u] -> cclbuffer[%u, %u]; cclbufferNo[%u] -> [%u]",
     136              :             __func__, sendServerId, localUsrInIndex[i], recvServerId, recvIntraRankIdx, needSendInputndex,
     137              :             recvCclbufferIndex);
     138            0 :     }
     139            0 :     HCCL_INFO("[%s] intra-server step[%u] run alltoall success", __func__, step);
     140            0 :     return HCCL_SUCCESS;
     141            0 : }
     142              : 
     143            0 : HcclResult MultiDeterPipeline::GroupTasksByStream(
     144              :     u32 activeCount, const std::vector<bool>& isReduceBlock, u32 retIndex,
     145              :     std::vector<std::vector<std::vector<std::pair<u32, u32>>>>& batchStreamTasks, // 输出:批次→流→任务
     146              :     std::vector<bool>& processed, std::vector<u32>& origIdxMap, u32& newActiveCount)
     147              : {
     148            0 :     batchStreamTasks.clear(); // 清空批次任务
     149            0 :     processed.assign(activeCount, false);
     150            0 :     newActiveCount = 0;
     151              : 
     152            0 :     const u32 mergeStep = 2;
     153            0 :     const u32 batchSize = MAX_REDUCE_STREAM_NUM;
     154            0 :     const u32 totalGroups = (activeCount + mergeStep - 1) / mergeStep;
     155            0 :     const u32 batchNum = (totalGroups + batchSize - 1) / batchSize;
     156              : 
     157              :     // 逐批生成流任务
     158            0 :     for (u32 batch = 0; batch < batchNum; batch++) {
     159              :         // 初始化当前批次的流任务(MAX_REDUCE_STREAM_NUM条流)
     160            0 :         std::vector<std::vector<std::pair<u32, u32>>> streamTasks(MAX_REDUCE_STREAM_NUM);
     161            0 :         u32 startGroup = batch * batchSize;
     162            0 :         u32 endGroup = std::min((batch + 1) * batchSize, totalGroups);
     163              : 
     164              :         // 处理当前批次的分组
     165            0 :         for (u32 group = startGroup; group < endGroup; group++) {
     166            0 :             u32 idx0 = group * mergeStep;
     167            0 :             u32 idx1 = idx0 + 1;
     168            0 :             if (idx1 >= activeCount) {
     169            0 :                 processed[idx0] = false;
     170            0 :                 continue;
     171              :             }
     172              : 
     173              :             // 选择dst/src(原有优先级逻辑不变)
     174              :             u32 dstIdx, srcIdx;
     175            0 :             if (origIdxMap[idx0] == retIndex) { // 当前idx0的原始索引是目标块,强制为dst
     176            0 :                 dstIdx = idx0;
     177            0 :                 srcIdx = idx1;
     178            0 :             } else if (origIdxMap[idx1] == retIndex) { // 当前idx1的原始索引是目标块,强制为dst
     179            0 :                 dstIdx = idx1;
     180            0 :                 srcIdx = idx0;
     181            0 :             } else if (isReduceBlock[idx0] && !isReduceBlock[idx1]) {
     182            0 :                 dstIdx = idx0;
     183            0 :                 srcIdx = idx1;
     184            0 :             } else if (!isReduceBlock[idx0] && isReduceBlock[idx1]) {
     185            0 :                 dstIdx = idx1;
     186            0 :                 srcIdx = idx0;
     187              :             } else {
     188            0 :                 dstIdx = std::max(idx0, idx1);
     189            0 :                 srcIdx = std::min(idx0, idx1);
     190              :             }
     191              : 
     192              :             // 分配到当前批次的流任务中
     193            0 :             u32 batchInnerGroupIdx = group - startGroup;
     194            0 :             u32 streamId = batchInnerGroupIdx % MAX_REDUCE_STREAM_NUM;
     195            0 :             streamTasks[streamId].emplace_back(srcIdx, dstIdx);
     196              : 
     197            0 :             processed[srcIdx] = true;
     198            0 :             processed[dstIdx] = false;
     199              : 
     200            0 :             HCCL_DEBUG(
     201              :                 "[%s] batch[%u] group[%u] merge src[%u] -> dst[%u] on stream[%u]", __func__, batch, group, srcIdx,
     202              :                 dstIdx, streamId + reduceStreamBegin_);
     203              :         }
     204              : 
     205              :         // 将当前批次的流任务加入总批次列表
     206            0 :         batchStreamTasks.push_back(streamTasks);
     207            0 :     }
     208              : 
     209            0 :     newActiveCount = std::count(processed.begin(), processed.end(), false);
     210            0 :     return HCCL_SUCCESS;
     211              : }
     212              : 
     213            0 : HcclResult MultiDeterPipeline::BatchPostNotifyForStreams(
     214              :     const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
     215              : {
     216            0 :     return HCCL_SUCCESS;
     217              : }
     218              : 
     219            0 : HcclResult MultiDeterPipeline::ExecuteStreamTasks(
     220              :     const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, const std::vector<DeviceMem>& validMem,
     221              :     std::vector<u32>& origIdxMap, bool useMainStream)
     222              : {
     223            0 :     for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
     224            0 :         if (streamTasks[s].empty())
     225            0 :             continue;
     226              : 
     227            0 :         u32 streamIdx = reduceStreamBegin_ + s;
     228            0 :         Stream& subStream = subStreams_[streamIdx];
     229            0 :         Stream& stream = useMainStream ? mainStream_ : subStream;
     230              : 
     231            0 :         for (const auto& task : streamTasks[s]) {
     232            0 :             u32 srcIdx = task.first;
     233            0 :             u32 dstIdx = task.second;
     234            0 :             const DeviceMem& dstMem = validMem[dstIdx];
     235            0 :             const DeviceMem& srcMem = validMem[srcIdx];
     236            0 :             u64 count = srcMem.size() / unitSize_;
     237            0 :             CHK_RET(HcclReduceAsync(
     238              :                 dispatcher_, srcMem.ptr(), count, dataType_, reductionOp_, stream, dstMem.ptr(), INVALID_VALUE_RANKID,
     239              :                 LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
     240            0 :             HCCL_DEBUG(
     241              :                 "[%s] stream[%u] execute task: merge src[%u] -> dst[%u], origSrc[%u] -> origDst[%u]", __func__,
     242              :                 useMainStream ? 0 : streamIdx, srcIdx, dstIdx, origIdxMap[srcIdx], origIdxMap[dstIdx]);
     243              :         }
     244              :     }
     245            0 :     return HCCL_SUCCESS;
     246              : }
     247              : 
     248            0 : void MultiDeterPipeline::CompressActiveSet(
     249              :     std::vector<DeviceMem>& validMem, std::vector<bool>& isReduceBlock, std::vector<u32>& origIdxMap,
     250              :     const std::vector<bool>& processed, u32& trackedTargetIdx, const u32 origRetIndex)
     251              : {
     252            0 :     std::vector<DeviceMem> newValidMem;
     253            0 :     std::vector<bool> newIsReduceBlock;
     254            0 :     std::vector<u32> newOrigIdxMap;
     255            0 :     u32 newTrackedTargetIdx = 0;
     256            0 :     bool foundTarget = false;
     257              : 
     258              :     // 保留未处理的块(processed=false,即dst块)
     259            0 :     for (u32 i = 0; i < validMem.size(); i++) {
     260            0 :         if (!processed[i]) { // 仅保留dst块,移除src块(processed=true)
     261            0 :             newValidMem.push_back(validMem[i]);
     262            0 :             newIsReduceBlock.push_back(isReduceBlock[i]);
     263            0 :             newOrigIdxMap.push_back(origIdxMap[i]);
     264              : 
     265              :             // 追踪原始目标块(retIndex)的新索引
     266            0 :             if (!foundTarget && origIdxMap[i] == origRetIndex) {
     267            0 :                 newTrackedTargetIdx = newValidMem.size() - 1;
     268            0 :                 foundTarget = true;
     269              :             }
     270              :         }
     271              :     }
     272              : 
     273            0 :     validMem.swap(newValidMem);
     274            0 :     isReduceBlock.swap(newIsReduceBlock);
     275            0 :     origIdxMap.swap(newOrigIdxMap);
     276              :     // 未找到目标块时,默认指向最后一个块
     277            0 :     trackedTargetIdx = foundTarget ? newTrackedTargetIdx : (validMem.empty() ? 0 : validMem.size() - 1);
     278            0 :     HCCL_DEBUG(
     279              :         "[%s] compressed: old size[%llu], new size[%llu], trackedTargetIdx[%u]", __func__, processed.size(),
     280              :         validMem.size(), trackedTargetIdx);
     281            0 : }
     282              : 
     283            0 : HcclResult MultiDeterPipeline::LocalReduce(
     284              :     std::vector<DeviceMem>& reduceMem, std::vector<bool>& isReduceBlock, u32 retIndex, bool useMainStream)
     285              : {
     286            0 :     const u32 totalBlockCount = reduceMem.size();
     287              :     // 校验1:容器大小匹配 + retIndex越界
     288            0 :     if (reduceMem.size() != isReduceBlock.size() || retIndex >= totalBlockCount) {
     289            0 :         HCCL_ERROR(
     290              :             "[%s] Invalid param (size mismatch: %llu vs %llu, retIndex: %u >= %u)", __func__, reduceMem.size(),
     291              :             isReduceBlock.size(), retIndex, totalBlockCount);
     292            0 :         return HCCL_E_PARA;
     293              :     }
     294              :     // 校验2:目标内存块有效
     295            0 :     const DeviceMem& targetCCLBuffer = reduceMem[retIndex];
     296            0 :     if (targetCCLBuffer.ptr() == nullptr || targetCCLBuffer.size() == 0) {
     297            0 :         HCCL_ERROR(
     298              :             "[%s] Target CCLBuffer invalid (ptr: %p, size: %llu)", __func__, targetCCLBuffer.ptr(),
     299              :             targetCCLBuffer.size());
     300            0 :         return HCCL_E_MEMORY;
     301              :     }
     302              : 
     303            0 :     std::vector<DeviceMem> validMem = std::move(reduceMem); // 外层不使用reduceMem
     304            0 :     std::vector<bool> validIsReduceBlock = std::move(isReduceBlock);
     305              :     // 1. 动态追踪目标块索引 2. 原索引→当前索引的映射表
     306            0 :     std::vector<u32> origIdxMap(validMem.size());
     307            0 :     for (size_t i = 0; i < origIdxMap.size(); ++i) {
     308            0 :         origIdxMap[i] = i;
     309              :     }
     310              : 
     311            0 :     if (validMem.size() == 1) {
     312            0 :         HCCL_ERROR("[%s] validMem size is one, only target block valid", __func__);
     313            0 :         return HCCL_E_PARA;
     314              :     }
     315              : 
     316            0 :     u32 trackedTargetIdx = retIndex;
     317            0 :     u32 activeCount = validMem.size();
     318            0 :     const u32 origRetIndex = retIndex;
     319              : 
     320            0 :     while (activeCount > 1) {
     321            0 :         std::vector<std::vector<std::vector<std::pair<u32, u32>>>> batchStreamTasks; // 批次→流→任务
     322            0 :         std::vector<bool> processed(activeCount, false);
     323            0 :         u32 newActiveCount = 0;
     324              : 
     325              :         // 2.1 生成按批次组织的流任务
     326            0 :         CHK_RET(GroupTasksByStream(
     327              :             activeCount, validIsReduceBlock, origRetIndex, batchStreamTasks, processed, origIdxMap, newActiveCount));
     328              : 
     329              :         // 2.2 逐批执行流任务(串行处理每批)
     330            0 :         for (const auto& streamTasks : batchStreamTasks) {
     331              :             // a. 执行start phase notify(仅处理有任务的流)
     332            0 :             CHK_RET(BatchPostNotifyForStreams(streamTasks, true, useMainStream));
     333              :             // b. 执行当前批次的流任务(src→dst归约)
     334            0 :             CHK_RET(ExecuteStreamTasks(streamTasks, validMem, origIdxMap, useMainStream));
     335              :             // c. 执行sync phase notify(等待当前批次完成)
     336            0 :             CHK_RET(BatchPostNotifyForStreams(streamTasks, false, useMainStream));
     337              :         }
     338              :         // 2.3 压缩活跃块(移除src块,保留dst块)
     339            0 :         CompressActiveSet(validMem, validIsReduceBlock, origIdxMap, processed, trackedTargetIdx, origRetIndex);
     340            0 :         activeCount = validMem.size();
     341            0 :         HCCL_DEBUG("[LocalReduce] round done: activeCount=%u -> %u", newActiveCount, activeCount);
     342            0 :     }
     343            0 :     HCCL_DEBUG(
     344              :         "[%s] Local reduce success (merge to retIndex[%u] CCLBuffer, final tracked idx[%u])", __func__, retIndex,
     345              :         trackedTargetIdx);
     346            0 :     return HCCL_SUCCESS;
     347            0 : }
     348              : 
     349            0 : HcclResult MultiDeterPipeline::RunIntraLocalReduce(u32 step) { return HCCL_SUCCESS; }
     350              : 
     351            0 : HcclResult MultiDeterPipeline::RunInterSend(u32 step) { return HCCL_SUCCESS; }
     352              : 
     353            0 : HcclResult MultiDeterPipeline::RunFinalReduce() { return HCCL_SUCCESS; }
     354              : 
     355            0 : HcclResult MultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase) { return HCCL_SUCCESS; }
     356              : 
     357            0 : HcclResult MultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase) { return HCCL_SUCCESS; }
     358              : 
     359            0 : HcclResult MultiDeterPipeline::AlltoallLocalReduceSync(u32 step, bool isStartPhase)
     360              : {
     361            0 :     bool alltoallStep = (step < allSteps_);
     362            0 :     bool localReduceStep = (step > 1 && step < allSteps_ + 1);
     363            0 :     if (alltoallStep) {
     364            0 :         CHK_RET(AlltoallSync(step, isStartPhase));
     365              :     }
     366            0 :     if (localReduceStep) {
     367            0 :         CHK_RET(LocalReduceSync(step, isStartPhase));
     368              :     }
     369            0 :     return HCCL_SUCCESS;
     370              : }
     371              : 
     372            0 : HcclResult MultiDeterPipeline::RunAsyncLocalReduceSerial()
     373              : {
     374            0 :     HCCL_INFO(
     375              :         "[MultiDeterPipeline] run begin: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p]", userRank_, userRankSize_,
     376              :         usrInMemPtr_, usrOutMemPtr_);
     377            0 :     CHK_SMART_PTR_NULL(dispatcher_);
     378              :     // 以机内8卡为例主流 + 从流 = 1 + 7 + 4 = 12
     379            0 :     allSteps_ = serverSize_;
     380              :     // #1 机间发送,#2 local reduce,#3 机内alltoall
     381              :     // 以机间的pairwise来分步,#n表示pairwise的第n步
     382            0 :     for (u32 step = 1; step < allSteps_ + 1; step++) {
     383            0 :         HCCL_DEBUG(
     384              :             "[%s] userRank[%u], intraRankId[%u], serverId[%u], step[%u/%u] begin", __func__, userRank_, intraRankId_,
     385              :             serverId_, step, allSteps_);
     386              :         // alltoall 主从流同步+前同步+拉齐
     387            0 :         CHK_RET(AlltoallSync(step, true));
     388            0 :         if (step < allSteps_ - 1) {
     389            0 :             CHK_RET(RunIntraAlltoallPreSync(step));
     390              :         }
     391            0 :         if (step == allSteps_ - 1) {
     392            0 :             CHK_RET(RunIntraAlltoallPreSync(0));
     393              :         }
     394            0 :         if (step == 1) {
     395            0 :             CHK_RET(RunLocalCopy());
     396              :         }
     397            0 :         if (step > 1) {
     398            0 :             CHK_RET(RunInterSend(step - 1));
     399              :         }
     400              :         // #1 机内RS alltoall + local reduce串行
     401            0 :         if (step < allSteps_) {
     402            0 :             CHK_RET(RunIntraAlltoall(step));
     403            0 :             CHK_RET(AlltoallSync(step, false));
     404            0 :             CHK_RET(LocalReduceSync(step, true));
     405            0 :             CHK_RET(RunIntraLocalReduce(step));
     406            0 :             CHK_RET(LocalReduceSync(step, false));
     407              :         } else {
     408              :             // #0 机内RS alltoall + local reduce串行,#0不需要向其他机发送数据
     409            0 :             CHK_RET(RunIntraAlltoall(0));
     410            0 :             CHK_RET(AlltoallSync(0, false));
     411            0 :             CHK_RET(LocalReduceSync(0, true));
     412            0 :             CHK_RET(RunIntraLocalReduce(0));
     413            0 :             CHK_RET(LocalReduceSync(0, false));
     414              :         }
     415              :     }
     416              :     // 总local reduce
     417            0 :     CHK_RET(RunFinalReduce());
     418            0 :     HCCL_INFO("[MultiDeterPipeline] MultiDeterPipeline success userRank[%u] ", userRank_);
     419            0 :     return HCCL_SUCCESS;
     420              : }
     421              : 
     422              : // 每个server内首先要进行alltoall full mesh收集数据,再进行机内local reduce,最后发送给指定server
     423            0 : HcclResult MultiDeterPipeline::RunAsyncReduceScatterPipeline()
     424              : {
     425            0 :     constexpr u64 HCCL_MEDIUM_COUNT_2_MB = 2 * 1024 * 1024;
     426              :     // 2机或者数据量小于2MB走localreduce串行算法
     427            0 :     if (serverSize_ <= LOCAL_REDUCE_SERIIAL_ALG_SERVER_NUM || GetLocalReduceSerialThresh() < HCCL_MEDIUM_COUNT_2_MB) {
     428            0 :         CHK_RET(RunAsyncLocalReduceSerial());
     429            0 :         return HCCL_SUCCESS;
     430              :     }
     431            0 :     HCCL_INFO("[MultiDeterPipeline] [%s] begin, userRank[%u]", __func__, userRank_);
     432              :     // 以机内8卡为例主流 + 从流 = 1 + 7 + 4 = 12
     433              :     // pairwise总共需要serverSize_步,#1机内alltoall只能自己执行,无法和其他步骤并行,所以总共需要serverSize_ + 1个步骤
     434            0 :     allSteps_ = serverSize_ + 1;
     435              :     // #1 机间发送,#2 local reduce,#3 机内alltoall
     436              :     // 以机间的pairwise来分步,#n表示pairwise的第n步
     437            0 :     for (u32 step = 1; step < allSteps_ + 1; step++) {
     438            0 :         HCCL_DEBUG(
     439              :             "[%s] userRank[%u], intraRankId[%u], serverId[%u], step[%u/%u] begin", __func__, userRank_, intraRankId_,
     440              :             serverId_, step, allSteps_);
     441            0 :         CHK_RET(AlltoallLocalReduceSync(step, true));
     442            0 :         if (step == 1) {
     443            0 :             CHK_RET(RunLocalCopy());
     444              :         }
     445            0 :         if (step < allSteps_ - 1) {
     446            0 :             CHK_RET(RunIntraAlltoallPreSync(step));
     447              :         }
     448            0 :         if (step == allSteps_ - 1) {
     449            0 :             CHK_RET(RunIntraAlltoallPreSync(0));
     450              :         }
     451            0 :         if (step > STEP_OFFSET_TWO) {
     452            0 :             CHK_RET(RunInterSend(step - STEP_OFFSET_TWO));
     453              :         }
     454              :         // allSteps_最小为3
     455            0 :         if (step > 1 && step < allSteps_) {
     456            0 :             CHK_RET(RunIntraLocalReduce(step - 1));
     457              :         }
     458              :         // 最后一步,需要额外执行 #0 机内RS local reduce,#0不需要向其他机发送数据
     459            0 :         if (step == allSteps_) {
     460            0 :             CHK_RET(RunIntraLocalReduce(0));
     461              :         }
     462            0 :         if (step < allSteps_ - 1) {
     463            0 :             CHK_RET(RunIntraAlltoall(step));
     464              :         }
     465            0 :         if (step == allSteps_ - 1) {
     466              :             // #0 机内RS alltoall
     467            0 :             CHK_RET(RunIntraAlltoall(0));
     468              :         }
     469            0 :         CHK_RET(AlltoallLocalReduceSync(step, false));
     470              :     }
     471              :     // 总local reduce
     472            0 :     CHK_RET(RunFinalReduce());
     473            0 :     HCCL_INFO("[MultiDeterPipeline] [%s] end, userRank[%u]", __func__, userRank_);
     474            0 :     return HCCL_SUCCESS;
     475              : }
     476              : 
     477              : // 遍历所有发送方rank和发送方块索引,预计算映射关系,目的是构造出下属矩阵
     478              : // srcRank\srcBlockIdx  0       1       2
     479              : //         0           MAX      0       0
     480              : //         1            0  MAX  1
     481              : //         2            1       1  MAX
     482              : // 例1,发送方是rank0,接收端是rank1,那么接收端非自身 rank 列表为[0, 2],那么rank0的索引为0,所以就发到rank1的第0块内存
     483              : // 例2,发送方是rank0,接收端是rank2,那么接收端非自身 rank 列表为[0, 1],那么rank0的索引为0,所以就发到rank2的第0块内存
     484              : // 例3,发送方是rank1,接收端是rank2,那么接收端非自身 rank 列表为[0, 1],那么rank1的索引为1,所以就发到rank2的第1块内存
     485            0 : void MultiDeterPipeline::InitAlltoallRecvBlockIdxMap()
     486              : {
     487            0 :     const u32 rankSize = intraRankSize_;
     488            0 :     alltoallRecvBlockIdxMap_.resize(rankSize, std::vector<u32>(rankSize, UINT32_MAX));
     489            0 :     if (rankSize <= 1) {
     490            0 :         return;
     491              :     }
     492              :     // 规则一:接收端的块索引 = 发送方 rank 在「接收端非自身 rank 列表」中的索引
     493            0 :     std::vector<std::vector<u32>> dstRankToRankIndex(rankSize, std::vector<u32>(rankSize, UINT32_MAX));
     494            0 :     for (u32 dstRank = 0; dstRank < rankSize; ++dstRank) {
     495            0 :         for (u32 srcRank = 0; srcRank < rankSize; ++srcRank) {
     496            0 :             if (srcRank == dstRank) {
     497            0 :                 continue;
     498              :             }
     499              :             // 若 srcRank < dstRank:索引 = srcRank, 若 srcRank > dstRank:索引 = srcRank - 1
     500            0 :             const u32 idx = (srcRank < dstRank) ? srcRank : (srcRank - 1);
     501            0 :             dstRankToRankIndex[dstRank][srcRank] = idx;
     502              :         }
     503              :     }
     504            0 :     for (u32 srcRank = 0; srcRank < rankSize; ++srcRank) {
     505            0 :         for (u32 srcBlockIdx = 0; srcBlockIdx < rankSize; ++srcBlockIdx) {
     506              :             // 规则二:发送方块索引 = 接收端 rank
     507            0 :             const u32 dstRank = srcBlockIdx;
     508              :             // 跳过自身发送的无效场景
     509            0 :             if (srcRank == dstRank) {
     510            0 :                 continue;
     511              :             }
     512              :             // 查找发送方rank在列表中的索引,存入映射表
     513            0 :             const u32 dstBlockIdx = dstRankToRankIndex[dstRank][srcRank];
     514            0 :             alltoallRecvBlockIdxMap_[srcRank][srcBlockIdx] = dstBlockIdx;
     515            0 :             HCCL_DEBUG(
     516              :                 "[%s] srcRank[%u], srcBlockIdx[%u] -> dstRank[%u], dstBlockIdx[%u]", __func__, srcRank, srcBlockIdx,
     517              :                 dstRank, dstBlockIdx);
     518              :         }
     519              :     }
     520            0 : }
     521              : 
     522            0 : HcclResult MultiDeterPipeline::PrepareTopoInfo(const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo)
     523              : {
     524            0 :     serverSize_ = level1CommInfo.localRankSize;
     525            0 :     CHK_PRT_RET(
     526              :         serverSize_ < MIN_SERVER_NUM,
     527              :         HCCL_ERROR("[%s] Unexpected inter rank size[%u], which should >= 2.", __func__, serverSize_), HCCL_E_PARA);
     528              : 
     529            0 :     intraRankSize_ = level0CommInfo.localRankSize;
     530            0 :     CHK_PRT_RET(
     531              :         intraRankSize_ < MIN_INTRA_RANK_NUM,
     532              :         HCCL_ERROR("[%s] Unexpected intra rank size[%u], which should >= 3.", __func__, intraRankSize_), HCCL_E_PARA);
     533            0 :     intraRankId_ = level0CommInfo.localRank;
     534            0 :     serverId_ = level1CommInfo.localRank;
     535            0 :     userRankSize_ = intraRankSize_ * serverSize_;
     536            0 :     userRank_ = intraRankId_ + serverId_ * intraRankSize_;
     537              : 
     538            0 :     intraLinks_ = level0CommInfo.links;  // 节点内
     539            0 :     serverLinks_ = level1CommInfo.links; // 节点间
     540            0 :     HCCL_INFO(
     541              :         "[%s] opInfo: dataType[%u], unitSize[%u], memSliceSize[%u], usrInMem[%p], usrOutMem[%p], reductionOp[%u]",
     542              :         __func__, dataType_, unitSize_, memSliceSize_, usrInMemPtr_, usrOutMemPtr_, reductionOp_);
     543            0 :     HCCL_INFO(
     544              :         "[%s] topoInfo: userRank[%u], intraRankId[%u], intraRankSize[%u], serverId[%u], interRankSize[%u]", __func__,
     545              :         userRank_, intraRankId_, intraRankSize_, serverId_, serverSize_);
     546            0 :     HCCL_INFO(
     547              :         "[%s] topoInfo: severLinksNum[%zu], intraLinksNum[%zu]", __func__, serverLinks_.size(), intraLinks_.size());
     548            0 :     InitAlltoallRecvBlockIdxMap();
     549            0 :     return HCCL_SUCCESS;
     550              : }
     551              : } // namespace hccl
        

Generated by: LCOV version 2.0-1