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