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 "transport.h"
12 : #include "dispatcher_pub.h"
13 : #include "hccl_primitive_remote.h"
14 : #include "hccl_primitive_local.h"
15 :
16 : using namespace hccl;
17 : extern HcclResult GetPubDispatcher(hccl::DispatcherPub** dispatcherPtr);
18 0 : HcclResult HcclRemoteWrite(StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* rmtBuf, HcclBuf* locBuf)
19 : {
20 0 : CHK_PTR_NULL(streamHandle);
21 0 : CHK_PTR_NULL(memTransport);
22 0 : CHK_PTR_NULL(rmtBuf);
23 0 : CHK_PTR_NULL(locBuf);
24 0 : HCCL_DEBUG(
25 : "[HcclRemoteWrite]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p], len[%llu].",
26 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len);
27 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
28 0 : struct Transport::Buffer localBuf(locBuf->addr, locBuf->len);
29 0 : struct Transport::Buffer remoteBuf(rmtBuf->addr, rmtBuf->len);
30 :
31 0 : return reinterpret_cast<Transport*>(memTransport)->WriteAsync(remoteBuf, localBuf, *stream);
32 : }
33 :
34 0 : HcclResult HcclRemoteRead(StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* locBuf, HcclBuf* rmtBuf)
35 : {
36 0 : CHK_PTR_NULL(streamHandle);
37 0 : CHK_PTR_NULL(memTransport);
38 0 : CHK_PTR_NULL(locBuf);
39 0 : CHK_PTR_NULL(rmtBuf);
40 0 : HCCL_DEBUG(
41 : "[HcclRemoteRead]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p], len[%llu].",
42 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len);
43 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
44 :
45 0 : struct Transport::Buffer localBuf(locBuf->addr, locBuf->len);
46 0 : struct Transport::Buffer remoteBuf(rmtBuf->addr, rmtBuf->len);
47 0 : return reinterpret_cast<Transport*>(memTransport)->ReadAsync(localBuf, remoteBuf, *stream);
48 : }
49 :
50 : constexpr uint32_t INVALID_COMPLETE_IDX = INVALID_UINT;
51 0 : HcclResult HcclRemoteWriteReduce(
52 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* rmtBuf, HcclBuf* locBuf,
53 : HcclReduceInfo reduceInfo)
54 : {
55 0 : CHK_PTR_NULL(streamHandle);
56 0 : CHK_PTR_NULL(memTransport);
57 0 : CHK_PTR_NULL(rmtBuf);
58 0 : CHK_PTR_NULL(locBuf);
59 0 : HCCL_DEBUG(
60 : "[HcclRemoteWriteReduce]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p], len[%llu],"
61 : " dataType[%d], reduceOp[%d].",
62 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len, reduceInfo.dataType, reduceInfo.reduceOp);
63 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
64 0 : struct Transport::Buffer localBuf(locBuf->addr, locBuf->len);
65 0 : struct Transport::Buffer remoteBuf(rmtBuf->addr, rmtBuf->len);
66 : return reinterpret_cast<Transport*>(memTransport)
67 0 : ->WriteReduceAsync(remoteBuf, localBuf, reduceInfo.dataType, reduceInfo.reduceOp, *stream);
68 : }
69 :
70 0 : HcclResult HcclRemoteReadReduce(
71 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* locBuf, HcclBuf* rmtBuf,
72 : HcclReduceInfo reduceInfo)
73 : {
74 0 : CHK_PTR_NULL(streamHandle);
75 0 : CHK_PTR_NULL(memTransport);
76 0 : CHK_PTR_NULL(locBuf);
77 0 : CHK_PTR_NULL(rmtBuf);
78 0 : HCCL_DEBUG(
79 : "[HcclRemoteReadReduce]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p], len[%llu],"
80 : " dataType[%d], reduceOp[%d].",
81 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len, reduceInfo.dataType, reduceInfo.reduceOp);
82 :
83 : // 后续使用transport,p2p支持,rdma不支持
84 0 : if (reinterpret_cast<Transport*>(memTransport)->GetLinkType() == LinkType::LINK_ROCE) {
85 0 : HCCL_ERROR("[HcclRemoteReadReduce]ROCE is not supported.");
86 0 : return HCCL_E_NOT_SUPPORT;
87 : }
88 :
89 0 : DispatcherPub* dispatcherPtr = nullptr;
90 0 : CHK_RET(GetPubDispatcher(&dispatcherPtr));
91 :
92 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
93 0 : return dispatcherPtr->InlineReduceAsync(
94 0 : rmtBuf->addr, rmtBuf->len / SIZE_TABLE[reduceInfo.dataType], reduceInfo.dataType, reduceInfo.reduceOp, *stream,
95 0 : locBuf->addr, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP);
96 : }
97 :
98 0 : HcclResult HcclRemoteNotifyRecord(StreamHandle streamHandle, HcclMemTransport memTransport, uint32_t notifyIndex)
99 : {
100 0 : CHK_PTR_NULL(streamHandle);
101 0 : CHK_PTR_NULL(memTransport);
102 0 : HCCL_DEBUG(
103 : "[HcclRemoteNotifyRecord]streamHandle[%p], memTransport[%p], notifyIndex[%d].", streamHandle, memTransport,
104 : notifyIndex);
105 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
106 0 : return reinterpret_cast<Transport*>(memTransport)->Post(notifyIndex, *stream);
107 : }
108 :
109 0 : HcclResult HcclRemoteNotifyWait(
110 : StreamHandle streamHandle, HcclMemTransport memTransport, uint32_t notifyIndex, const uint32_t timeOut)
111 : {
112 0 : CHK_PTR_NULL(streamHandle);
113 0 : CHK_PTR_NULL(memTransport);
114 0 : HCCL_DEBUG(
115 : "[HcclRemoteNotifyWait]streamHandle[%p], memTransport[%p], notifyIndex[%u], timeOut[%u s].", streamHandle,
116 : memTransport, notifyIndex, timeOut);
117 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
118 0 : return reinterpret_cast<Transport*>(memTransport)->Wait(notifyIndex, *stream, timeOut);
119 : }
120 :
121 0 : HcclResult HcclRemoteWriteWithNotify(
122 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* rmtBuf, HcclBuf* locBuf, uint32_t notifyIndex)
123 : {
124 0 : CHK_PTR_NULL(streamHandle);
125 0 : CHK_PTR_NULL(memTransport);
126 0 : CHK_PTR_NULL(locBuf);
127 0 : CHK_PTR_NULL(rmtBuf);
128 0 : HCCL_DEBUG(
129 : "[HcclRemoteWriteWithNotify]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p]"
130 : ", len[%llu], notifyIndex[%u].",
131 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len, notifyIndex);
132 0 : return HCCL_E_NOT_SUPPORT;
133 : }
134 :
135 0 : HcclResult HcclRemoteWriteReduceWithNotify(
136 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBuf* rmtBuf, HcclBuf* locBuf,
137 : HcclReduceInfo reduceInfo, uint32_t notifyIndex)
138 : {
139 0 : CHK_PTR_NULL(streamHandle);
140 0 : CHK_PTR_NULL(memTransport);
141 0 : CHK_PTR_NULL(locBuf);
142 0 : CHK_PTR_NULL(rmtBuf);
143 0 : HCCL_DEBUG(
144 : "[HcclRemoteWriteReduceWithNotify]streamHandle[%p], memTransport[%p], locBuf addr[%p], rmtBuf addr[%p],"
145 : " len[%llu], dataType[%d], reduceOp[%d], notifyIndex[%u].",
146 : streamHandle, memTransport, locBuf->addr, rmtBuf->addr, rmtBuf->len, reduceInfo.dataType, reduceInfo.reduceOp,
147 : notifyIndex);
148 0 : return HCCL_E_NOT_SUPPORT;
149 : }
150 :
151 0 : HcclResult HcclRemoteFence(StreamHandle streamHandle, HcclMemTransport memTransport, uint32_t orderFlag)
152 : {
153 0 : CHK_PTR_NULL(streamHandle);
154 0 : CHK_PTR_NULL(memTransport);
155 0 : HCCL_DEBUG(
156 : "[HcclRemoteFence]streamHandle[%p], memTransport[%p], orderFlag[%u].", streamHandle, memTransport, orderFlag);
157 0 : return reinterpret_cast<Transport*>(memTransport)->Fence();
158 : }
159 :
160 0 : HcclResult HcclRemoteBatchWrite(
161 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBufPair* bufPairs, uint32_t bufPairNum)
162 : {
163 0 : CHK_PTR_NULL(streamHandle);
164 0 : CHK_PTR_NULL(memTransport);
165 0 : CHK_PTR_NULL(bufPairs);
166 0 : HCCL_DEBUG(
167 : "[HcclRemoteBatchWrite]streamHandle[%p], memTransport[%p], bufPairNum[%u].", streamHandle, memTransport,
168 : bufPairNum);
169 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
170 0 : Transport* transport = reinterpret_cast<Transport*>(memTransport);
171 0 : for (uint32_t i = 0; i < bufPairNum; i++) {
172 0 : CHK_PTR_NULL(bufPairs[i].loc.addr);
173 0 : CHK_PTR_NULL(bufPairs[i].rmt.addr);
174 0 : struct Transport::Buffer localBuf(bufPairs[i].loc.addr, bufPairs[i].loc.len);
175 0 : struct Transport::Buffer remoteBuf(bufPairs[i].rmt.addr, bufPairs[i].rmt.len);
176 0 : CHK_RET(transport->WriteAsync(remoteBuf, localBuf, *stream));
177 : }
178 :
179 0 : return HCCL_SUCCESS;
180 : }
181 :
182 4 : HcclResult HcclRemoteBatchRead(
183 : StreamHandle streamHandle, HcclMemTransport memTransport, HcclBufPair* bufPairs, uint32_t bufPairNum)
184 : {
185 4 : CHK_PTR_NULL(streamHandle);
186 3 : CHK_PTR_NULL(memTransport);
187 2 : CHK_PTR_NULL(bufPairs);
188 1 : CHK_PRT_RET(bufPairNum == 0, HCCL_ERROR("[HcclRemoteBatchRead]bufPairsNum is 0."), HCCL_E_PARA);
189 0 : HCCL_DEBUG(
190 : "[HcclRemoteBatchRead]streamHandle[%p], memTransport[%p], bufPairNum[%u].", streamHandle, memTransport,
191 : bufPairNum);
192 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
193 0 : Transport* transport = reinterpret_cast<Transport*>(memTransport);
194 :
195 0 : for (uint32_t i = 0; i < bufPairNum; i++) {
196 0 : CHK_PTR_NULL(bufPairs[i].loc.addr);
197 0 : CHK_PTR_NULL(bufPairs[i].rmt.addr);
198 0 : struct Transport::Buffer localBuf(bufPairs[i].loc.addr, bufPairs[i].loc.len);
199 0 : struct Transport::Buffer remoteBuf(bufPairs[i].rmt.addr, bufPairs[i].rmt.len);
200 0 : CHK_RET(transport->ReadAsync(localBuf, remoteBuf, *stream));
201 : }
202 :
203 0 : return HCCL_SUCCESS;
204 : }
205 :
206 0 : HcclResult HcclRemoteBatchTransfer(
207 : StreamHandle streamHandle, HcclMemTransport memTransport, const HcclBatchTransferInfo* transferInfo,
208 : [[maybe_unused]] uint32_t bufPairNum)
209 : {
210 0 : CHK_PTR_NULL(streamHandle);
211 0 : CHK_PTR_NULL(memTransport);
212 0 : CHK_PTR_NULL(transferInfo);
213 0 : return HCCL_E_NOT_SUPPORT;
214 : }
215 :
216 0 : HcclResult HcclRemoteDrain(StreamHandle streamHandle, HcclMemTransport memTransport)
217 : {
218 0 : CHK_PTR_NULL(streamHandle);
219 0 : CHK_PTR_NULL(memTransport);
220 0 : Stream* stream = reinterpret_cast<Stream*>(streamHandle);
221 0 : Transport* transport = reinterpret_cast<Transport*>(memTransport);
222 0 : return transport->Drain(*stream);
223 : }
|