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