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 "alg_data_trans_wrapper.h"
12 : #include "log.h"
13 :
14 : namespace Hccl {
15 0 : HcclResult Send(const DataInfo& sendInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
16 : {
17 0 : CHK_RET(TxReady(sendInfo.link_, queue, topicId, dmaMode));
18 0 : CHK_RET(TxDataWithFin(sendInfo.link_, queue, sendInfo.slices_, topicId, dmaMode));
19 0 : if (needNetFinAck) {
20 0 : CHK_RET(TxFinAck(sendInfo.link_, queue, topicId, dmaMode));
21 : }
22 :
23 0 : return HcclResult::HCCL_SUCCESS;
24 : }
25 :
26 0 : HcclResult Recv(const DataInfo& recvInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
27 : {
28 0 : CHK_RET(RxReady(recvInfo.link_, queue, topicId, dmaMode));
29 0 : CHK_RET(RxDataWithFin(recvInfo.link_, queue, recvInfo.slices_, topicId, dmaMode));
30 0 : if (needNetFinAck) {
31 0 : CHK_RET(RxFinAck(recvInfo.link_, queue, topicId, dmaMode));
32 : }
33 :
34 0 : return HcclResult::HCCL_SUCCESS;
35 : }
36 :
37 0 : HcclResult SendRecv(const SendRecvInfo& sendRecvInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
38 : {
39 0 : CHK_RET(TxRxReady(sendRecvInfo.sendRecvLinks_, queue, topicId, dmaMode));
40 0 : CHK_RET(TxRxDataWithFin(sendRecvInfo.sendRecvLinks_, queue, sendRecvInfo.sendRecvSlices_, topicId, dmaMode));
41 0 : if (needNetFinAck) {
42 0 : CHK_RET(TxRxFinAck(sendRecvInfo.sendRecvLinks_, queue, topicId, dmaMode));
43 : }
44 :
45 0 : return HcclResult::HCCL_SUCCESS;
46 : }
47 :
48 : HcclResult
49 0 : SendReduce(const DataReduceInfo& sendReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
50 : {
51 0 : CHK_RET(TxReady(sendReduceInfo.link_, queue, topicId, dmaMode));
52 0 : CHK_RET(TxReduceWithFin(
53 : sendReduceInfo.link_, queue, {sendReduceInfo.slices_, sendReduceInfo.dataType_, sendReduceInfo.reduceOp_},
54 : topicId, dmaMode));
55 0 : if (needNetFinAck) {
56 0 : CHK_RET(TxFinAck(sendReduceInfo.link_, queue, topicId, dmaMode));
57 : }
58 :
59 0 : return HcclResult::HCCL_SUCCESS;
60 : }
61 :
62 : HcclResult
63 0 : RecvReduce(const DataReduceInfo& recvReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
64 : {
65 0 : CHK_RET(RxReady(recvReduceInfo.link_, queue, topicId, dmaMode));
66 0 : CHK_RET(RxReduceWithFin(
67 : recvReduceInfo.link_, queue, {recvReduceInfo.slices_, recvReduceInfo.dataType_, recvReduceInfo.reduceOp_},
68 : topicId, dmaMode));
69 0 : if (needNetFinAck) {
70 0 : CHK_RET(RxFinAck(recvReduceInfo.link_, queue, topicId, dmaMode));
71 : }
72 :
73 0 : return HcclResult::HCCL_SUCCESS;
74 : }
75 :
76 0 : HcclResult SendRecvReduce(
77 : const SendRecvReduceInfo& sendRecvReduceInfo, InsQuePtr queue, u32 topicId, bool needNetFinAck, DmaMode dmaMode)
78 : {
79 0 : CHK_RET(TxRxReady(sendRecvReduceInfo.sendRecvLinks_, queue, topicId, dmaMode));
80 0 : CHK_RET(TxRxReduceWithFin(
81 : sendRecvReduceInfo.sendRecvLinks_, queue,
82 : {sendRecvReduceInfo.sendRecvSlices_, sendRecvReduceInfo.dataType_, sendRecvReduceInfo.reduceOp_}, topicId,
83 : dmaMode));
84 0 : if (needNetFinAck) {
85 0 : CHK_RET(TxRxFinAck(sendRecvReduceInfo.sendRecvLinks_, queue, topicId, dmaMode));
86 : }
87 :
88 0 : return HcclResult::HCCL_SUCCESS;
89 : }
90 :
91 0 : HcclResult MultiSendCounter(const MultiDataInfo& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
92 : {
93 0 : if (sendInfo.links_.size() == 0) {
94 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: link size equals 0, do nothing.");
95 0 : return HcclResult::HCCL_SUCCESS;
96 : }
97 :
98 0 : CHK_PRT_RET(
99 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
100 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: inter-rank CounterNotify is "
101 : "supported only when device supports StarsPollNetCq."),
102 : HcclResult::HCCL_E_INTERNAL);
103 :
104 0 : CHK_PRT_RET(
105 : (sendInfo.links_.size() != queues.size()) || (sendInfo.slices_.size() != queues.size()),
106 : HCCL_ERROR(
107 : "[InsCollAlgFactory] [AlgDataTrans] MultiSendCounter: invalid input with link num [%zu], "
108 : "slice num [%zu], queue num [%zu].",
109 : sendInfo.links_.size(), sendInfo.slices_.size(), queues.size()),
110 : HcclResult::HCCL_E_INTERNAL);
111 :
112 0 : auto linkIter = sendInfo.links_.begin();
113 0 : auto queIter = queues.begin();
114 :
115 0 : for (; linkIter != sendInfo.links_.end(); linkIter++, queIter++) {
116 0 : CHK_RET(TxReady((*linkIter), (*queIter), topicId, dmaMode));
117 : }
118 :
119 0 : CHK_RET(MultiTxDataWithFinCounter(sendInfo.links_, queues, sendInfo.slices_, topicId, dmaMode));
120 :
121 0 : return HcclResult::HCCL_SUCCESS;
122 : }
123 :
124 0 : HcclResult MultiRecvCounter(const MultiDataInfo& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
125 : {
126 0 : if (recvInfo.links_.size() == 0) {
127 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: link size equals 0, do nothing.");
128 0 : return HcclResult::HCCL_SUCCESS;
129 : }
130 :
131 0 : CHK_PRT_RET(
132 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
133 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: inter-rank CounterNotify is "
134 : "supported only when device supports StarsPollNetCq."),
135 : HcclResult::HCCL_E_INTERNAL);
136 :
137 0 : CHK_PRT_RET(
138 : (recvInfo.links_.size() != queues.size()) || (recvInfo.slices_.size() != queues.size()),
139 : HCCL_ERROR(
140 : "[InsCollAlgFactory] [AlgDataTrans] MultiRecvCounter: invalid input with link num [%zu], "
141 : "slice num [%zu], queue num [%zu].",
142 : recvInfo.links_.size(), recvInfo.slices_.size(), queues.size()),
143 : HcclResult::HCCL_E_INTERNAL);
144 :
145 0 : auto linkIter = recvInfo.links_.begin();
146 0 : auto queIter = queues.begin();
147 :
148 0 : for (; linkIter != recvInfo.links_.end(); linkIter++, queIter++) {
149 0 : CHK_RET(RxReady((*linkIter), (*queIter), topicId, dmaMode));
150 : }
151 :
152 0 : CHK_RET(MultiRxDataWithFinCounter(recvInfo.links_, queues, recvInfo.slices_, topicId, dmaMode));
153 :
154 0 : return HcclResult::HCCL_SUCCESS;
155 : }
156 :
157 0 : HcclResult MultiSendRecvCounter(
158 : const MultiSendRecvInfo& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
159 : {
160 0 : if (sendRecvInfo.txRxLinks_.size() == 0) {
161 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: link size equals 0, do nothing.");
162 0 : return HcclResult::HCCL_SUCCESS;
163 : }
164 :
165 0 : CHK_PRT_RET(
166 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
167 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: inter-rank CounterNotify is "
168 : "supported only when device supports StarsPollNetCq."),
169 : HcclResult::HCCL_E_INTERNAL);
170 :
171 0 : CHK_PRT_RET(
172 : (sendRecvInfo.txRxLinks_.size() != queues.size()) || (sendRecvInfo.txRxSlices_.size() != queues.size()),
173 : HCCL_ERROR(
174 : "[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvCounter: invalid input with link num [%zu], "
175 : "slice num [%zu], queue num [%zu].",
176 : sendRecvInfo.txRxLinks_.size(), sendRecvInfo.txRxSlices_.size(), queues.size()),
177 : HcclResult::HCCL_E_INTERNAL);
178 :
179 0 : auto linkIter = sendRecvInfo.txRxLinks_.begin();
180 0 : auto queIter = queues.begin();
181 :
182 0 : for (; linkIter != sendRecvInfo.txRxLinks_.end(); linkIter++, queIter++) {
183 0 : CHK_RET(TxRxReady((*linkIter), (*queIter), topicId, dmaMode));
184 : }
185 :
186 0 : CHK_RET(MultiTxRxDataWithFinCounter(sendRecvInfo.txRxLinks_, queues, sendRecvInfo.txRxSlices_, topicId, dmaMode));
187 :
188 0 : return HcclResult::HCCL_SUCCESS;
189 : }
190 :
191 0 : HcclResult MultiSendReduceCounter(
192 : const MultiDataReduceInfo& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
193 : {
194 0 : if (sendInfo.links_.size() == 0) {
195 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: link size equals 0, do nothing.");
196 0 : return HcclResult::HCCL_SUCCESS;
197 : }
198 :
199 0 : CHK_PRT_RET(
200 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
201 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: inter-rank CounterNotify is "
202 : "supported only when device supports StarsPollNetCq."),
203 : HcclResult::HCCL_E_INTERNAL);
204 :
205 0 : CHK_PRT_RET(
206 : (sendInfo.links_.size() != queues.size()) || (sendInfo.slices_.size() != queues.size()),
207 : HCCL_ERROR(
208 : "[InsCollAlgFactory] [AlgDataTrans] MultiSendReduceCounter: invalid input with link num [%zu], "
209 : "slice num [%zu], queue num [%zu].",
210 : sendInfo.links_.size(), sendInfo.slices_.size(), queues.size()),
211 : HcclResult::HCCL_E_INTERNAL);
212 :
213 0 : auto linkIter = sendInfo.links_.begin();
214 0 : auto queIter = queues.begin();
215 :
216 0 : for (; linkIter != sendInfo.links_.end(); linkIter++, queIter++) {
217 0 : CHK_RET(TxReady((*linkIter), (*queIter), topicId, dmaMode));
218 : }
219 :
220 0 : CHK_RET(MultiTxReduceWithFinCounter(sendInfo.links_, queues, sendInfo.slices_, topicId, dmaMode));
221 :
222 0 : return HcclResult::HCCL_SUCCESS;
223 : }
224 :
225 0 : HcclResult MultiRecvReduceCounter(
226 : const MultiDataReduceInfo& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
227 : {
228 0 : if (recvInfo.links_.size() == 0) {
229 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: link size equals 0, do nothing.");
230 0 : return HcclResult::HCCL_SUCCESS;
231 : }
232 :
233 0 : CHK_PRT_RET(
234 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
235 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: inter-rank CounterNotify is "
236 : "supported only when device supports StarsPollNetCq."),
237 : HcclResult::HCCL_E_INTERNAL);
238 :
239 0 : CHK_PRT_RET(
240 : (recvInfo.links_.size() != queues.size()) || (recvInfo.slices_.size() != queues.size()),
241 : HCCL_ERROR(
242 : "[InsCollAlgFactory] [AlgDataTrans] MultiRecvReduceCounter: invalid input with link num [%zu], "
243 : "slice num [%zu], queue num [%zu].",
244 : recvInfo.links_.size(), recvInfo.slices_.size(), queues.size()),
245 : HcclResult::HCCL_E_INTERNAL);
246 :
247 0 : auto linkIter = recvInfo.links_.begin();
248 0 : auto queIter = queues.begin();
249 :
250 0 : for (; linkIter != recvInfo.links_.end(); linkIter++, queIter++) {
251 0 : CHK_RET(RxReady((*linkIter), (*queIter), topicId, dmaMode));
252 : }
253 :
254 0 : CHK_RET(MultiRxReduceWithFinCounter(recvInfo.links_, queues, recvInfo.slices_, topicId, dmaMode));
255 :
256 0 : return HcclResult::HCCL_SUCCESS;
257 : }
258 :
259 0 : HcclResult MultiSendRecvReduceCounter(
260 : const MultiSendRecvReduceInfo& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, DmaMode dmaMode)
261 : {
262 0 : if (sendRecvInfo.txRxLinks_.size() == 0) {
263 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: link size equals 0, do nothing.");
264 0 : return HcclResult::HCCL_SUCCESS;
265 : }
266 :
267 0 : CHK_PRT_RET(
268 : !DevCapability::GetInstance().IsSupportStarsPollNetCq(),
269 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: inter-rank CounterNotify is "
270 : "supported only when device supports StarsPollNetCq."),
271 : HcclResult::HCCL_E_INTERNAL);
272 :
273 0 : CHK_PRT_RET(
274 : (sendRecvInfo.txRxLinks_.size() != queues.size()) || (sendRecvInfo.txRxSlices_.size() != queues.size()),
275 : HCCL_ERROR(
276 : "[InsCollAlgFactory] [AlgDataTrans] MultiSendRecvReduceCounter: invalid input with link num [%zu], "
277 : "slice num [%zu], queue num [%zu].",
278 : sendRecvInfo.txRxLinks_.size(), sendRecvInfo.txRxSlices_.size(), queues.size()),
279 : HcclResult::HCCL_E_INTERNAL);
280 :
281 0 : auto linkIter = sendRecvInfo.txRxLinks_.begin();
282 0 : auto queIter = queues.begin();
283 :
284 0 : for (; linkIter != sendRecvInfo.txRxLinks_.end(); linkIter++, queIter++) {
285 0 : CHK_RET(TxRxReady((*linkIter), (*queIter), topicId, dmaMode));
286 : }
287 :
288 0 : CHK_RET(MultiTxRxReduceWithFinCounter(sendRecvInfo.txRxLinks_, queues, sendRecvInfo.txRxSlices_, topicId, dmaMode));
289 :
290 0 : return HcclResult::HCCL_SUCCESS;
291 : }
292 :
293 0 : HcclResult SendThruMultiLinks(
294 : const std::vector<DataInfo>& sendInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
295 : DmaMode dmaMode)
296 : {
297 0 : if (sendInfo.size() == 0) {
298 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: sendInfo size equals 0, do nothing.");
299 0 : return HcclResult::HCCL_SUCCESS;
300 : }
301 :
302 0 : CHK_PRT_RET(
303 : sendInfo.size() != queues.size(),
304 : HCCL_ERROR(
305 : "[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: sendInfo size [%zu] is non-equal to queue num "
306 : "[%zu].",
307 : sendInfo.size(), queues.size()),
308 : HcclResult::HCCL_E_INTERNAL);
309 :
310 : // only those worker queues required to be sync: put mode in send
311 0 : std::vector<InsQuePtr> syncQues = {queues[0]};
312 0 : bool hasDiffDmaMode = false;
313 :
314 0 : CHK_RET(ProceedMultiLinks(
315 : sendInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::PUT, dmaMode), syncQues, hasDiffDmaMode));
316 :
317 0 : if (hasDiffDmaMode) {
318 0 : HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] SendThruMultiLinks: current send links have two DmaMode.");
319 0 : CHK_RET(TxRxReady({sendInfo[0].link_, sendInfo[0].link_}, queues[0], topicId, dmaMode));
320 : } else {
321 0 : CHK_RET(TxReady(sendInfo[0].link_, queues[0], topicId, dmaMode));
322 : }
323 :
324 0 : CHK_RET(PreSyncQues(syncQues, 0));
325 :
326 0 : auto dataInfoIter = sendInfo.begin();
327 0 : auto queIter = queues.begin();
328 0 : for (; dataInfoIter != sendInfo.end(); dataInfoIter++, queIter++) {
329 0 : if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
330 0 : CHK_RET(TxData(dataInfoIter->link_, (*queIter), dataInfoIter->slices_, dmaMode));
331 : }
332 : }
333 :
334 0 : CHK_RET(PostSyncQues(syncQues, 0));
335 :
336 0 : if (hasDiffDmaMode) {
337 0 : CHK_RET(TxRxFin({sendInfo[0].link_, sendInfo[0].link_}, queues[0], topicId, dmaMode));
338 : } else {
339 0 : CHK_RET(TxFin(sendInfo[0].link_, queues[0], topicId, dmaMode));
340 : }
341 :
342 0 : if (needNetFinAck) {
343 0 : TxFinAck(sendInfo[0].link_, queues[0], topicId, dmaMode);
344 : }
345 :
346 0 : return HcclResult::HCCL_SUCCESS;
347 0 : }
348 :
349 0 : HcclResult RecvThruMultiLinks(
350 : const std::vector<DataInfo>& recvInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
351 : DmaMode dmaMode)
352 : {
353 0 : if (recvInfo.size() == 0) {
354 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: recvInfo size equals 0, do nothing.");
355 0 : return HcclResult::HCCL_SUCCESS;
356 : }
357 :
358 0 : CHK_PRT_RET(
359 : recvInfo.size() != queues.size(),
360 : HCCL_ERROR(
361 : "[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: invalid input with recvInfo size [%u], "
362 : "queue num [%u].",
363 : recvInfo.size(), queues.size()),
364 : HcclResult::HCCL_E_INTERNAL);
365 :
366 : // only those worker queues required to be sync: put mode in send
367 0 : std::vector<InsQuePtr> syncQues = {queues[0]};
368 0 : bool hasDiffDmaMode = false;
369 :
370 0 : CHK_RET(ProceedMultiLinks(
371 : recvInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::GET, dmaMode), syncQues,
372 : hasDiffDmaMode)); // Get mode should be sync for Recv
373 :
374 0 : if (hasDiffDmaMode) {
375 0 : HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] RecvThruMultiLinks: current recv links have two DmaMode.");
376 0 : CHK_RET(TxRxReady({recvInfo[0].link_, recvInfo[0].link_}, queues[0], topicId, dmaMode));
377 : } else {
378 0 : CHK_RET(RxReady(recvInfo[0].link_, queues[0], topicId, dmaMode));
379 : }
380 :
381 0 : CHK_RET(PreSyncQues(syncQues, 0));
382 :
383 0 : auto dataInfoIter = recvInfo.begin();
384 0 : auto queIter = queues.begin();
385 0 : for (; dataInfoIter != recvInfo.end(); dataInfoIter++, queIter++) {
386 0 : if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
387 0 : CHK_RET(RxData(dataInfoIter->link_, (*queIter), dataInfoIter->slices_, dmaMode));
388 : }
389 : }
390 :
391 0 : CHK_RET(PostSyncQues(syncQues, 0));
392 :
393 0 : if (hasDiffDmaMode) {
394 0 : CHK_RET(TxRxFin({recvInfo[0].link_, recvInfo[0].link_}, queues[0], topicId, dmaMode));
395 : } else {
396 0 : CHK_RET(RxFin(recvInfo[0].link_, queues[0], topicId, dmaMode));
397 : }
398 :
399 0 : if (needNetFinAck) {
400 0 : RxFinAck(recvInfo[0].link_, queues[0], topicId, dmaMode);
401 : }
402 :
403 0 : return HcclResult::HCCL_SUCCESS;
404 0 : }
405 :
406 0 : HcclResult SendRecvThruMultiLinks(
407 : const std::vector<SendRecvInfo>& sendRecvInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
408 : DmaMode dmaMode)
409 : {
410 0 : if (sendRecvInfo.size() == 0) {
411 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: empty sendRecvInfo, do nothing.");
412 0 : return HcclResult::HCCL_SUCCESS;
413 : }
414 :
415 0 : CHK_PRT_RET(
416 : sendRecvInfo.size() != queues.size(),
417 : HCCL_ERROR(
418 : "[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: invalid input with recvInfo size [%zu], "
419 : "queue num [%zu].",
420 : sendRecvInfo.size(), queues.size()),
421 : HcclResult::HCCL_E_INTERNAL);
422 :
423 0 : auto sendRecvInfoIter = sendRecvInfo.begin();
424 0 : auto queIter = queues.begin();
425 0 : u32 netTxLinksNum = 0;
426 0 : u32 netRxLinksNum = 0;
427 :
428 0 : CHK_RET(TxRxReady(sendRecvInfoIter->sendRecvLinks_, (*queIter), topicId, dmaMode));
429 :
430 0 : u32 mainQueIdx = 0;
431 0 : CHK_RET(PreSyncQues(queues, mainQueIdx));
432 :
433 0 : for (; sendRecvInfoIter != sendRecvInfo.end(); sendRecvInfoIter++, queIter++) {
434 0 : if (((sendRecvInfoIter->sendRecvLinks_).txLink_).GetType() == PortDeploymentType::DEV_NET) {
435 0 : netTxLinksNum++;
436 : }
437 0 : if (((sendRecvInfoIter->sendRecvLinks_).rxLink_).GetType() == PortDeploymentType::DEV_NET) {
438 0 : netRxLinksNum++;
439 : }
440 0 : CHK_RET(TxRxData(sendRecvInfoIter->sendRecvLinks_, (*queIter), sendRecvInfoIter->sendRecvSlices_, dmaMode));
441 : }
442 :
443 0 : CHK_PRT_RET(
444 : ((netTxLinksNum > 1) || (netRxLinksNum > 1)),
445 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: multi net links is not "
446 : "supported as NET operations are async, use mid-level wrapper instead."),
447 : HcclResult::HCCL_E_INTERNAL);
448 :
449 0 : CHK_PRT_RET(
450 : (((netTxLinksNum == 1) && (sendRecvInfo[0].sendRecvLinks_.txLink_.GetType() == PortDeploymentType::P2P))
451 : || ((netRxLinksNum == 1) && (sendRecvInfo[0].sendRecvLinks_.rxLink_.GetType() == PortDeploymentType::P2P))),
452 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: first link must be NET when there "
453 : "exists NET links."),
454 : HcclResult::HCCL_E_INTERNAL);
455 :
456 0 : CHK_RET(PostSyncQues(queues, mainQueIdx));
457 :
458 0 : CHK_RET(TxRxFin(sendRecvInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
459 :
460 0 : if (needNetFinAck) {
461 0 : CHK_RET(TxRxFinAck(sendRecvInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
462 : }
463 :
464 0 : return HcclResult::HCCL_SUCCESS;
465 : }
466 :
467 0 : HcclResult SendReduceThruMultiLinks(
468 : const std::vector<DataReduceInfo>& sendReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
469 : DmaMode dmaMode)
470 : {
471 0 : if (sendReduceInfo.size() == 0) {
472 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: empty sendReduceInfo, do nothing.");
473 0 : return HcclResult::HCCL_SUCCESS;
474 : }
475 :
476 0 : CHK_PRT_RET(
477 : sendReduceInfo.size() != queues.size(),
478 : HCCL_ERROR(
479 : "[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: sendReduceInfo size [%zu] is "
480 : "non-equal to queue num [%zu].",
481 : sendReduceInfo.size(), queues.size()),
482 : HcclResult::HCCL_E_INTERNAL);
483 :
484 : // only those worker queues required to be sync: put mode in send
485 0 : std::vector<InsQuePtr> syncQues = {queues[0]};
486 0 : bool hasDiffDmaMode = false;
487 :
488 0 : CHK_RET(ProceedMultiLinks(
489 : sendReduceInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::PUT, dmaMode), syncQues, hasDiffDmaMode));
490 :
491 0 : if (hasDiffDmaMode) {
492 0 : HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] SendReduceThruMultiLinks: current send links have two DmaMode.");
493 0 : CHK_RET(TxRxReady({sendReduceInfo[0].link_, sendReduceInfo[0].link_}, queues[0], topicId, dmaMode));
494 : } else {
495 0 : CHK_RET(TxReady(sendReduceInfo[0].link_, queues[0], topicId, dmaMode));
496 : }
497 :
498 0 : CHK_RET(PreSyncQues(syncQues, 0));
499 :
500 0 : auto dataInfoIter = sendReduceInfo.begin();
501 0 : auto queIter = queues.begin();
502 0 : for (; dataInfoIter != sendReduceInfo.end(); dataInfoIter++, queIter++) {
503 0 : if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
504 0 : CHK_RET(TxReduce(
505 : dataInfoIter->link_, (*queIter),
506 : {dataInfoIter->slices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
507 : }
508 : }
509 :
510 0 : CHK_RET(PostSyncQues(syncQues, 0));
511 :
512 0 : if (hasDiffDmaMode) {
513 0 : CHK_RET(TxRxFin({sendReduceInfo[0].link_, sendReduceInfo[0].link_}, queues[0], topicId, dmaMode));
514 : } else {
515 0 : CHK_RET(TxFin(sendReduceInfo[0].link_, queues[0], topicId, dmaMode));
516 : }
517 :
518 0 : if (needNetFinAck) {
519 0 : TxFinAck(sendReduceInfo[0].link_, queues[0], topicId, dmaMode);
520 : }
521 :
522 0 : return HcclResult::HCCL_SUCCESS;
523 0 : }
524 :
525 0 : HcclResult RecvReduceThruMultiLinks(
526 : const std::vector<DataReduceInfo>& recvReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId, bool needNetFinAck,
527 : DmaMode dmaMode)
528 : {
529 0 : if (recvReduceInfo.size() == 0) {
530 0 : HCCL_WARNING("[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: empty recvReduceInfo, do nothing.");
531 0 : return HcclResult::HCCL_SUCCESS;
532 : }
533 :
534 0 : CHK_PRT_RET(
535 : recvReduceInfo.size() != queues.size(),
536 : HCCL_ERROR(
537 : "[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: invalid input with recvReduceInfo size "
538 : "[%zu], "
539 : "queue num [%zu].",
540 : recvReduceInfo.size(), queues.size()),
541 : HcclResult::HCCL_E_INTERNAL);
542 :
543 : // only those worker queues required to be sync: put mode in send
544 0 : std::vector<InsQuePtr> syncQues = {queues[0]};
545 0 : bool hasDiffDmaMode = false;
546 :
547 0 : CHK_RET(ProceedMultiLinks(
548 : recvReduceInfo, queues, MultiDataLinksDmaModeInfo(DmaMode::GET, dmaMode), syncQues,
549 : hasDiffDmaMode)); // Get mode should be sync for Recv
550 :
551 0 : if (hasDiffDmaMode) {
552 0 : HCCL_DEBUG("[InsCollAlgFactory] [AlgDataTrans] RecvReduceThruMultiLinks: current recv links have two DmaMode.");
553 0 : CHK_RET(TxRxReady({recvReduceInfo[0].link_, recvReduceInfo[0].link_}, queues[0], topicId, dmaMode));
554 : } else {
555 0 : CHK_RET(RxReady(recvReduceInfo[0].link_, queues[0], topicId, dmaMode));
556 : }
557 :
558 0 : CHK_RET(PreSyncQues(syncQues, 0));
559 :
560 0 : auto dataInfoIter = recvReduceInfo.begin();
561 0 : auto queIter = queues.begin();
562 0 : for (; dataInfoIter != recvReduceInfo.end(); dataInfoIter++, queIter++) {
563 0 : if (std::find(syncQues.begin(), syncQues.end(), (*queIter)) != syncQues.end()) {
564 0 : CHK_RET(RxReduce(
565 : dataInfoIter->link_, (*queIter),
566 : {dataInfoIter->slices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
567 : }
568 : }
569 :
570 0 : CHK_RET(PostSyncQues(syncQues, 0));
571 :
572 0 : if (hasDiffDmaMode) {
573 0 : CHK_RET(TxRxFin({recvReduceInfo[0].link_, recvReduceInfo[0].link_}, queues[0], topicId, dmaMode));
574 : } else {
575 0 : CHK_RET(RxFin(recvReduceInfo[0].link_, queues[0], topicId, dmaMode));
576 : }
577 :
578 0 : if (needNetFinAck) {
579 0 : RxFinAck(recvReduceInfo[0].link_, queues[0], topicId, dmaMode);
580 : }
581 :
582 0 : return HcclResult::HCCL_SUCCESS;
583 0 : }
584 :
585 0 : HcclResult SendRecvReduceThruMultiLinks(
586 : const std::vector<SendRecvReduceInfo>& sendRecvReduceInfo, std::vector<InsQuePtr>& queues, u32 topicId,
587 : bool needNetFinAck, DmaMode dmaMode)
588 : {
589 0 : if (sendRecvReduceInfo.size() == 0) {
590 0 : HCCL_WARNING(
591 : "[InsCollAlgFactory] [AlgDataTrans] SendRecvReduceThruMultiLinks: empty sendRecvReduceInfo, do nothing.");
592 0 : return HcclResult::HCCL_SUCCESS;
593 : }
594 :
595 0 : CHK_PRT_RET(
596 : sendRecvReduceInfo.size() != queues.size(),
597 : HCCL_ERROR(
598 : "[InsCollAlgFactory] [AlgDataTrans] SendRecvReduceThruMultiLinks: sendRecvReduceInfo size [%u] is "
599 : "non-equal to queue num [%u].",
600 : sendRecvReduceInfo.size(), queues.size()),
601 : HcclResult::HCCL_E_INTERNAL);
602 :
603 0 : auto dataInfoIter = sendRecvReduceInfo.begin();
604 0 : auto queIter = queues.begin();
605 0 : u32 netTxLinksNum = 0;
606 0 : u32 netRxLinksNum = 0;
607 :
608 0 : CHK_RET(TxRxReady(dataInfoIter->sendRecvLinks_, (*queIter), topicId, dmaMode));
609 :
610 0 : u32 mainQueIdx = 0;
611 0 : CHK_RET(PreSyncQues(queues, mainQueIdx));
612 :
613 0 : for (; dataInfoIter != sendRecvReduceInfo.end(); dataInfoIter++, queIter++) {
614 0 : if (((dataInfoIter->sendRecvLinks_).txLink_).GetType() == PortDeploymentType::DEV_NET) {
615 0 : netTxLinksNum++;
616 : }
617 0 : if (((dataInfoIter->sendRecvLinks_).rxLink_).GetType() == PortDeploymentType::DEV_NET) {
618 0 : netRxLinksNum++;
619 : }
620 0 : CHK_RET(TxRxReduce(
621 : dataInfoIter->sendRecvLinks_, (*queIter),
622 : {dataInfoIter->sendRecvSlices_, dataInfoIter->dataType_, dataInfoIter->reduceOp_}, dmaMode));
623 : }
624 :
625 0 : CHK_PRT_RET(
626 : ((netTxLinksNum > 1) || (netRxLinksNum > 1)),
627 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: multi net links is not "
628 : "supported as NET operations are async, use mid-level wrapper instead."),
629 : HcclResult::HCCL_E_INTERNAL);
630 :
631 0 : CHK_PRT_RET(
632 : (((netTxLinksNum == 1) && (sendRecvReduceInfo[0].sendRecvLinks_.txLink_.GetType() == PortDeploymentType::P2P))
633 : || ((netRxLinksNum == 1)
634 : && (sendRecvReduceInfo[0].sendRecvLinks_.rxLink_.GetType() == PortDeploymentType::P2P))),
635 : HCCL_ERROR("[InsCollAlgFactory] [AlgDataTrans] SendRecvThruMultiLinks: first link must be NET when there "
636 : "exists NET links."),
637 : HcclResult::HCCL_E_INTERNAL);
638 :
639 0 : CHK_RET(PostSyncQues(queues, mainQueIdx));
640 :
641 0 : CHK_RET(TxRxFin(sendRecvReduceInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
642 :
643 0 : if (needNetFinAck) {
644 0 : CHK_RET(TxRxFinAck(sendRecvReduceInfo[0].sendRecvLinks_, queues[0], topicId, dmaMode));
645 : }
646 :
647 0 : return HcclResult::HCCL_SUCCESS;
648 : }
649 :
650 : } // namespace Hccl
|