LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_executor/ins_alg_executor/all_reduce - ins_all_reduce_parallel_executor_opt.h (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 51 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 5 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              : #ifndef HCCLV2_INS_ALL_REDUCE_PARALLEL_EXECUTOR_OPT_H
      12              : #define HCCLV2_INS_ALL_REDUCE_PARALLEL_EXECUTOR_OPT_H
      13              : #include "ins_coll_alg_base.h"
      14              : 
      15              : namespace Hccl {
      16              : 
      17              : template <
      18              :     typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1, typename InsAlgTemplate2,
      19              :     typename InsAlgTemplate3>
      20              : class InsAllReduceParallelExecutorV2 : public InsCollAlgBase {
      21              : public:
      22              :     explicit InsAllReduceParallelExecutorV2();
      23              :     ~InsAllReduceParallelExecutorV2() override;
      24              : 
      25            0 :     std::string Describe() const override { return "Instruction based All Reduce Parallel Executor."; }
      26              : 
      27              :     // HOST 接口
      28              :     HcclResult Orchestrate(
      29              :         const RankGraph* rankGraph, const CollAlgOperator& op, const CollAlgParams& params, InsQuePtr insQue) override;
      30              :     // AICPU 接口
      31              :     HcclResult Orchestrate(
      32              :         const AlgTopoInfo& topoInfo, const CollAlgOperator& op, const CollAlgParams& params, ConnectedLinkMgr* linkMgr,
      33              :         InsQuePtr insQue) override;
      34              : 
      35              :     HcclResult CalcResOffload(const RankGraph* rankGraph, const u64& dataSize, CollOffloadOpResReq& resReq) override;
      36              : 
      37              :     HcclResult CalcRes(const RankGraph* rankGraph, CollAlgResReq& algResReq) override;
      38              : 
      39              : private:
      40              :     HcclResult CalcLocalRankSize();
      41              :     HcclResult GenInsQues(
      42              :         InsAlgTemplate0& tempAlgIntraRS, InsAlgTemplate1& tempAlgInterRS, InsAlgTemplate2& tempAlgIntraAG,
      43              :         InsAlgTemplate3& tempAlgInterAG);
      44              :     void GetParallelDataSplitRate(std::vector<float>& splitDataSize) const;
      45              :     HcclResult PrepareResForTemplate(
      46              :         const RankGraph* rankGraph, InsAlgTemplate0& tempAlgIntraRS, InsAlgTemplate1& tempAlgInterRS,
      47              :         InsAlgTemplate2& tempAlgIntraAG, InsAlgTemplate3& tempAlgInterAG);
      48              :     HcclResult PrepareResForTemplate(
      49              :         ConnectedLinkMgr* linkMgr, InsAlgTemplate0& tempAlgIntraRS, InsAlgTemplate1& tempAlgInterRS,
      50              :         InsAlgTemplate2& tempAlgIntraAG, InsAlgTemplate3& tempAlgInterAG);
      51              : 
      52              :     void GenRSIntraParams0(
      53              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      54              : 
      55              :     void GenRSInterParams0(
      56              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      57              : 
      58              :     void GenAGInterParams0(
      59              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      60              : 
      61              :     void GenAGIntraParams0(
      62              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      63              : 
      64              :     void GenRSInterParams1(
      65              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      66              : 
      67              :     void GenRSIntraParams1(
      68              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      69              : 
      70              :     void GenAGIntraParams1(
      71              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      72              : 
      73              :     void GenAGInterParams1(
      74              :         const u64 dataOffset, const u64 dataCount, const u64 scratchOff, TemplateDataParams& params) const;
      75              : 
      76            0 :     inline void InitAlgCommonParams(
      77              :         InsAlgTemplate0& tempAlgIntraRS, InsAlgTemplate1& tempAlgInterRS, InsAlgTemplate2& tempAlgIntraAG,
      78              :         InsAlgTemplate3& tempAlgInterAG, const CollAlgOperator& op) const
      79              :     {
      80            0 :         tempAlgIntraRS.SetDmaMode(dmaMode_);
      81            0 :         tempAlgIntraRS.InitReduceInfo(redOp_, dataType_);
      82            0 :         tempAlgIntraRS.SetCollOp(op);
      83              : 
      84            0 :         tempAlgInterRS.SetDmaMode(dmaMode_);
      85            0 :         tempAlgInterRS.InitReduceInfo(redOp_, dataType_);
      86            0 :         tempAlgInterRS.SetCollOp(op);
      87              : 
      88            0 :         tempAlgIntraAG.SetDmaMode(dmaMode_);
      89            0 :         tempAlgIntraAG.SetCollOp(op);
      90            0 :         tempAlgIntraAG.SetDataType(dataType_);
      91              : 
      92            0 :         tempAlgInterAG.SetDmaMode(dmaMode_);
      93            0 :         tempAlgInterAG.SetCollOp(op);
      94            0 :         tempAlgInterAG.SetDataType(dataType_);
      95            0 :     }
      96              : 
      97              :     // 统一设置 TemplateDataParams 的公共字段
      98            0 :     inline void SetTemplateDataParams(
      99              :         TemplateDataParams& params, BufferType inBuffType, BufferType outBuffType, u64 inBuffBaseOff,
     100              :         u64 outBuffBaseOff, u64 scratchBuffBaseOff, u64 sliceSize, u64 inputSliceStride, u64 outputSliceStride,
     101              :         u32 repeatNum, u64 inputRepeatStride, u64 outputRepeatStride, u64 tailSize) const
     102              :     {
     103            0 :         params.buffInfo.inBuffType = inBuffType;
     104            0 :         params.buffInfo.outBuffType = outBuffType;
     105            0 :         params.buffInfo.scratBuffType = BufferType::SCRATCH;
     106            0 :         params.buffInfo.inBuffBaseOff = inBuffBaseOff;
     107            0 :         params.buffInfo.outBuffBaseOff = outBuffBaseOff;
     108            0 :         params.buffInfo.scratchBuffBaseOff = scratchBuffBaseOff;
     109            0 :         params.sliceSize = sliceSize;
     110            0 :         params.inputSliceStride = inputSliceStride;
     111            0 :         params.outputSliceStride = outputSliceStride;
     112            0 :         params.repeatNum = repeatNum;
     113            0 :         params.inputRepeatStride = inputRepeatStride;
     114            0 :         params.outputRepeatStride = outputRepeatStride;
     115            0 :         params.tailSize = tailSize;
     116            0 :     }
     117              : 
     118              :     // 计算 Gen*Params1 中 intra 函数的公共 dataCountTmp
     119            0 :     inline u64 CalcDataCountTmp1(u64 dataCount) const
     120              :     {
     121            0 :         return (rankIdxLevel1_ != rankSizeLevel1_ - 1) ?
     122            0 :                    dataCount / rankSize_ * rankSizeLevel0_ :
     123            0 :                    dataCount - dataCount / rankSize_ * rankSizeLevel0_ * (rankSizeLevel1_ - 1);
     124              :     }
     125              : 
     126            0 :     inline HcclResult CalcQue(
     127              :         AlgTempResReq& resReqIntraRS, AlgTempResReq& resReqInterRS, AlgTempResReq& resReqIntraAG,
     128              :         AlgTempResReq& resReqInterAG)
     129              :     {
     130              :         // 申请算法模板所需资源
     131            0 :         if (!(resReqIntraRS.queNum > 0 && resReqInterRS.queNum > 0 && resReqIntraAG.queNum > 0
     132            0 :               && resReqInterAG.queNum > 0)) {
     133            0 :             HCCL_ERROR("resReqIntra.queNum and resReqInter.queNum must larger than 0.");
     134            0 :             return HcclResult::HCCL_E_INTERNAL;
     135              :         }
     136            0 :         u32 intraQueNum = std::max(resReqIntraRS.queNum, resReqIntraAG.queNum);
     137            0 :         u32 interQueNum = std::max(resReqInterRS.queNum, resReqInterAG.queNum);
     138            0 :         u32 totalQueNum = intraQueNum + interQueNum;
     139            0 :         CHK_RET(InitQueue(totalQueNum, requiredQue_));
     140            0 :         for (u32 i = 0; i < requiredQue_.size(); i++) {
     141            0 :             if (i < intraQueNum) {
     142            0 :                 intraQue_.push_back(requiredQue_[i]);
     143              :             } else {
     144            0 :                 interQue_.push_back(requiredQue_[i]);
     145              :             }
     146              :         }
     147            0 :         HCCL_INFO(
     148              :             "LGC requiredQue_.size is [%llu]. intraQue_.size is [%llu]. interQue_.size is [%llu].", requiredQue_.size(),
     149              :             intraQue_.size(), interQue_.size());
     150            0 :         syncQueues_.emplace_back(intraQue_[0]);
     151            0 :         syncQueues_.emplace_back(interQue_[0]);
     152            0 :         return HCCL_SUCCESS;
     153              :     }
     154              : 
     155              :     uint64_t rankSizeLevel0_{0};
     156              :     uint64_t rankSizeLevel1_{0};
     157              :     uint64_t rankSize_{0};
     158              : 
     159              :     uint64_t rankIdxLevel0_{0};
     160              :     uint64_t rankIdxLevel1_{0};
     161              : 
     162              :     u64 sliceCount_;
     163              : 
     164              :     std::vector<std::vector<RankId>> virtRanks_;
     165              :     std::vector<std::map<RankId, u32>> virtRankMap_; // map<virtRank, virtRankOrder>
     166              :     std::vector<std::vector<std::vector<RankId>>> vTopo_;
     167              : 
     168              :     std::vector<InsQuePtr> requiredQue_;
     169              :     std::vector<InsQuePtr> intraQue_;
     170              :     std::vector<InsQuePtr> interQue_;
     171              :     std::vector<InsQuePtr> syncQueues_;
     172              : 
     173              :     ResLinks intraRSLinks_;
     174              :     ResLinks interRSLinks_;
     175              :     ResLinks intraAGLinks_;
     176              :     ResLinks interAGLinks_;
     177              : };
     178              : 
     179              : } // namespace Hccl
     180              : 
     181              : #endif // HCCLV2_INS_ALL_REDUCE_PARALLEL_EXECUTOR_H
        

Generated by: LCOV version 2.0-1