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