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