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 : #ifndef HCCLV2_ALG_DATA_TRANS_WRAPPER
12 : #define HCCLV2_ALG_DATA_TRANS_WRAPPER
13 :
14 : #include <vector>
15 : #include "data_type.h"
16 : #include "reduce_op.h"
17 :
18 : #include "coll_alg_params.h"
19 : #include "virtual_topo.h"
20 : #include "connected_link_mgr.h"
21 : #include "dev_capability.h"
22 : #include "primitive.h"
23 : #include "prim_queue.h"
24 : #include "instruction.h"
25 : #include "ins_queue.h"
26 :
27 : namespace Hccl {
28 : using InsQuePtr = std::shared_ptr<InsQueue>;
29 :
30 : using SlicesList = struct SlicesListDef {
31 : std::vector<DataSlice> srcSlices_;
32 : std::vector<DataSlice> dstSlices_;
33 :
34 0 : SlicesListDef(const std::vector<DataSlice> &srcSlices, const std::vector<DataSlice> &dstSlices)
35 0 : : srcSlices_(srcSlices), dstSlices_(dstSlices)
36 : {
37 0 : }
38 : };
39 :
40 : using ReduceSlicesList = struct ReduceSlicesListDef {
41 : std::vector<DataSlice> srcSlices_;
42 : std::vector<DataSlice> dstSlices_;
43 : DataType dataType_;
44 : ReduceOp reduceOp_;
45 :
46 : ReduceSlicesListDef(const std::vector<DataSlice> &srcSlices, const std::vector<DataSlice> &dstSlices,
47 : const DataType &dataType, const ReduceOp &reduceOp)
48 : : srcSlices_(srcSlices), dstSlices_(dstSlices), dataType_(dataType), reduceOp_(reduceOp)
49 : {
50 : }
51 0 : ReduceSlicesListDef(const SlicesList &slices, const DataType &dataType, const ReduceOp &reduceOp)
52 0 : : srcSlices_(slices.srcSlices_), dstSlices_(slices.dstSlices_), dataType_(dataType), reduceOp_(reduceOp)
53 : {
54 0 : }
55 : };
56 :
57 : using TxRxSlicesList = struct TxRxSlicesListDef {
58 : SlicesList txSlicesList_;
59 : SlicesList rxSlicesList_;
60 :
61 0 : TxRxSlicesListDef(const SlicesList &txSlicesList, const SlicesList &rxSlicesList)
62 0 : : txSlicesList_(txSlicesList), rxSlicesList_(rxSlicesList)
63 : {
64 0 : }
65 : };
66 :
67 : using TxRxReduceSlicesList = struct TxRxReduceSliceInfoDef {
68 : SlicesList txSlicesList_;
69 : SlicesList rxSlicesList_;
70 : DataType dataType_;
71 : ReduceOp reduceOp_;
72 :
73 : TxRxReduceSliceInfoDef(const SlicesList &txSlicesList, const SlicesList &rxSlicesList, const DataType dataType,
74 : const ReduceOp reduceOp)
75 : : txSlicesList_(txSlicesList), rxSlicesList_(rxSlicesList), dataType_(dataType), reduceOp_(reduceOp)
76 : {
77 : }
78 :
79 0 : TxRxReduceSliceInfoDef(const TxRxSlicesList &slices, const DataType dataType, const ReduceOp reduceOp)
80 0 : : txSlicesList_(slices.txSlicesList_), rxSlicesList_(slices.rxSlicesList_), dataType_(dataType), reduceOp_(reduceOp)
81 : {
82 0 : }
83 : };
84 :
85 : using TxRxLinks = struct TxRxLinkDef {
86 : LinkData txLink_;
87 : LinkData rxLink_;
88 :
89 0 : TxRxLinkDef(const LinkData &txLink, const LinkData &rxLink) : txLink_(txLink), rxLink_(rxLink)
90 : {
91 0 : }
92 : };
93 :
94 : using DataInfo = struct DataInfoDef {
95 : LinkData link_;
96 : SlicesList slices_;
97 :
98 0 : DataInfoDef(const LinkData &link, const SlicesList &slices) : link_(link), slices_(slices)
99 : {
100 0 : }
101 : };
102 :
103 : using SendRecvInfo = struct SendRecvInfoDef {
104 : TxRxLinks sendRecvLinks_;
105 : TxRxSlicesList sendRecvSlices_;
106 :
107 0 : SendRecvInfoDef(const TxRxLinks &sendRecvLinks, const TxRxSlicesList &sendRecvSlices)
108 0 : : sendRecvLinks_(sendRecvLinks), sendRecvSlices_(sendRecvSlices)
109 : {
110 0 : }
111 : };
112 :
113 : using DataReduceInfo = struct DataReduceInfoDef {
114 : LinkData link_;
115 : SlicesList slices_;
116 : DataType dataType_;
117 : ReduceOp reduceOp_;
118 :
119 : DataReduceInfoDef(const LinkData &link, const SlicesList &slices, const DataType dataType, const ReduceOp reduceOp)
120 : : link_(link), slices_(slices), dataType_(dataType), reduceOp_(reduceOp)
121 : {
122 : }
123 : };
124 :
125 : using SendRecvReduceInfo = struct SendRecvReduceInfoDef {
126 : TxRxLinks sendRecvLinks_;
127 : TxRxSlicesList sendRecvSlices_;
128 : DataType dataType_;
129 : ReduceOp reduceOp_;
130 :
131 0 : SendRecvReduceInfoDef(const TxRxLinks &sendRecvLinks, const TxRxSlicesList &sendRecvSlices, const DataType dataType,
132 : const ReduceOp reduceOp)
133 0 : : sendRecvLinks_(sendRecvLinks), sendRecvSlices_(sendRecvSlices), dataType_(dataType), reduceOp_(reduceOp)
134 : {
135 0 : }
136 : };
137 :
138 : using MultiDataInfo = struct MultiDataInfoDef {
139 : std::vector<LinkData> &links_;
140 : std::vector<SlicesList> &slices_;
141 :
142 : MultiDataInfoDef(std::vector<LinkData> &links, std::vector<SlicesList> &slices) : links_(links), slices_(slices)
143 : {
144 : }
145 : };
146 :
147 : using MultiSendRecvInfo = struct MultiSendRecvInfoDef {
148 : std::vector<TxRxLinks> &txRxLinks_;
149 : std::vector<TxRxSlicesList> &txRxSlices_;
150 :
151 : MultiSendRecvInfoDef(std::vector<TxRxLinks> &txRxLinks, std::vector<TxRxSlicesList> &txRxSlices)
152 : : txRxLinks_(txRxLinks), txRxSlices_(txRxSlices)
153 : {
154 : }
155 : };
156 :
157 : using MultiDataReduceInfo = struct MultiDataReduceInfoDef {
158 : std::vector<LinkData> &links_;
159 : std::vector<ReduceSlicesList> &slices_;
160 :
161 : MultiDataReduceInfoDef(std::vector<LinkData> &links, std::vector<ReduceSlicesList> &slices)
162 : : links_(links), slices_(slices)
163 : {
164 : }
165 : };
166 :
167 : using MultiSendRecvReduceInfo = struct MultiSendRecvReduceInfoDef {
168 : std::vector<TxRxLinks> &txRxLinks_;
169 : std::vector<TxRxReduceSlicesList> &txRxSlices_;
170 :
171 : MultiSendRecvReduceInfoDef(std::vector<TxRxLinks> &txRxLinks, std::vector<TxRxReduceSlicesList> &txRxSlices)
172 : : txRxLinks_(txRxLinks), txRxSlices_(txRxSlices)
173 : {
174 : }
175 : };
176 :
177 : using SlicePair = struct SlicePairDef {
178 : DataSlice srcSlice_;
179 : DataSlice dstSlice_;
180 : DataType dataType_;
181 : ReduceOp reduceOp_;
182 :
183 0 : SlicePairDef(const DataSlice &srcSlice, const DataSlice &dstSlice) : srcSlice_(srcSlice), dstSlice_(dstSlice)
184 : {
185 0 : }
186 : };
187 :
188 : using TransSlicesInfo = struct TransSlicesInfoDef {
189 : bool reduceFlag;
190 :
191 : std::vector<DataSlice> srcSlices;
192 : std::vector<DataSlice> dstSlices;
193 : DataType dataType_;
194 : ReduceOp reduceOp_;
195 :
196 : bool enableCounterNotify_;
197 :
198 0 : explicit TransSlicesInfoDef(const SlicesList &slices, const bool enableCounterNotify = false)
199 0 : : reduceFlag(false), srcSlices(slices.srcSlices_), dstSlices(slices.dstSlices_),
200 0 : enableCounterNotify_(enableCounterNotify)
201 : {
202 0 : }
203 :
204 0 : TransSlicesInfoDef(const SlicesList &slices, const DataType dataType, const ReduceOp reduceOp,
205 : const bool enableCounterNotify = false)
206 0 : : reduceFlag(true), srcSlices(slices.srcSlices_), dstSlices(slices.dstSlices_), dataType_(dataType),
207 0 : reduceOp_(reduceOp), enableCounterNotify_(enableCounterNotify)
208 : {
209 0 : }
210 :
211 0 : explicit TransSlicesInfoDef(const ReduceSlicesList &slices, const bool enableCounterNotify = false)
212 0 : : reduceFlag(true), srcSlices(slices.srcSlices_), dstSlices(slices.dstSlices_), dataType_(slices.dataType_),
213 0 : reduceOp_(slices.reduceOp_), enableCounterNotify_(enableCounterNotify)
214 : {
215 0 : }
216 : };
217 :
218 : using MultiDataLinksDmaModeInfo = struct MultiDataLinksDmaModeInfoDef {
219 : DmaMode modeNeedSync_;
220 : DmaMode modeSet_;
221 :
222 0 : MultiDataLinksDmaModeInfoDef(const DmaMode modeNeedSync, const DmaMode modeSet)
223 0 : : modeNeedSync_(modeNeedSync), modeSet_(modeSet)
224 : {
225 0 : }
226 : };
227 :
228 : // mid-level
229 : // sync
230 : HcclResult TxReady(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
231 : HcclResult RxReady(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
232 : HcclResult TxFin(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
233 : HcclResult RxFin(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
234 : HcclResult TxFinAck(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
235 : HcclResult RxFinAck(const LinkData &link, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
236 :
237 : // data
238 : HcclResult TxData(const LinkData &link, InsQuePtr queue, const SlicesList &slices, DmaMode dmaMode = DmaMode::DEFAULT);
239 : HcclResult RxData(const LinkData &link, InsQuePtr queue, const SlicesList &slices, DmaMode dmaMode = DmaMode::DEFAULT);
240 : HcclResult TxReduce(const LinkData &link, InsQuePtr queue, const ReduceSlicesList &slices,
241 : DmaMode dmaMode = DmaMode::DEFAULT);
242 : HcclResult RxReduce(const LinkData &link, InsQuePtr queue, const ReduceSlicesList &slices,
243 : DmaMode dmaMode = DmaMode::DEFAULT);
244 :
245 : // data with sync
246 : HcclResult TxDataWithFin(const LinkData &link, InsQuePtr queue, const SlicesList &slices, u32 topicId = 0,
247 : DmaMode dmaMode = DmaMode::DEFAULT);
248 : HcclResult RxDataWithFin(const LinkData &link, InsQuePtr queue, const SlicesList &slices, u32 topicId = 0,
249 : DmaMode dmaMode = DmaMode::DEFAULT);
250 : HcclResult TxReduceWithFin(const LinkData &link, InsQuePtr queue, const ReduceSlicesList &slices, u32 topicId = 0,
251 : DmaMode dmaMode = DmaMode::DEFAULT);
252 : HcclResult RxReduceWithFin(const LinkData &link, InsQuePtr queue, const ReduceSlicesList &slices, u32 topicId = 0,
253 : DmaMode dmaMode = DmaMode::DEFAULT);
254 :
255 : // data with sync in counterNotify mode
256 : HcclResult MultiTxDataWithFinCounter(const std::vector<LinkData> &links, const std::vector<InsQuePtr> &queues,
257 : const std::vector<SlicesList> &slices, u32 topicId = 0,
258 : DmaMode dmaMode = DmaMode::DEFAULT);
259 : HcclResult MultiRxDataWithFinCounter(const std::vector<LinkData> &links, const std::vector<InsQuePtr> &queues,
260 : const std::vector<SlicesList> &slices, u32 topicId = 0,
261 : DmaMode dmaMode = DmaMode::DEFAULT);
262 : HcclResult MultiTxReduceWithFinCounter(const std::vector<LinkData> &links, const std::vector<InsQuePtr> &queues,
263 : const std::vector<ReduceSlicesList> &slices, u32 topicId = 0,
264 : DmaMode dmaMode = DmaMode::DEFAULT);
265 : HcclResult MultiRxReduceWithFinCounter(const std::vector<LinkData> &links, const std::vector<InsQuePtr> &queues,
266 : const std::vector<ReduceSlicesList> &slices, u32 topicId = 0,
267 : DmaMode dmaMode = DmaMode::DEFAULT);
268 :
269 : // sync
270 : HcclResult TxRxReady(const TxRxLinks &txRxlinks, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
271 : HcclResult TxRxFin(const TxRxLinks &txRxlinks, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
272 : HcclResult TxRxFinAck(const TxRxLinks &txRxlinks, InsQuePtr queue, u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
273 :
274 : // data
275 : HcclResult TxRxData(const TxRxLinks &txRxlinks, InsQuePtr queue, const TxRxSlicesList &txRxSlices,
276 : DmaMode dmaMode = DmaMode::DEFAULT);
277 : HcclResult TxRxReduce(const TxRxLinks &txRxlinks, InsQuePtr queue, const TxRxReduceSlicesList &txRxSlices,
278 : DmaMode dmaMode = DmaMode::DEFAULT);
279 :
280 : // data with sync
281 : HcclResult TxRxDataWithFin(const TxRxLinks &txRxlinks, InsQuePtr queue, const TxRxSlicesList &txRxSlices,
282 : u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
283 : HcclResult TxRxReduceWithFin(const TxRxLinks &txRxlinks, InsQuePtr queue, const TxRxReduceSlicesList &txRxSlices,
284 : u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
285 :
286 : HcclResult MultiTxRxDataWithFinCounter(const std::vector<TxRxLinks> &links, const std::vector<InsQuePtr> &queues,
287 : const std::vector<TxRxSlicesList> &slices, u32 topicId = 0,
288 : DmaMode dmaMode = DmaMode::DEFAULT);
289 : HcclResult MultiTxRxReduceWithFinCounter(const std::vector<TxRxLinks> &links, const std::vector<InsQuePtr> &queues,
290 : const std::vector<TxRxReduceSlicesList> &slices, u32 topicId = 0,
291 : DmaMode dmaMode = DmaMode::DEFAULT);
292 :
293 : // else
294 : HcclResult LocalCopy(InsQuePtr queue, const DataSlice &srcSlice, const DataSlice &dstSlice);
295 : HcclResult LocalCopySlices(InsQuePtr queue, const std::vector<DataSlice> &srcSlices,
296 : const std::vector<DataSlice> &dstSlices); // 支持连续数据片融合,未来支持strideCount
297 : HcclResult LocalReduce(InsQuePtr queue, const DataSlice &srcSlice, const DataSlice &dstSlice, const DataType dataType,
298 : const ReduceOp reduceOp);
299 : HcclResult LocalReduceSlices(InsQuePtr queue, const std::vector<DataSlice> &srcSlices,
300 : const std::vector<DataSlice> &dstSlices, const DataType dataType, const ReduceOp reduceOp);
301 : HcclResult AicpuReduce(InsQuePtr queue, const DataSlice &srcSlice, const DataSlice &dstSlice, const DataType dataType,
302 : const ReduceOp reduceOp);
303 : HcclResult AicpuReduceSlices(InsQuePtr queue, const std::vector<DataSlice> &srcSlices,
304 : const std::vector<DataSlice> &dstSlices, const DataType dataType, const ReduceOp reduceOp);
305 :
306 : HcclResult StreamSync(std::vector<InsQuePtr> &queues);
307 :
308 : HcclResult PreSyncQues(const std::vector<InsQuePtr> &syncQueues, const u32 postQueIdx, u32 topicId = 0,
309 : bool enableCounterNotify = false);
310 : HcclResult PostSyncQues(const std::vector<InsQuePtr> &syncQueues, const u32 waitQueIdx, u32 topicId = 0,
311 : bool enableCounterNotify = false);
312 :
313 : // high-level
314 : HcclResult Send(const DataInfo &sendInfo, InsQuePtr queue, u32 topicId = 0, bool needNetFinAck = true,
315 : DmaMode dmaMode = DmaMode::DEFAULT);
316 : HcclResult Recv(const DataInfo &recvInfo, InsQuePtr queue, u32 topicId = 0, bool needNetFinAck = true,
317 : DmaMode dmaMode = DmaMode::DEFAULT);
318 : HcclResult SendRecv(const SendRecvInfo &sendRecvInfo, InsQuePtr queue, u32 topicId = 0, bool needNetFinAck = true,
319 : DmaMode dmaMode = DmaMode::DEFAULT);
320 :
321 : HcclResult SendReduce(const DataReduceInfo &sendReduceInfo, InsQuePtr queue, u32 topicId = 0, bool needNetFinAck = true,
322 : DmaMode dmaMode = DmaMode::DEFAULT);
323 : HcclResult RecvReduce(const DataReduceInfo &recvReduceInfo, InsQuePtr queue, u32 topicId = 0, bool needNetFinAck = true,
324 : DmaMode dmaMode = DmaMode::DEFAULT);
325 : HcclResult SendRecvReduce(const SendRecvReduceInfo &sendRecvReduceInfo, InsQuePtr queue, u32 topicId = 0,
326 : bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
327 :
328 : // Inter-rank CounterNotify is supported when device supports poll cqe
329 : HcclResult MultiSendCounter(const MultiDataInfo &sendInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
330 : DmaMode dmaMode = DmaMode::DEFAULT);
331 : HcclResult MultiRecvCounter(const MultiDataInfo &recvInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
332 : DmaMode dmaMode = DmaMode::DEFAULT);
333 : HcclResult MultiSendRecvCounter(const MultiSendRecvInfo &sendRecvInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
334 : DmaMode dmaMode = DmaMode::DEFAULT);
335 :
336 : HcclResult MultiSendReduceCounter(const MultiDataReduceInfo &sendInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
337 : DmaMode dmaMode = DmaMode::DEFAULT);
338 : HcclResult MultiRecvReduceCounter(const MultiDataReduceInfo &recvInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
339 : DmaMode dmaMode = DmaMode::DEFAULT);
340 : HcclResult MultiSendRecvReduceCounter(const MultiSendRecvReduceInfo &sendRecvInfo, std::vector<InsQuePtr> &queues,
341 : u32 topicId = 0, DmaMode dmaMode = DmaMode::DEFAULT);
342 :
343 : // send/recv through multi links (support detour case)
344 : HcclResult SendThruMultiLinks(const std::vector<DataInfo> &sendInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
345 : bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
346 : HcclResult RecvThruMultiLinks(const std::vector<DataInfo> &recvInfo, std::vector<InsQuePtr> &queues, u32 topicId = 0,
347 : bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
348 : HcclResult SendRecvThruMultiLinks(const std::vector<SendRecvInfo> &sendRecvInfo, std::vector<InsQuePtr> &queues,
349 : u32 topicId = 0, bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
350 :
351 : HcclResult SendReduceThruMultiLinks(const std::vector<DataReduceInfo> &sendReduceInfo, std::vector<InsQuePtr> &queues,
352 : u32 topicId = 0, bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
353 : HcclResult RecvReduceThruMultiLinks(const std::vector<DataReduceInfo> &recvReduceInfo, std::vector<InsQuePtr> &queues,
354 : u32 topicId = 0, bool needNetFinAck = true, DmaMode dmaMode = DmaMode::DEFAULT);
355 : HcclResult SendRecvReduceThruMultiLinks(const std::vector<SendRecvReduceInfo> &sendRecvReduceInfo,
356 : std::vector<InsQuePtr> &queues, u32 topicId = 0, bool needNetFinAck = true,
357 : DmaMode dmaMode = DmaMode::DEFAULT);
358 :
359 : // auxiliary functions
360 : HcclResult GetDMAMode(const DmaMode setMode, const PortDeploymentType linkPortType, DmaMode &mode);
361 : bool IsContinuousSlice(const DataSlice &nxtSlice, const DataSlice &currSlice);
362 :
363 : void TransSlice(const LinkData &link, InsQuePtr queue, const SlicePair &txRxSlice, DmaMode dmaMode, bool reduceFlag);
364 : HcclResult TransSlicesLists(const LinkData &link, InsQuePtr queue, const TransSlicesInfo &slices, DmaMode dmaMode);
365 :
366 : HcclResult WriteSlicesListsWithFin(const LinkData &link, InsQuePtr queue, const TransSlicesInfo &slices, u32 topicId);
367 :
368 : // for high-level wrapper function
369 : HcclResult ProceedMultiLinks(const std::vector<DataInfo> &dataInfo, const std::vector<InsQuePtr> &queues,
370 : const MultiDataLinksDmaModeInfo &dmaModeInfo, std::vector<InsQuePtr> &syncQues,
371 : bool &hasDiffDmaMode);
372 : HcclResult ProceedMultiLinks(const std::vector<DataReduceInfo> &dataInfo, const std::vector<InsQuePtr> &queues,
373 : const MultiDataLinksDmaModeInfo &dmaModeInfo, std::vector<InsQuePtr> &syncQues,
374 : bool &hasDiffDmaMode);
375 : } // namespace Hccl
376 :
377 : #endif // !HCCLV2_ALG_DATA_TRANS_WRAPPER
|