Line data Source code
1 : /**
2 : * Copyright (c) 2026 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_BROADCAST_PARALLEL_APICPU_EXECUTOR_H
12 : #define HCCLV2_INS_BROADCAST_PARALLEL_APICPU_EXECUTOR_H
13 :
14 : #include "ins_coll_alg_base.h"
15 :
16 : namespace Hccl {
17 :
18 : template <
19 : typename AlgTopoMatch, typename InsAlgTemplate0, typename InsAlgTemplate1, typename InsAlgTemplate2,
20 : typename InsAlgTemplate3>
21 : class InsBroadcastParallelAiCpuExecutor : public InsCollAlgBase {
22 : public:
23 0 : explicit InsBroadcastParallelAiCpuExecutor() = default;
24 0 : ~InsBroadcastParallelAiCpuExecutor() override = default;
25 :
26 0 : std::string Describe() const override
27 : {
28 0 : return "Instruction based BroadCast Parallel AICPU Executor.";
29 : }
30 :
31 : HcclResult CalcRes(const RankGraph* rankGraph, CollAlgResReq& algResReq) override;
32 :
33 : HcclResult CalcResOffload(const RankGraph* rankGraph, const u64& dataSize, CollOffloadOpResReq& resReq) override;
34 : // HOST 接口
35 : HcclResult Orchestrate(
36 : const RankGraph* rankGraph, const CollAlgOperator& op, const CollAlgParams& params, InsQuePtr insQue) override;
37 : // AICPU 接口
38 : HcclResult Orchestrate(
39 : const AlgTopoInfo& topoInfo, const CollAlgOperator& op, const CollAlgParams& params, ConnectedLinkMgr* linkMgr,
40 : InsQuePtr insQue) override;
41 :
42 : private:
43 : struct ScratchMultiple {
44 : u32 interScatter;
45 : u32 intraScatter;
46 : u32 interAllGather;
47 : u32 intraAllGather;
48 : float maxMultiple;
49 : };
50 : struct SliceConfig {
51 : u32 loopTimes = 0;
52 : u64 sliceCount = 0; // 正常切分个数
53 : u64 sliceCountPart0 = 0; // 正常切块第一部分个数
54 : u64 sliceCountPart1 = 0; // 正常切块第二部分个数
55 : u64 finalSliceCount = 0; // 尾块切分个数
56 : u64 finalSliceCountPart0 = 0; // 尾块第一部分切分个数
57 : u64 finalSliceCountPart1 = 0; // 尾块第二部分切分个数
58 : u64 finalTailCountPart0 = 0; // 尾块非卡整数倍尾巴
59 : u64 finalTailCountPart1 = 0; // 尾块非卡整数倍尾巴
60 : };
61 : struct ScratchOffset {
62 : u64 interScatterStage0;
63 : u64 intraScatterStage0;
64 : u64 intraScatterStage1;
65 : u64 interScatterStage1;
66 : u64 intraAllGatherStage2;
67 : u64 interAllGatherStage2;
68 : u64 interAllGatherStage3;
69 : u64 intraAllGatherStage3;
70 : };
71 : struct DataParameters {
72 : u64 dataOffset[2] = {0, 0}; // 每个part数据偏移
73 : std::vector<std::vector<u64>> sliceSize{2}; // 正常分块part每个阶段数据大小
74 : std::vector<std::vector<u64>> inputStride{2}; // 正常分块partInputStride大小
75 : std::vector<std::vector<u64>> scratchOffset{2}; // 每个分块scratchoffset
76 : std::vector<std::vector<u64>> tailSize{2}; // 尾片的整数倍数据大小
77 : };
78 :
79 : struct StageProcAlgPara {
80 : std::function<HcclResult(TempFuncs&, TemplateDataParams&, ResLinks&, std::vector<InsQuePtr>&)> part0FuncPtr;
81 : ResLinks part0links;
82 : std::vector<InsQuePtr> part0Que;
83 : std::function<HcclResult(TempFuncs&, TemplateDataParams&, ResLinks&, std::vector<InsQuePtr>&)> part1FuncPtr;
84 : ResLinks part1links;
85 : std::vector<InsQuePtr> part1Que;
86 : };
87 0 : HcclResult WrapPrepResLinks(const RankGraph* type, const LinkReq& linkReq, ResLinks& resLinks)
88 : {
89 0 : return PrepResLinks(myRank_, type, linkPriority_, linkReq, resLinks);
90 : }
91 0 : HcclResult WrapPrepResLinks(ConnectedLinkMgr* type, const LinkReq& linkReq, ResLinks& resLinks) const
92 : {
93 0 : return PrepResLinks(myRank_, linkReq, type, resLinks);
94 : }
95 : HcclResult PreCalcRes(
96 : const RankGraph* rankGraph, AlgTempResReq& resReqIntraScatter, AlgTempResReq& resReqInterScatter,
97 : AlgTempResReq& resReqIntraAllGather, AlgTempResReq& resReqInterAllGather);
98 : template <typename T>
99 : HcclResult CalcSingleAlgRes(
100 : InsAlgTemplate0& intraScatter, InsAlgTemplate1& interScatter, InsAlgTemplate2& intraAllGather,
101 : InsAlgTemplate3& interAllGather, T* type, AlgTempResReq& resReqIntraScatter, AlgTempResReq& resReqInterScatter,
102 : AlgTempResReq& resReqIntraAllGather, AlgTempResReq& resReqInterAllGather) const;
103 :
104 : template <typename T>
105 : HcclResult PrepareRes(
106 : T* type, AlgTempResReq& resReqIntraScatter, AlgTempResReq& resReqInterScatter,
107 : AlgTempResReq& resReqIntraAllGather, AlgTempResReq& resReqInterAllGather);
108 :
109 0 : HcclResult CalcLocalRankSize()
110 : {
111 0 : uint64_t virtRanks_2 = 2;
112 0 : CHK_PRT_RET(
113 : virtRanks_.size() < virtRanks_2, HCCL_ERROR("[CalcLocalRankSize] virtRanks level num is smaller than 2."),
114 : HcclResult::HCCL_E_INTERNAL);
115 :
116 0 : intraLocalRankSize_ = virtRanks_.at(0).size();
117 0 : interLocalRankSize_ = virtRanks_.at(1).size();
118 :
119 0 : HCCL_INFO(
120 : "[CalcLocalRankSize] localRankSize: myRank[%d] intraLocalRankSize[%u] interLocalRankSize[%u]", myRank_,
121 : intraLocalRankSize_, interLocalRankSize_);
122 0 : return HcclResult::HCCL_SUCCESS;
123 : };
124 0 : void GetParallelDataSplit(std::vector<double>& splitDataSize) const
125 : {
126 : // to do 先做等分,后续根据性能做调整
127 0 : double splitData = 0.5;
128 0 : splitDataSize.push_back(splitData);
129 0 : splitDataSize.push_back(splitData);
130 0 : return;
131 : }
132 0 : void InitDataParameters(SliceConfig& slice, ScratchMultiple& scratchMultiple, DataParameters& dataParameters) const
133 : {
134 0 : dataParameters.sliceSize.at(0) = {
135 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_,
136 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_,
137 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_,
138 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_};
139 0 : dataParameters.sliceSize.at(1) = {
140 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_,
141 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_,
142 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_,
143 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_};
144 0 : dataParameters.inputStride.at(0) = {
145 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_,
146 0 : slice.sliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_, 0, 0};
147 0 : dataParameters.inputStride.at(1) = {
148 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_,
149 0 : slice.sliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_, 0, 0};
150 : // 计算Scratch偏移,数据尾块必然小于常规块,不用额外计算尾块时的Scratch偏移
151 0 : dataParameters.scratchOffset.at(0) = {0, 0, 0, 0};
152 0 : dataParameters.scratchOffset.at(1) = {
153 0 : slice.sliceCountPart0 * scratchMultiple.interScatter * dataTypeSize_,
154 0 : (slice.sliceCountPart0 / interLocalRankSize_) * scratchMultiple.intraScatter * dataTypeSize_,
155 0 : (slice.sliceCountPart0 / interLocalRankSize_ / intraLocalRankSize_) * scratchMultiple.intraAllGather *
156 0 : dataTypeSize_,
157 0 : (slice.sliceCountPart0 / interLocalRankSize_) * scratchMultiple.interAllGather * dataTypeSize_};
158 0 : dataParameters.tailSize = dataParameters.sliceSize;
159 0 : return;
160 : }
161 0 : void InitFinalSliceDataParameters(
162 : SliceConfig& slice, ScratchMultiple& scratchMultiple, DataParameters& dataParameters) const
163 : {
164 0 : dataParameters.sliceSize.at(0) = {
165 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_,
166 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_,
167 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_,
168 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_};
169 0 : dataParameters.sliceSize.at(1) = {
170 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_,
171 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_,
172 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_,
173 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_};
174 0 : dataParameters.inputStride.at(0) = {
175 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_,
176 0 : slice.finalSliceCountPart0 * dataTypeSize_ / interLocalRankSize_ / intraLocalRankSize_, 0, 0};
177 0 : dataParameters.inputStride.at(1) = {
178 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_,
179 0 : slice.finalSliceCountPart1 * dataTypeSize_ / intraLocalRankSize_ / interLocalRankSize_, 0, 0};
180 0 : dataParameters.scratchOffset.at(0) = {0, 0, 0, 0};
181 0 : dataParameters.scratchOffset.at(1) = {
182 0 : slice.finalSliceCountPart0 * scratchMultiple.interScatter * dataTypeSize_,
183 0 : (slice.finalSliceCountPart0 / interLocalRankSize_) * scratchMultiple.intraScatter * dataTypeSize_,
184 0 : (slice.finalSliceCountPart0 / interLocalRankSize_ / intraLocalRankSize_) * scratchMultiple.intraAllGather *
185 0 : dataTypeSize_,
186 0 : (slice.finalSliceCountPart0 / interLocalRankSize_) * scratchMultiple.interAllGather * dataTypeSize_};
187 : // 只有最后一片数据的part1部分存在尾片数据,scatter算子和allgather算子都需要支持该数据收集
188 0 : for (size_t i = 0; i < dataParameters.sliceSize.at(0).size(); i++) {
189 0 : dataParameters.tailSize.at(0).at(i) =
190 0 : dataParameters.sliceSize.at(0).at(i) + slice.finalTailCountPart0 * dataTypeSize_;
191 : }
192 0 : for (size_t i = 0; i < dataParameters.sliceSize.at(1).size(); i++) {
193 0 : dataParameters.tailSize.at(1).at(i) =
194 0 : dataParameters.sliceSize.at(1).at(i) + slice.finalTailCountPart1 * dataTypeSize_;
195 : }
196 0 : return;
197 : }
198 0 : HcclResult CalcLocalRoot()
199 : {
200 0 : CHK_PRT_RET(
201 : root_ >= rankSize_, HCCL_ERROR("[CalcLocalRoot] root[%u] is out of rankSize[%u]", root_, rankSize_),
202 : HcclResult::HCCL_E_INTERNAL);
203 :
204 0 : u32 intraLocalRootIdx = root_ % intraLocalRankSize_;
205 0 : intraLocalRoot_ = static_cast<u32>(vTopo_.at(0).at(0).at(intraLocalRootIdx));
206 0 : u32 interLocalRootIdx = root_ / intraLocalRankSize_;
207 0 : interLocalRoot_ = static_cast<u32>(vTopo_.at(1).at(0).at(interLocalRootIdx));
208 :
209 0 : HCCL_INFO(
210 : "[CalcLocalRoot] localRoot: myRank[%d] intraLocalRoot[%u] interLocalRoot[%u]", myRank_, intraLocalRoot_,
211 : interLocalRoot_);
212 0 : return HcclResult::HCCL_SUCCESS;
213 : }
214 :
215 0 : void CalcScratchMultiple(
216 : std::vector<double>& splitDataSize, ScratchMultiple& scratchMultiple, InsAlgTemplate0& intraScatterTempAlg,
217 : InsAlgTemplate1& interScatterTempAlg, InsAlgTemplate2& intraAllGatherTempAlg,
218 : InsAlgTemplate3& interAllGatherTempAlg) const
219 : {
220 0 : scratchMultiple.intraScatter = intraScatterTempAlg.CalcScratchMultiple(BufferType::INPUT, BufferType::INPUT);
221 0 : scratchMultiple.interScatter = interScatterTempAlg.CalcScratchMultiple(BufferType::INPUT, BufferType::INPUT);
222 0 : scratchMultiple.intraAllGather =
223 0 : intraAllGatherTempAlg.CalcScratchMultiple(BufferType::INPUT, BufferType::INPUT);
224 0 : scratchMultiple.interAllGather =
225 0 : interAllGatherTempAlg.CalcScratchMultiple(BufferType::INPUT, BufferType::INPUT);
226 : // 计算第一步需要的倍数和最后一步所需要的数据缓存倍数,取multiple最大需求
227 0 : float multiple0 = splitDataSize.at(0) * float(scratchMultiple.interScatter) +
228 0 : splitDataSize.at(1) * float(scratchMultiple.intraScatter);
229 0 : float multiple1 = splitDataSize.at(0) * float(scratchMultiple.interAllGather / interLocalRankSize_) +
230 0 : splitDataSize.at(1) * float(scratchMultiple.intraAllGather / intraLocalRankSize_);
231 0 : scratchMultiple.maxMultiple = std::max(multiple0, multiple1);
232 0 : return;
233 : }
234 : void CalcSlice(std::vector<double>& splitDataSize, float scratchMaxMultiple, SliceConfig& slice);
235 0 : void LogAlgInfo(
236 : InsAlgTemplate0& intraScatterTempAlg, InsAlgTemplate1& interScatterTempAlg,
237 : InsAlgTemplate2& intraAllGatherTempAlg, InsAlgTemplate3& interAllGatherTempAlg) const
238 : {
239 0 : HCCL_INFO("[InsBroadcastParallelAiCpuExecutor] Alg0 is [%s]", intraScatterTempAlg.Describe().c_str());
240 0 : HCCL_INFO("[InsBroadcastParallelAiCpuExecutor] Alg1 is [%s]", interScatterTempAlg.Describe().c_str());
241 0 : HCCL_INFO("[InsBroadcastParallelAiCpuExecutor] Alg2 is [%s]", intraAllGatherTempAlg.Describe().c_str());
242 0 : HCCL_INFO("[InsBroadcastParallelAiCpuExecutor] Alg3 is [%s]", interAllGatherTempAlg.Describe().c_str());
243 0 : return;
244 : }
245 : HcclResult StageProcess(DataParameters& dataParameters, std::vector<StageProcAlgPara>& algParaVec);
246 0 : void AlgTemplateInitPara(
247 : const CollAlgOperator& op, InsAlgTemplate0& intraScatterTempAlg, InsAlgTemplate1& interScatterTempAlg,
248 : InsAlgTemplate2& intraAllGatherTempAlg, InsAlgTemplate3& interAllGatherTempAlg)
249 : {
250 0 : intraScatterTempAlg.SetDmaMode(dmaMode_);
251 0 : intraScatterTempAlg.SetCollOp(op);
252 0 : intraScatterTempAlg.SetDataType(dataType_);
253 0 : intraScatterTempAlg.SetRoot(intraLocalRoot_);
254 :
255 0 : interScatterTempAlg.SetDmaMode(dmaMode_);
256 0 : interScatterTempAlg.SetCollOp(op);
257 0 : interScatterTempAlg.SetDataType(dataType_);
258 0 : interScatterTempAlg.SetRoot(interLocalRoot_);
259 :
260 0 : intraAllGatherTempAlg.SetDmaMode(dmaMode_);
261 0 : intraAllGatherTempAlg.SetCollOp(op);
262 0 : intraAllGatherTempAlg.SetDataType(dataType_);
263 0 : intraAllGatherTempAlg.SetRoot(intraLocalRoot_);
264 :
265 0 : interAllGatherTempAlg.SetDmaMode(dmaMode_);
266 0 : interAllGatherTempAlg.SetCollOp(op);
267 0 : interAllGatherTempAlg.SetDataType(dataType_);
268 0 : interAllGatherTempAlg.SetRoot(intraLocalRoot_);
269 0 : return;
270 : }
271 : // Host
272 : HcclResult PrepareResForTemplate(
273 : const RankGraph* rankGraph, InsAlgTemplate0& intraScatterTempAlg, InsAlgTemplate1& interScatterTempAlg,
274 : InsAlgTemplate2& intraAllGatherTempAlg, InsAlgTemplate3& interAllGatherTempAlg);
275 : // Aicpu
276 : HcclResult PrepareResForTemplate(
277 : ConnectedLinkMgr* linkMgr, InsAlgTemplate0& intraScatterTempAlg, InsAlgTemplate1& interScatterTempAlg,
278 : InsAlgTemplate2& intraAllGatherTempAlg, InsAlgTemplate3& interAllGatherTempAlg);
279 :
280 0 : void GenDataParamsStage(
281 : const u32 part, const u32 stage, DataParameters& dataParameters, TemplateDataParams& dataParams) const
282 : {
283 0 : dataParams.buffInfo.inBuffType = BufferType::INPUT;
284 0 : dataParams.buffInfo.outBuffType = BufferType::INPUT;
285 0 : dataParams.buffInfo.scratBuffType = BufferType::SCRATCH;
286 0 : dataParams.buffInfo.inBuffBaseOff = dataParameters.dataOffset[part];
287 0 : dataParams.buffInfo.outBuffBaseOff = dataParameters.dataOffset[part];
288 0 : dataParams.buffInfo.scratchBuffBaseOff = dataParameters.scratchOffset.at(part).at(stage);
289 0 : dataParams.sliceSize = dataParameters.sliceSize.at(part).at(stage);
290 0 : dataParams.inputSliceStride = dataParameters.inputStride.at(part).at(stage);
291 0 : dataParams.outputSliceStride = dataParameters.sliceSize.at(part).at(stage);
292 0 : dataParams.repeatNum = 1;
293 0 : dataParams.inputRepeatStride = 0;
294 0 : dataParams.outputRepeatStride = 0;
295 0 : dataParams.tailSize = dataParameters.tailSize.at(part).at(stage);
296 0 : return;
297 : }
298 0 : void InitStageProcAlgParaVec(
299 : std::vector<StageProcAlgPara>& stageProcAlgParaVec, InsAlgTemplate0& intraScatterTempAlg,
300 : InsAlgTemplate1& interScatterTempAlg, InsAlgTemplate2& intraAllGatherTempAlg,
301 : InsAlgTemplate3& interAllGatherTempAlg)
302 : {
303 : // 以此输入第1个阶段的part0 GenExtIns, scratchOffset, links,que以及 part1部分对应信息
304 0 : stageProcAlgParaVec = {
305 0 : {[&](auto&... args) { return interScatterTempAlg.GenExtIns(args...); }, scatterInterLinks_,
306 0 : interQue_, // stage0 part0
307 0 : [&](auto&... args) { return intraScatterTempAlg.GenExtIns(args...); }, scatterIntraLinks_,
308 0 : intraQue_}, // stage0 part1
309 0 : {[&](auto&... args) { return intraScatterTempAlg.GenExtIns(args...); }, scatterIntraLinks_,
310 0 : intraQue_, // stage1 part0
311 0 : [&](auto&... args) { return interScatterTempAlg.GenExtIns(args...); }, scatterInterLinks_,
312 0 : interQue_}, // stage1 part1
313 0 : {[&](auto&... args) { return intraAllGatherTempAlg.GenExtIns(args...); }, allGatherIntraLinks_,
314 0 : intraQue_, // stage2 part0
315 0 : [&](auto&... args) { return interAllGatherTempAlg.GenExtIns(args...); }, allGatherInterLinks_,
316 0 : interQue_}, // stage2 part1
317 0 : {[&](auto&... args) { return interAllGatherTempAlg.GenExtIns(args...); }, allGatherInterLinks_,
318 0 : interQue_, // stage3 part0
319 0 : [&](auto&... args) { return intraAllGatherTempAlg.GenExtIns(args...); }, allGatherIntraLinks_,
320 0 : intraQue_}, // stage3 part1
321 : };
322 0 : }
323 : HcclResult GenInsQues(
324 : InsAlgTemplate0& intraScatterTempAlg, InsAlgTemplate1& interScatterTempAlg,
325 : InsAlgTemplate2& intraAllGatherTempAlg, InsAlgTemplate3& interAllGatherTempAlg);
326 :
327 : u32 intraLocalRankSize_{0}; // server内算法rankSize
328 : u32 interLocalRankSize_{0}; // server间算法rankSize
329 :
330 : RankId intraLocalRank_{INVALID_RANKID}; // server内算法rank
331 : RankId interLocalRank_{INVALID_RANKID}; // server间算法rank
332 :
333 : u32 intraLocalRoot_{0}; // server内算法root
334 : u32 interLocalRoot_{0}; // server间算法root
335 :
336 : const RankGraph* rankGraph_ = nullptr;
337 :
338 : std::vector<std::vector<std::vector<RankId>>> vTopo_;
339 : std::vector<std::vector<RankId>> virtRanks_;
340 : std::vector<std::map<RankId, u32>> virtRankMap_; // map<virtRank, virtRankOrder>
341 :
342 : std::vector<InsQuePtr> requiredQue_;
343 : std::vector<InsQuePtr> intraQue_;
344 : std::vector<InsQuePtr> interQue_;
345 : std::vector<InsQuePtr> syncQueues_;
346 : ResLinks scatterIntraLinks_;
347 : ResLinks scatterInterLinks_;
348 : ResLinks allGatherIntraLinks_;
349 : ResLinks allGatherInterLinks_;
350 :
351 : const RankGraph* rankGraphPtr_ = nullptr;
352 : };
353 :
354 : } // namespace Hccl
355 :
356 : #endif
|