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 "asymmetric_hierarchical_concatenate_alg_template_base.h"
12 :
13 : #include <iostream>
14 : #include <fstream>
15 :
16 : namespace hccl {
17 :
18 0 : AHCAlgTemplateBase::AHCAlgTemplateBase(const HcclDispatcher dispatcher)
19 : : AlgTemplateBase(dispatcher),
20 0 : needTraslateSliceAddr_(false),
21 0 : rankSize_(1),
22 0 : extendFlag_(false)
23 0 : {}
24 :
25 0 : AHCAlgTemplateBase::~AHCAlgTemplateBase() {}
26 :
27 0 : HcclResult AHCAlgTemplateBase::Prepare(u64 reduceAttrBitMap, [[maybe_unused]] HcomCollOpInfo* opInfo)
28 : {
29 0 : reduceAttr_ = reduceAttrBitMap;
30 0 : return HCCL_SUCCESS;
31 : }
32 :
33 0 : HcclResult AHCAlgTemplateBase::Prepare(
34 : u64 totalCount, const std::vector<std::vector<std::vector<u32>>>& globalSubGroups,
35 : std::map<AHCConcOpType, TemplateType>& ahcAlgOption, bool extendFlag, AHCExtendPreparePara extendPara)
36 : {
37 0 : globalSubGroups_ = globalSubGroups;
38 0 : totalCount_ = totalCount;
39 0 : ahcAlgOption_ = ahcAlgOption;
40 0 : extendFlag_ = extendFlag;
41 0 : ahcExtendPreparePara_ = extendPara;
42 0 : return HCCL_SUCCESS;
43 : }
44 :
45 0 : HcclResult AHCAlgTemplateBase::DisposeSubGroups([[maybe_unused]] const u32 rank) { return HCCL_SUCCESS; }
46 :
47 0 : HcclResult AHCAlgTemplateBase::CommAHCInfoInit() { return HCCL_SUCCESS; }
48 :
49 0 : HcclResult AHCAlgTemplateBase::PrepareRunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
50 : {
51 0 : HcclResult ret = HCCL_SUCCESS;
52 0 : CHK_SMART_PTR_NULL(dispatcher_);
53 0 : CHK_PTR_NULL(stream_.ptr());
54 0 : CHK_PRT_RET(
55 : !outputMem_ || !inputMem_,
56 : HCCL_ERROR("[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] run_async inputmem or outputmem is null", rank),
57 : HCCL_E_PTR);
58 :
59 0 : HCCL_INFO(
60 : "AHCAlgTemplateBase run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
61 : inputMem_.ptr(), outputMem_.ptr(), count_);
62 :
63 0 : CHK_PRT_RET(
64 : links.size() < rankSize,
65 : HCCL_ERROR(
66 : "[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] linksize[%llu] is less "
67 : "than rankSize[%u]",
68 : rank, links.size(), rankSize),
69 : HCCL_E_INTERNAL);
70 :
71 : // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
72 0 : if (rankSize == 1) {
73 0 : if (inputMem_ != outputMem_) {
74 0 : ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
75 0 : CHK_PRT_RET(
76 : ret != HCCL_SUCCESS,
77 : HCCL_ERROR("[AHCAlgTemplateBase][PrepareRunAsync]rank[%u] memcpy async failed", rank), ret);
78 : }
79 0 : return ret;
80 : }
81 :
82 0 : DisposeSubGroups(rank);
83 :
84 0 : rankSize_ = rankSize;
85 :
86 0 : CommAHCInfoInit();
87 :
88 : // 保存物理slice,可能非连续
89 0 : physicalSlices_ = slices_;
90 :
91 : // 检查、并清空逻辑slices_
92 0 : if (slices_.size() != 0) {
93 0 : HCCL_DEBUG("[AHCAlgTemplateBase][PrepareRunAsync] clear logic slice_");
94 0 : slices_.clear();
95 : }
96 :
97 0 : return HCCL_SUCCESS;
98 : }
99 :
100 0 : HcclResult AHCAlgTemplateBase::GetNslbAdjInfoPro(
101 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
102 : {
103 0 : HCCL_DEBUG("[NSLB-AHC] entry GetNslbAdjInfoPro");
104 0 : DisposeSubGroups(rank);
105 0 : CommAHCInfoInit();
106 :
107 0 : if (rankSize == 1 || links.size() < rankSize) {
108 0 : return HCCL_SUCCESS;
109 : }
110 0 : u32 nSteps = 0;
111 0 : std::vector<u32> dstRanks;
112 0 : HCCL_DEBUG("[NSLB-AHC] try to GetNslbDstRanks, rank = %u, ranksize = %u", rank, rankSize);
113 0 : CHK_RET(commAHCBaseInfo_->GetNslbDstRanks(rank, dstRanks));
114 0 : if (dstRanks.size() == 0 || dstRanks.size() > NSLBDP_MAX_PHASE) {
115 0 : HCCL_DEBUG("[NSLB-AHC] dstRanks size not support");
116 0 : return HCCL_SUCCESS;
117 : }
118 0 : for (u32 nextRank : dstRanks) {
119 0 : LINK linkRight = links[nextRank];
120 0 : CHK_SMART_PTR_NULL(linkRight);
121 0 : NslbDpAdjInfo adjInfoStep = {0, 0, 0};
122 0 : adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
123 0 : adjInfoStep.phaseId = nSteps + 1;
124 0 : adjInfoStep.rev = 0;
125 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
126 0 : nSteps++;
127 0 : }
128 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
129 0 : return HCCL_SUCCESS;
130 0 : }
131 :
132 0 : HcclResult AHCAlgTemplateBase::PrepareAlgTemplate(
133 : std::unique_ptr<AlgTemplateBase>& tempAlg, const std::vector<Slice>& slices, AHCOpType opType)
134 : {
135 0 : HcclResult ret = HCCL_SUCCESS;
136 0 : switch (opType) {
137 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER: {
138 0 : ret = tempAlg->Prepare(
139 0 : inputMem_, inputMem_, scratchMem_, count_, dataType_, stream_, reductionOp_, root_, slices,
140 : baseOffset_);
141 0 : break;
142 : }
143 0 : case AHCOpType::AHC_OP_TYPE_ALLGATHER: {
144 0 : ret = tempAlg->Prepare(
145 0 : outputMem_, outputMem_, scratchMem_, count_, dataType_, stream_, reductionOp_, root_, slices,
146 : baseOffset_);
147 0 : break;
148 : }
149 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE: {
150 0 : ret = tempAlg->Prepare(
151 0 : inputMem_, outputMem_, scratchMem_, count_, dataType_, stream_, reductionOp_, root_, slices,
152 : baseOffset_);
153 0 : break;
154 : }
155 0 : case AHCOpType::AHC_OP_TYPE_RESERVED: {
156 : // 其他算子不支持,直接返回
157 0 : ret = HCCL_E_PARA;
158 0 : break;
159 : }
160 : }
161 :
162 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AHCAlgTemplateBase][PrepareAlgTemplate] prepare step error"), ret);
163 :
164 0 : return ret;
165 : }
166 :
167 0 : HcclResult AHCAlgTemplateBase::MemcpyForSingleOp(const u32 rank, AHCOpType opType)
168 : {
169 0 : HcclResult ret = HCCL_SUCCESS;
170 0 : u32 commRank = commAHCBaseInfo_->GetCommRank(rank);
171 0 : HCCL_DEBUG("[AHCAlgTemplateBase][MemcpyForSingleOp] rank[%u] commRank[%u]", rank, commRank);
172 0 : switch (opType) {
173 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER: {
174 : u64 srcSize
175 0 : = (inputMem_.size() - commRank * count_ * DataUnitSize(dataType_)) > count_ * DataUnitSize(dataType_) ?
176 0 : count_ * DataUnitSize(dataType_) :
177 0 : (inputMem_.size() - commRank * count_ * DataUnitSize(dataType_));
178 0 : DeviceMem srcMem = inputMem_.range(commRank * count_ * DataUnitSize(dataType_), srcSize);
179 0 : ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, srcMem, stream_);
180 0 : break;
181 0 : }
182 0 : case AHCOpType::AHC_OP_TYPE_ALLGATHER: {
183 : u64 dstSize
184 0 : = (outputMem_.size() - commRank * count_ * DataUnitSize(dataType_)) > count_ * DataUnitSize(dataType_) ?
185 0 : count_ * DataUnitSize(dataType_) :
186 0 : (outputMem_.size() - commRank * count_ * DataUnitSize(dataType_));
187 0 : DeviceMem dstMem = outputMem_.range(commRank * count_ * DataUnitSize(dataType_), dstSize);
188 0 : ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, inputMem_, stream_);
189 0 : break;
190 0 : }
191 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE: {
192 : // 使用 RS+AG 实现 AR 时,需要在 RS 完成时进行一次额外的数据搬运
193 0 : ret = HcclD2DMemcpyAsync(dispatcher_, inputMem_, outputMem_, stream_);
194 0 : break;
195 : }
196 0 : case AHCOpType::AHC_OP_TYPE_RESERVED: {
197 : // 其他算子不支持,无需copy,直接返回
198 0 : break;
199 : }
200 : }
201 0 : return ret;
202 : }
203 :
204 0 : HcclResult AHCAlgTemplateBase::RunInstance(
205 : const u32 rank, const std::vector<LINK>& links, std::vector<Slice>& slices,
206 : std::unique_ptr<AlgTemplateBase>& tempAlg, AHCOpType opType)
207 : {
208 0 : HcclResult ret = HCCL_SUCCESS;
209 :
210 : // 判断是否关闭reducescatter的barrier
211 0 : if (!barrierSwitchOn_) {
212 0 : tempAlg->CloseBarrier();
213 : }
214 :
215 : // 地址映射
216 0 : if (needTraslateSliceAddr_) {
217 0 : ret = commAHCBaseInfo_->TrasLogicSliceToPhysical(slices, physicalSlices_);
218 0 : CHK_PRT_RET(
219 : ret != HCCL_SUCCESS,
220 : HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] optype[%d] translate slice failed", rank, opType),
221 : ret);
222 : }
223 :
224 : // 调用算法执行
225 0 : ret = PrepareAlgTemplate(tempAlg, slices, opType);
226 0 : CHK_PRT_RET(
227 : ret != HCCL_SUCCESS,
228 : HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] prepare optype[%d] failed", rank, opType), ret);
229 :
230 0 : ret = tempAlg->RegisterProfiler(profilerInput_.planeID, profilerInput_.stage, profilerInput_.step, stream_);
231 0 : CHK_PRT_RET(
232 : ret != HCCL_SUCCESS,
233 : HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] registerProfiler optype[%d] failed", rank, opType), ret);
234 :
235 0 : ret = tempAlg->RunAsync(rank, links.size(), links);
236 0 : CHK_PRT_RET(
237 : ret != HCCL_SUCCESS,
238 : HCCL_ERROR("[AHCAlgTemplateBase][RunInstance]rank[%u] run optype[%d] failed", rank, opType), ret);
239 :
240 0 : return ret;
241 : }
242 :
243 0 : ReduceScatterAHCBase::ReduceScatterAHCBase(const HcclDispatcher dispatcher) : AHCAlgTemplateBase(dispatcher) {}
244 :
245 0 : ReduceScatterAHCBase::~ReduceScatterAHCBase() {}
246 :
247 0 : HcclResult ReduceScatterAHCBase::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
248 : {
249 0 : HCCL_INFO("[ReduceScatterAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
250 :
251 0 : HcclResult ret = HCCL_SUCCESS;
252 0 : ret = PrepareRunAsync(rank, rankSize, links);
253 0 : HCCL_DEBUG(
254 : "[ReduceScatterAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(),
255 : outputMem_.size(), count_);
256 :
257 : // 设置地址翻译标记有效,并计算逻辑的totalsize
258 0 : needTraslateSliceAddr_ = true;
259 0 : commAHCBaseInfo_->ParseInputSlice(physicalSlices_);
260 :
261 0 : CHK_PRT_RET(
262 : ret != HCCL_SUCCESS,
263 : HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_),
264 : ret);
265 :
266 0 : CHK_PRT_RET(
267 : rankSize == 1, HCCL_INFO("[ReduceScatterAHCBase][RunAsync] rankSize[%u], do nothing.", rankSize), HCCL_SUCCESS);
268 :
269 0 : HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] begin intra rs", rank);
270 :
271 : // 做组内 reduce-scatter
272 0 : ret = RunIntraReduceScatter(rank, links, commAHCBaseInfo_);
273 0 : CHK_PRT_RET(
274 : ret != HCCL_SUCCESS,
275 : HCCL_ERROR(
276 : "[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in "
277 : "RunIntraReduceScatter step",
278 : rank, count_),
279 : ret);
280 :
281 0 : HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] end intra rs begin inter", rank);
282 :
283 : // 做组间 reduce-scatter
284 0 : ret = RunInterReduceScatter(rank, links, commAHCBaseInfo_);
285 0 : CHK_PRT_RET(
286 : ret != HCCL_SUCCESS,
287 : HCCL_ERROR(
288 : "[ReduceScatterAHCBase][RunAsync]rank[%u] count[%llu] failed in "
289 : "RunInterReduceScatter step",
290 : rank, count_),
291 : ret);
292 :
293 : // 对于单独的 Reduce-scatter 算子,在运算结束时进行数据搬运
294 0 : if (inputMem_ != outputMem_) {
295 0 : ret = MemcpyForSingleOp(rank, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER);
296 0 : CHK_PRT_RET(
297 : ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterAHCBase][RunAsync]rank[%u] memcpy failed", rank), ret);
298 : }
299 :
300 0 : HCCL_DEBUG("[ReduceScatterAHCBase][RunAsync] rank[%u] end inter rs", rank);
301 :
302 0 : HCCL_INFO("[ReduceScatterAHCBase][RunAsync] finished: rank[%u]", rank);
303 0 : return HCCL_SUCCESS;
304 : }
305 :
306 0 : HcclResult ReduceScatterAHCBase::GetNslbAdjInfo(
307 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
308 : {
309 0 : return GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo);
310 : }
311 :
312 0 : HcclResult ReduceScatterAHCBase::RunIntraReduceScatter(
313 : const u32 rank, const std::vector<LINK>& links, const std::unique_ptr<CommAHCBaseInfo>& commAHCBaseInfo)
314 : {
315 : // 获取当前rank的组内rank
316 0 : HcclResult ret = HCCL_SUCCESS;
317 0 : HCCL_INFO("[ReduceScatterAHC][RunIntraReduceScatter] begin intra ReduceScatter rank[%u] count[%llu]", rank, count_);
318 :
319 0 : u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
320 :
321 : // 创建执行算子实列
322 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
323 0 : commAHCBaseInfo->GetIntraAlgTemplateOpInstance(
324 0 : AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER, tempAlg, dispatcher_, reduceAttr_, extendFlag_, ahcExtendPreparePara_);
325 :
326 0 : std::vector<std::vector<Slice>> intraSlicesVector;
327 0 : std::vector<std::vector<LINK>> intraLinksVector;
328 0 : CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(
329 : rank, DataUnitSize(dataType_), count_, links, intraLinksVector, intraSlicesVector));
330 :
331 0 : HCCL_DEBUG("[ReduceScatterAHCBase][RunIntraReduceScatter] run inst rank[%u] intraRank[%u]", rank, intraRank);
332 :
333 0 : for (u32 i = 0; i < intraLinksVector.size(); i++) {
334 0 : std::vector<Slice> intraSlices = intraSlicesVector[i];
335 0 : std::vector<LINK> intraLinks = intraLinksVector[i];
336 0 : if (intraLinks.size() <= 1) {
337 0 : continue;
338 : }
339 0 : CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER));
340 0 : }
341 :
342 0 : HCCL_DEBUG("[ReduceScatterAHCBase][RunIntraReduceScatter] end intra ReduceScatter rank[%u]", rank);
343 :
344 0 : return ret;
345 0 : }
346 :
347 0 : AllGatherAHCBase::AllGatherAHCBase(const HcclDispatcher dispatcher) : AHCAlgTemplateBase(dispatcher) {}
348 :
349 0 : AllGatherAHCBase::~AllGatherAHCBase() {}
350 :
351 0 : HcclResult AllGatherAHCBase::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
352 : {
353 0 : HCCL_INFO("[AllGatherAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
354 :
355 0 : HcclResult ret = HCCL_SUCCESS;
356 0 : ret = PrepareRunAsync(rank, rankSize, links);
357 0 : HCCL_DEBUG(
358 : "[AllGatherAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(),
359 : outputMem_.size(), count_);
360 :
361 : // 设置地址翻译标记有效,并计算逻辑的totalsize
362 0 : needTraslateSliceAddr_ = true;
363 0 : commAHCBaseInfo_->ParseInputSlice(physicalSlices_);
364 :
365 0 : CHK_PRT_RET(
366 : ret != HCCL_SUCCESS,
367 : HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_),
368 : ret);
369 :
370 0 : CHK_PRT_RET(
371 : rankSize == 1, HCCL_INFO("[AllGatherAHCBase][RunAsync] rankSize[%u], do nothing.", rankSize), HCCL_SUCCESS);
372 :
373 0 : HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] begin intra ag", rank);
374 :
375 : // 对于单独的 All-gather 算子,在运算开始时进行数据搬运
376 0 : if (inputMem_ != outputMem_) {
377 0 : ret = MemcpyForSingleOp(rank, AHCOpType::AHC_OP_TYPE_ALLGATHER);
378 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherAHCBase][RunAsync]rank[%u] memcpy failed", rank), ret);
379 : }
380 :
381 : // 做组间 all-gather
382 0 : ret = RunInterAllGather(rank, links, commAHCBaseInfo_);
383 0 : CHK_PRT_RET(
384 : ret != HCCL_SUCCESS,
385 : HCCL_ERROR(
386 : "[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in "
387 : "RunInterAllGather step",
388 : rank, count_),
389 : ret);
390 :
391 0 : HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] end inter ag", rank);
392 :
393 : // 做组内 allgather
394 0 : ret = RunIntraAllGather(rank, links, commAHCBaseInfo_);
395 0 : CHK_PRT_RET(
396 : ret != HCCL_SUCCESS,
397 : HCCL_ERROR(
398 : "[AllGatherAHCBase][RunAsync]rank[%u] count[%llu] failed in "
399 : "RunIntraAllGather step",
400 : rank, count_),
401 : ret);
402 :
403 0 : HCCL_DEBUG("[AllGatherAHCBase][RunAsync] rank[%u] end intra ag begin inter", rank);
404 :
405 0 : HCCL_INFO("[AllGatherAHCBase][RunAsync] finished: rank[%u]", rank);
406 0 : return HCCL_SUCCESS;
407 : }
408 :
409 0 : HcclResult AllGatherAHCBase::GetNslbAdjInfo(
410 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
411 : {
412 0 : return GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo);
413 : }
414 :
415 0 : HcclResult AllGatherAHCBase::RunIntraAllGather(
416 : const u32 rank, const std::vector<LINK>& links, const std::unique_ptr<CommAHCBaseInfo>& commAHCBaseInfo)
417 : {
418 : // 获取当前rank的组内rank
419 0 : HCCL_INFO("[AllGatherAHCBase][RunIntraAllGather] begin intra AllGather rank[%u]", rank);
420 :
421 0 : u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
422 :
423 : // 创建执行算子实列
424 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
425 0 : commAHCBaseInfo->GetIntraAlgTemplateOpInstance(
426 0 : AHCOpType::AHC_OP_TYPE_ALLGATHER, tempAlg, dispatcher_, reduceAttr_, extendFlag_, ahcExtendPreparePara_);
427 :
428 0 : std::vector<std::vector<Slice>> intraSlicesVector;
429 0 : std::vector<std::vector<LINK>> intraLinksVector;
430 0 : CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(
431 : rank, DataUnitSize(dataType_), count_, links, intraLinksVector, intraSlicesVector));
432 :
433 0 : HCCL_DEBUG("[AllGatherAHCBase][RunIntraAllGather] run inst rank[%u] intraRank[%u]", rank, intraRank);
434 :
435 0 : for (u32 i = 0; i < intraLinksVector.size(); i++) {
436 0 : std::vector<Slice> intraSlices = intraSlicesVector[i];
437 0 : std::vector<LINK> intraLinks = intraLinksVector[i];
438 0 : if (intraLinks.size() <= 1) {
439 0 : continue;
440 : }
441 0 : CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_ALLGATHER));
442 0 : }
443 :
444 0 : HCCL_DEBUG("[AllGatherAHCBase][RunIntraAllGather] end intra AllGather rank[%u]", rank);
445 :
446 0 : return HCCL_SUCCESS;
447 0 : }
448 :
449 0 : AllReduceAHCBase::AllReduceAHCBase(const HcclDispatcher dispatcher) : AHCAlgTemplateBase(dispatcher) {}
450 :
451 0 : AllReduceAHCBase::~AllReduceAHCBase() {}
452 :
453 0 : HcclResult AllReduceAHCBase::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
454 : {
455 0 : HCCL_INFO("[AllReduceAHCBase][RunAsync] start rank[%u] rankSize[%u]", rank, rankSize);
456 :
457 0 : HcclResult ret = HCCL_SUCCESS;
458 0 : ret = PrepareRunAsync(rank, rankSize, links);
459 0 : HCCL_DEBUG(
460 : "[AllReduceAHCBase][RunAsync] inputmem.size[%llu] outputmem.size[%llu] count[%llu]", inputMem_.size(),
461 : outputMem_.size(), count_);
462 :
463 0 : CHK_PRT_RET(
464 : ret != HCCL_SUCCESS,
465 : HCCL_ERROR("[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in PrepareRunAsync step", rank, count_),
466 : ret);
467 :
468 0 : CHK_PRT_RET(
469 : rankSize == 1, HCCL_INFO("[AllReduceAHCBase][RunAsync] rankSize[%u], do nothing.", rankSize), HCCL_SUCCESS);
470 :
471 0 : CHK_PRT_RET(count_ == 0, HCCL_INFO("[AllReduceAHCBase][RunAsync] count_[%llu], do nothing.", count_), HCCL_SUCCESS);
472 :
473 0 : HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] begin intra rs", rank);
474 :
475 0 : ret = RunIntraReduceScatter(rank, links, commAHCBaseInfo_);
476 0 : CHK_PRT_RET(
477 : ret != HCCL_SUCCESS,
478 : HCCL_ERROR(
479 : "[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "
480 : "RunIntraReduceScatter step",
481 : rank, count_),
482 : ret);
483 :
484 0 : HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end intra rs begin inter", rank);
485 :
486 : // 垂直方向做allreduce ring
487 0 : ret = RunInterAllReduce(rank, links, commAHCBaseInfo_);
488 0 : CHK_PRT_RET(
489 : ret != HCCL_SUCCESS,
490 : HCCL_ERROR(
491 : "[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "
492 : "RunInterAllReduce step",
493 : rank, count_),
494 : ret);
495 :
496 0 : HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end inter begin intra ag", rank);
497 :
498 : // 水平方向做broken allgather ring
499 0 : ret = RunIntraAllGather(rank, links, commAHCBaseInfo_);
500 0 : CHK_PRT_RET(
501 : ret != HCCL_SUCCESS,
502 : HCCL_ERROR(
503 : "[AllReduceAHCBase][RunAsync]rank[%u] count[%llu] failed in "
504 : "RunIntraAllGather step",
505 : rank, count_),
506 : ret);
507 :
508 0 : HCCL_DEBUG("[AllReduceAHCBase][RunAsync] rank[%u] end intra ag", rank);
509 :
510 0 : HCCL_INFO("[AllReduceAHCBase][RunAsync] finished: rank[%u]", rank);
511 0 : return HCCL_SUCCESS;
512 : }
513 :
514 0 : HcclResult AllReduceAHCBase::GetNslbAdjInfo(
515 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
516 : {
517 : // 获取reducescatter的部分
518 0 : CHK_RET(GetNslbAdjInfoPro(rank, rankSize, links, nslbAdjInfo));
519 : // 后续模拟all_gather的部分
520 0 : HCCL_INFO("[NSLB-AHC]try to get allgather part");
521 0 : if (nslbAdjInfo.dstRankNum == 0 || nslbAdjInfo.nsAdjInfo.size() == 0) {
522 0 : HCCL_INFO("[NSLB-AHC] get reducescatter part is null");
523 0 : return HCCL_SUCCESS;
524 : }
525 0 : uint16_t nsteps = nslbAdjInfo.nsAdjInfo.size();
526 :
527 0 : for (size_t index = 0; index < nsteps; index++) {
528 0 : NslbDpAdjInfo adjInfoStep = {0, 0, 0};
529 0 : adjInfoStep.dstLocalRankId = nslbAdjInfo.nsAdjInfo[nsteps - index - 1].dstLocalRankId;
530 0 : adjInfoStep.phaseId = nsteps + index + 1;
531 0 : adjInfoStep.rev = 0;
532 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
533 0 : nsteps++;
534 : }
535 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
536 0 : HCCL_INFO("[NSLB-AHC]success to get allgather part");
537 0 : return HCCL_SUCCESS;
538 : }
539 :
540 0 : HcclResult AllReduceAHCBase::RunIntraReduceScatter(
541 : const u32 rank, const std::vector<LINK>& links, const std::unique_ptr<CommAHCBaseInfo>& commAHCBaseInfo)
542 : {
543 : // 获取当前rank的组内rank
544 0 : HCCL_INFO("[AllReduceAHCBase][RunIntraReduceScatter] begin intra ReduceScatter rank[%u]", rank);
545 :
546 0 : u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
547 :
548 : // 创建执行算子实列
549 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
550 0 : commAHCBaseInfo->GetIntraAlgTemplateOpInstance(
551 0 : AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER, tempAlg, dispatcher_, reduceAttr_, extendFlag_, ahcExtendPreparePara_);
552 :
553 0 : std::vector<Slice> intraSlices;
554 0 : std::vector<LINK> intraLinks;
555 0 : CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(
556 : rank, DataUnitSize(dataType_), count_, links, intraLinks, intraSlices));
557 :
558 : // 长度不足2,直接跳过
559 0 : if (intraLinks.size() <= 1) {
560 0 : return HCCL_SUCCESS;
561 : }
562 :
563 0 : HCCL_DEBUG(
564 : "[AllReduceAHCBase][RunIntraReduceScatter] run inst rank[%u] intraRank[%u], IntraSize=%u", rank, intraRank,
565 : intraLinks.size());
566 :
567 0 : CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER));
568 :
569 0 : HCCL_DEBUG("[AllReduceAHCBase][RunIntraReduceScatter] end intra ReduceScatter rank[%u]", rank);
570 :
571 0 : return HCCL_SUCCESS;
572 0 : }
573 :
574 0 : HcclResult AllReduceAHCBase::RunIntraAllGather(
575 : const u32 rank, const std::vector<LINK>& links, const std::unique_ptr<CommAHCBaseInfo>& commAHCBaseInfo)
576 : {
577 0 : HCCL_INFO("[AllReduceAHCBase][RunIntraAllGather] begin intra allgather rank[%u]", rank);
578 :
579 : // 获取当前rank的组内rank
580 0 : u32 intraRank = commAHCBaseInfo->GetIntraRank(rank);
581 :
582 : // 创建执行算子实列
583 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
584 0 : commAHCBaseInfo->GetIntraAlgTemplateOpInstance(
585 0 : AHCOpType::AHC_OP_TYPE_ALLGATHER, tempAlg, dispatcher_, reduceAttr_, extendFlag_, ahcExtendPreparePara_);
586 :
587 0 : std::vector<Slice> intraSlices;
588 0 : std::vector<LINK> intraLinks;
589 :
590 0 : CHK_RET(commAHCBaseInfo->CalcIntraSlicesAndLinks(
591 : rank, DataUnitSize(dataType_), count_, links, intraLinks, intraSlices));
592 :
593 : // 长度不足2,直接跳过
594 0 : if (intraLinks.size() <= 1) {
595 0 : return HCCL_SUCCESS;
596 : }
597 :
598 0 : HCCL_DEBUG(
599 : "[AllReduceAHCBase][RunIntraAllGather] run inst rank[%u] intraRank[%u], IntraSize=%u", rank, intraRank,
600 : intraLinks.size());
601 :
602 0 : CHK_RET(RunInstance(intraRank, intraLinks, intraSlices, tempAlg, AHCOpType::AHC_OP_TYPE_ALLGATHER));
603 :
604 0 : HCCL_DEBUG("[AllReduceAHCBase][RunIntraAllGather] end intra allgather rank[%u]", rank);
605 0 : return HCCL_SUCCESS;
606 0 : }
607 :
608 : } // namespace hccl
|