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 "all_gather_recursive_hd.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : AllGatherRecursiveHalvingDoubling::AllGatherRecursiveHalvingDoubling(const HcclDispatcher dispatcher)
16 0 : : RecursiveHalvingDoublingBase(dispatcher)
17 0 : {}
18 :
19 0 : AllGatherRecursiveHalvingDoubling::~AllGatherRecursiveHalvingDoubling() {}
20 :
21 : // 服务器间allreduce的入口函数
22 0 : HcclResult AllGatherRecursiveHalvingDoubling::RunAsync(
23 : const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
24 : {
25 0 : CHK_SMART_PTR_NULL(dispatcher_);
26 0 : CHK_PTR_NULL(stream_.ptr());
27 0 : HCCL_INFO(
28 : "AllGatherRecursiveHalvingDoubling run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank,
29 : rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
30 :
31 0 : HcclResult ret = HCCL_SUCCESS;
32 :
33 0 : if (rankSize == 1) {
34 0 : if (inputMem_ != outputMem_) {
35 0 : ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
36 : }
37 0 : return ret;
38 : }
39 :
40 0 : if (links.size() < rankSize) {
41 0 : HCCL_ERROR(
42 : "[AllGatherRecursiveHalvingDoubling][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank,
43 : links.size(), rankSize);
44 0 : return HCCL_E_INTERNAL;
45 : }
46 :
47 0 : ret = CalcPartOneSizeAndBlockSize(rankSize);
48 0 : CHK_PRT_RET(
49 : ret != HCCL_SUCCESS,
50 : HCCL_ERROR(
51 : "[AllGatherRecursiveHalvingDoubling][RunAsync]Calculate Par1Size[%u] "
52 : "And BlockSize[%u] Failed! rankSize[%u]",
53 : part1Size_, blockSize_, rankSize),
54 : ret);
55 :
56 0 : ret = CalculateSlices(dataBytes_, rankSize);
57 0 : CHK_PRT_RET(
58 : ret != HCCL_SUCCESS,
59 : HCCL_ERROR(
60 : "[AllGatherRecursiveHalvingDoubling][RunAsync]Calculate slices failed, "
61 : "dataBytes[%llu], rankSize[%u]",
62 : dataBytes_, rankSize),
63 : ret);
64 :
65 0 : CHK_RET(GatherInPartOneToEven(rank, links));
66 :
67 0 : CHK_RET(AllGatherInBlock(rank, rankSize, links));
68 :
69 0 : CHK_RET(GatherInPartOneToOdd(rank, links));
70 :
71 0 : HCCL_INFO("AllGatherRecursiveHalvingDoubling finished: rank[%u] finished", rank);
72 0 : return HCCL_SUCCESS;
73 : }
74 :
75 0 : HcclResult AllGatherRecursiveHalvingDoubling::CalculateSlices(u64 dataBytes, const u32 rankSize) const
76 : {
77 0 : slices_.resize(blockSize_);
78 0 : u64 bytesPerSlice = dataBytes;
79 0 : u64 totalBytes = dataBytes * rankSize;
80 0 : u64 bytesLeft = totalBytes;
81 0 : u32 i = 0;
82 0 : while (bytesLeft > 0 && i < part1Size_ / 2) { // 除2计算part1在做完操作后block内slice数
83 0 : slices_[i].size = 2 * bytesPerSlice < bytesLeft ? 2 * bytesPerSlice : bytesLeft; // 乘2表示slice为part2两倍
84 0 : slices_[i].offset = totalBytes - bytesLeft;
85 0 : bytesLeft -= slices_[i].size;
86 0 : i++;
87 : }
88 :
89 0 : while (bytesLeft > 0) {
90 0 : slices_[i].size = bytesPerSlice < bytesLeft ? bytesPerSlice : bytesLeft;
91 0 : slices_[i].offset = totalBytes - bytesLeft;
92 0 : bytesLeft -= slices_[i].size;
93 0 : i++;
94 : }
95 0 : return HCCL_SUCCESS;
96 : }
97 :
98 0 : HcclResult AllGatherRecursiveHalvingDoubling::GatherInPartOneToEven(u32 rank, const std::vector<LINK>& links)
99 : {
100 0 : if (rank < part1Size_ && rank % 2 == 0) { // 模2判断奇偶性,从下一个rank的output收数据到output
101 0 : u32 peerRank = rank + 1; // 加1计算下一个rank号
102 0 : if (peerRank < links.size()) {
103 0 : CHK_SMART_PTR_NULL(links[peerRank]);
104 :
105 0 : HcclResult ret = links[peerRank]->TxAck(stream_);
106 0 : CHK_PRT_RET(
107 : ret != HCCL_SUCCESS,
108 : HCCL_ERROR("[Gather][InPartOneToEven]rank[%u] tx ack from peerank[%u] failed", rank, peerRank), ret);
109 0 : ret = links[peerRank]->RxAck(stream_);
110 0 : CHK_PRT_RET(
111 : ret != HCCL_SUCCESS,
112 : HCCL_ERROR("[Gather][InPartOneToEven]rank[%u] rx ack from peerank[%u] failed", rank, peerRank), ret);
113 0 : DeviceMem gatherOutputMem = outputMem_.range(dataBytes_ * rank, dataBytes_);
114 : // 接收数据到本端的 output
115 0 : HCCL_DEBUG(
116 : "rank[%u] outputMem[%p] receive from PeerRank[%u] outputMem, Offset[%llu], Size[%llu]", rank,
117 : gatherOutputMem.ptr(), peerRank, baseOffset_ + dataBytes_ * rank, gatherOutputMem.size());
118 :
119 0 : ret = ExecuteRxSync(
120 0 : links[peerRank], UserMemType::OUTPUT_MEM, dataBytes_ * rank, gatherOutputMem.ptr(), dataBytes_,
121 0 : stream_);
122 0 : CHK_PRT_RET(
123 : ret != HCCL_SUCCESS,
124 : HCCL_ERROR(
125 : "[Gather][InPartOneToEven]rank[%u] rx sync from PeerRank[%u] "
126 : "failed",
127 : rank, peerRank),
128 : ret);
129 0 : ret = links[peerRank]->RxWaitDone(stream_);
130 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Gather][InPartOneToEven]RxWaitDone failed"), ret);
131 0 : }
132 0 : } else if (rank < part1Size_ && rank % 2 == 1) { // 模2判断奇偶性,向上一个rank发送数据
133 0 : u32 peerRank = rank - 1; // 减1计算上一个rank号
134 : // 发送到对端的output
135 0 : if (peerRank < links.size()) {
136 0 : CHK_SMART_PTR_NULL(links[peerRank]);
137 0 : HcclResult ret = links[peerRank]->TxAck(stream_);
138 0 : CHK_PRT_RET(
139 : ret != HCCL_SUCCESS,
140 : HCCL_ERROR("[Gather][InPartOneToEven]rank[%u] tx ack from peerank[%u] failed", rank, peerRank), ret);
141 0 : ret = links[peerRank]->RxAck(stream_);
142 0 : HCCL_DEBUG("[AllGatherRecursiveHalvingDoubling][GatherInPartOneToEven]peerRank is %u", peerRank);
143 : // 等待对端可以接收数据
144 0 : CHK_PRT_RET(
145 : ret != HCCL_SUCCESS,
146 : HCCL_ERROR("[Gather][InPartOneToEven]rank[%u] rx ack from peerank[%u] failed", rank, peerRank), ret);
147 : // 设置gather的发送内存范围
148 0 : DeviceMem gatherOutputMem = outputMem_.range(dataBytes_ * rank, dataBytes_);
149 : // 发送数据到对端的 output
150 0 : HCCL_DEBUG(
151 : "rank[%u] outputMem[%p] sends to PeerRank[%u] outputMem, Offset[%llu], Size[%llu]", rank,
152 : gatherOutputMem.ptr(), peerRank, baseOffset_ + dataBytes_ * rank, gatherOutputMem.size());
153 :
154 0 : ret = ExecuteTxSync(
155 0 : links[peerRank], UserMemType::OUTPUT_MEM, dataBytes_ * rank, gatherOutputMem.ptr(), dataBytes_,
156 0 : stream_);
157 0 : CHK_PRT_RET(
158 : ret != HCCL_SUCCESS,
159 : HCCL_ERROR("[Gather][InPartOneToEven]rank[%u] tx sync to PeerRank[%u] failed", rank, peerRank), ret);
160 0 : ret = links[peerRank]->TxWaitDone(stream_);
161 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Gather][InPartOneToEven]TxWaitDone failed"), ret);
162 0 : }
163 : }
164 0 : return HCCL_SUCCESS;
165 : }
166 :
167 0 : HcclResult AllGatherRecursiveHalvingDoubling::GatherInPartOneToOdd(u32 rank, const std::vector<LINK>& links)
168 : {
169 0 : if (rank < part1Size_ && rank % 2 == 0) { // 模2判断奇偶性,向下一个rank发送数据
170 0 : u32 peerRank = rank + 1; // 加1计算下一个rank号
171 : // 发送到对端的output
172 0 : if (peerRank < links.size()) {
173 0 : CHK_SMART_PTR_NULL(links[peerRank]);
174 0 : HcclResult ret = links[peerRank]->TxAck(stream_);
175 0 : CHK_PRT_RET(
176 : ret != HCCL_SUCCESS,
177 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] tx ack from peerank[%u] failed.", rank, peerRank), ret);
178 0 : ret = links[peerRank]->RxAck(stream_);
179 : // 等待对端可以接收数据
180 0 : CHK_PRT_RET(
181 : ret != HCCL_SUCCESS,
182 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] rx ack from peerank[%u] failed", rank, peerRank), ret);
183 :
184 0 : HCCL_DEBUG(
185 : "rank[%u] outputMem[%p] sends to PeerRank[%u] outputMem, Offset[%llu], Size[%llu]", rank,
186 : outputMem_.ptr(), peerRank, baseOffset_, outputMem_.size());
187 0 : ret = ExecuteTxSync(
188 0 : links[peerRank], UserMemType::OUTPUT_MEM, baseOffset_, outputMem_.ptr(), outputMem_.size(), stream_);
189 0 : CHK_PRT_RET(
190 : ret != HCCL_SUCCESS,
191 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] tx sync to PeerRank[%u] failed", rank, peerRank), ret);
192 0 : ret = links[peerRank]->TxWaitDone(stream_);
193 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Gather][InPartOneToOdd]TxWaitDone failed"), ret);
194 : }
195 0 : } else if (rank < part1Size_ && rank % 2 == 1) { // 模2判断奇偶性,从上一个rank的output收数据到output
196 0 : u32 peerRank = rank - 1; // 减1计算上一个rank号
197 0 : if (peerRank < links.size()) {
198 0 : CHK_SMART_PTR_NULL(links[peerRank]);
199 : // 知会对端本人可以接收数据
200 0 : HcclResult ret = links[peerRank]->TxAck(stream_);
201 : // 等待对端可以接收数据
202 0 : CHK_PRT_RET(
203 : ret != HCCL_SUCCESS,
204 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] tx ack from peerank[%u] failed", rank, peerRank), ret);
205 0 : ret = links[peerRank]->RxAck(stream_);
206 0 : CHK_PRT_RET(
207 : ret != HCCL_SUCCESS,
208 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] rx ack from peerank[%u] failed", rank, peerRank), ret);
209 : // 接收数据到本端的 output
210 0 : HCCL_DEBUG(
211 : "rank[%u] outputMem[%p] receive from PeerRank[%u] outputMem, Offset[%llu], "
212 : "Size[%llu]",
213 : rank, outputMem_.ptr(), peerRank, baseOffset_, outputMem_.size());
214 0 : ret = ExecuteRxSync(
215 0 : links[peerRank], UserMemType::OUTPUT_MEM, baseOffset_, outputMem_.ptr(), outputMem_.size(), stream_);
216 0 : CHK_PRT_RET(
217 : ret != HCCL_SUCCESS,
218 : HCCL_ERROR("[Gather][InPartOneToOdd]rank[%u] rx sync from PeerRank[%u] failed", rank, peerRank), ret);
219 0 : ret = links[peerRank]->RxWaitDone(stream_);
220 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Gather][InPartOneToOdd]RxWaitDone failed"), ret);
221 : }
222 : }
223 0 : return HCCL_SUCCESS;
224 : }
225 :
226 0 : HcclResult AllGatherRecursiveHalvingDoubling::AllGatherInBlock(u32 rank, u32 rankSize, const std::vector<LINK>& links)
227 : {
228 0 : u32 rankInBlock = 0;
229 0 : if (rank < part1Size_ && (rank % 2) == 1) { // 模2余1代表当前rank在part1的奇数位置上,不参与block内的计算
230 0 : return HCCL_SUCCESS;
231 0 : } else if (rank < part1Size_ && (rank % 2) == 0) { // 模2余0代表当前rank在part1的偶数位置上,参与block内的计算
232 0 : rankInBlock = rank / 2; // 除2计算出在block内的rank号
233 : } else {
234 0 : rankInBlock = rank - part1Size_ / 2; // rank在part2中,用原始rank减part1除2,计算出在block内的rank号
235 : }
236 :
237 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
238 0 : TemplateType::TEMPLATE_ALL_GATHER_HALVING_DOUBLING, dispatcher_);
239 0 : CHK_SMART_PTR_NULL(tempAlg);
240 0 : CHK_RET(tempAlg->Prepare(blockSize_, UserMemType::OUTPUT_MEM, UserMemType::OUTPUT_MEM));
241 0 : CHK_RET(tempAlg->Prepare(
242 : outputMem_, outputMem_, count_, dataType_, stream_, reductionOp_, root_, slices_, baseOffset_));
243 :
244 0 : CHK_RET(tempAlg->RegisterProfiler(profilerInput_.planeID, profilerInput_.stage, profilerInput_.step, stream_));
245 :
246 0 : std::vector<LINK> subLinks;
247 0 : CHK_RET(BuildSubLinks(links, subLinks, rankSize));
248 :
249 0 : CHK_PRT_RET(
250 : subLinks.size() == 0,
251 : HCCL_ERROR("[AllGatherRecursiveHalvingDoubling][AllGatherInBlock]rank[%u] BuildSubLinks failed", rank),
252 : HCCL_E_PARA);
253 :
254 0 : CHK_RET(tempAlg->RunAsync(rankInBlock, blockSize_, subLinks));
255 :
256 0 : return HCCL_SUCCESS;
257 0 : }
258 :
259 0 : HcclResult AllGatherRecursiveHalvingDoubling::GetNslbAdjInfo(
260 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
261 : {
262 0 : u32 nslbRound = 0;
263 0 : u32 base = 1;
264 0 : const u32 minExponent = 1;
265 0 : HCCL_DEBUG("[AllGatherRecursiveHalvingDoubling]GetNslbAdjInfo begins");
266 0 : while ((base << nslbRound) <= rankSize) {
267 0 : nslbRound++;
268 : }
269 0 : if (nslbRound >= minExponent) {
270 0 : nslbRound = nslbRound - minExponent;
271 : }
272 0 : u32 nslbBlockSize = base << nslbRound;
273 : // 获取第一部分:rank数减block数乘2
274 0 : u32 nslbPart1Size = (rankSize - nslbBlockSize) * NSLBDP_ALL_GATHER_MOLD2;
275 : // 2的次幂场景下处理流程
276 0 : if (nslbPart1Size == 0) {
277 0 : u32 stepNum = 0;
278 0 : while ((rankSize >> (stepNum + 1)) != 0) {
279 0 : stepNum++;
280 : }
281 0 : for (u32 step = 0; step < stepNum; step++) {
282 0 : HCCL_DEBUG("[AllGatherRecursiveHalvingDoubling]current step is %u", step);
283 0 : u32 peerRankBitmask = 1 << (stepNum - step - 1);
284 0 : u32 peerRank = rank ^ peerRankBitmask;
285 0 : NslbDpAdjInfo adjInfoStep = {};
286 0 : u32 remoteuserRank = links[peerRank]->GetRemoteRank();
287 0 : adjInfoStep.dstLocalRankId = remoteuserRank;
288 0 : adjInfoStep.phaseId = step + 1;
289 0 : adjInfoStep.rev = 0;
290 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
291 0 : HCCL_DEBUG("[AllGatherRecursiveHalvingDoubling]current step %u success", step);
292 : }
293 0 : nslbAdjInfo.dstRankNum = stepNum;
294 0 : return HCCL_SUCCESS;
295 : }
296 : // 非2的次幂场景下,被合并部分的奇数rank处理流程
297 0 : if (rank < nslbPart1Size && rank % NSLBDP_ALL_GATHER_MOLD2 == 1) {
298 0 : u32 peerRank = rank - 1;
299 0 : if (peerRank < links.size()) {
300 0 : NslbDpAdjInfo adjInfoStep = {};
301 0 : adjInfoStep.dstLocalRankId = links[peerRank]->GetRemoteRank();
302 0 : adjInfoStep.phaseId = 1;
303 0 : adjInfoStep.rev = 0;
304 0 : HCCL_INFO("AllGatherHDR-nslb: peerRank[%u]", peerRank);
305 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
306 0 : nslbAdjInfo.dstRankNum = 1;
307 : }
308 0 : return HCCL_SUCCESS;
309 : }
310 : // 针对合并后映射成2的次幂场景处理
311 0 : u32 rankInBlock = 0;
312 0 : if (rank < nslbPart1Size && (rank % NSLBDP_ALL_GATHER_MOLD2) == 0) {
313 0 : rankInBlock = rank / NSLBDP_ALL_GATHER_MOLD2; // 直接除以2即为本rank的在block内的排序
314 : } else {
315 : rankInBlock
316 0 : = rank
317 0 : - nslbPart1Size / NSLBDP_ALL_GATHER_MOLD2; // 通过rank减去part1除2的大小即不处于第一部分的block内rank号
318 : }
319 0 : std::vector<LINK> subLinks;
320 0 : std::vector<LINK>::const_iterator iter = links.begin();
321 0 : subLinks.resize(nslbBlockSize);
322 0 : for (u32 i = 0; i < rankSize; i++) {
323 0 : if (i < nslbPart1Size
324 0 : && (i % NSLBDP_ALL_GATHER_MOLD2) == 1) { // 模2余1代表当前rank在part1的奇数位置上,不参与block内的建链
325 0 : continue;
326 0 : } else if (i < nslbPart1Size && (i % NSLBDP_ALL_GATHER_MOLD2) == 0) { // 模2余0代表当前rank在part1的偶数位置上
327 0 : std::vector<LINK>::const_iterator niter = std::next(iter, i);
328 0 : if (niter != links.end()) {
329 0 : subLinks[i / NSLBDP_ALL_GATHER_MOLD2] = *niter;
330 : }
331 0 : } else {
332 0 : std::vector<LINK>::const_iterator niter = std::next(iter, i);
333 0 : if (niter != links.end()) {
334 0 : subLinks[i - nslbPart1Size / NSLBDP_ALL_GATHER_MOLD2] = *niter;
335 : }
336 : }
337 : }
338 0 : u32 stepNum = 0;
339 0 : while ((rankSize >> (stepNum + 1)) != 0) {
340 0 : stepNum++;
341 : }
342 : // 映射完成后针对以新的通信域进行邻接表获取
343 0 : for (u32 step = 0; step < stepNum; step++) {
344 0 : u32 peerRankBitmask = (1 << step);
345 0 : u32 peerRank = rankInBlock ^ peerRankBitmask;
346 0 : if (subLinks[peerRank] == nullptr) {
347 0 : continue;
348 : }
349 0 : NslbDpAdjInfo adjInfoStep = {};
350 0 : u32 remoteuserRank = subLinks[peerRank]->GetRemoteRank();
351 0 : HCCL_DEBUG("[AllGatherRecursiveHalvingDoubling][GetNslbAdjInfo]remoteuserRank is %u", remoteuserRank);
352 0 : adjInfoStep.dstLocalRankId = remoteuserRank;
353 0 : adjInfoStep.phaseId = step + 1;
354 0 : adjInfoStep.rev = 0;
355 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
356 : }
357 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
358 :
359 0 : if (nslbAdjInfo.nsAdjInfo.size() == 0) {
360 0 : return HCCL_SUCCESS;
361 : }
362 : // 上面处理完成后,紧接着处理合并部分的偶数rank同步到奇数rank增加phaseId
363 0 : if (rank < nslbPart1Size && rank % NSLBDP_ALL_GATHER_MOLD2 == 0) {
364 0 : u32 peerRank = rank + 1;
365 0 : uint16_t phaseSize = nslbAdjInfo.nsAdjInfo.size();
366 0 : if (peerRank < links.size()) {
367 0 : NslbDpAdjInfo adjInfoStep = {};
368 0 : adjInfoStep.dstLocalRankId = links[peerRank]->GetRemoteRank();
369 0 : adjInfoStep.phaseId = nslbAdjInfo.nsAdjInfo[phaseSize - 1].phaseId + 1;
370 0 : adjInfoStep.rev = 0;
371 0 : HCCL_INFO("Scatter-nslb: peerRank[%u]", peerRank);
372 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
373 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
374 : }
375 0 : return HCCL_SUCCESS;
376 : }
377 0 : return HCCL_SUCCESS;
378 0 : }
379 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING, AllGatherRecursiveHalvingDoubling);
380 : } // namespace hccl
|