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