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 <arpa/inet.h>
12 : #include <securec.h>
13 : #include <chrono>
14 : #include <memory>
15 : #include "network/hccp_common.h"
16 : #include "device_capacity.h"
17 : #include "dlhns_function.h"
18 : #include "adapter_verbs.h"
19 : #include "transport_device_ibverbs.h"
20 : #include "new/hccl_dispatcher_ctx.h"
21 :
22 : constexpr u32 RDMA_QP_EXPECT_STATUS_PAUSE = 5;
23 : constexpr u32 RDMA_QP_EXPECT_STATUS_CONNECTED = 1;
24 : constexpr s32 RDMA_QP_NO_MEM = -12;
25 :
26 : constexpr u32 RDMA_WRITE_NOTIFY_OFFSET_MASK = 0xffffff;
27 : constexpr u32 RDMA_WRITE_NOTIFY_VALUE_RECORD = 0x1000000;
28 :
29 : // 内存屏障,确保wqe下到HBM里
30 : #if defined(__x86_64__)
31 : #define HCOMM_DSB() asm volatile("" ::: "memory")
32 : #elif defined(__aarch64__)
33 : #define HCOMM_DSB() asm volatile("dsb st" ::: "memory")
34 : #else
35 : #define HCOMM_DSB()
36 : #endif
37 :
38 : namespace hccl {
39 : namespace {
40 11 : inline BufferKey<uintptr_t, u64> MakeMemLookupKey(u64 logicalStartVa, u64 size)
41 : {
42 11 : return BufferKey<uintptr_t, u64>(static_cast<uintptr_t>(logicalStartVa), size);
43 : }
44 :
45 3 : inline BufferKey<uintptr_t, u64> MakeMemLookupKey(const void *logicalPtr, u64 size)
46 : {
47 3 : return MakeMemLookupKey(static_cast<u64>(reinterpret_cast<uintptr_t>(logicalPtr)), size);
48 : }
49 :
50 2 : inline void *LogicalPtrToDevPtr(const RoceMemDetails &md, const void *logicalPtr)
51 : {
52 2 : const u64 logicalVa = static_cast<u64>(reinterpret_cast<uintptr_t>(logicalPtr));
53 2 : const u64 offset = logicalVa - md.addr;
54 2 : const u64 devVa = md.devAddr + offset;
55 2 : return reinterpret_cast<void *>(static_cast<uintptr_t>(devVa));
56 : }
57 : } // namespace
58 :
59 : std::atomic<u64> TransportDeviceIbverbs::wrIdOffset_ = {0};
60 :
61 :
62 10 : TransportDeviceIbverbs::TransportDeviceIbverbs(DispatcherPub *dispatcher,
63 : const std::unique_ptr<NotifyPool> ¬ifyPool,
64 : MachinePara &machinePara,
65 : std::chrono::milliseconds timeout,
66 10 : const TransportDeviceIbverbsData &transDevIbverbsData)
67 : : TransportIbverbs(dispatcher, notifyPool, machinePara, timeout),
68 10 : transDevIbverbsData_(transDevIbverbsData)
69 : {
70 10 : }
71 :
72 11 : TransportDeviceIbverbs::~TransportDeviceIbverbs()
73 : {
74 10 : HCCL_DEBUG("~TransportDeviceIbverbs Enter!");
75 :
76 10 : (void)DeInit();
77 :
78 10 : if (machinePara_.deviceLogicId >= 0 && (static_cast<u32>(machinePara_.deviceLogicId) < MAX_MODULE_DEVICE_NUM)) {
79 10 : if ( instanceRef_[machinePara_.deviceLogicId].Unref() == 0) {
80 3 : std::unique_lock<std::mutex> lock(notifyValueMutex_[machinePara_.deviceLogicId]);
81 3 : notifyValueMem_[machinePara_.deviceLogicId].free();
82 3 : }
83 : }
84 10 : HCCL_DEBUG("~TransportDeviceIbverbs Success!");
85 11 : }
86 :
87 1 : HcclResult TransportDeviceIbverbs::InitDrainNotifyInfo()
88 : {
89 1 : HCCL_DEBUG("[%s] RemoteNotifyAddr[%llu], remoteNotifyKey[%u], localDataNotifyAddr[%llu], localDataNotifyKey[%u]," \
90 : "notifySize[%u]", __func__, transDevIbverbsData_.remoteNotifyValueAddr,
91 : transDevIbverbsData_.remoteNotifyValueKey, transDevIbverbsData_.localDataNotifyAddr,
92 : transDevIbverbsData_.localDataNotifyKey, transDevIbverbsData_.notifySize);
93 1 : CHK_PRT_RET((transDevIbverbsData_.localDataNotifyAddr == 0 || transDevIbverbsData_.remoteNotifyValueAddr == 0),
94 : HCCL_ERROR("[%s] Notify addr is nullptr, RemoteNotifyAddr[%llu], localDataNotifyAddr[%llu]",
95 : transDevIbverbsData_.remoteNotifyValueAddr, transDevIbverbsData_.localDataNotifyAddr), HCCL_E_PTR);
96 1 : memMsg_[MemType::DATA_NOTIFY_MEM].addr = reinterpret_cast<void *>(transDevIbverbsData_.localDataNotifyAddr);
97 1 : memMsg_[MemType::DATA_NOTIFY_MEM].lkey = transDevIbverbsData_.localDataNotifyKey;
98 1 : memMsg_[MemType::DATA_NOTIFY_MEM].len = transDevIbverbsData_.notifySize;
99 1 : remoteMemMsg_[MemType::NOTIFY_SRC_MEM].addr = reinterpret_cast<void *>(transDevIbverbsData_.remoteNotifyValueAddr);
100 1 : remoteMemMsg_[MemType::NOTIFY_SRC_MEM].lkey = transDevIbverbsData_.remoteNotifyValueKey;
101 1 : remoteMemMsg_[MemType::NOTIFY_SRC_MEM].len = transDevIbverbsData_.notifySize;
102 1 : CHK_RET(SignalInit(transDevIbverbsData_.dataNotify, dataNotify_));
103 1 : return HCCL_SUCCESS;
104 : }
105 :
106 6 : HcclResult TransportDeviceIbverbs::Init()
107 : {
108 6 : HCCL_DEBUG("TransportDeviceIbverbs Init Enter! notifyNum[%u]", machinePara_.notifyNum);
109 6 : if (transDevIbverbsData_.useMemDetailsMgr) {
110 6 : return InitMemDetails();
111 : }
112 0 : CHK_RET(SignalInit(transDevIbverbsData_.ackNotify, ackNotify_));
113 0 : CHK_RET(SignalInit(transDevIbverbsData_.dataNotify, dataNotify_));
114 0 : CHK_RET(SignalInit(transDevIbverbsData_.dataAckNotify, dataAckNotify_));
115 0 : constexpr u32 QPINFO_SIZE_MAX = 33;
116 0 : constexpr u32 QPINFO_SIZE_MIN = 1;
117 0 : constexpr u32 QP_PERCONNECTION_MAX = 32;
118 0 : constexpr u32 QP_PERCONNECTION_MIN = 1;
119 0 : u32 qpInfoSize = transDevIbverbsData_.qpInfo.size();
120 0 : if (transDevIbverbsData_.qpsPerConnection + static_cast<u32>(qpInfoSize > 1) != qpInfoSize ||
121 0 : qpInfoSize > QPINFO_SIZE_MAX || qpInfoSize < QPINFO_SIZE_MIN ||
122 0 : transDevIbverbsData_.qpsPerConnection > QP_PERCONNECTION_MAX ||
123 0 : transDevIbverbsData_.qpsPerConnection < QP_PERCONNECTION_MIN) {
124 0 : HCCL_ERROR("[TransportDeviceIbverbs][Init]QPNum[%d] or qpInfos size[%u] is invalid",
125 : transDevIbverbsData_.qpsPerConnection,
126 : qpInfoSize);
127 0 : return HCCL_E_INTERNAL;
128 : }
129 0 : combineAiQpInfo_.aiQpInfo.aiQpAddr = transDevIbverbsData_.qpInfo[0].qpPtr;
130 0 : combineAiQpInfo_.aiQpInfo.sqIndex = transDevIbverbsData_.qpInfo[0].sqIndex;
131 0 : combineAiQpInfo_.aiQpInfo.dbIndex = transDevIbverbsData_.qpInfo[0].dbIndex;
132 0 : combineAiQpInfos_.resize(transDevIbverbsData_.qpsPerConnection);
133 0 : for (u32 i = 1, j = 0; i < qpInfoSize; i++, j++) {
134 0 : combineAiQpInfos_[j].aiQpInfo.aiQpAddr = transDevIbverbsData_.qpInfo[i].qpPtr;
135 0 : combineAiQpInfos_[j].aiQpInfo.sqIndex = transDevIbverbsData_.qpInfo[i].sqIndex;
136 0 : combineAiQpInfos_[j].aiQpInfo.dbIndex = transDevIbverbsData_.qpInfo[i].dbIndex;
137 0 : HCCL_DEBUG("TransportDeviceIbverbs Init multiQp[%u], aiQpAddr[%llu] sqIndex[%u] dbIndex[%u]",
138 : j,
139 : transDevIbverbsData_.qpInfo[i].qpPtr,
140 : transDevIbverbsData_.qpInfo[i].sqIndex,
141 : transDevIbverbsData_.qpInfo[i].dbIndex);
142 : }
143 0 : notifySize_ = transDevIbverbsData_.notifySize;
144 0 : remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].addr = transDevIbverbsData_.inputBufferPtr;
145 0 : remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].lkey = transDevIbverbsData_.remoteInputKey;
146 :
147 0 : remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].addr = transDevIbverbsData_.outputBufferPtr;
148 0 : remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].lkey = transDevIbverbsData_.remoteOutputKey;
149 :
150 0 : u32 ackNotifyIdx = static_cast<u32>(MemType::ACK_NOTIFY_MEM);
151 0 : remoteMemMsg_[ackNotifyIdx].addr = reinterpret_cast<void *>(transDevIbverbsData_.remoteAckNotifyDetails.addr);
152 0 : remoteMemMsg_[ackNotifyIdx].notifyId = transDevIbverbsData_.remoteAckNotifyDetails.notifyId;
153 0 : remoteMemMsg_[ackNotifyIdx].lkey = transDevIbverbsData_.remoteAckNotifyDetails.key;
154 :
155 0 : u32 dataNotifyIdx = static_cast<u32>(MemType::DATA_NOTIFY_MEM);
156 0 : remoteMemMsg_[dataNotifyIdx].addr = reinterpret_cast<void *>(transDevIbverbsData_.remoteDataNotifyDetails.addr);
157 0 : remoteMemMsg_[dataNotifyIdx].notifyId = transDevIbverbsData_.remoteDataNotifyDetails.notifyId;
158 0 : remoteMemMsg_[dataNotifyIdx].lkey = transDevIbverbsData_.remoteDataNotifyDetails.key;
159 :
160 0 : u32 dataAckIdx = static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM);
161 0 : remoteMemMsg_[dataAckIdx].addr = reinterpret_cast<void *>(transDevIbverbsData_.remoteDataAckNotifyDetails.addr);
162 0 : remoteMemMsg_[dataAckIdx].notifyId = transDevIbverbsData_.remoteDataAckNotifyDetails.notifyId;
163 0 : remoteMemMsg_[dataAckIdx].lkey = transDevIbverbsData_.remoteDataAckNotifyDetails.key;
164 :
165 0 : HCCL_INFO("%s ACK:addr[0x%llx] notifyId[%d] lkey[%u], DATA:addr[0x%llx] notifyId[%d] lkey[%u], "\
166 : "DATA_ACK:addr[0x%llx] notifyId[%d] lkey[%u]", __func__,
167 : remoteMemMsg_[ackNotifyIdx].addr, remoteMemMsg_[ackNotifyIdx].notifyId, remoteMemMsg_[ackNotifyIdx].lkey,
168 : remoteMemMsg_[dataNotifyIdx].addr, remoteMemMsg_[dataNotifyIdx].notifyId, remoteMemMsg_[dataNotifyIdx].lkey,
169 : remoteMemMsg_[dataAckIdx].addr, remoteMemMsg_[dataAckIdx].notifyId, remoteMemMsg_[dataAckIdx].lkey);
170 :
171 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr =
172 0 : reinterpret_cast<void *>(transDevIbverbsData_.localNotifyValueAddr);
173 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey = transDevIbverbsData_.notifyValueKey;
174 0 : localInputMem_ = transDevIbverbsData_.localInputMem;
175 0 : memMsg_[MemType::USER_INPUT_MEM].addr = reinterpret_cast<void *>(transDevIbverbsData_.localInputMem.addr);
176 0 : memMsg_[MemType::USER_INPUT_MEM].len = transDevIbverbsData_.localInputMem.size;
177 0 : memMsg_[MemType::USER_INPUT_MEM].lkey = transDevIbverbsData_.localInputMem.key;
178 :
179 0 : localOutputMem_ = transDevIbverbsData_.localOutputMem;
180 0 : memMsg_[MemType::USER_OUTPUT_MEM].addr = reinterpret_cast<void *>(transDevIbverbsData_.localOutputMem.addr);
181 0 : memMsg_[MemType::USER_OUTPUT_MEM].len = transDevIbverbsData_.localOutputMem.size;
182 0 : memMsg_[MemType::USER_OUTPUT_MEM].lkey = transDevIbverbsData_.localOutputMem.key;
183 :
184 0 : notifyValueAddr_ = reinterpret_cast<void *>(transDevIbverbsData_.localNotifyValueAddr);
185 0 : CHK_RET(CheckDeviceId());
186 0 : CHK_RET(DlHnsFunction::GetInstance().DlHnsFunctionInit());
187 0 : transportAttr_.linkType = LinkType::LINK_ROCE;
188 0 : multiQpThreshold_ = transDevIbverbsData_.multiQpThreshold;
189 0 : qpsPerConnection_ = transDevIbverbsData_.qpsPerConnection;
190 0 : if (transDevIbverbsData_.userLocalNotify.size() != qpsPerConnection_ ||
191 0 : transDevIbverbsData_.userRemoteNotifyDetails.size() != qpsPerConnection_) {
192 0 : HCCL_ERROR("[TransportDeviceIbverbs][Init]userLocalNotify size[%u] is not equal to qpsPerConnection[%u]",
193 : transDevIbverbsData_.userLocalNotify.size(),
194 : qpsPerConnection_);
195 0 : return HCCL_E_INTERNAL;
196 : }
197 :
198 0 : userMultiQpLocalNotify_.resize(transDevIbverbsData_.qpsPerConnection);
199 0 : u32 multiQpExtNotifyLength = transDevIbverbsData_.qpsPerConnection > 1 ? transDevIbverbsData_.qpsPerConnection: 0;
200 0 : multiQpDataNotify_.resize(multiQpExtNotifyLength);
201 0 : for (u32 i = 0; i < transDevIbverbsData_.qpsPerConnection; ++i) {
202 0 : CHK_PRT_RET(transDevIbverbsData_.userLocalNotify[i].empty() && transDevIbverbsData_.qpsPerConnection > 1,
203 : HCCL_ERROR("[TransportDeviceIbverbs][Init]userLocalNotify[%u] is empty, qpsPerConnection[%u]",
204 : i,
205 : transDevIbverbsData_.qpsPerConnection),
206 : HCCL_E_INTERNAL);
207 0 : u32 singleQpNotifyNum = transDevIbverbsData_.qpsPerConnection > 1
208 0 : ? transDevIbverbsData_.userLocalNotify[i].size() - 1
209 0 : : transDevIbverbsData_.userLocalNotify[i].size();
210 0 : CHK_PRT_RET(singleQpNotifyNum != notifyNum_,
211 : HCCL_ERROR(
212 : "[TransportDeviceIbverbs][Init] qpIdx[%u] userLocalNotify notifynum[%u] is not equal to notifyNum_[%u]",
213 : i,
214 : singleQpNotifyNum,
215 : notifyNum_),
216 : HCCL_E_INTERNAL);
217 0 : userMultiQpLocalNotify_[i].resize(singleQpNotifyNum);
218 0 : for (u32 j = 0; j < singleQpNotifyNum; ++j) {
219 0 : CHK_RET(SignalInit(transDevIbverbsData_.userLocalNotify[i][j], userMultiQpLocalNotify_[i][j]));
220 : }
221 0 : if (transDevIbverbsData_.qpsPerConnection > 1) {
222 0 : CHK_RET(SignalInit(transDevIbverbsData_.userLocalNotify[i][singleQpNotifyNum], multiQpDataNotify_[i]));
223 : }
224 : }
225 :
226 0 : userMultiQpRemoteNotifyMsg_.resize(transDevIbverbsData_.qpsPerConnection);
227 0 : multiQpDataNotifyRemoteMemMsg_.resize(multiQpExtNotifyLength);
228 0 : for (u32 i = 0; i < transDevIbverbsData_.qpsPerConnection; ++i) {
229 0 : CHK_PRT_RET(transDevIbverbsData_.userRemoteNotifyDetails[i].empty() && transDevIbverbsData_.qpsPerConnection > 1,
230 : HCCL_ERROR("[TransportDeviceIbverbs][Init]userLocalNotify[%u] is empty, qpsPerConnection[%u]",
231 : i,
232 : transDevIbverbsData_.qpsPerConnection),
233 : HCCL_E_INTERNAL);
234 0 : u32 singleQpNotifyNum = transDevIbverbsData_.qpsPerConnection > 1
235 0 : ? transDevIbverbsData_.userRemoteNotifyDetails[i].size() - 1
236 0 : : transDevIbverbsData_.userRemoteNotifyDetails[i].size();
237 0 : CHK_PRT_RET(singleQpNotifyNum != notifyNum_,
238 : HCCL_ERROR(
239 : "[TransportDeviceIbverbs][Init] qpIdx[%u] userLocalNotify notifynum[%u] is not equal to notifyNum_[%u]",
240 : i,
241 : singleQpNotifyNum,
242 : notifyNum_),
243 : HCCL_E_INTERNAL);
244 0 : userMultiQpRemoteNotifyMsg_[i].resize(singleQpNotifyNum);
245 0 : u32 j = 0;
246 0 : for (; j < singleQpNotifyNum; ++j) {
247 0 : userMultiQpRemoteNotifyMsg_[i][j].addr =
248 0 : reinterpret_cast<void *>(transDevIbverbsData_.userRemoteNotifyDetails[i][j].addr);
249 0 : userMultiQpRemoteNotifyMsg_[i][j].notifyId = transDevIbverbsData_.userRemoteNotifyDetails[i][j].notifyId;
250 0 : userMultiQpRemoteNotifyMsg_[i][j].lkey = transDevIbverbsData_.userRemoteNotifyDetails[i][j].key;
251 0 : HCCL_INFO("userMultiQpRemoteNotifyMsg_[%u][%u] addr[0x%llx] notifyId[%u] lkey[%u]", i, j,
252 : userMultiQpRemoteNotifyMsg_[i][j].addr, userMultiQpRemoteNotifyMsg_[i][j].notifyId,
253 : userMultiQpRemoteNotifyMsg_[i][j].lkey);
254 : }
255 0 : if (transDevIbverbsData_.qpsPerConnection > 1) {
256 0 : multiQpDataNotifyRemoteMemMsg_[i].addr =
257 0 : reinterpret_cast<void *>(transDevIbverbsData_.userRemoteNotifyDetails[i][j].addr);
258 0 : multiQpDataNotifyRemoteMemMsg_[i].notifyId = transDevIbverbsData_.userRemoteNotifyDetails[i][j].notifyId;
259 0 : multiQpDataNotifyRemoteMemMsg_[i].lkey = transDevIbverbsData_.userRemoteNotifyDetails[i][j].key;
260 0 : HCCL_INFO("multiQpDataNotifyRemoteMemMsg_[%u] addr[0x%llx] notifyId[%u] lkey[%u]",
261 : i, multiQpDataNotifyRemoteMemMsg_[i].addr, multiQpDataNotifyRemoteMemMsg_[i].notifyId,
262 : multiQpDataNotifyRemoteMemMsg_[i].lkey);
263 : }
264 : }
265 0 : useAtomicWrite_ = transDevIbverbsData_.useAtomicWrite;
266 0 : HCCL_USER_CRITICAL_LOG("create hccl transport:communicator[%s], local rank[%u], remote rank[%u],"\
267 : "transporttype[%s], atomicWrite[%d]", machinePara_.tag.c_str(), machinePara_.localUserrank,
268 : machinePara_.remoteUserrank, GetLinkTypeEnumStr(GetLinkType()).c_str(), useAtomicWrite_);
269 :
270 0 : return HCCL_SUCCESS;
271 : }
272 :
273 6 : HcclResult TransportDeviceIbverbs::InitMemDetails()
274 : {
275 6 : constexpr u32 QPINFO_SIZE_MAX = 33;
276 6 : constexpr u32 QPINFO_SIZE_MIN = 1;
277 6 : constexpr u32 QP_PERCONNECTION_MAX = 32;
278 6 : constexpr u32 QP_PERCONNECTION_MIN = 1;
279 6 : u32 qpSize = transDevIbverbsData_.qpInfo.size();
280 6 : if (transDevIbverbsData_.qpsPerConnection + static_cast<u32>(qpSize > 1) != qpSize ||
281 5 : qpSize > QPINFO_SIZE_MAX || qpSize < QPINFO_SIZE_MIN ||
282 5 : transDevIbverbsData_.qpsPerConnection > QP_PERCONNECTION_MAX ||
283 5 : transDevIbverbsData_.qpsPerConnection < QP_PERCONNECTION_MIN) {
284 1 : HCCL_ERROR("[TransportDeviceIbverbs][InitMemDetails]QPNum[%d] or qpInfos size[%u] is invalid",
285 : transDevIbverbsData_.qpsPerConnection,
286 : qpSize);
287 1 : return HCCL_E_INTERNAL;
288 : }
289 5 : combineAiQpInfo_.aiQpInfo.aiQpAddr = transDevIbverbsData_.qpInfo[0].qpPtr;
290 5 : combineAiQpInfo_.aiQpInfo.sqIndex = transDevIbverbsData_.qpInfo[0].sqIndex;
291 5 : combineAiQpInfo_.aiQpInfo.dbIndex = transDevIbverbsData_.qpInfo[0].dbIndex;
292 5 : combineAiQpInfos_.resize(transDevIbverbsData_.qpsPerConnection);
293 5 : for (u32 i = 1, j = 0; i < qpSize; i++, j++) {
294 0 : combineAiQpInfos_[j].aiQpInfo.aiQpAddr = transDevIbverbsData_.qpInfo[i].qpPtr;
295 0 : combineAiQpInfos_[j].aiQpInfo.sqIndex = transDevIbverbsData_.qpInfo[i].sqIndex;
296 0 : combineAiQpInfos_[j].aiQpInfo.dbIndex = transDevIbverbsData_.qpInfo[i].dbIndex;
297 : }
298 :
299 5 : CHK_RET(CheckDeviceId());
300 5 : CHK_RET(DlHnsFunction::GetInstance().DlHnsFunctionInit());
301 5 : transportAttr_.linkType = LinkType::LINK_ROCE;
302 5 : multiQpThreshold_ = transDevIbverbsData_.multiQpThreshold;
303 5 : qpsPerConnection_ = transDevIbverbsData_.qpsPerConnection;
304 5 : useAtomicWrite_ = transDevIbverbsData_.useAtomicWrite;
305 5 : HCCL_USER_CRITICAL_LOG("create hccl transport:communicator[%s], local rank[%u], remote rank[%u],"\
306 : "transporttype[%s], atomicWrite[%d]", machinePara_.tag.c_str(), machinePara_.localUserrank,
307 : machinePara_.remoteUserrank, GetLinkTypeEnumStr(GetLinkType()).c_str(), useAtomicWrite_);
308 5 : CHK_RET(BuildMemDetailsRmaMgrs());
309 5 : return HCCL_SUCCESS;
310 : }
311 :
312 5 : HcclResult TransportDeviceIbverbs::BuildMemDetailsRmaMgrs()
313 : {
314 5 : localMemDetailsRmaMgr_.reset();
315 5 : remoteMemDetailsRmaMgr_.reset();
316 5 : useMemDetailsLookup_ = false;
317 5 : localMemDetailsRmaMgr_ = std::make_unique<DeviceMemDetailsRmaMgr>();
318 5 : remoteMemDetailsRmaMgr_ = std::make_unique<DeviceMemDetailsRmaMgr>();
319 9 : for (const auto &md : transDevIbverbsData_.localRoceMemDetailsList) {
320 4 : if (md.size == 0U) {
321 0 : continue;
322 : }
323 4 : auto ent = std::make_shared<RoceMemDetails>(md);
324 4 : auto pr = localMemDetailsRmaMgr_->Add(MakeMemLookupKey(md.addr, md.size), ent);
325 4 : if (pr.first == localMemDetailsRmaMgr_->End()) {
326 0 : HCCL_ERROR("[TransportDeviceIbverbs][BuildMemDetailsRmaMgrs] add local mem range failed, "
327 : "logical[0x%llx, +%llu) devBase[0x%llx] key[%u]",
328 : static_cast<unsigned long long>(md.addr), static_cast<unsigned long long>(md.size),
329 : static_cast<unsigned long long>(md.devAddr), md.key);
330 0 : return HCCL_E_INTERNAL;
331 : }
332 4 : HCCL_DEBUG("[TransportDeviceIbverbs][BuildMemDetailsRmaMgrs] add local MR logical[0x%llx, +%llu) "
333 : "devBase[0x%llx] key[%u]",
334 : static_cast<unsigned long long>(md.addr), static_cast<unsigned long long>(md.size),
335 : static_cast<unsigned long long>(md.devAddr), md.key);
336 4 : }
337 9 : for (const auto &md : transDevIbverbsData_.remoteRoceMemDetailsList) {
338 4 : if (md.size == 0U) {
339 0 : continue;
340 : }
341 4 : auto ent = std::make_shared<RoceMemDetails>(md);
342 4 : auto pr = remoteMemDetailsRmaMgr_->Add(MakeMemLookupKey(md.addr, md.size), ent);
343 4 : if (pr.first == remoteMemDetailsRmaMgr_->End()) {
344 0 : HCCL_ERROR("[TransportDeviceIbverbs][BuildMemDetailsRmaMgrs] add remote mem range failed, "
345 : "logical[0x%llx, +%llu) devBase[0x%llx] key[%u]",
346 : static_cast<unsigned long long>(md.addr), static_cast<unsigned long long>(md.size),
347 : static_cast<unsigned long long>(md.devAddr), md.key);
348 0 : return HCCL_E_INTERNAL;
349 : }
350 4 : HCCL_DEBUG("[TransportDeviceIbverbs][BuildMemDetailsRmaMgrs] add remote MR logical[0x%llx, +%llu) "
351 : "devBase[0x%llx] key[%u]",
352 : static_cast<unsigned long long>(md.addr), static_cast<unsigned long long>(md.size),
353 : static_cast<unsigned long long>(md.devAddr), md.key);
354 4 : }
355 5 : HCCL_INFO("[TransportDeviceIbverbs][BuildMemDetailsRmaMgrs] indexed localMR[%zu] remoteMR[%zu]",
356 : localMemDetailsRmaMgr_->size(), remoteMemDetailsRmaMgr_->size());
357 5 : useMemDetailsLookup_ = true;
358 5 : return HCCL_SUCCESS;
359 : }
360 :
361 0 : HcclResult TransportDeviceIbverbs::AddWrList(void *dstMemPtr, const void *srcMemPtr, u64 srcMemSize,
362 : u32 srcKey, u32 dstKey, WqeType wqeType, WrAuxInfo &aux, std::vector<WrInformation> &wrInfoVec)
363 : {
364 0 : HCCL_DEBUG("TransportDeviceIbverbs AddWrList start");
365 0 : if (srcMemSize == 0) {
366 0 : return HCCL_SUCCESS;
367 : }
368 0 : WrInformation wrInfoTmp;
369 0 : wrInfoTmp.wrData.dstAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(dstMemPtr));
370 0 : wrInfoTmp.wrData.rkey = dstKey;
371 0 : wrInfoTmp.wrData.sendFlags = fence_ ? (RA_SEND_SIGNALED | RA_SEND_FENCE) : RA_SEND_SIGNALED;
372 0 : fence_ = false;
373 0 : wrInfoTmp.wrData.immData = 0;
374 0 : wrInfoTmp.wrData.wrId = 0;
375 0 : wrInfoTmp.wrData.memList.addr = static_cast<u64>(reinterpret_cast<uintptr_t>(srcMemPtr));
376 0 : wrInfoTmp.wrData.memList.len = srcMemSize;
377 0 : wrInfoTmp.wrData.memList.lkey = srcKey;
378 :
379 0 : switch (wqeType) {
380 0 : case WqeType::WQE_TYPE_DATA:
381 : case WqeType::WQE_TYPE_DATA_NOTIFY:
382 : case WqeType::WQE_TYPE_ACK_NOTIFY:
383 : case WqeType::WQE_TYPE_DATA_ACK_NOTIFY:
384 : case WqeType::WQE_TYPE_DATA_WITH_NOTIFY:
385 0 : wrInfoTmp.wrData.op = RA_WR_RDMA_WRITE;
386 0 : wrInfoTmp.type = static_cast<u64>(wqeType);
387 0 : break;
388 0 : case WqeType::WQE_TYPE_DATA_WITH_REDUCE:
389 0 : wrInfoTmp.wrData.op = RA_WR_RDMA_REDUCE_WRITE;
390 0 : wrInfoTmp.wrData.aux = aux;
391 : // REDUCE WRITE 作为特殊的DATA
392 0 : wrInfoTmp.type = static_cast<u64>(WqeType::WQE_TYPE_DATA);
393 0 : break;
394 0 : case WqeType::WQE_TYPE_READ_DATA:
395 0 : wrInfoTmp.wrData.op = RA_WR_RDMA_READ;
396 0 : wrInfoTmp.type = static_cast<u64>(wqeType);
397 0 : break;
398 0 : default:
399 0 : HCCL_ERROR("error wqeType[%d]", wqeType);
400 0 : return HCCL_E_INTERNAL;
401 : }
402 0 : CHK_RET(GetWrDataAddr(dstMemPtr, wqeType, wrInfoTmp.wrDataAddr, wrInfoTmp.notifyId));
403 0 : HCCL_DEBUG("wrInfoTmp dst_addr[0x%llx] memList addr[0x%llx] len[%llu]", wrInfoTmp.wrData.dstAddr,
404 : wrInfoTmp.wrData.memList.addr, srcMemSize);
405 0 : wrInfoVec.push_back(wrInfoTmp);
406 0 : HCCL_DEBUG("TransportDeviceIbverbs AddWrList end");
407 0 : return HCCL_SUCCESS;
408 : }
409 :
410 0 : HcclResult TransportDeviceIbverbs::GetMemInfo(UserMemType memType, void **dstMemPtr, unsigned int *dstKey,
411 : u64 &dstMemSize)
412 : {
413 0 : CHK_PTR_NULL(dstMemPtr);
414 0 : CHK_PTR_NULL(dstKey);
415 :
416 0 : switch (memType) {
417 0 : case UserMemType::INPUT_MEM: {
418 0 : *dstMemPtr = remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].addr;
419 0 : dstMemSize = remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].len;
420 0 : *dstKey = remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].lkey;
421 0 : break;
422 : }
423 :
424 0 : case UserMemType::OUTPUT_MEM: {
425 0 : *dstMemPtr = remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].addr;
426 0 : dstMemSize = remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].len;
427 0 : *dstKey = remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].lkey;
428 0 : break;
429 : }
430 :
431 0 : default: {
432 0 : HCCL_ERROR("[Get][MemInfo]not support dst_mem_type=%d", memType);
433 0 : return HCCL_E_NOT_SUPPORT;
434 : }
435 : }
436 0 : return HCCL_SUCCESS;
437 : }
438 :
439 0 : HcclResult TransportDeviceIbverbs::ConstructPayLoadWqe(void *dstMemPtr, u32 dstKey, const void *src,
440 : u32 srcKey, u64 len, WqeType wqeType, WrAuxInfo &aux, std::vector<WrInformation> &wrInfoVec,
441 : u32 txSendDataTimes)
442 : {
443 : HcclResult ret;
444 : // 发送数据Wqe
445 0 : for (u32 txSendDataIdx = 0; txSendDataIdx < txSendDataTimes; txSendDataIdx++) {
446 0 : u64 txSendDataOffset = txSendDataIdx * RDMA_SEND_MAX_SIZE;
447 0 : u64 txSendDataSize = (txSendDataIdx == (txSendDataTimes - 1)) ? len - txSendDataOffset : RDMA_SEND_MAX_SIZE;
448 :
449 0 : void* txdstMemPtr = reinterpret_cast<void *>(reinterpret_cast<char *>(dstMemPtr) +
450 : txSendDataOffset);
451 :
452 0 : const void* txsrcMemPtr = reinterpret_cast<const void *>(reinterpret_cast<const char *>(src) +
453 : txSendDataOffset);
454 0 : ret = AddWrList(txdstMemPtr, txsrcMemPtr, txSendDataSize, srcKey, dstKey, wqeType, aux, wrInfoVec);
455 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
456 : HCCL_ERROR("[TransportDeviceIbverbs][TxAsync]errNo[0x%016llx] In lbv exp, add wqe list failed."\
457 : "srcMemSize[%llu]", HCCL_ERROR_CODE(ret), txSendDataSize), ret);
458 : }
459 0 : HCCL_DEBUG("TransportDeviceIbverbs TxPayLoad end");
460 :
461 0 : return HCCL_SUCCESS;
462 : }
463 :
464 0 : HcclResult TransportDeviceIbverbs::TxPayLoad(UserMemType dstMemType, u64 dstOffset, const void *src, u64 len,
465 : WqeType wqeType, WrAuxInfo &aux, std::vector<WrInformation>& wrInfoVec)
466 : {
467 0 : HCCL_DEBUG("TransportDeviceIbverbs TxPayLoad start");
468 0 : void *dstMemPtr = nullptr;
469 : unsigned int dstKey;
470 : unsigned int srcKey;
471 0 : u64 dstMemSize = 0;
472 : // 为保证单算子下不同数据量下子图的结构相同,zero byte message 时也需要下发task
473 0 : u32 txSendDataTimes = (len + RDMA_SEND_MAX_SIZE - 1) / RDMA_SEND_MAX_SIZE;
474 :
475 : // 当前len不可用,无法校验dstOffset > dstMemSize
476 0 : CHK_RET(GetMemInfo(dstMemType, &dstMemPtr, &dstKey, dstMemSize));
477 :
478 0 : u64 srcAddr = reinterpret_cast<u64>(src);
479 0 : if (srcAddr >= localInputMem_.addr && srcAddr < localInputMem_.addr + localInputMem_.size) {
480 0 : srcKey = localInputMem_.key;
481 0 : } else if (srcAddr >= localOutputMem_.addr && srcAddr <= localOutputMem_.addr + localOutputMem_.size) {
482 0 : srcKey = localOutputMem_.key;
483 : } else {
484 0 : HCCL_ERROR("[TransportDeviceIbverbs][TxAsync]src_ptr=%p is out of range, inputmem src[%p], size[%llu];"
485 : " outputmem src[%p] size[%llu]", src, localInputMem_.addr, localInputMem_.size,
486 : localOutputMem_.addr, localOutputMem_.size);
487 0 : return HCCL_E_INTERNAL;
488 : }
489 :
490 0 : dstMemPtr = reinterpret_cast<void *>(reinterpret_cast<char *>(dstMemPtr) + dstOffset);
491 0 : CHK_RET(ConstructPayLoadWqe(dstMemPtr, dstKey, src, srcKey, len, wqeType, aux, wrInfoVec, txSendDataTimes));
492 :
493 0 : return HCCL_SUCCESS;
494 : }
495 :
496 0 : HcclResult TransportDeviceIbverbs::TxAsync(UserMemType dstMemType, u64 dstOffset,
497 : const void *src, u64 len, Stream &stream)
498 : {
499 0 : CHK_SMART_PTR_NULL(stream);
500 0 : std::vector<WrInformation> wrInfoVec;
501 0 : struct WrAuxInfo aux = {0};
502 0 : HCCL_DEBUG("TX src[%p] len[%llu] dstOffset[%llu]", src, len, dstOffset);
503 :
504 0 : if (len > 0) {
505 0 : CHK_PTR_NULL(src);
506 0 : CHK_RET(TxPayLoad(dstMemType, dstOffset, src, len, WqeType::WQE_TYPE_DATA, aux, wrInfoVec));
507 : }
508 :
509 0 : CHK_RET(TxSendDataAndNotify(wrInfoVec, stream, GetUseOneDoorbellValue()));
510 0 : return HCCL_SUCCESS;
511 0 : }
512 :
513 0 : HcclResult TransportDeviceIbverbs::TxWithReduce(UserMemType dstMemType, u64 dstOffset, const void *src, u64 len,
514 : const HcclDataType datatype, HcclReduceOp redOp, Stream &stream)
515 : {
516 0 : CHK_SMART_PTR_NULL(stream);
517 0 : std::vector<WrInformation> wrInfoVec;
518 0 : struct WrAuxInfo aux = {0};
519 0 : aux.dataType = RDMA_REDUCE_DATA_TYPE_TABLE[datatype];
520 0 : aux.reduceType = RDMA_REDUCE_OP_TYPE_TABLE[redOp];
521 0 : if (aux.dataType == static_cast<uint8_t>(RdmaReduceDataType::RDMA_REDUCE_DATA_INVALID) ||
522 0 : aux.reduceType == static_cast<uint8_t>(RdmaReduceOpType::RDMA_REDUCE_OP_INVALID)) {
523 0 : HCCL_ERROR("unsupported data type [%s] or Reduce type [%s]",
524 : GetDataTypeEnumStr(datatype).c_str(), GetReduceOpEnumStr(redOp).c_str());
525 0 : return HCCL_E_INTERNAL;
526 : }
527 0 : if (len > 0) {
528 0 : CHK_PTR_NULL(src);
529 0 : CHK_RET(TxPayLoad(dstMemType, dstOffset, src, len, WqeType::WQE_TYPE_DATA_WITH_REDUCE, aux, wrInfoVec));
530 : }
531 :
532 0 : CHK_RET(TxSendDataAndNotify(wrInfoVec, stream, GetUseOneDoorbellValue()));
533 0 : return HCCL_SUCCESS;
534 0 : }
535 :
536 0 : HcclResult TransportDeviceIbverbs::TxWithReduce(const std::vector<TxMemoryInfo> &txWithReduceMems,
537 : const HcclDataType datatype, HcclReduceOp redOp, Stream &stream)
538 : {
539 0 : CHK_SMART_PTR_NULL(stream);
540 0 : std::vector<WrInformation> wrInfoVec;
541 0 : struct WrAuxInfo aux = {0};
542 0 : aux.dataType = RDMA_REDUCE_DATA_TYPE_TABLE[datatype];
543 0 : aux.reduceType = RDMA_REDUCE_OP_TYPE_TABLE[redOp];
544 0 : if (aux.dataType == static_cast<uint8_t>(RdmaReduceDataType::RDMA_REDUCE_DATA_INVALID) ||
545 0 : aux.reduceType == static_cast<uint8_t>(RdmaReduceOpType::RDMA_REDUCE_OP_INVALID)) {
546 0 : HCCL_ERROR("unsupported data type [%s] or Reduce type [%s]",
547 : GetDataTypeEnumStr(datatype).c_str(), GetReduceOpEnumStr(redOp).c_str());
548 0 : return HCCL_E_INTERNAL;
549 : }
550 :
551 0 : for (const TxMemoryInfo &txWithReduceMem : txWithReduceMems) {
552 0 : if (txWithReduceMem.len == 0) {
553 0 : continue;
554 : }
555 0 : CHK_PTR_NULL(txWithReduceMem.src);
556 0 : CHK_RET(TxPayLoad(txWithReduceMem.dstMemType, txWithReduceMem.dstOffset, txWithReduceMem.src,
557 : txWithReduceMem.len, WqeType::WQE_TYPE_DATA_WITH_REDUCE, aux, wrInfoVec));
558 : }
559 :
560 0 : CHK_RET(TxSendDataAndNotify(wrInfoVec, stream, GetUseOneDoorbellValue()));
561 0 : return HCCL_SUCCESS;
562 0 : }
563 :
564 0 : HcclResult TransportDeviceIbverbs::TxSendDataAndNotifyWithSingleQP(
565 : std::vector<WrInformation> &wrInfoVec, Stream &stream, bool useOneDoorbell)
566 : {
567 : // 发送data notify同步信息
568 0 : struct WrAuxInfo aux = {0};
569 0 : void *remoteNotifyaddr = remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].addr;;
570 0 : CHK_RET(AddWrList(remoteNotifyaddr, notifyValueAddr_, notifySize_,
571 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
572 : remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].lkey,
573 : WqeType::WQE_TYPE_DATA_NOTIFY, aux, wrInfoVec));
574 :
575 0 : CHK_RET(RdmaSendAsync(wrInfoVec, stream, useOneDoorbell));
576 0 : return HCCL_SUCCESS;
577 : }
578 :
579 0 : HcclResult TransportDeviceIbverbs::TxSendDataAndNotify(std::vector<WrInformation> &wrInfoVec,
580 : Stream &stream, bool useOneDoorbell)
581 : {
582 0 : u32 maxLength = 0;
583 0 : for (u32 i = 0; i < wrInfoVec.size(); i++) {
584 0 : if (wrInfoVec[i].wrData.memList.len > maxLength) {
585 0 : maxLength = wrInfoVec[i].wrData.memList.len;
586 : }
587 : }
588 0 : u32 actualMultiQpNum = GetActualQpNum(maxLength);
589 0 : HCCL_DEBUG("[TransportDeviceIbverbs][TxSendDataAndNotify] UseMultiQp[%d] MultiQpNum[%u] actualMultiQpNum[%u] "
590 : "maxLength[%u]",
591 : UseMultiQp(),
592 : qpsPerConnection_,
593 : actualMultiQpNum,
594 : maxLength);
595 0 : if (UseMultiQp() && actualMultiQpNum != 1 && actualMultiQpNum <= qpsPerConnection_ && maxLength != 0) {
596 0 : CHK_RET(TxSendDataAndNotifyWithMultiQP(wrInfoVec, actualMultiQpNum, stream, useOneDoorbell));
597 : } else {
598 0 : CHK_RET(TxSendDataAndNotifyWithSingleQP(wrInfoVec, stream, useOneDoorbell));
599 : }
600 0 : return HCCL_SUCCESS;
601 : }
602 :
603 0 : HcclResult TransportDeviceIbverbs::TxAsync(std::vector<TxMemoryInfo>& txMems, Stream &stream)
604 : {
605 0 : CHK_SMART_PTR_NULL(stream);
606 :
607 0 : std::vector<WrInformation> wrInfoVec;
608 0 : struct WrAuxInfo aux = {0};
609 :
610 0 : for (auto& mem : txMems) {
611 0 : HCCL_DEBUG("TX src[%p] len[%llu] dstOffset[%llu]", mem.src, mem.len, mem.dstOffset);
612 0 : if (mem.len == 0) {
613 0 : continue;
614 : }
615 0 : CHK_PTR_NULL(mem.src);
616 0 : CHK_RET(TxPayLoad(mem.dstMemType, mem.dstOffset, mem.src, mem.len, WqeType::WQE_TYPE_DATA, aux, wrInfoVec));
617 : }
618 :
619 0 : CHK_RET(TxSendDataAndNotify(wrInfoVec, stream, GetUseOneDoorbellValue()));
620 0 : return HCCL_SUCCESS;
621 0 : }
622 :
623 0 : HcclResult TransportDeviceIbverbs::TxWrList(std::vector<WrInformation> &wrInfoVec, Stream &stream,
624 : std::vector<struct SendWrRsp> &opRspVec, u32 multiQpIndex)
625 : {
626 : (void)stream;
627 :
628 0 : u32 totalWqeCount = wrInfoVec.size();
629 0 : WrInformation *wrlist = wrInfoVec.data();
630 0 : struct SendWrRsp *opRsp = opRspVec.data();
631 :
632 : // HCCP会校验 zero byte messages 的内存地址是否已注册MR。对于 zero byte messages 不下发WR,将opRsp设置为特殊值。
633 : // 下发rdmasend task时检查该特殊值,如果zero byte message则不下发rdmasend task。
634 0 : bool batchSendWr = true;
635 0 : for (u32 i = 0; i < totalWqeCount; i++) {
636 0 : if (wrInfoVec[i].wrData.memList.len == 0) {
637 0 : batchSendWr = false;
638 0 : break;
639 : }
640 : }
641 :
642 0 : if (batchSendWr) {
643 0 : CHK_RET(SendWrList(totalWqeCount, wrlist, opRsp, multiQpIndex));
644 : } else {
645 0 : for (u32 i = 0; i < totalWqeCount; i++) {
646 0 : if (wrInfoVec[i].wrData.memList.len > 0) {
647 0 : CHK_RET(SendWrList(1U, &wrlist[i], &opRsp[i], multiQpIndex));
648 : } else {
649 0 : opRsp[i].wqeTmp.sqIndex = INVALID_UINT;
650 0 : opRsp[i].wqeTmp.wqeIndex = INVALID_UINT;
651 0 : opRsp[i].db.dbIndex = INVALID_UINT;
652 0 : opRsp[i].db.dbInfo = INVALID_U64;
653 : }
654 : }
655 : }
656 :
657 0 : return HCCL_SUCCESS;
658 : }
659 :
660 1 : HcclResult TransportDeviceIbverbs::SendWrList(
661 : u32 wrNum, WrInformation *wrlist, struct SendWrRsp *opRsp, u32 multiQpIndex)
662 : {
663 1 : unsigned int completeNum = 0;
664 1 : HcclResult ret = SendWrlistExt(wrlist, opRsp, wrNum, &completeNum, multiQpIndex);
665 1 : CHK_PRT_RET(ret != HCCL_SUCCESS,
666 : HCCL_ERROR("[TransportDeviceIbverbs][SendWrList]In ibv send wq list, SendWrlistExt failed.ret[%d]", ret),
667 : HCCL_E_NETWORK);
668 1 : return HCCL_SUCCESS;
669 : }
670 :
671 0 : HcclResult TransportDeviceIbverbs::SendWrlistExt(WrInformation wr[], struct SendWrRsp opRsp[], unsigned int sendNum,
672 : unsigned int *completeNum, u32 multiQpIndex)
673 : {
674 0 : HcclResult ret = HCCL_SUCCESS;
675 0 : auto startTime = std::chrono::steady_clock::now();
676 0 : u32 remainNum = sendNum;
677 0 : unsigned int completeNumLocal = 0;
678 0 : *completeNum = 0;
679 : while (true) {
680 0 : if (remainNum > sendNum) {
681 0 : HCCL_ERROR("[Aicpu][Send][Wr]wr list send async fail. return[%d], remainNum[%u], "\
682 : "sendNum[%u].", HCCL_E_ROCE_TRANSFER, remainNum, sendNum);
683 0 : return HCCL_E_ROCE_TRANSFER;
684 : }
685 0 : if (remainNum == 0) {
686 0 : break;
687 : }
688 0 : ret = TxSendWrlistExt(
689 0 : wr + (sendNum - remainNum), remainNum, opRsp + (sendNum - remainNum),
690 : &completeNumLocal, multiQpIndex);
691 0 : *completeNum += completeNumLocal;
692 0 : if (ret == HCCL_SUCCESS && *completeNum == sendNum) {
693 0 : break; // 成功跳出
694 : }
695 :
696 0 : if (ret == HCCL_E_AGAIN || *completeNum < sendNum) {
697 0 : remainNum -= completeNumLocal;
698 0 : bool bTimeout = ((std::chrono::steady_clock::now() - startTime) >= timeout_);
699 0 : CHK_PRT_RET(bTimeout, HCCL_ERROR("[Aicpu][Send][Wr]errNo[0x%016llx] wrlist send async timeout[%d]ms. "\
700 : "return[%d], params: send_wrAddr[%p], opRspAddr[%p]",
701 : HCCL_ERROR_CODE(HCCL_E_ROCE_TRANSFER), timeout_, ret, wr, opRsp), HCCL_E_ROCE_TRANSFER);
702 0 : SaluSleep(ONE_MILLISECOND_OF_USLEEP);
703 0 : } else {
704 0 : HCCL_ERROR("[Aicpu][Send][Wr]wrlist send async fail. return[%d], para: send_wrAddr[%p], "\
705 : "opRspAddr[%p].", ret, wr, opRsp);
706 0 : return HCCL_E_ROCE_TRANSFER; // 非-2/-11场景错误,不轮询,直接退出
707 : }
708 0 : }
709 0 : return HCCL_SUCCESS;
710 : }
711 :
712 0 : HcclResult TransportDeviceIbverbs::TxSendWrlistExt(WrInformation wrList[], u32 sendNum,
713 : struct SendWrRsp opRsp[], unsigned int *completeNum, u32 multiQpIndex)
714 : {
715 0 : u32 i = 0;
716 0 : s32 ret = 0;
717 0 : struct ibv_send_wr ib_wr = {0};
718 0 : struct ibv_sge list = {0};
719 0 : struct ibv_send_wr *bad_wr = nullptr;
720 0 : struct WrExpRsp exp_rsp = {0};
721 0 : struct IbvPostSendExtResp ext_rsp = {0};
722 0 : struct IbvPostSendExtAddt ext_attr = {0};
723 0 : for (; i < sendNum; i++) {
724 0 : if (wrList[i].wrData.memList.len > IBV_SGLIST_LEN_MAX) {
725 0 : HCCL_ERROR("sg list len is more than 2G, len[%u]", wrList[i].wrData.memList.len);
726 0 : return HCCL_E_PARA;
727 : }
728 0 : u64 wrIdoffset = wrIdOffset_++;
729 :
730 : // 910B和910_93,reduce的下一个notify要设置为atomic write
731 : u32& preWrOpcode = multiQpIndex == RDMA_INVALID_QP_INDEX ?
732 0 : combineAiQpInfo_.preWrOpcode : combineAiQpInfos_[multiQpIndex].preWrOpcode;
733 0 : ModifyAtomicWriteAfterReduce(preWrOpcode, wrList[i].type, wrList[i].wrData.op, wrList[i].wrData.immData);
734 :
735 0 : if (wrList[i].wrData.op != RA_WR_SEND && wrList[i].wrData.op != RA_WR_SEND_WITH_IMM) {
736 0 : HCCL_DEBUG("remote wr dst addr is 0x%llx", wrList[i].wrData.dstAddr);
737 0 : list.addr = wrList[i].wrData.memList.addr;
738 0 : list.length = wrList[i].wrData.memList.len;
739 0 : list.lkey = wrList[i].wrData.memList.lkey;
740 :
741 0 : ib_wr.sg_list = &list;
742 0 : ib_wr.opcode = static_cast<enum ibv_wr_opcode>(wrList[i].wrData.op);
743 0 : ib_wr.send_flags = static_cast<unsigned int>(wrList[i].wrData.sendFlags);
744 0 : ib_wr.imm_data = wrList[i].wrData.immData;
745 :
746 0 : ib_wr.num_sge = 1; /* only support one sge */
747 0 : ib_wr.wr_id = wrList[i].wrData.wrId += wrIdoffset;
748 0 : ib_wr.wr.rdma.rkey = wrList[i].wrData.rkey;
749 0 : ib_wr.wr.rdma.remote_addr = wrList[i].wrData.dstAddr;
750 : } else {
751 0 : list.addr = wrList[i].wrData.memList.addr;
752 0 : list.length = wrList[i].wrData.memList.len;
753 0 : list.lkey = wrList[i].wrData.memList.lkey;
754 :
755 0 : ib_wr.sg_list = &list;
756 0 : ib_wr.opcode = static_cast<enum ibv_wr_opcode>(wrList[i].wrData.op);
757 0 : ib_wr.send_flags = static_cast<unsigned int>(wrList[i].wrData.sendFlags);
758 0 : ib_wr.imm_data = wrList[i].wrData.immData;
759 :
760 0 : ib_wr.num_sge = 1; /* only support one sge */
761 0 : ib_wr.wr_id = wrList[i].wrData.wrId += wrIdoffset;
762 : }
763 0 : unsigned long long aiQpAddr = multiQpIndex == RDMA_INVALID_QP_INDEX ?
764 0 : combineAiQpInfo_.aiQpInfo.aiQpAddr : combineAiQpInfos_[multiQpIndex].aiQpInfo.aiQpAddr;
765 0 : struct ibv_qp *qp = reinterpret_cast<struct ibv_qp *>(aiQpAddr);
766 0 : HCCL_DEBUG("ib_wr.sglist[%u].addr[%p], ib_wr.sglist[%u].length[%u], ib_wr.sglist[%u], "
767 : "ib_wr.wr_id[%llu], raddr[%p], opcode[%d], imm_data[0x%llx]", i, list.addr, i, list.length, i, ib_wr.wr_id,
768 : ib_wr.wr.rdma.remote_addr, wrList[i].wrData.op, ib_wr.imm_data);
769 0 : if (wrList[i].wrData.op == RA_WR_RDMA_ATOMIC_WRITE) {
770 0 : ext_attr.reduce_op = wrList[i].wrData.aux.reduceType;
771 0 : ext_attr.reduce_type = wrList[i].wrData.aux.dataType;
772 0 : ret = DlHnsFunction::GetInstance().dlHnsIbvExtPostSend(qp, &ib_wr, &bad_wr, &ext_attr, &ext_rsp);
773 0 : HCOMM_DSB();
774 0 : exp_rsp.wqe_index = ext_rsp.wqe_index;
775 0 : exp_rsp.db_info = ext_rsp.db_info;
776 0 : HCCL_DEBUG("ibv_ext_post_send, op = [0x%x], imm_data = [0x%lx], reduce_op = [%d], reduceType = [%d]",
777 : wrList[i].wrData.op, ib_wr.imm_data, ext_attr.reduce_op, ext_attr.reduce_type);
778 0 : } else if (wrList[i].wrData.op == RA_WR_RDMA_WRITE_WITH_NOTIFY ||
779 0 : wrList[i].wrData.op == RA_WR_RDMA_REDUCE_WRITE ||
780 0 : wrList[i].wrData.op == RA_WR_RDMA_REDUCE_WRITE_WITH_NOTIFY) {
781 0 : ib_wr.imm_data = htobe32((wrList[i].wrData.aux.notifyOffset & RDMA_WRITE_NOTIFY_OFFSET_MASK) |
782 : RDMA_WRITE_NOTIFY_VALUE_RECORD);
783 0 : ext_attr.reduce_op = wrList[i].wrData.aux.reduceType;
784 0 : ext_attr.reduce_type = wrList[i].wrData.aux.dataType;
785 0 : ret = DlHnsFunction::GetInstance().dlHnsIbvExtPostSend(qp, &ib_wr, &bad_wr, &ext_attr, &ext_rsp);
786 0 : HCOMM_DSB();
787 0 : exp_rsp.wqe_index = ext_rsp.wqe_index;
788 0 : exp_rsp.db_info = ext_rsp.db_info;
789 0 : HCCL_DEBUG("ibv_ext_post_send, op = [0x%x], imm_data = [0x%lx], reduce_op = [%d],reduceType = [%d]",
790 : wrList[i].wrData.op, ib_wr.imm_data, ext_attr.reduce_op, ext_attr.reduce_type);
791 : } else {
792 0 : ret = DlHnsFunction::GetInstance().dlHnsIbvExpPostSend(qp, &ib_wr, &bad_wr, &exp_rsp);
793 0 : HCOMM_DSB();
794 0 : HCCL_DEBUG("ibv_exp_post_send, op = [0x%x], remote_addr = [0x%llx], size = [%d]",
795 : wrList[i].wrData.op, ib_wr.wr.rdma.remote_addr, ib_wr.sg_list->length);
796 : }
797 0 : if (ret) {
798 0 : HCCL_WARNING("[TxSendWrlistExt]ibv_post_send failed ret %d, i[%u]", ret, i);
799 0 : break;
800 : }
801 0 : unsigned long long dbIndex = multiQpIndex == RDMA_INVALID_QP_INDEX ?
802 0 : combineAiQpInfo_.aiQpInfo.dbIndex : combineAiQpInfos_[multiQpIndex].aiQpInfo.dbIndex;
803 0 : opRsp[i].db.dbIndex = (unsigned int)dbIndex;
804 0 : HCCL_DEBUG("opRsp.db.dbIndex = [%d]", opRsp[i].db.dbIndex);
805 0 : opRsp[i].db.dbInfo = exp_rsp.db_info;
806 : }
807 :
808 0 : HCCL_DEBUG("completeNum[%d], ret[%d]", i, ret);
809 0 : *completeNum = i;
810 0 : if ((ret == SOCK_ENOENT) || (ret == ROCE_EAGAIN) ||
811 0 : (workFlowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE && ret == ROCE_ENOMEM)) {
812 0 : return HCCL_E_AGAIN;
813 0 : } else if (!ret) {
814 0 : return HCCL_SUCCESS;
815 0 : } else if (ret == RDMA_QP_NO_MEM) { // 表示qp已满,内存不足,需要重发
816 0 : ib_wr.wr_id = wrList[i].wrData.wrId -= wrIdOffset_;
817 : // 可能出现主流没有launch,但从流一直下发导致卡死超时的问题,所以这里将所有流都下发
818 0 : CHK_RET(dispatcher_->LaunchAllTasks());
819 0 : return HCCL_E_AGAIN;
820 : } else {
821 0 : return HCCL_E_ROCE_TRANSFER;
822 : }
823 :
824 : return HCCL_SUCCESS;
825 : }
826 :
827 0 : HcclResult TransportDeviceIbverbs::RdmaSendAsync(
828 : std::vector<WrInformation> &wrInfoVec, Stream &stream, bool useOneDoorbell, u32 multiQpIndex)
829 : {
830 : HcclResult ret;
831 :
832 0 : std::vector<struct SendWrRsp> opRspVec(wrInfoVec.size());
833 0 : CHK_RET(TxWrList(wrInfoVec, stream, opRspVec, multiQpIndex));
834 :
835 0 : for (u32 i = 0; i < wrInfoVec.size(); i++) {
836 0 : if (useOneDoorbell && i != wrInfoVec.size() - 1) {
837 : // 如果useOneDoorbell为true,只敲最后一次doorbell
838 0 : continue;
839 : }
840 :
841 0 : RdmaTaskInfo taskInfo = {};
842 0 : taskInfo.remoteRank = machinePara_.remoteWorldRank;
843 0 : taskInfo.rdmaType = (wrInfoVec[i].type == static_cast<u64>(WqeType::WQE_TYPE_DATA)) ?
844 : RdmaType::RDMA_SEND_PAYLOAD : RdmaType::RDMA_SEND_NOTIFY;
845 :
846 0 : if (useOneDoorbell) {
847 : // 如果useOneDoorbell为true,一次性传入所有wr
848 0 : taskInfo.wrInfos = wrInfoVec;
849 : } else {
850 0 : taskInfo.wrInfos.push_back(wrInfoVec[i]);
851 : }
852 :
853 0 : u32 dbIndex = static_cast<u32>(opRspVec[i].db.dbIndex);
854 0 : HCCL_DEBUG("dbIndex = [%d]", dbIndex);
855 0 : u64 dbInfo = static_cast<u64>(opRspVec[i].db.dbInfo);
856 :
857 0 : ret = dispatcher_->RdmaSend(dbIndex, dbInfo, stream, taskInfo);
858 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
859 : HCCL_ERROR("[TransportDeviceIbverbs][RdmaSendAsync]errNo[0x%016llx] In lbv exp op base mode, "\
860 : "rdma send failed. dbIndex[%u] dbInfo[%llu] wqe type[%llu] addr[%llu]", HCCL_ERROR_CODE(ret), dbIndex,
861 : dbInfo, wrInfoVec[i].type, wrInfoVec[i].wrDataAddr), ret);
862 0 : }
863 0 : return HCCL_SUCCESS;
864 0 : }
865 :
866 1 : HcclResult TransportDeviceIbverbs::RdmaSendAsync(struct SendWr &wr, Stream &stream, WqeType wqeType, u64 notifyAddr,
867 : u32 notifyId)
868 : {
869 : HcclResult ret;
870 1 : WrInformation wrInfoTmp;
871 1 : struct SendWrRsp opRsp = {0};
872 1 : struct WrAuxInfo aux = {0};
873 1 : wrInfoTmp.wrData.memList = wr.bufList[0];
874 1 : wrInfoTmp.wrData.dstAddr = wr.dstAddr;
875 1 : wrInfoTmp.wrData.op = wr.op;
876 1 : wrInfoTmp.wrData.sendFlags = wr.sendFlag;
877 1 : wrInfoTmp.wrData.immData = 0;
878 1 : wrInfoTmp.wrData.wrId = 0xFF;
879 1 : wrInfoTmp.wrData.rkey = wr.rkey;
880 1 : wrInfoTmp.wrData.aux = aux;
881 :
882 1 : CHK_RET(SendWrList(1U, &wrInfoTmp, &opRsp));
883 1 : u32 dbIndex = static_cast<u32>(opRsp.db.dbIndex);
884 1 : u64 dbInfo = static_cast<u64>(opRsp.db.dbInfo);
885 1 : HCCL_DEBUG("dbIndex = [%d]", dbIndex);
886 1 : RdmaTaskInfo taskInfo = {};
887 1 : taskInfo.remoteRank = machinePara_.remoteWorldRank;
888 1 : taskInfo.rdmaType = (wqeType == WqeType::WQE_TYPE_DATA) ? RdmaType::RDMA_SEND_PAYLOAD : RdmaType::RDMA_SEND_NOTIFY;
889 1 : wrInfoTmp.type = static_cast<u64>(wqeType);
890 1 : wrInfoTmp.wrDataAddr = notifyAddr;
891 1 : wrInfoTmp.notifyId = notifyId;
892 1 : taskInfo.wrInfos.push_back(wrInfoTmp);
893 :
894 1 : ret = dispatcher_->RdmaSend(dbIndex, dbInfo, stream, taskInfo);
895 1 : CHK_PRT_RET(ret != HCCL_SUCCESS,
896 : HCCL_ERROR("[TransportDeviceIbverbs][RdmaSendAsync]errNo[0x%016llx] In lbv exp op base mode, "\
897 : "rdma send failed. dbIndex[%u] dbInfo[%llu], addr[%llu]", HCCL_ERROR_CODE(ret), dbIndex, dbInfo,
898 : notifyAddr), ret);
899 1 : return HCCL_SUCCESS;
900 1 : }
901 :
902 0 : HcclResult TransportDeviceIbverbs::GetWrDataAddr(void *dstAddr, WqeType wqeType, u64 &wrDataAddr, u32 ¬ifyId)
903 : {
904 0 : switch (wqeType) {
905 0 : case WqeType::WQE_TYPE_DATA:
906 : case WqeType::WQE_TYPE_DATA_WITH_NOTIFY:
907 : case WqeType::WQE_TYPE_DATA_WITH_REDUCE:
908 : case WqeType::WQE_TYPE_READ_DATA:
909 0 : wrDataAddr = reinterpret_cast<u64>(dstAddr);
910 0 : notifyId = INVALID_UINT;
911 0 : break;
912 0 : case WqeType::WQE_TYPE_DATA_NOTIFY:
913 0 : wrDataAddr = reinterpret_cast<u64>(remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].addr);
914 0 : notifyId = remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].notifyId;
915 0 : break;
916 0 : case WqeType::WQE_TYPE_ACK_NOTIFY:
917 0 : wrDataAddr = reinterpret_cast<u64>(remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].addr);
918 0 : notifyId = remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].notifyId;
919 0 : break;
920 0 : case WqeType::WQE_TYPE_DATA_ACK_NOTIFY:
921 0 : wrDataAddr = reinterpret_cast<u64>(remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].addr);
922 0 : notifyId = remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].notifyId;
923 0 : break;
924 0 : default:
925 0 : HCCL_ERROR("[Get][WrDataAddr]error wqeType[%d]", wqeType);
926 0 : return HCCL_E_INTERNAL;
927 : }
928 0 : HCCL_DEBUG("%s dstAddr:%p, wqeType:%d, wrDataAddr:%llu, notifyId:%u",
929 : __func__, dstAddr, wqeType, wrDataAddr, notifyId);
930 0 : return HCCL_SUCCESS;
931 : }
932 :
933 0 : HcclResult TransportDeviceIbverbs::TxSendWqe(void *dstMemPtr, u32 dstKey, const void *srcMemPtr, u32 srcKey,
934 : u64 srcMemSize, Stream &stream, WqeType wqeType)
935 : {
936 0 : struct SgList list = {0};
937 0 : struct SendWr wr = {nullptr};
938 : // 构造wr信息
939 0 : list.addr = static_cast<u64>(reinterpret_cast<uintptr_t>(srcMemPtr));
940 0 : list.len = srcMemSize;
941 0 : list.lkey = srcKey;
942 :
943 0 : wr.bufList = &list;
944 0 : wr.bufNum = 1; /* 此处list只有一个,设置为1 */
945 0 : wr.dstAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(dstMemPtr));
946 0 : wr.rkey = dstKey;
947 0 : wr.op = 0; /* RDMA_WRITE: 0 */
948 0 : wr.sendFlag = fence_ ? (RA_SEND_SIGNALED | RA_SEND_FENCE) : RA_SEND_SIGNALED;
949 0 : fence_ = false;
950 :
951 : // 获取notify偏移地址,对于发送数据时,偏移地址为0
952 0 : u64 wrDataAddr = 0;
953 0 : u32 notifyId = INVALID_UINT;
954 0 : CHK_RET(GetWrDataAddr(dstMemPtr, wqeType, wrDataAddr, notifyId));
955 :
956 : // RDMA异步发送
957 0 : CHK_RET(RdmaSendAsync(wr, stream, wqeType, wrDataAddr, notifyId));
958 0 : return HCCL_SUCCESS;
959 : }
960 :
961 0 : HcclResult TransportDeviceIbverbs::RxAsync(UserMemType srcMemType, u64 srcOffset, void *dst, u64 len, Stream &stream)
962 : {
963 0 : CHK_SMART_PTR_NULL(stream);
964 : // 等待TS把任务处理完成
965 0 : HCCL_DEBUG("RX dst[%p] len[%llu] srcOffset[%llu]", dst, len, srcOffset);
966 0 : u32 actualMultiQpNum = 1;
967 0 : const u32 KByteToByte = 1024; // 1024 多QP阈值单位是KB
968 0 : if (len / qpsPerConnection_ > multiQpThreshold_ * KByteToByte) {
969 0 : actualMultiQpNum = qpsPerConnection_;
970 : } else {
971 0 : u32 quotient = len / (multiQpThreshold_ * KByteToByte);
972 0 : u32 remainder = len % (multiQpThreshold_ * KByteToByte);
973 0 : actualMultiQpNum = quotient + (remainder != 0 ? 1 : 0);
974 : }
975 0 : if (UseMultiQp() && actualMultiQpNum != 1 && actualMultiQpNum <= qpsPerConnection_ && len != 0) {
976 0 : for (u32 i = 0; i < actualMultiQpNum; i++) {
977 0 : CHK_RET(dispatcher_->SignalWait(multiQpDataNotify_[i]->ptr(),
978 : stream,
979 : machinePara_.localUserrank,
980 : machinePara_.remoteWorldRank,
981 : INVALID_VALUE_STAGE,
982 : false,
983 : multiQpDataNotify_[i]->notifyId_));
984 : }
985 : } else {
986 0 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(),
987 : stream,
988 : machinePara_.localUserrank,
989 : machinePara_.remoteWorldRank,
990 : INVALID_VALUE_STAGE,
991 : false,
992 : dataNotify_->notifyId_));
993 : }
994 0 : return HCCL_SUCCESS;
995 : }
996 :
997 0 : HcclResult TransportDeviceIbverbs::RxAsync(std::vector<RxMemoryInfo>& rxMems, Stream &stream)
998 : {
999 0 : CHK_PRT_RET(rxMems.size() == 0, HCCL_ERROR("Invalid rxMem size[%u]", rxMems.size()), HCCL_E_PARA);
1000 0 : CHK_SMART_PTR_NULL(stream);
1001 0 : for (auto& mem : rxMems) {
1002 0 : HCCL_DEBUG("RX dst[%p] len[%llu] dstOffset[%llu]", mem.dst, mem.len, mem.srcOffset);
1003 : }
1004 0 : u32 maxLength = 0;
1005 0 : for (u32 i = 0; i < rxMems.size(); i++) {
1006 0 : if (rxMems[i].len > maxLength) {
1007 0 : maxLength = rxMems[i].len;
1008 : }
1009 : }
1010 :
1011 0 : CHK_RET(RxAsync(rxMems[0].srcMemType, rxMems[0].srcOffset, rxMems[0].dst, maxLength, stream));
1012 0 : return HCCL_SUCCESS;
1013 : }
1014 :
1015 0 : HcclResult TransportDeviceIbverbs::DataReceivedAck(Stream &stream)
1016 : {
1017 0 : CHK_RET(PostFinAck(stream));
1018 0 : CHK_RET(WaitFinAck(stream));
1019 :
1020 0 : return HCCL_SUCCESS;
1021 : }
1022 :
1023 0 : HcclResult TransportDeviceIbverbs::TxWaitDone(Stream &stream)
1024 : {
1025 0 : return HCCL_SUCCESS;
1026 : }
1027 :
1028 : /* 发送ack消息(同步模式) */
1029 0 : HcclResult TransportDeviceIbverbs::TxAck(Stream &stream)
1030 : {
1031 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].addr,
1032 : remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].lkey,
1033 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1034 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1035 : notifySize_, stream, WqeType::WQE_TYPE_ACK_NOTIFY));
1036 0 : return HCCL_SUCCESS;
1037 : }
1038 :
1039 : /* 接收ack消息(同步模式) */
1040 0 : HcclResult TransportDeviceIbverbs::RxAck(Stream &stream)
1041 : {
1042 0 : CHK_RET(dispatcher_->SignalWait(ackNotify_->ptr(), stream, machinePara_.localUserrank,
1043 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, ackNotify_->notifyId_));
1044 0 : return HCCL_SUCCESS;
1045 : }
1046 :
1047 0 : HcclResult TransportDeviceIbverbs::TxDataSignal(Stream &stream)
1048 : {
1049 : // 发送data notify同步信息
1050 0 : void *remoteNotifyaddr = remoteDataNotifyMsg_.addr;
1051 0 : HcclResult ret = TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].addr,
1052 0 : remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].lkey,
1053 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1054 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1055 0 : notifySize_, stream, WqeType::WQE_TYPE_DATA_NOTIFY);
1056 :
1057 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
1058 : HCCL_ERROR("[TransportDeviceIbverbs][TxDataSignal]errNo[0x%016llx] In ibv tx data signal, send notify "\
1059 : "wqe failed. dstMemPtr[%p], srcMemPtr[%p], srcMemSize[%llu]", HCCL_ERROR_CODE(ret), remoteNotifyaddr,
1060 : notifyValueAddr_, notifySize_), ret);
1061 : // 每发送一个data notify wqe, count 自增
1062 0 : return HCCL_SUCCESS;
1063 : }
1064 :
1065 0 : HcclResult TransportDeviceIbverbs::RxDataSignal(Stream &stream)
1066 : {
1067 : /* 等待send_ready_event事件 */
1068 0 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(), stream, machinePara_.localUserrank,
1069 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, dataNotify_->notifyId_));
1070 0 : return HCCL_SUCCESS;
1071 : }
1072 :
1073 : /* 发送ack消息(同步模式) */
1074 0 : HcclResult TransportDeviceIbverbs::TxPrepare(Stream &stream)
1075 : {
1076 0 : CHK_RET(dispatcher_->SignalWait(ackNotify_->ptr(), stream, machinePara_.localUserrank,
1077 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, ackNotify_->notifyId_));
1078 0 : return HCCL_SUCCESS;
1079 : }
1080 :
1081 : /* 接收ack消息(同步模式) */
1082 0 : HcclResult TransportDeviceIbverbs::RxPrepare(Stream &stream)
1083 : {
1084 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].addr,
1085 : remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].lkey,
1086 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1087 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1088 : notifySize_, stream, WqeType::WQE_TYPE_ACK_NOTIFY));
1089 0 : return HCCL_SUCCESS;
1090 : }
1091 :
1092 0 : HcclResult TransportDeviceIbverbs::TxData(UserMemType dstMemType, u64 dstOffset, const void *src, u64 len, Stream &stream)
1093 : {
1094 0 : CHK_SMART_PTR_NULL(stream);
1095 0 : std::vector<WrInformation> wrInfoVec;
1096 0 : struct WrAuxInfo aux = {0};
1097 0 : HCCL_DEBUG("TX src[%p] len[%llu] dstOffset[%llu]", src, len, dstOffset);
1098 :
1099 0 : if (len > 0) {
1100 0 : CHK_PTR_NULL(src);
1101 0 : CHK_RET(TxPayLoad(dstMemType, dstOffset, src, len, WqeType::WQE_TYPE_DATA, aux, wrInfoVec));
1102 : }
1103 :
1104 0 : CHK_RET(RdmaSendAsync(wrInfoVec, stream, false));
1105 0 : return HCCL_SUCCESS;
1106 0 : }
1107 :
1108 0 : HcclResult TransportDeviceIbverbs::RxData(UserMemType srcMemType, u64 srcOffset, void *dst, u64 len, Stream &stream)
1109 : {
1110 0 : return HCCL_SUCCESS;
1111 : }
1112 :
1113 0 : HcclResult TransportDeviceIbverbs::TxDone(Stream &stream)
1114 : {
1115 : // 发送数据接收确认notify
1116 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].addr,
1117 : remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].lkey,
1118 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1119 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1120 : notifySize_, stream, WqeType::WQE_TYPE_DATA_NOTIFY));
1121 : // 接收数据接收确认notify
1122 0 : CHK_RET(dispatcher_->SignalWait(dataAckNotify_->ptr(), stream, machinePara_.localUserrank,
1123 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, dataAckNotify_->notifyId_));
1124 0 : return HCCL_SUCCESS;
1125 : }
1126 :
1127 0 : HcclResult TransportDeviceIbverbs::RxDone(Stream &stream)
1128 : {
1129 : // 接收数据接收确认notify
1130 0 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(), stream, machinePara_.localUserrank,
1131 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, dataNotify_->notifyId_));
1132 :
1133 : // 发送数据接收确认notify
1134 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].addr,
1135 : remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].lkey,
1136 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1137 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1138 : notifySize_, stream, WqeType::WQE_TYPE_DATA_ACK_NOTIFY));
1139 0 : return HCCL_SUCCESS;
1140 : }
1141 :
1142 0 : HcclResult TransportDeviceIbverbs::PostReady(Stream &stream)
1143 : {
1144 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].addr,
1145 : remoteMemMsg_[static_cast<u32>(MemType::ACK_NOTIFY_MEM)].lkey,
1146 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1147 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1148 : notifySize_, stream, WqeType::WQE_TYPE_ACK_NOTIFY));
1149 0 : return HCCL_SUCCESS;
1150 : }
1151 :
1152 0 : HcclResult TransportDeviceIbverbs::WaitReady(Stream &stream)
1153 : {
1154 0 : CHK_RET(dispatcher_->SignalWait(ackNotify_->ptr(), stream, machinePara_.localUserrank,
1155 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, ackNotify_->notifyId_));
1156 0 : return HCCL_SUCCESS;
1157 : }
1158 :
1159 0 : HcclResult TransportDeviceIbverbs::PostFin(Stream &stream)
1160 : {
1161 : // 发送data notify同步信息
1162 0 : void *remoteNotifyaddr = remoteDataNotifyMsg_.addr;
1163 0 : HcclResult ret = TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].addr,
1164 0 : remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].lkey,
1165 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1166 0 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1167 0 : notifySize_, stream, WqeType::WQE_TYPE_DATA_NOTIFY);
1168 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
1169 : HCCL_ERROR("[TransportDeviceIbverbs][PostFin]errNo[0x%016llx] In ibv tx data signal, send notify "\
1170 : "wqe failed. dstMemPtr[%p], srcMemPtr[%p], srcMemSize[%llu]", HCCL_ERROR_CODE(ret), remoteNotifyaddr,
1171 : notifyValueAddr_, notifySize_), ret);
1172 : // 每发送一个data notify wqe, count 自增
1173 0 : return HCCL_SUCCESS;
1174 : }
1175 :
1176 0 : HcclResult TransportDeviceIbverbs::WaitFin(Stream &stream)
1177 : {
1178 0 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(), stream, machinePara_.localUserrank,
1179 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, dataNotify_->notifyId_));
1180 0 : return HCCL_SUCCESS;
1181 : }
1182 :
1183 0 : HcclResult TransportDeviceIbverbs::PostFinAck(Stream &stream)
1184 : {
1185 0 : CHK_RET(TxSendWqe(remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].addr,
1186 : remoteMemMsg_[static_cast<u32>(MemType::DATA_ACK_NOTIFY_MEM)].lkey,
1187 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1188 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1189 : notifySize_, stream, WqeType::WQE_TYPE_DATA_ACK_NOTIFY));
1190 0 : return HCCL_SUCCESS;
1191 : }
1192 :
1193 0 : HcclResult TransportDeviceIbverbs::WaitFinAck(Stream &stream)
1194 : {
1195 0 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(), stream, machinePara_.localUserrank,
1196 : machinePara_.remoteWorldRank, INVALID_VALUE_STAGE, false, dataAckNotify_->notifyId_));
1197 0 : return HCCL_SUCCESS;
1198 : }
1199 :
1200 5 : HcclResult TransportDeviceIbverbs::ResolveRdmaAddrsAndKeys(RdmaAddrKeyResolveParam ¶m)
1201 : {
1202 5 : param.transLocalAddr = const_cast<void *>(param.localAddr);
1203 5 : param.transRemoteAddr = const_cast<void *>(param.remoteAddr);
1204 5 : if (useMemDetailsLookup_) {
1205 2 : return ResolveRdmaKeysFromMemDetails(param);
1206 : }
1207 3 : return ResolveRdmaKeysFromIoMemRanges(param);
1208 : }
1209 :
1210 2 : HcclResult TransportDeviceIbverbs::ResolveRdmaKeysFromMemDetails(RdmaAddrKeyResolveParam ¶m)
1211 : {
1212 2 : auto rf = remoteMemDetailsRmaMgr_->Find(MakeMemLookupKey(param.remoteAddr, param.length));
1213 2 : if (!rf.first || rf.second == nullptr) {
1214 1 : HCCL_ERROR("[TransportDeviceIbverbs]Can't find remoteBuffer key by addr and size {%p, %llu}, "
1215 : "registered remote MR count[%zu]",
1216 : param.remoteAddr, param.length, remoteMemDetailsRmaMgr_->size());
1217 1 : return HCCL_E_INTERNAL;
1218 : }
1219 1 : param.dstKey = rf.second->key;
1220 1 : param.transRemoteAddr = LogicalPtrToDevPtr(*rf.second, param.remoteAddr);
1221 :
1222 1 : auto lf = localMemDetailsRmaMgr_->Find(MakeMemLookupKey(param.localAddr, param.length));
1223 1 : CHK_PRT_RET(!lf.first || lf.second == nullptr,
1224 : HCCL_ERROR("[TransportDeviceIbverbs]Can't find localBuffer key by addr and size {%p, %llu}, "
1225 : "registered local MR count[%zu]",
1226 : param.localAddr, param.length, localMemDetailsRmaMgr_->size()),
1227 : HCCL_E_INTERNAL);
1228 1 : param.srcKey = lf.second->key;
1229 1 : param.transLocalAddr = LogicalPtrToDevPtr(*lf.second, param.localAddr);
1230 1 : return HCCL_SUCCESS;
1231 2 : }
1232 :
1233 5 : HcclResult TransportDeviceIbverbs::ResolveRdmaKeysFromIoMemRanges(RdmaAddrKeyResolveParam ¶m)
1234 : {
1235 5 : u64 dstAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(param.remoteAddr));
1236 : u64 remoteInputAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(
1237 5 : remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].addr));
1238 5 : u64 remoteInputSize = localInputMem_.size;
1239 5 : u32 remoteInputKey = remoteMemMsg_[static_cast<u32>(MemType::USER_INPUT_MEM)].lkey;
1240 : u64 remoteOutputAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(
1241 5 : remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].addr));
1242 5 : u64 remoteOutputSize = localOutputMem_.size;
1243 5 : u32 remoteOutputKey = remoteMemMsg_[static_cast<u32>(MemType::USER_OUTPUT_MEM)].lkey;
1244 5 : if (dstAddr >= remoteInputAddr && dstAddr < remoteInputAddr + remoteInputSize) {
1245 2 : param.dstKey = remoteInputKey;
1246 3 : } else if (dstAddr >= remoteOutputAddr && dstAddr <= remoteOutputAddr + remoteOutputSize) {
1247 1 : param.dstKey = remoteOutputKey;
1248 : } else {
1249 2 : HCCL_ERROR("[TransportDeviceIbverbs][TxAsync]src_ptr=%p is out of range, inputmem src[%p], size[%llu];"
1250 : " outputmem src[%p] size[%llu]", param.remoteAddr, remoteInputAddr, remoteInputSize,
1251 : remoteOutputAddr, remoteOutputSize);
1252 2 : return HCCL_E_INTERNAL;
1253 : }
1254 :
1255 3 : u64 srcAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(param.localAddr));
1256 3 : if (srcAddr >= localInputMem_.addr && srcAddr < localInputMem_.addr + localInputMem_.size) {
1257 2 : param.srcKey = localInputMem_.key;
1258 1 : } else if (srcAddr >= localOutputMem_.addr && srcAddr <= localOutputMem_.addr + localOutputMem_.size) {
1259 1 : param.srcKey = localOutputMem_.key;
1260 : } else {
1261 0 : HCCL_ERROR("[TransportDeviceIbverbs][TxAsync]src_ptr=%p is out of range, inputmem src[%p], size[%llu];"
1262 : " outputmem src[%p] size[%llu]", param.localAddr, localInputMem_.addr, localInputMem_.size,
1263 : localOutputMem_.addr, localOutputMem_.size);
1264 0 : return HCCL_E_INTERNAL;
1265 : }
1266 3 : return HCCL_SUCCESS;
1267 : }
1268 :
1269 0 : HcclResult TransportDeviceIbverbs::WriteCommon(const void *remoteAddr, const void *localAddr, u64 length, Stream &stream,
1270 : WqeType wqeType, struct WrAuxInfo &aux)
1271 : {
1272 0 : if (machinePara_.dctxPtr != nullptr) {
1273 0 : CHK_RET(SetDispatcherCtx(static_cast<DispatcherCtxPtr>(machinePara_.dctxPtr)));
1274 : }
1275 :
1276 0 : std::vector<WrInformation> wrInfoVec;
1277 0 : HCCL_DEBUG("write localAddr[%p] remoteAddr[%p] len[%llu]",
1278 : localAddr, remoteAddr, length);
1279 :
1280 0 : if (localAddr != nullptr) {
1281 : // 为保证单算子下不同数据量下子图的结构相同,zero byte message 时也需要下发task
1282 0 : u32 txSendDataTimes = (length == 0) ? 1 : (length + RDMA_SEND_MAX_SIZE - 1) / RDMA_SEND_MAX_SIZE;
1283 :
1284 0 : RdmaAddrKeyResolveParam resolve{};
1285 0 : resolve.remoteAddr = remoteAddr;
1286 0 : resolve.localAddr = localAddr;
1287 0 : resolve.length = length;
1288 0 : CHK_RET(ResolveRdmaAddrsAndKeys(resolve));
1289 :
1290 0 : CHK_RET(ConstructPayLoadWqe(resolve.transRemoteAddr, resolve.dstKey,
1291 : resolve.transLocalAddr, resolve.srcKey,
1292 : length,
1293 : wqeType,
1294 : aux,
1295 : wrInfoVec,
1296 : txSendDataTimes));
1297 : }
1298 0 : u32 maxLength = 0;
1299 0 : for (u32 i = 0; i < wrInfoVec.size(); i++) {
1300 0 : if (wrInfoVec[i].wrData.memList.len > maxLength) {
1301 0 : maxLength = wrInfoVec[i].wrData.memList.len;
1302 : }
1303 : }
1304 :
1305 0 : u32 actualMultiQpNum = GetActualQpNum(maxLength);
1306 :
1307 0 : HCCL_DEBUG("[TransportDeviceIbverbs][TxSendDataAndNotify] UseMultiQp[%d] MultiQpNum[%u] actualMultiQpNum[%u] "
1308 : "maxLength[%u]",
1309 : UseMultiQp(),
1310 : qpsPerConnection_,
1311 : actualMultiQpNum,
1312 : maxLength);
1313 0 : if (UseMultiQp() && actualMultiQpNum != 1 && actualMultiQpNum <= qpsPerConnection_ && maxLength != 0) {
1314 0 : std::vector<std::vector<WrInformation>> multiQpWqeInfoVct(actualMultiQpNum, wrInfoVec);
1315 0 : for (u32 i = 0; i < wrInfoVec.size(); i++) {
1316 0 : WrInformation tmpWqeInfo = wrInfoVec[i];
1317 0 : u32 curLen = tmpWqeInfo.wrData.memList.len;
1318 0 : std::vector<u32> splittedLen = RdmaLengthSplit(curLen, actualMultiQpNum);
1319 0 : uint64_t curSrcAddr = tmpWqeInfo.wrData.memList.addr;
1320 0 : uint64_t curDstAddr = tmpWqeInfo.wrData.dstAddr;
1321 0 : for (u32 qpIndex = 0; qpIndex < actualMultiQpNum; qpIndex++) {
1322 0 : multiQpWqeInfoVct[qpIndex][i].wrData.memList.len = splittedLen[qpIndex];
1323 0 : multiQpWqeInfoVct[qpIndex][i].wrData.memList.addr = curSrcAddr;
1324 0 : multiQpWqeInfoVct[qpIndex][i].wrData.dstAddr = curDstAddr;
1325 0 : curSrcAddr += splittedLen[qpIndex];
1326 0 : curDstAddr += splittedLen[qpIndex];
1327 : }
1328 0 : }
1329 :
1330 : // useOneDoorbell 配置成true。最后一个payload去按doorbell
1331 0 : for (u32 qpIndex = 0; qpIndex < actualMultiQpNum; qpIndex++) {
1332 0 : CHK_RET(RdmaSendAsync(multiQpWqeInfoVct[qpIndex], stream, true, qpIndex)); // 多QP使用同一个stream异步doorbell触发
1333 : }
1334 0 : } else {
1335 0 : CHK_RET(RdmaSendAsync(wrInfoVec, stream, GetUseOneDoorbellValue()));
1336 : }
1337 0 : return HCCL_SUCCESS;
1338 :
1339 : #ifndef CCL_LLT
1340 : CHK_RET(RdmaSendAsync(wrInfoVec, stream, GetUseOneDoorbellValue()));
1341 : #endif
1342 : return HCCL_SUCCESS;
1343 0 : }
1344 :
1345 0 : HcclResult TransportDeviceIbverbs::WriteAsync(
1346 : struct Transport::Buffer &remoteBuf, struct Transport::Buffer &localBuf, Stream &stream)
1347 : {
1348 0 : struct WrAuxInfo aux = {0};
1349 0 : return WriteCommon(remoteBuf.addr, localBuf.addr, remoteBuf.size, stream, WqeType::WQE_TYPE_DATA, aux);
1350 : }
1351 :
1352 0 : HcclResult TransportDeviceIbverbs::ReadAsync(
1353 : struct Transport::Buffer &localBuf, struct Transport::Buffer &remoteBuf, Stream &stream)
1354 : {
1355 0 : HCCL_DEBUG("[TransportDeviceIbverbs][ReadAsync]");
1356 0 : struct WrAuxInfo aux = {0};
1357 0 : return WriteCommon(remoteBuf.addr, localBuf.addr, remoteBuf.size, stream, WqeType::WQE_TYPE_READ_DATA, aux);
1358 : }
1359 :
1360 4 : HcclResult TransportDeviceIbverbs::ResolveTransferDesc(
1361 : const HcommBatchTransferDesc &desc, const void *&remoteAddr,
1362 : const void *&localAddr, u64 &length, WqeType &wqeType, struct WrAuxInfo &aux)
1363 : {
1364 4 : if (desc.transType == HCOMM_TRANSFER_TYPE_WRITE) {
1365 3 : CHK_PTR_NULL(desc.transferInfo.write.dst);
1366 1 : CHK_PTR_NULL(desc.transferInfo.write.src);
1367 1 : remoteAddr = desc.transferInfo.write.dst;
1368 1 : localAddr = desc.transferInfo.write.src;
1369 1 : length = desc.transferInfo.write.len;
1370 1 : wqeType = WqeType::WQE_TYPE_DATA;
1371 1 : } else if (desc.transType == HCOMM_TRANSFER_TYPE_READ) {
1372 1 : CHK_PTR_NULL(desc.transferInfo.read.dst);
1373 1 : CHK_PTR_NULL(desc.transferInfo.read.src);
1374 1 : remoteAddr = desc.transferInfo.read.src;
1375 1 : localAddr = desc.transferInfo.read.dst;
1376 1 : length = desc.transferInfo.read.len;
1377 1 : wqeType = WqeType::WQE_TYPE_READ_DATA;
1378 : } else {
1379 0 : HCCL_ERROR("[ResolveTransferDesc] Unsupported transType[%d].", desc.transType);
1380 0 : return HCCL_E_NOT_SUPPORT;
1381 : }
1382 2 : return HCCL_SUCCESS;
1383 : }
1384 :
1385 0 : HcclResult TransportDeviceIbverbs::SubmitWqeBatch(
1386 : std::vector<WrInformation> &wrInfoVec, Stream &stream)
1387 : {
1388 0 : u32 maxLength = 0;
1389 0 : for (u32 i = 0; i < wrInfoVec.size(); i++) {
1390 0 : if (wrInfoVec[i].wrData.memList.len > maxLength) {
1391 0 : maxLength = wrInfoVec[i].wrData.memList.len;
1392 : }
1393 : }
1394 :
1395 0 : u32 actualMultiQpNum = GetActualQpNum(maxLength);
1396 0 : if (UseMultiQp() && actualMultiQpNum != 1 && actualMultiQpNum <= qpsPerConnection_ && maxLength != 0) {
1397 0 : CHK_RET(TxSendDataAndNotifyWithMultiQP(wrInfoVec, actualMultiQpNum, stream, true));
1398 : } else {
1399 0 : CHK_RET(TxSendDataAndNotifyWithSingleQP(wrInfoVec, stream, true));
1400 : }
1401 0 : return HCCL_SUCCESS;
1402 : }
1403 :
1404 4 : HcclResult TransportDeviceIbverbs::BatchTransferImpl(
1405 : const HcommBatchTransferDesc *transferDescs, uint32_t descNum, Stream &stream)
1406 : {
1407 4 : if (machinePara_.dctxPtr != nullptr) {
1408 0 : CHK_RET(SetDispatcherCtx(static_cast<DispatcherCtxPtr>(machinePara_.dctxPtr)));
1409 : }
1410 :
1411 4 : CHK_PTR_NULL(transferDescs);
1412 4 : std::vector<WrInformation> wrInfoVec;
1413 4 : for (uint32_t i = 0; i < descNum; i++) {
1414 4 : const void *localAddr = nullptr;
1415 4 : const void *remoteAddr = nullptr;
1416 4 : u64 length = 0;
1417 4 : WqeType wqeType = WqeType::WQE_TYPE_DATA;
1418 4 : struct WrAuxInfo aux = {0};
1419 6 : CHK_RET(ResolveTransferDesc(transferDescs[i], remoteAddr, localAddr, length, wqeType, aux));
1420 :
1421 2 : HCCL_DEBUG("[BatchTransferImpl] index[%u] localAddr[%p] remoteAddr[%p] len[%llu] wqeType[%d]",
1422 : i, localAddr, remoteAddr, length, static_cast<int>(wqeType));
1423 :
1424 2 : if (localAddr != nullptr) {
1425 2 : u32 txSendDataTimes = (length == 0) ? 1 :
1426 2 : (length + RDMA_SEND_MAX_SIZE - 1) / RDMA_SEND_MAX_SIZE;
1427 2 : RdmaAddrKeyResolveParam resolve{};
1428 2 : resolve.remoteAddr = remoteAddr;
1429 2 : resolve.localAddr = localAddr;
1430 2 : resolve.length = length;
1431 2 : CHK_RET(ResolveRdmaAddrsAndKeys(resolve));
1432 0 : CHK_RET(ConstructPayLoadWqe(resolve.transRemoteAddr, resolve.dstKey,
1433 : resolve.transLocalAddr, resolve.srcKey,
1434 : length, wqeType, aux, wrInfoVec, txSendDataTimes));
1435 : }
1436 : }
1437 :
1438 0 : if (!wrInfoVec.empty()) {
1439 0 : return SubmitWqeBatch(wrInfoVec, stream);
1440 : }
1441 0 : return HCCL_SUCCESS;
1442 4 : }
1443 :
1444 4 : HcclResult TransportDeviceIbverbs::BatchTransferAsync(
1445 : const HcommBatchTransferDesc *transferDescs, uint32_t descNum, Stream &stream)
1446 : {
1447 4 : return BatchTransferImpl(transferDescs, descNum, stream);
1448 : }
1449 :
1450 0 : HcclResult TransportDeviceIbverbs::WriteReduceAsync(struct Transport::Buffer &remoteBuf,
1451 : struct Transport::Buffer &localBuf, const HcclDataType datatype, HcclReduceOp redOp, Stream &stream)
1452 : {
1453 0 : struct WrAuxInfo aux = {0};
1454 0 : aux.dataType = RDMA_REDUCE_DATA_TYPE_TABLE[datatype];
1455 0 : aux.reduceType = RDMA_REDUCE_OP_TYPE_TABLE[redOp];
1456 0 : if (aux.dataType == static_cast<uint8_t>(RdmaReduceDataType::RDMA_REDUCE_DATA_INVALID) ||
1457 0 : aux.reduceType == static_cast<uint8_t>(RdmaReduceOpType::RDMA_REDUCE_OP_INVALID)) {
1458 0 : HCCL_ERROR("unsupported data type [%s] or Reduce type [%s]",
1459 : GetDataTypeEnumStr(datatype).c_str(), GetReduceOpEnumStr(redOp).c_str());
1460 0 : return HCCL_E_INTERNAL;
1461 : }
1462 :
1463 0 : return WriteCommon(remoteBuf.addr, localBuf.addr, remoteBuf.size, stream, WqeType::WQE_TYPE_DATA_WITH_REDUCE, aux);
1464 : }
1465 :
1466 0 : HcclResult TransportDeviceIbverbs::Post(u32 notifyIdx, Stream &stream)
1467 : {
1468 : // 校验notifyIdx有效性
1469 0 : bool bRet = (notifyIdx >= notifyNum_);
1470 0 : CHK_PRT_RET(bRet,
1471 : HCCL_ERROR("[TransportDeviceIbverbs][Post]notifyNum[%u], notifyIdx[%u] out of range[0, %u]", \
1472 : notifyNum_, notifyIdx, notifyNum_-1), HCCL_E_INTERNAL);
1473 :
1474 : // 每个QP发送一个指定idx的notify
1475 0 : for (u32 i = 0; i < qpsPerConnection_; i++) {
1476 0 : CHK_RET(TxSendWqe(userMultiQpRemoteNotifyMsg_[i][notifyIdx].addr,
1477 : userMultiQpRemoteNotifyMsg_[i][notifyIdx].lkey,
1478 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].addr,
1479 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1480 : notifySize_, stream, WqeType::WQE_TYPE_DATA_WITH_NOTIFY));
1481 : }
1482 0 : return HCCL_SUCCESS;
1483 : }
1484 :
1485 0 : HcclResult TransportDeviceIbverbs::Wait(u32 notifyIdx, Stream &stream, const u32 timeOut)
1486 : {
1487 : // 校验notifyIdx有效性
1488 0 : bool bRet = (notifyIdx >= notifyNum_);
1489 0 : CHK_PRT_RET(bRet,
1490 : HCCL_ERROR("[TransportDeviceIbverbs][Wait]notifyNum[%u], notifyIdx[%u] out of range[0, %u]", \
1491 : notifyNum_, notifyIdx, notifyNum_-1), HCCL_E_INTERNAL);
1492 :
1493 : // 单QP接收一个指定idx的notify
1494 : // 每个qp接收一个指定idx的notify
1495 0 : for (u32 i = 0; i < qpsPerConnection_; i++) {
1496 0 : CHK_RET(dispatcher_->SignalWait(userMultiQpLocalNotify_[i][notifyIdx]->ptr(),
1497 : stream,
1498 : machinePara_.localUserrank,
1499 : machinePara_.remoteWorldRank,
1500 : INVALID_VALUE_STAGE,
1501 : false,
1502 : userMultiQpLocalNotify_[i][notifyIdx]->notifyId_, timeOut));
1503 : }
1504 0 : return HCCL_SUCCESS;
1505 : }
1506 :
1507 0 : bool TransportDeviceIbverbs::UseMultiQp()
1508 : {
1509 0 : return qpsPerConnection_ != 1;
1510 : }
1511 :
1512 0 : u32 TransportDeviceIbverbs::GetActualQpNum(u32 maxLength)
1513 : {
1514 0 : u32 actualMultiQpNum = 1;
1515 0 : const u32 KByteToByte = 1024; // 1024 多QP阈值单位是KB
1516 0 : if (maxLength / qpsPerConnection_ >= multiQpThreshold_ * KByteToByte) {
1517 0 : actualMultiQpNum = qpsPerConnection_;
1518 : } else {
1519 0 : u32 quotient = maxLength / (multiQpThreshold_ * KByteToByte);
1520 0 : u32 remainder = maxLength % (multiQpThreshold_ * KByteToByte);
1521 0 : actualMultiQpNum = quotient + (remainder != 0 ? 1 : 0);
1522 : }
1523 :
1524 0 : return actualMultiQpNum;
1525 : }
1526 :
1527 0 : HcclResult TransportDeviceIbverbs::TxSendDataAndNotifyWithMultiQP(std::vector<WrInformation> &wqeInfoVec,
1528 : u32 actualMultiQpNum, Stream &stream, bool useOneDoorbell)
1529 : {
1530 : // vector<WrInformation> 是一个vector的原因是 单个wqe只能发2GB数据,如果超过2GB,就拆分到多个WqeInfo中了
1531 : // 多QP下,对每个WqeInfo都进行多QP切分,然后在收发每一个QP的数据
1532 0 : std::vector<std::vector<WrInformation>> multiQpWqeInfoVct(actualMultiQpNum, wqeInfoVec);
1533 0 : for (u32 i = 0; i < wqeInfoVec.size(); i++) {
1534 0 : WrInformation tmpWqeInfo = wqeInfoVec[i];
1535 0 : u32 curLen = tmpWqeInfo.wrData.memList.len;
1536 0 : std::vector<u32> splittedLen = RdmaLengthSplit(curLen, actualMultiQpNum);
1537 0 : uint64_t curSrcAddr = tmpWqeInfo.wrData.memList.addr;
1538 0 : uint64_t curDstAddr = tmpWqeInfo.wrData.dstAddr;
1539 0 : for (u32 qpIndex = 0; qpIndex < actualMultiQpNum; qpIndex++) {
1540 0 : multiQpWqeInfoVct[qpIndex][i].wrData.memList.len = splittedLen[qpIndex];
1541 0 : multiQpWqeInfoVct[qpIndex][i].wrData.memList.addr = curSrcAddr;
1542 0 : multiQpWqeInfoVct[qpIndex][i].wrData.dstAddr = curDstAddr;
1543 0 : curSrcAddr += splittedLen[qpIndex];
1544 0 : curDstAddr += splittedLen[qpIndex];
1545 : }
1546 0 : }
1547 : // 给每个QP最后增加一个属于该QP的DataNotify
1548 0 : for (u32 qpIndex = 0; qpIndex < actualMultiQpNum; qpIndex++) {
1549 0 : struct WrAuxInfo aux = {0};
1550 0 : void *remoteNotifyaddr = multiQpDataNotifyRemoteMemMsg_[qpIndex].addr;
1551 0 : CHK_RET(AddWrList(remoteNotifyaddr,
1552 : notifyValueAddr_,
1553 : notifySize_,
1554 : memMsg_[static_cast<u32>(MemType::NOTIFY_SRC_MEM)].lkey,
1555 : remoteMemMsg_[static_cast<u32>(MemType::DATA_NOTIFY_MEM)].lkey,
1556 : WqeType::WQE_TYPE_DATA_NOTIFY,
1557 : aux,
1558 : multiQpWqeInfoVct[qpIndex]));
1559 : }
1560 : // useOneDoorbell 配置成true。最后一个payload去按doorbell
1561 0 : for (u32 qpIndex = 0; qpIndex < actualMultiQpNum; qpIndex++) {
1562 0 : CHK_RET(
1563 : RdmaSendAsync(multiQpWqeInfoVct[qpIndex], stream, true, qpIndex)); // 多QP使用同一个stream异步doorbell触发
1564 : }
1565 0 : return HCCL_SUCCESS;
1566 0 : }
1567 0 : HcclResult TransportDeviceIbverbs::GetTransportId(u32 &id)
1568 : {
1569 0 : struct ibv_qp *qp = reinterpret_cast<struct ibv_qp *>(combineAiQpInfo_.aiQpInfo.aiQpAddr);
1570 0 : if (nullptr != qp)
1571 : {
1572 0 : id = qp->qp_num;
1573 : }
1574 0 : return HCCL_SUCCESS;
1575 : }
1576 :
1577 0 : HcclResult TransportDeviceIbverbs::HnsPostSend(const TransportDeviceNormalData &ibvData, struct MemDetails *localMems,
1578 : struct MemDetails *remoteMems, u32 memNum, HcclWrOpCode opCode, u64 &dbInfo, bool fence)
1579 : {
1580 0 : CHK_PTR_NULL(localMems);
1581 0 : CHK_PTR_NULL(remoteMems);
1582 :
1583 0 : const uint32_t SEND_WR_LEN = 8;
1584 0 : uint32_t last = memNum - 1;
1585 0 : CHK_PRT_RET(memNum > SEND_WR_LEN,
1586 : HCCL_ERROR("[TransportDeviceIbverbs][HnsPostSend] buffer size is:%u over SEND_WR_LEN: %u", memNum, SEND_WR_LEN),
1587 : HCCL_E_PARA);
1588 0 : struct ibv_send_wr sendWr[SEND_WR_LEN] = {0};
1589 0 : struct ibv_sge sge[SEND_WR_LEN] = {0};
1590 :
1591 0 : for (uint32_t index = 0; index < memNum; index++) {
1592 : // 设置WR的SGE
1593 0 : sge[index].addr = reinterpret_cast<u64>(localMems[index].addr);
1594 0 : sge[index].length = remoteMems[index].size;
1595 0 : sge[index].lkey = localMems[index].key;
1596 :
1597 : // 设置WR属性
1598 0 : sendWr[index].wr_id = wrIdOffset_.fetch_add(1, std::memory_order_relaxed);
1599 0 : sendWr[index].num_sge = 1; // 只有一个SGE
1600 0 : sendWr[index].sg_list = &sge[index];
1601 0 : sendWr[index].wr.rdma.remote_addr = reinterpret_cast<u64>(remoteMems[index].addr);
1602 0 : sendWr[index].wr.rdma.rkey = remoteMems[index].key;
1603 0 : sendWr[index].next = (index == last) ? nullptr : &sendWr[index + 1]; // 第一个WR指向第二个WR
1604 0 : sendWr[index].send_flags = (index == last) ?
1605 : (fence ? (IBV_SEND_SIGNALED | IBV_SEND_FENCE) : IBV_SEND_SIGNALED) : 0; // 最后一个WR才需要回复CQE
1606 0 : sendWr[index].opcode = static_cast<enum ibv_wr_opcode>(opCode);
1607 0 : HCCL_DEBUG("[TransportDeviceIbverbs][HnsPostSend] Direct ibv_post_send[%llu], opcode=[0x%x], "
1608 : "remote_addr=[0x%llx], size=[%u], fence[%u]", wrIdOffset_.load(), sendWr[index].opcode,
1609 : sendWr[index].wr.rdma.remote_addr, sendWr[index].sg_list->length, fence);
1610 : }
1611 :
1612 0 : struct ibv_send_wr *badWr = nullptr;
1613 0 : struct WrExpRsp exp_rsp = {0};
1614 0 : struct ibv_qp *qp = reinterpret_cast<struct ibv_qp *>(ibvData.qpInfo.qpPtr);
1615 0 : CHK_PTR_NULL(qp);
1616 0 : HCCL_DEBUG("[TransportDeviceIbverbs][HnsPostSend] qp=%p, handle=%u, qp_num=%u, qp_type=%d, qp_stat=%d", qp,
1617 : qp->handle, qp->qp_num, qp->qp_type, qp->state);
1618 0 : HcclResult ret = HrtHnsIbvExpPostSend(qp, &sendWr[0], &badWr, &exp_rsp);
1619 0 : HCOMM_DSB();
1620 0 : CHK_PRT_RET(ret != HCCL_SUCCESS && ret != HCCL_E_AGAIN,
1621 : HCCL_ERROR("[TransportDeviceIbverbs][HnsPostSend] failed, qp=%p, handle=%u, qp_num=%u, qp_type=%d, qp_stat=%d",
1622 : qp, qp->handle, qp->qp_num, qp->qp_type, qp->state),
1623 : ret);
1624 0 : if (ret == HCCL_SUCCESS) {
1625 0 : dbInfo = exp_rsp.db_info;
1626 : }
1627 0 : return ret;
1628 : }
1629 :
1630 1 : HcclResult TransportDeviceIbverbs::Drain(Stream &stream)
1631 : {
1632 1 : CHK_PTR_NULL(dataNotify_);
1633 1 : CHK_PTR_NULL(remoteMemMsg_[MemType::NOTIFY_SRC_MEM].addr);
1634 1 : CHK_PTR_NULL(memMsg_[MemType::DATA_NOTIFY_MEM].addr);
1635 :
1636 1 : struct SgList list = {0};
1637 1 : struct SendWr wr = {nullptr};
1638 1 : list.addr = static_cast<u64>(reinterpret_cast<uintptr_t>(memMsg_[MemType::DATA_NOTIFY_MEM].addr));
1639 1 : list.len = static_cast<u32>(memMsg_[MemType::DATA_NOTIFY_MEM].len);
1640 1 : list.lkey = memMsg_[MemType::DATA_NOTIFY_MEM].lkey;
1641 :
1642 1 : wr.bufList = &list;
1643 1 : wr.bufNum = 1; /* 此处list只有一个,设置为1 */
1644 1 : wr.dstAddr = reinterpret_cast<u64>(remoteMemMsg_[MemType::NOTIFY_SRC_MEM].addr);
1645 1 : wr.rkey = remoteMemMsg_[MemType::NOTIFY_SRC_MEM].lkey;
1646 1 : wr.op = RaWrOpcode::RA_WR_RDMA_READ;
1647 1 : wr.sendFlag = RA_SEND_SIGNALED | RA_SEND_FENCE; // fence
1648 :
1649 1 : CHK_RET(RdmaSendAsync(wr, stream, WqeType::WQE_TYPE_DATA_NOTIFY, wr.dstAddr, INVALID_UINT));
1650 1 : CHK_RET(dispatcher_->SignalWait(dataNotify_->ptr(),
1651 : stream, machinePara_.localUserrank, machinePara_.remoteWorldRank,
1652 : INVALID_VALUE_STAGE, false, dataNotify_->notifyId_));
1653 1 : return HCCL_SUCCESS;
1654 : }
1655 : } // namespace hccl
|