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_heterog_roce.h"
12 : #include "log.h"
13 : #include "externalinput_pub.h"
14 : #include "mr_manager.h"
15 : #include "adapter_hccp.h"
16 : #include "adapter_rts.h"
17 : #include "network_manager_pub.h"
18 : #include "transport_ibverbs_pub.h"
19 : #include "hccl_socket.h"
20 :
21 : using namespace std;
22 : namespace hccl {
23 : constexpr u32 MAX_COSTTIME_COUNT = 1000;
24 : constexpr u32 MAX_TOTALCOST_COUNT = 10000; // 总耗时预警门槛 10ms
25 : constexpr u32 BLOCK_ALLOCATOR_POOL_SIZE = 4096;
26 : constexpr s32 PROTOCOL_TYPE = 0;
27 : constexpr s32 LINK_NUM = 3;
28 : constexpr u32 SOCKET_FOR_TAG_QP = 0;
29 : constexpr u32 SOCKET_FOR_DATA_QP = 1;
30 : constexpr u32 SOCKET_FOR_SENDRECV_QP = 2;
31 :
32 : constexpr u32 RECV_WQE_HDC_BATCH_NUM = 128;
33 : constexpr u32 RECV_WQE_HDC_BATCH_SUPPLEMENT = 1;
34 : constexpr u32 RECV_WQE_NUM_THRESHOLD = 96;
35 : constexpr u32 RECV_WQE_BATCH_SUPPLEMENT = 96;
36 : constexpr u32 SMALL_PAGE_SIZE = 4096;
37 : constexpr u32 EIGHE_BIT = 8;
38 : constexpr u32 WAIT_SLEEP_TIME_US = 50;
39 :
40 : enum class RdmaOp { OP_WRITE = 0, OP_SEND = 2, OP_READ = 4 };
41 :
42 0 : TransportHeterogRoce::TransportHeterogRoce(
43 : const std::string& transTag, HcclIpAddress& selfIp, HcclIpAddress& peerIp, u32 peerPort, u32 selfPort,
44 0 : const TransportResourceInfo& transportResourceInfo)
45 : : TransportHeterog(transTag, selfIp, peerIp, peerPort, selfPort, transportResourceInfo),
46 0 : nicRdmaHandle_(nullptr),
47 0 : mrManager_(transportResourceInfo.mrManager.get()),
48 0 : blockMemLkey_(transportResourceInfo.lkey),
49 0 : recvWqeBatchNum_(RECV_WQE_BATCH_NUM),
50 0 : recvWqeBatchThreshold_(RECV_WQE_NUM_THRESHOLD),
51 0 : recvWqeBatchSupplement_(RECV_WQE_BATCH_SUPPLEMENT),
52 0 : access_(RA_ACCESS_LOCAL_WRITE | RA_ACCESS_REMOTE_WRITE),
53 0 : tagRecvWqeNum_(0),
54 0 : dataRecvWqeNum_(0),
55 0 : dataRecvWqeExpNum_(0),
56 0 : memBlocksManager_(transportResourceInfo.memBlocksManager),
57 0 : pRecvWrInfosMem_(transportResourceInfo.pRecvWrInfosMem),
58 0 : deviceEvePtr_(nullptr),
59 0 : deviceEveLkey_(0),
60 0 : useDevMem_(true),
61 0 : isRawConn_(transportResourceInfo.isRawConn)
62 : {
63 0 : GetTransportResourceInfo(transportResourceInfo);
64 0 : }
65 :
66 0 : TransportHeterogRoce::TransportHeterogRoce(const TransportResourceInfo& transportResourceInfo)
67 : : TransportHeterog(transportResourceInfo),
68 0 : nicRdmaHandle_(nullptr),
69 0 : mrManager_(transportResourceInfo.mrManager.get()),
70 0 : blockMemLkey_(transportResourceInfo.lkey),
71 0 : recvWqeBatchNum_(RECV_WQE_BATCH_NUM),
72 0 : recvWqeBatchThreshold_(RECV_WQE_NUM_THRESHOLD),
73 0 : recvWqeBatchSupplement_(RECV_WQE_BATCH_SUPPLEMENT),
74 0 : access_(RA_ACCESS_LOCAL_WRITE | RA_ACCESS_REMOTE_WRITE),
75 0 : tagRecvWqeNum_(0),
76 0 : dataRecvWqeNum_(0),
77 0 : dataRecvWqeExpNum_(0),
78 0 : memBlocksManager_(transportResourceInfo.memBlocksManager),
79 0 : pRecvWrInfosMem_(transportResourceInfo.pRecvWrInfosMem),
80 0 : deviceEvePtr_(nullptr),
81 0 : deviceEveLkey_(0),
82 0 : useDevMem_(true),
83 0 : isRawConn_(transportResourceInfo.isRawConn)
84 : {
85 0 : GetTransportResourceInfo(transportResourceInfo);
86 0 : }
87 :
88 0 : TransportHeterogRoce::~TransportHeterogRoce() {}
89 :
90 0 : u64 HostAddrToDev(const u64& hostAddr, u64 hostAddrBegin, u64 devAddrBegin)
91 : {
92 0 : u64 devAddr = hostAddr - hostAddrBegin + devAddrBegin;
93 0 : return devAddr;
94 : }
95 :
96 0 : HcclResult TransportHeterogRoce::Init()
97 : {
98 0 : if (!isHdcMode_ && (remoteIsHdc_ && (deviceLogicId_ == HOST_DEVICE_ID))) {
99 0 : HCCL_INFO("TransportHeterogRoce no useDevMem_");
100 0 : useDevMem_ = false;
101 : }
102 :
103 0 : CHK_RET(CheckRecvMsgAndRequestBuffer());
104 :
105 0 : CHK_RET(GetNetworkResource());
106 :
107 0 : CHK_RET(PreQpConnect());
108 :
109 0 : CHK_RET(InitTransportConnect(PROTOCOL_TYPE, LINK_NUM));
110 :
111 0 : CHK_RET(ConnectAsync());
112 :
113 0 : return HCCL_SUCCESS;
114 : }
115 :
116 0 : HcclResult TransportHeterogRoce::Deinit()
117 : {
118 0 : if (isDeinited_ == true) {
119 0 : return HCCL_SUCCESS;
120 : }
121 :
122 0 : if (isHdcMode_) {
123 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
124 0 : CHK_PRT(MemBlocksManagerDeInit());
125 : }
126 0 : CHK_PRT(MrManagerDeInit());
127 0 : if (deviceEvePtr_ != nullptr) {
128 : #ifndef CCL_KERNEL
129 : CHK_RET(HrtDevFree(deviceEvePtr_));
130 : #endif
131 : }
132 : }
133 :
134 0 : CHK_RET(DeleteNotifyValueBuffer());
135 :
136 0 : CHK_RET(DestroyCqAndQp());
137 :
138 0 : CHK_RET(SocketClose());
139 :
140 0 : isDeinited_ = true;
141 0 : return HCCL_SUCCESS;
142 : }
143 :
144 : HcclResult
145 0 : TransportHeterogRoce::Isend(const TransData& sendData, const TransportEndPointParam& epParam, HcclRequestInfo*& request)
146 : {
147 0 : CHK_RET(GenerateSendRequest(sendData, epParam, request));
148 :
149 0 : u32 lkey = 0;
150 0 : CHK_RET(RegMr(
151 : reinterpret_cast<void*>(sendData.srcBuf), static_cast<u64>(sendData.count * SIZE_TABLE[sendData.dataType]),
152 : lkey));
153 0 : HCCL_DEBUG(
154 : "addr[%llu] count[%d] datatype[%s]", sendData.srcBuf, sendData.count,
155 : GetDataTypeEnumStr(sendData.dataType).c_str());
156 :
157 : HcclEnvelope envelope(
158 0 : request->transportRequest.protocol, request->transportRequest.transData, request->transportRequest.epParam,
159 0 : lkey, request->transportRequest.msn);
160 :
161 : // 如果建链未完成,或者积压的信封未发送完成,则Isend不进行信封发送。
162 : // Test接口中推动积压信封发送完成后,Isend接口才启动信封发送。
163 0 : std::unique_lock<std::mutex> lock(envelopeBacklogQueueLock_);
164 0 : if (GetState() != ConnState::CONN_STATE_COMPLETE) {
165 0 : envelopeBacklogQueue_.push(envelope);
166 0 : return HCCL_SUCCESS;
167 : }
168 :
169 0 : return SendEnvelope(envelope);
170 0 : }
171 :
172 0 : HcclResult TransportHeterogRoce::Send(const TransData& sendData, const TransportEndPointParam& epParam)
173 : {
174 0 : HCCL_ERROR("TransportHeterogRoce::Send is not supported.");
175 0 : return HCCL_E_NOT_SUPPORT;
176 : }
177 :
178 0 : HcclResult TransportHeterogRoce::Improbe(
179 : const TransportEndPointParam& epParam, s32& matched, HcclMessageInfo*& msg, HcclStatus& status, bool& flag)
180 : {
181 0 : return Improbe(epParam, matched, msg, status);
182 : }
183 :
184 0 : HcclResult TransportHeterogRoce::Improbe(
185 : const TransportEndPointParam& epParam, s32& matched, HcclMessageInfo*& msg, HcclStatus& status)
186 : {
187 : // 建链未完成时,返回未匹配到
188 0 : if (GetState() != ConnState::CONN_STATE_COMPLETE) {
189 0 : CHK_RET(ConnectAsync());
190 0 : return ProbeNothing(matched, msg, status);
191 : }
192 :
193 : // 先检查本地能否匹配
194 0 : HcclEnvelopeSummary envelopInfo;
195 0 : bool envelopeExist = GetSavedEnvelope(envelopInfo);
196 :
197 0 : auto probeSomething = [&]() -> HcclResult {
198 0 : CHK_RET(GenerateRecvMessage(envelopInfo, msg, status));
199 0 : matched = HCCL_IMPROBE_COMPLETED;
200 0 : return HCCL_SUCCESS;
201 0 : };
202 :
203 0 : if (envelopeExist) {
204 0 : return probeSomething();
205 : }
206 : // 再拉取roce cqe,检查是否能匹配
207 0 : CHK_RET(PullRecvRequestStatus());
208 :
209 0 : envelopeExist = GetSavedEnvelope(envelopInfo);
210 0 : if (envelopeExist) {
211 0 : return probeSomething();
212 : } else {
213 0 : return ProbeNothing(matched, msg, status);
214 : }
215 : }
216 :
217 : HcclResult
218 0 : TransportHeterogRoce::Iwrite(const TransData& sendData, const HcclEnvelope& envelope, HcclRequestInfo*& request)
219 : {
220 0 : if (isHdcMode_ && dataQpInfo_.qpMode != NORMAL_QP_MODE) {
221 0 : CHK_RET(TransportHeterog::WaitBuildLinkComplete());
222 : }
223 :
224 0 : TransportEndPointParam epParam{};
225 0 : CHK_RET(GenerateSendRequest(sendData, epParam, request));
226 0 : request->transportRequest.requestType = HcclRequestType::HCCL_REQUEST_RECV;
227 :
228 0 : u32 lkey = 0;
229 0 : CHK_RET(RegMr(
230 : reinterpret_cast<void*>(sendData.srcBuf), static_cast<u64>(sendData.count * SIZE_TABLE[sendData.dataType]),
231 : lkey, false));
232 :
233 0 : bool tmp = true;
234 0 : CHK_RET(GetQpStatus(tmp));
235 :
236 0 : if (!isHdcMode_ || dataQpInfo_.qpMode == NORMAL_QP_MODE) {
237 0 : dataWriteSge_.addr = static_cast<uint64_t>(sendData.srcBuf);
238 0 : dataWriteSge_.length = envelope.transData.count * SIZE_TABLE[envelope.transData.dataType];
239 0 : dataWriteSge_.lkey = lkey;
240 0 : dataWriteWr_.wr_id = reinterpret_cast<uint64_t>(request);
241 0 : dataWriteWr_.wr.rdma.remote_addr = static_cast<uint64_t>(envelope.transData.dstBuf);
242 0 : dataWriteWr_.wr.rdma.rkey = envelope.key;
243 :
244 0 : struct ibv_send_wr* badWr = nullptr;
245 0 : HCCL_INFO(
246 : "rdma write: remote addr[%llu] count[%d] datatype[%s] wrId[%llu]",
247 : reinterpret_cast<u64>(envelope.transData.dstBuf), envelope.transData.count,
248 : GetDataTypeEnumStr(envelope.transData.dataType).c_str(), dataWriteWr_.wr_id);
249 0 : CHK_RET(hrtIbvPostSend(dataQpInfo_.qp, &dataWriteWr_, &badWr));
250 : // 写notify
251 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
252 0 : Stream tmpStream(nullptr);
253 0 : CHK_RET(RecordNotify(tmpStream, RdmaNotifyOp::SEND_NOTIFY, dataWriteWr_.wr_id));
254 0 : }
255 0 : } else {
256 0 : struct SgList list = {};
257 0 : u64 srcBufDevAddr = 0;
258 0 : CHK_RET(dataQpMrManager_->GetDevVirAddr(
259 : reinterpret_cast<void*>(sendData.srcBuf),
260 : static_cast<u64>(envelope.transData.count * SIZE_TABLE[envelope.transData.dataType]), srcBufDevAddr));
261 :
262 0 : list.addr = srcBufDevAddr;
263 0 : list.len = envelope.transData.count * SIZE_TABLE[envelope.transData.dataType];
264 0 : list.lkey = lkey;
265 :
266 0 : struct SendWrV2 wr {};
267 0 : wr.wrId = reinterpret_cast<uint64_t>(request);
268 0 : HCCL_INFO("iwrite wr.wrId[%llu]", wr.wrId);
269 0 : wr.bufList = &list;
270 0 : wr.bufNum = 1;
271 0 : wr.dstAddr = static_cast<uint64_t>(envelope.transData.dstBuf);
272 0 : wr.rkey = envelope.key;
273 0 : wr.op = static_cast<u32>(RdmaOp::OP_WRITE);
274 0 : wr.sendFlag = RA_SEND_FENCE;
275 0 : struct SendWrRsp opRsp = {0};
276 0 : CHK_RET(HrtRaSendWrV2(dataQpInfo_.qpHandle, &wr, &opRsp, GetWorkflowMode()));
277 0 : CHK_RET(DoorBellSend(dataQpInfo_.qpMode, opRsp));
278 :
279 : // 写notify
280 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
281 0 : Stream tmpStream(nullptr);
282 0 : CHK_RET(RecordNotify(tmpStream, RdmaNotifyOp::SEND_NOTIFY, wr.wrId));
283 0 : }
284 :
285 0 : s32 writeAndNotifyFlag = HCCL_TEST_INCOMPLETED;
286 0 : TIME_PRINT(CHK_RET(this->Wait(*request, writeAndNotifyFlag)));
287 : }
288 :
289 0 : return HCCL_SUCCESS;
290 : }
291 :
292 0 : HcclResult TransportHeterogRoce::Imrecv(
293 : const TransData& recvData, HcclMessageInfo& msg, HcclRequestInfo*& request, bool flag, bool needRecordFlag)
294 : {
295 0 : HcclResult ret = Imrecv(recvData, msg, request);
296 0 : return ret;
297 : }
298 :
299 0 : HcclResult TransportHeterogRoce::Imrecv(const TransData& recvData, HcclMessageInfo& msg, HcclRequestInfo*& request)
300 : {
301 0 : CHK_RET(GenerateRecvRequest(recvData, msg, request));
302 :
303 0 : u32 lkey = 0;
304 0 : CHK_RET(RegMr(
305 : reinterpret_cast<void*>(recvData.dstBuf), static_cast<u64>(recvData.count * SIZE_TABLE[recvData.dataType]),
306 : lkey, false));
307 :
308 0 : HcclEnvelope& envelope = msg.envelope.envelope;
309 :
310 0 : if (!isHdcMode_ || dataQpInfo_.qpMode == NORMAL_QP_MODE) {
311 0 : dataReadSge_.addr = static_cast<uint64_t>(recvData.dstBuf);
312 0 : dataReadSge_.length = envelope.transData.count * SIZE_TABLE[envelope.transData.dataType];
313 0 : dataReadSge_.lkey = lkey;
314 0 : dataReadWr_.wr_id = reinterpret_cast<uint64_t>(request);
315 0 : dataReadWr_.wr.rdma.remote_addr = static_cast<uint64_t>(envelope.transData.srcBuf);
316 0 : dataReadWr_.wr.rdma.rkey = envelope.key;
317 0 : dataReadWr_.next = nullptr;
318 :
319 0 : if (!(remoteIsHdc_ && (deviceLogicId_ == HOST_DEVICE_ID))) {
320 0 : HCCL_INFO("general server not load ack ");
321 0 : dataReadWr_.next = &dataAckWr_;
322 0 : dataAckSge_.addr = reinterpret_cast<uint64_t>(&envelope.msn);
323 0 : dataAckSge_.length = sizeof(uint64_t);
324 0 : dataAckSge_.lkey = 0;
325 0 : dataAckWr_.wr_id = 0;
326 : }
327 :
328 0 : struct ibv_send_wr* badWr = nullptr;
329 0 : HCCL_INFO(
330 : "rdma read: remote addr[%llx] count[%d] datatype[%s] wrId[%llu]",
331 : reinterpret_cast<u64>(envelope.transData.srcBuf), envelope.transData.count,
332 : GetDataTypeEnumStr(envelope.transData.dataType).c_str(), dataReadWr_.wr_id);
333 0 : CHK_RET(hrtIbvPostSend(dataQpInfo_.qp, &dataReadWr_, &badWr));
334 0 : } else {
335 0 : if (envelope.transData.count == 0) {
336 0 : request->transportRequest.transData.count = 0;
337 0 : CHK_RET(FreeRecvMessage(msg));
338 0 : return HCCL_SUCCESS;
339 : }
340 :
341 0 : struct SgList list = {};
342 0 : u64 devAddr = 0;
343 0 : CHK_RET(dataQpMrManager_->GetDevVirAddr(
344 : reinterpret_cast<void*>(recvData.dstBuf), static_cast<u64>(recvData.count * SIZE_TABLE[recvData.dataType]),
345 : devAddr));
346 0 : list.addr = static_cast<uint64_t>(devAddr);
347 0 : list.len = envelope.transData.count * SIZE_TABLE[envelope.transData.dataType];
348 0 : list.lkey = lkey;
349 :
350 0 : struct SendWrV2 wr = {};
351 0 : wr.wrId = reinterpret_cast<uint64_t>(request);
352 :
353 0 : HCCL_INFO("Imrecv wr.wrId[%llu]", wr.wrId);
354 0 : wr.bufList = &list;
355 0 : wr.bufNum = 1; /* 此处list只有一个,设置为1 */
356 0 : wr.dstAddr = static_cast<uint64_t>(envelope.transData.srcBuf);
357 0 : wr.rkey = envelope.key;
358 0 : wr.op = static_cast<u32>(RdmaOp::OP_READ);
359 0 : wr.sendFlag = RA_SEND_SIGNALED | RA_SEND_FENCE;
360 0 : struct SendWrRsp opRsp = {};
361 0 : CHK_RET(HrtRaSendWrV2(dataQpInfo_.qpHandle, &wr, &opRsp, GetWorkflowMode()));
362 0 : CHK_RET(DoorBellSend(dataQpInfo_.qpMode, opRsp));
363 :
364 0 : s32 imrecvFlag = HCCL_TEST_INCOMPLETED;
365 0 : TIME_PRINT(CHK_RET(this->Wait(*request, imrecvFlag)));
366 : }
367 :
368 0 : CHK_RET(FreeRecvMessage(msg));
369 :
370 0 : return HCCL_SUCCESS;
371 : }
372 :
373 0 : HcclResult TransportHeterogRoce::Test(HcclRequestInfo& request, s32& flag, HcclStatus& compState)
374 : {
375 0 : if (isHdcMode_ && dataQpInfo_.qpMode != NORMAL_QP_MODE) {
376 0 : flag = HCCL_TEST_COMPLETED;
377 0 : HCCL_INFO("TransportHeterogRoce QueryRequestStatus: flag [%d]", flag);
378 0 : compState.error = 0;
379 0 : CHK_RET(FreeRequest(request));
380 0 : return HCCL_SUCCESS;
381 : }
382 :
383 : // 建链未完成时,继续推进建链流程;
384 0 : if (GetState() != ConnState::CONN_STATE_COMPLETE) {
385 0 : CHK_RET(ConnectAsync());
386 : }
387 :
388 0 : CHK_RET(PullSendOrRecvStatus(request));
389 :
390 0 : return QueryRequestStatus(request, flag, compState);
391 : }
392 :
393 0 : HcclResult TransportHeterogRoce::PullSendOrRecvStatus(const HcclRequestInfo& request)
394 : {
395 0 : if ((GetState() != ConnState::CONN_STATE_COMPLETE) && (GetState() != ConnState::CONN_STATE_FLUSH_QUEUE)) {
396 0 : return HCCL_SUCCESS;
397 : }
398 :
399 0 : if (request.transportRequest.requestType == HcclRequestType::HCCL_REQUEST_SEND) {
400 0 : CHK_RET(PullSendStatus());
401 0 : } else if (request.transportRequest.requestType == HcclRequestType::HCCL_REQUEST_RECV) {
402 0 : CHK_RET(PullRecvStatus());
403 : } else {
404 0 : HCCL_ERROR("[HcclTest] requestType[%u] is invalid", request.transportRequest.requestType);
405 0 : return HCCL_E_PARA;
406 : }
407 :
408 0 : return HCCL_SUCCESS;
409 : }
410 :
411 0 : HcclResult TransportHeterogRoce::Wait(HcclRequestInfo& request, s32& flag)
412 : {
413 : // 建链未完成时,继续推进建链流程;
414 0 : if (GetState() != ConnState::CONN_STATE_COMPLETE) {
415 0 : CHK_RET(ConnectAsync());
416 : }
417 :
418 0 : auto startTime = chrono::steady_clock::now();
419 0 : auto timeout = chrono::seconds(GetExternalInputHcclLinkTimeOut());
420 :
421 0 : while ((flag != HCCL_TEST_COMPLETED) && ((chrono::steady_clock::now() - startTime) < timeout)) {
422 0 : CHK_RET(PullSendOrRecvStatus(request));
423 :
424 0 : if (request.transportRequest.status >= 0) {
425 0 : flag = HCCL_TEST_COMPLETED;
426 0 : HCCL_INFO("QueryRequestStatus: flag[%d]", flag);
427 0 : return HCCL_SUCCESS;
428 : }
429 :
430 0 : SaluSleep(WAIT_SLEEP_TIME_US);
431 : }
432 :
433 0 : HCCL_ERROR("Wait Cqe timeOut[%d] s, State[%d]", GetExternalInputHcclLinkTimeOut(), GetState());
434 :
435 0 : return HCCL_E_TIMEOUT;
436 : }
437 :
438 0 : HcclResult TransportHeterogRoce::QueryRequestStatus(HcclRequestInfo& request, s32& flag, HcclStatus& compState)
439 : {
440 0 : if (request.transportRequest.status >= 0) {
441 : // 该request已完成
442 0 : flag = HCCL_TEST_COMPLETED;
443 0 : HCCL_INFO("QueryRequestStatus: flag [%d]", flag);
444 0 : compState.tag = request.transportRequest.epParam.src.tag;
445 0 : compState.srcRank = request.transportRequest.epParam.src.rank;
446 0 : compState.error = request.transportRequest.status;
447 0 : CHK_RET(FreeRequest(request));
448 : }
449 0 : return HCCL_SUCCESS;
450 : }
451 :
452 0 : HcclResult TransportHeterogRoce::PullSendStatus(bool allowNotify)
453 : {
454 0 : if (isHdcMode_ && (tagQpInfo_.qpMode == OFFLINE_QP_MODE || tagQpInfo_.qpMode == OFFLINE_QP_MODE_EXT)) {
455 0 : return HCCL_SUCCESS;
456 : }
457 :
458 : struct ibv_wc wcTagCq[HCCL_POLL_CQ_DEPTH];
459 0 : s32 tagCqNum = 0;
460 0 : CHK_RET(PollCq(tagQpInfo_, true, tagCqNum, wcTagCq));
461 0 : for (int i = 0; i < tagCqNum; i++) {
462 0 : if (wcTagCq[i].status != 0) {
463 0 : CHK_RET(ParseErrorTagSqe(wcTagCq, i));
464 0 : HCCL_ERROR("rdma poll tag sq failed, cqe status[%u]", wcTagCq[i].status);
465 0 : return HCCL_E_NETWORK;
466 : }
467 : }
468 : struct ibv_wc wcDataRq[HCCL_POLL_CQ_DEPTH];
469 0 : s32 dataRqNum = 0;
470 :
471 0 : CHK_RET(PollCq(dataQpInfo_, false, dataRqNum, wcDataRq));
472 0 : if ((dataRqNum == 0) && allowNotify) {
473 0 : CHK_RET(hrtIbvReqNotifyCq(dataQpInfo_.recvCq, 0));
474 0 : } else {
475 0 : HCCL_DEBUG("data rq: poll cq num:%d", dataRqNum);
476 0 : CHK_RET(ParseDataRqes(wcDataRq, dataRqNum));
477 : }
478 0 : return HCCL_SUCCESS;
479 : }
480 :
481 0 : HcclResult TransportHeterogRoce::PullRecvRequestStatus(bool allowNotify)
482 : {
483 : struct ibv_wc wc[HCCL_POLL_CQ_DEPTH];
484 0 : s32 num = 0;
485 0 : CHK_RET(PollCq(tagQpInfo_, false, num, wc));
486 0 : if ((num == 0) && allowNotify) {
487 0 : CHK_RET(hrtIbvReqNotifyCq(tagQpInfo_.recvCq, 0));
488 0 : } else {
489 0 : HCCL_DEBUG("tag rq: poll cq num:%d", num);
490 0 : CHK_RET(ParseTagRqes(wc, num));
491 : }
492 0 : return HCCL_SUCCESS;
493 : }
494 :
495 0 : HcclResult TransportHeterogRoce::PullRecvStatus(bool allowNotify)
496 : {
497 0 : HCCL_INFO("Pull dataQp RecvStatus");
498 : struct ibv_wc wc[HCCL_POLL_CQ_DEPTH];
499 0 : s32 num = 0;
500 0 : CHK_RET(PollCq(dataQpInfo_, true, num, wc));
501 0 : if ((num == 0) && allowNotify) {
502 0 : CHK_RET(hrtIbvReqNotifyCq(dataQpInfo_.sendCq, 0));
503 0 : } else {
504 0 : HCCL_DEBUG("data sq: poll cq num:%d", num);
505 0 : CHK_RET(ParseDataSqes(wc, num));
506 : }
507 0 : return HCCL_SUCCESS;
508 : }
509 :
510 0 : HcclResult TransportHeterogRoce::ParseTagRqes(const struct ibv_wc* wc, int num)
511 : {
512 0 : for (int i = 0; i < num; i++) {
513 0 : HCCL_INFO("rq cqe info: wrId[%llu] status[%u] opcode[%u]", wc[i].wr_id, wc[i].status, wc[i].opcode);
514 0 : CHK_PRT_RET(
515 : wc[i].status != 0,
516 : HCCL_ERROR(
517 : "rdma send failed, cqe status[%u] wrId[%llu] opcode[%u]", wc[i].status, wc[i].wr_id, wc[i].opcode),
518 : HCCL_E_INTERNAL);
519 0 : RecvWrInfo* info = reinterpret_cast<RecvWrInfo*>(wc[i].wr_id);
520 0 : CHK_PTR_NULL(info);
521 :
522 0 : TransportHeterogRoce* transportPtr = reinterpret_cast<TransportHeterogRoce*>(info->transportHandle);
523 0 : CHK_PTR_NULL(transportPtr);
524 0 : CHK_RET(transportPtr->SupplyTagRecvWqe());
525 0 : HcclEnvelope* envelope = nullptr;
526 0 : if (useDevMem_ && (deviceLogicId_ == HOST_DEVICE_ID)) {
527 : #ifndef CCL_KERNEL
528 : // 根据device内存求host内存
529 : CHK_RET(hrtSetDevice(index_));
530 : u64 uDevPtr = reinterpret_cast<u64>(info->buf);
531 : void* devPtr = reinterpret_cast<void*>(uDevPtr);
532 : u64 uHostPtr = uDevPtr - reinterpret_cast<uint64_t>(deviceEvePtr_) + hostAddrBegin_;
533 : void* hostPtr = reinterpret_cast<void*>(uHostPtr);
534 : HCCL_DEBUG("ParseTagRqes devPtr[%p][%llu] hostPtr[%p][%llu]", devPtr, uDevPtr, hostPtr, uHostPtr);
535 : CHK_RET(hrtMemcpy(
536 : hostPtr, MEM_BLOCK_SIZE, devPtr, MEM_BLOCK_SIZE, HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_DEVICE_TO_HOST));
537 : envelope = reinterpret_cast<HcclEnvelope*>(hostPtr);
538 : // ps侧DbSend对当前线程SetDevice后,改变了原GE通信域初始化时setdevice 0
539 : // 若后面使用本线程save会导致getctx失败获取不到通信域句柄,所以需要在此处重新set回默认
540 : CHK_RET(hrtSetDevice(0));
541 : #endif
542 : } else {
543 0 : envelope = reinterpret_cast<HcclEnvelope*>(info->buf);
544 : }
545 0 : CHK_PTR_NULL(envelope);
546 :
547 0 : HCCL_INFO(
548 : "recv request: tag:%d srcRank:%u dstRank:%u status:%u msn:0x%016llx count:%d", envelope->epParam.src.tag,
549 : envelope->epParam.src.rank, envelope->epParam.dst.rank, wc[i].status, envelope->msn,
550 : envelope->transData.count);
551 0 : HcclEnvelopeSummary envelopSummary(*envelope, wc[i].status);
552 0 : transportPtr->SaveEnvelope(envelopSummary);
553 0 : CHK_RET(transportPtr->FreeMemBlock(envelope));
554 0 : CHK_RET(transportPtr->FreeRecvWrId(wc[i].wr_id));
555 : }
556 0 : return HCCL_SUCCESS;
557 : }
558 :
559 0 : void TransportHeterogRoce::SaveEnvelope(HcclEnvelopeSummary& envelope)
560 : {
561 0 : unique_lock<mutex> lock(envelopeQueMutex_);
562 0 : envelopeQue_.push(envelope);
563 0 : }
564 :
565 0 : bool TransportHeterogRoce::GetSavedEnvelope(HcclEnvelopeSummary& envelope)
566 : {
567 0 : unique_lock<mutex> lock(envelopeQueMutex_);
568 0 : if (envelopeQue_.empty()) {
569 0 : return false;
570 : }
571 0 : envelope = envelopeQue_.front();
572 0 : envelopeQue_.pop();
573 0 : return true;
574 0 : }
575 :
576 0 : HcclResult TransportHeterogRoce::ParseErrorTagSqe(const struct ibv_wc* wc, int index)
577 : {
578 : // wr_id内容即信封中的msn
579 0 : HcclRequestInfo* wrPtr = reinterpret_cast<HcclRequestInfo*>(wc[index].wr_id);
580 0 : CHK_PTR_NULL(wrPtr);
581 0 : wrPtr->transportRequest.status = wc[index].status;
582 0 : TransportHeterogRoce* transportPtr = reinterpret_cast<TransportHeterogRoce*>(wrPtr->transportHandle);
583 0 : CHK_PTR_NULL(transportPtr);
584 :
585 0 : HCCL_INFO(
586 : "exception send msg: tag:%d srcRank:%u dstRank:%u status:%d msn:0x%016llx request:%p",
587 : wrPtr->transportRequest.epParam.src.tag, wrPtr->transportRequest.epParam.src.rank,
588 : wrPtr->transportRequest.epParam.dst.rank, wrPtr->transportRequest.status, wrPtr->transportRequest.msn, wrPtr);
589 :
590 0 : CHK_RET(transportPtr->DeregMr(
591 : reinterpret_cast<void*>(wrPtr->transportRequest.transData.srcBuf),
592 : static_cast<u64>(
593 : wrPtr->transportRequest.transData.count * SIZE_TABLE[wrPtr->transportRequest.transData.dataType])));
594 0 : return HCCL_SUCCESS;
595 : }
596 :
597 0 : HcclResult TransportHeterogRoce::ParseDataRqes(const struct ibv_wc* wc, int num)
598 : {
599 0 : for (int i = 0; i < num; i++) {
600 0 : HCCL_INFO("rq cqe info: wrId[%llu] status[%u] opcode[%u]", wc[i].wr_id, wc[i].status, wc[i].opcode);
601 0 : CHK_PRT_RET(
602 : wc[i].status != 0,
603 : HCCL_ERROR(
604 : "rdma poll data rq failed, cqe status[%u] wrId[%llu] opcode[%u]", wc[i].status, wc[i].wr_id,
605 : wc[i].opcode),
606 : HCCL_E_NETWORK);
607 0 : RecvWrInfo* info = reinterpret_cast<RecvWrInfo*>(wc[i].wr_id);
608 0 : CHK_PTR_NULL(info);
609 0 : HcclRequestInfo* wrPtr = reinterpret_cast<HcclRequestInfo*>(*reinterpret_cast<u64*>(info->buf));
610 0 : CHK_PRT_RET(
611 : wrPtr == nullptr, HCCL_ERROR("wrId[%llu] status[%u] opcode[%u]", wc[i].wr_id, wc[i].status, wc[i].opcode),
612 : HCCL_E_PTR);
613 0 : TransportHeterogRoce* transportPtr = reinterpret_cast<TransportHeterogRoce*>(wrPtr->transportHandle);
614 0 : CHK_PRT_RET(
615 : transportPtr == nullptr,
616 : HCCL_ERROR(
617 : "wrId[%llu] opcode[%u] tag:%d peerRank:%u status:%d msn:0x%016llx request:%p", wc[i].wr_id,
618 : wc[i].opcode, wrPtr->transportRequest.epParam.src.tag, wrPtr->transportRequest.epParam.src.rank,
619 : wrPtr->transportRequest.status, wrPtr->transportRequest.msn, wrPtr),
620 : HCCL_E_PTR);
621 0 : CHK_RET(transportPtr->SupplyDataRecvWqe());
622 0 : wrPtr->transportRequest.status = wc[i].status;
623 :
624 0 : CHK_RET(transportPtr->DeregMr(
625 : reinterpret_cast<void*>(wrPtr->transportRequest.transData.srcBuf),
626 : static_cast<u64>(
627 : wrPtr->transportRequest.transData.count * SIZE_TABLE[wrPtr->transportRequest.transData.dataType]),
628 : false));
629 0 : CHK_RET(transportPtr->FreeMemBlock(info->buf));
630 0 : CHK_RET(transportPtr->FreeRecvWrId(wc[i].wr_id));
631 0 : HCCL_INFO(
632 : "send completion: tag:%d peerRank:%u status:%d msn:0x%016llx request:%p",
633 : wrPtr->transportRequest.epParam.src.tag, wrPtr->transportRequest.epParam.src.rank,
634 : wrPtr->transportRequest.status, wrPtr->transportRequest.msn, wrPtr);
635 : }
636 0 : return HCCL_SUCCESS;
637 : }
638 :
639 0 : HcclResult TransportHeterogRoce::ParseDataSqes(const struct ibv_wc* wc, int num)
640 : {
641 0 : for (int i = 0; i < num; i++) {
642 0 : HCCL_INFO("sq cqe info: wrId[%llu] status[%u] opcode[%u]", wc[i].wr_id, wc[i].status, wc[i].opcode);
643 0 : CHK_PRT_RET(
644 : wc[i].status != 0,
645 : HCCL_ERROR(
646 : "rdma poll data sq failed, cqe status[%u] wrId[%llu] opcode[%u]", wc[i].status, wc[i].wr_id,
647 : wc[i].opcode),
648 : HCCL_E_NETWORK);
649 0 : HcclRequestInfo* wrPtr = reinterpret_cast<HcclRequestInfo*>(wc[i].wr_id);
650 :
651 0 : CHK_PRT_RET(
652 : wrPtr == nullptr, HCCL_ERROR("wrId[%llu] status[%u] opcode[%u]", wc[i].wr_id, wc[i].status, wc[i].opcode),
653 : HCCL_E_PTR);
654 0 : TransportHeterogRoce* transportPtr = reinterpret_cast<TransportHeterogRoce*>(wrPtr->transportHandle);
655 0 : CHK_PRT_RET(
656 : transportPtr == nullptr,
657 : HCCL_ERROR(
658 : "wrId[%llu] opcode[%u] tag:%d peerRank:%u status:%d msn:0x%016llx request:%p", wc[i].wr_id,
659 : wc[i].opcode, wrPtr->transportRequest.epParam.src.tag, wrPtr->transportRequest.epParam.src.rank,
660 : wrPtr->transportRequest.status, wrPtr->transportRequest.msn, wrPtr),
661 : HCCL_E_PTR);
662 0 : wrPtr->transportRequest.status = wc[i].status;
663 :
664 0 : if (!isHdcMode_ && !(remoteIsHdc_ && (deviceLogicId_ == HOST_DEVICE_ID))) {
665 0 : CHK_RET(transportPtr->DeregMr(
666 : reinterpret_cast<void*>(wrPtr->transportRequest.transData.dstBuf),
667 : static_cast<u64>(
668 : wrPtr->transportRequest.transData.count * SIZE_TABLE[wrPtr->transportRequest.transData.dataType]),
669 : false));
670 : }
671 0 : HCCL_INFO(
672 : "recv completion: tag:%d peerRank:%u status:%d msn:0x%016llx", wrPtr->transportRequest.epParam.src.tag,
673 : wrPtr->transportRequest.epParam.src.rank, wrPtr->transportRequest.status, wrPtr->transportRequest.msn);
674 : }
675 0 : return HCCL_SUCCESS;
676 : }
677 :
678 0 : HcclResult TransportHeterogRoce::SendEnvelope(HcclEnvelope& envelopInfo, void* stream)
679 : {
680 0 : if (isHdcMode_ && tagQpInfo_.qpMode != NORMAL_QP_MODE) {
681 0 : CHK_RET(TransportHeterog::WaitBuildLinkComplete());
682 : }
683 :
684 0 : if (!isHdcMode_ || tagQpInfo_.qpMode == NORMAL_QP_MODE) {
685 0 : CHK_RET(SendFlowControl());
686 : }
687 :
688 0 : if (!isHdcMode_ || tagQpInfo_.qpMode == NORMAL_QP_MODE) {
689 0 : envelopeSge_.addr = reinterpret_cast<uint64_t>(&envelopInfo);
690 0 : envelopeSge_.length = sizeof(envelopInfo);
691 0 : envelopeSge_.lkey = 0;
692 0 : envelopeWr_.wr_id = envelopInfo.msn;
693 :
694 0 : struct ibv_send_wr* badWr = nullptr;
695 0 : HCCL_INFO(
696 : "rdma send: srcRank[%u] dstRank[%u] tag[%d]: addr[%llu] count[%d] dtype[%s] msn[%llu]",
697 : envelopInfo.epParam.src.rank, envelopInfo.epParam.dst.rank, envelopInfo.epParam.src.tag,
698 : reinterpret_cast<u64>(envelopInfo.transData.srcBuf), envelopInfo.transData.count,
699 : GetDataTypeEnumStr(envelopInfo.transData.dataType).c_str(), envelopInfo.msn);
700 0 : CHK_RET(hrtIbvPostSend(tagQpInfo_.qp, &envelopeWr_, &badWr));
701 0 : } else {
702 0 : struct SgList list = {};
703 0 : list.addr = reinterpret_cast<uint64_t>(&envelopInfo);
704 0 : list.len = sizeof(envelopInfo);
705 0 : list.lkey = 0;
706 :
707 0 : struct SendWr wr = {};
708 0 : wr.bufList = &list;
709 0 : wr.bufNum = 1;
710 0 : wr.op = static_cast<u32>(RdmaOp::OP_SEND);
711 0 : wr.sendFlag = RA_SEND_SIGNALED;
712 0 : struct SendWrRsp opRsp = {};
713 0 : CHK_RET(HrtRaSendWr(tagQpInfo_.qpHandle, &wr, &opRsp));
714 0 : CHK_RET(DoorBellSend(tagQpInfo_.qpMode, opRsp, stream));
715 : }
716 :
717 0 : return HCCL_SUCCESS;
718 : }
719 :
720 0 : HcclResult TransportHeterogRoce::InitTagRecvWqe()
721 : {
722 0 : CHK_RET(IssueRecvWqe(tagQpInfo_.qp, recvWqeBatchNum_));
723 0 : tagRecvWqeNum_ = recvWqeBatchNum_;
724 0 : return HCCL_SUCCESS;
725 : }
726 :
727 0 : HcclResult TransportHeterogRoce::InitDataRecvWqe()
728 : {
729 0 : CHK_RET(IssueRecvWqe(dataQpInfo_.qp, recvWqeBatchNum_));
730 0 : dataRecvWqeNum_ = recvWqeBatchNum_;
731 0 : dataRecvWqeExpNum_ = recvWqeBatchNum_;
732 0 : return HCCL_SUCCESS;
733 : }
734 :
735 0 : HcclResult TransportHeterogRoce::SendFlowControl()
736 : {
737 0 : if (dataRecvWqeNum_ <= recvWqeBatchThreshold_) {
738 0 : CHK_RET(IssueRecvWqe(dataQpInfo_.qp, recvWqeBatchSupplement_));
739 0 : dataRecvWqeNum_ += recvWqeBatchSupplement_;
740 0 : dataRecvWqeExpNum_ += recvWqeBatchSupplement_;
741 : }
742 :
743 0 : u32 dataRecvWqeNum = dataRecvWqeNum_.load();
744 0 : u32 dataRecvWqeExpNum = dataRecvWqeExpNum_.load();
745 0 : if (dataRecvWqeNum - dataRecvWqeExpNum >= recvWqeBatchSupplement_) {
746 0 : CHK_RET(PullSendStatus());
747 0 : HCCL_RUN_INFO(
748 : "Flow control is activated, because dataRecvWqeNum[%u] - dataRecvWqeExpNum[%u] >="
749 : " recvWqeBatchSupplement[%u]",
750 : dataRecvWqeNum, dataRecvWqeExpNum, recvWqeBatchSupplement_);
751 :
752 0 : return HCCL_E_AGAIN;
753 : }
754 :
755 0 : dataRecvWqeExpNum_--;
756 0 : return HCCL_SUCCESS;
757 : }
758 :
759 0 : HcclResult TransportHeterogRoce::SupplyTagRecvWqe()
760 : {
761 0 : tagRecvWqeNum_--;
762 0 : if (tagRecvWqeNum_ <= recvWqeBatchThreshold_) {
763 0 : CHK_RET(IssueRecvWqe(tagQpInfo_.qp, recvWqeBatchSupplement_));
764 0 : tagRecvWqeNum_ += recvWqeBatchSupplement_;
765 : }
766 :
767 0 : return HCCL_SUCCESS;
768 : }
769 :
770 0 : HcclResult TransportHeterogRoce::SupplyDataRecvWqe()
771 : {
772 0 : dataRecvWqeNum_--;
773 0 : if (dataRecvWqeNum_ <= recvWqeBatchThreshold_) {
774 0 : CHK_RET(IssueRecvWqe(dataQpInfo_.qp, recvWqeBatchSupplement_));
775 0 : dataRecvWqeNum_ += recvWqeBatchSupplement_;
776 0 : dataRecvWqeExpNum_ += recvWqeBatchSupplement_;
777 : }
778 0 : return HCCL_SUCCESS;
779 : }
780 :
781 0 : HcclResult TransportHeterogRoce::IssueRecvWqe(struct ibv_qp* qp, u32 num)
782 : {
783 0 : if (isHdcMode_ && (tagQpInfo_.qpMode == OFFLINE_QP_MODE || tagQpInfo_.qpMode == OFFLINE_QP_MODE_EXT)) {
784 0 : return HCCL_SUCCESS;
785 : }
786 :
787 0 : list<void*> blockList(num, nullptr);
788 0 : CHK_RET(AllocMemBlocks(blockList));
789 :
790 0 : auto iter = blockList.begin();
791 0 : struct ibv_recv_wr* nextRqWr = nullptr;
792 0 : struct ibv_recv_wr rqWr[num];
793 0 : struct ibv_sge sgeList[num];
794 :
795 0 : std::vector<struct RecvWrlistData> recvWrVec(num);
796 0 : struct RecvWrlistData* recvWr = recvWrVec.data();
797 :
798 0 : if (!isHdcMode_ || tagQpInfo_.qpMode == NORMAL_QP_MODE) {
799 0 : for (int i = num - 1; i >= 0; i--) {
800 0 : CHK_PTR_NULL(*iter);
801 0 : u64 wrId = 0;
802 0 : CHK_RET(GenerateRecvWrId(*iter, wrId));
803 0 : rqWr[i].wr_id = wrId;
804 0 : rqWr[i].next = nextRqWr;
805 0 : rqWr[i].sg_list = &sgeList[i];
806 0 : rqWr[i].num_sge = 1;
807 0 : sgeList[i].addr = reinterpret_cast<uint64_t>(*iter);
808 0 : sgeList[i].length = MEM_BLOCK_SIZE;
809 0 : sgeList[i].lkey = blockMemLkey_;
810 0 : nextRqWr = &rqWr[i];
811 0 : iter++;
812 : }
813 :
814 0 : struct ibv_recv_wr* badRqWr = nullptr;
815 0 : CHK_RET(hrtIbvPostRecv(qp, &rqWr[0], &badRqWr));
816 0 : return HCCL_SUCCESS;
817 : } else {
818 0 : for (int i = num - 1; i >= 0; i--) {
819 0 : CHK_PTR_NULL(*iter);
820 0 : u64 wrId = 0;
821 0 : if (useDevMem_) {
822 : // 根据host内存地址计算出device内存地址
823 : u64 uDevPtr
824 0 : = reinterpret_cast<uint64_t>(*iter) - hostAddrBegin_ + reinterpret_cast<uint64_t>(deviceEvePtr_);
825 0 : CHK_RET(GenerateRecvWrId(reinterpret_cast<void*>(uDevPtr), wrId));
826 0 : recvWr[i].memList.addr = uDevPtr;
827 0 : recvWr[i].memList.lkey = deviceEveLkey_;
828 : } else {
829 0 : CHK_RET(GenerateRecvWrId(*iter, wrId));
830 0 : recvWr[i].memList.addr
831 0 : = HostAddrToDev(reinterpret_cast<uint64_t>(*iter), hostAddrBegin_, devAddrBegin_);
832 0 : recvWr[i].memList.lkey = blockMemLkey_;
833 : }
834 0 : recvWr[i].wrId = wrId;
835 0 : recvWr[i].memList.len = MEM_BLOCK_SIZE;
836 0 : iter++;
837 : }
838 : }
839 :
840 0 : u32 completeNum = 0;
841 0 : s32 ret = hrtRaRecvWrlist(tagQpInfo_.qpHandle, recvWr, num, &completeNum);
842 0 : if (ret == HCCL_SUCCESS && completeNum == num) {
843 0 : HCCL_INFO("hrtRaRecvWrlist success ");
844 0 : return HCCL_SUCCESS;
845 : } else {
846 0 : HCCL_ERROR("[Transport][RdmaData]In RdmaDataTransport, hrtRaRecvWrlist failed. ret[%d]", ret);
847 0 : return HCCL_E_NETWORK;
848 : }
849 :
850 : return HCCL_SUCCESS;
851 0 : }
852 :
853 0 : HcclResult TransportHeterogRoce::GetQpStatus(bool& completed)
854 : {
855 0 : int qpStatus = 0;
856 0 : s32 ret = 0;
857 :
858 0 : ret = hrtGetRaQpStatus(tagQpInfo_.qpHandle, &qpStatus);
859 0 : if (ret != 0) {
860 0 : HCCL_ERROR("get tag qp status fail. qpStatus[%d] ret[%d]", qpStatus, ret);
861 0 : return HCCL_E_INTERNAL;
862 0 : } else if (ret == 0 && qpStatus != 1) { // 为1时,qp 建链成功
863 0 : return HCCL_E_AGAIN;
864 : }
865 :
866 0 : ret = hrtGetRaQpStatus(dataQpInfo_.qpHandle, &qpStatus);
867 0 : if (ret != 0) {
868 0 : HCCL_ERROR("get data qp status fail. qpStatus[%d] ret[%d]", qpStatus, ret);
869 0 : return HCCL_E_INTERNAL;
870 0 : } else if (ret == 0 && qpStatus != 1) { // 为1时,qp 建链成功
871 0 : return HCCL_E_AGAIN;
872 : }
873 :
874 0 : completed = true;
875 0 : return HCCL_SUCCESS;
876 : }
877 :
878 0 : HcclResult TransportHeterogRoce::AllocMemBlocks(list<void*>& blockList)
879 : {
880 : const std::unique_ptr<HeterogMemBlocksManager>& memBlocksManagerPtr
881 0 : = (IsRamdHandleLevelMr()) ? memBlocksManager_ : tagMemBlocksManager_;
882 0 : CHK_PTR_NULL(memBlocksManagerPtr);
883 0 : CHK_RET(memBlocksManagerPtr->Alloc(blockList));
884 0 : if (isHdcMode_) {
885 0 : for (auto iter : blockList) {
886 0 : wqeBlockLists_.push_back(iter);
887 : }
888 : }
889 0 : return HCCL_SUCCESS;
890 : }
891 :
892 0 : HcclResult TransportHeterogRoce::FreeMemBlock(void* block)
893 : {
894 : const std::unique_ptr<HeterogMemBlocksManager>& memBlocksManagerPtr
895 0 : = ((IsRamdHandleLevelMr())) ? memBlocksManager_ : tagMemBlocksManager_;
896 0 : CHK_PTR_NULL(memBlocksManagerPtr);
897 0 : CHK_RET(memBlocksManagerPtr->Free(block));
898 0 : if (isHdcMode_) {
899 0 : auto iter = std::find(wqeBlockLists_.begin(), wqeBlockLists_.end(), block);
900 0 : if (iter != wqeBlockLists_.end()) {
901 0 : wqeBlockLists_.erase(iter);
902 : }
903 : }
904 0 : return HCCL_SUCCESS;
905 : }
906 :
907 0 : HcclResult TransportHeterogRoce::FreeRecvWrId(u64 wrId)
908 : {
909 0 : pRecvWrInfosMem_->Free(reinterpret_cast<RecvWrInfo*>(wrId));
910 0 : return HCCL_SUCCESS;
911 : }
912 :
913 0 : HcclResult TransportHeterogRoce::GenerateRecvWrId(void* recvBuf, u64& wrId)
914 : {
915 0 : RecvWrInfo* data = pRecvWrInfosMem_->Alloc();
916 0 : CHK_PTR_NULL(data);
917 0 : data->buf = recvBuf;
918 0 : data->transportHandle = reinterpret_cast<void*>(this);
919 0 : CHK_PTR_NULL(data->transportHandle);
920 0 : wrId = reinterpret_cast<uint64_t>(data);
921 0 : return HCCL_SUCCESS;
922 : }
923 :
924 0 : HcclResult TransportHeterogRoce::GetNetworkResource()
925 : {
926 0 : RaResourceInfo raResourceInfo;
927 0 : CHK_RET(NetworkManager::GetInstance(index_).GetRaResourceInfo(raResourceInfo));
928 0 : auto it = raResourceInfo.nicSocketMap.find(selfIp_);
929 0 : if (it == raResourceInfo.nicSocketMap.end()) {
930 0 : HCCL_ERROR("[TransportHeterogRoce][Init]nic socket handle did not found");
931 0 : return HCCL_E_PARA;
932 : }
933 0 : nicSocketHandle_ = it->second.nicSocketHandle;
934 0 : CHK_PTR_NULL(nicSocketHandle_);
935 0 : nicRdmaHandle_ = it->second.nicRdmaHandle;
936 0 : CHK_PTR_NULL(nicRdmaHandle_);
937 0 : HCCL_INFO(
938 : "TransportHeterogRoce GetNetworkResource index_[%d] nicSocketHandle_[%p] nicRdmaHandle_[%p]", index_,
939 : nicSocketHandle_, nicRdmaHandle_);
940 0 : return HCCL_SUCCESS;
941 0 : }
942 :
943 0 : HcclResult TransportHeterogRoce::PreQpConnect()
944 : {
945 : // 创建QP及CQ,多个QP可共享CQ
946 0 : CHK_RET(CreateCqAndQp());
947 :
948 0 : if (isHdcMode_) { // 不是HDC模式,走的peer,但是训练时,wqe下发情况和hdc相同
949 0 : CHK_RET(PreHdcResource());
950 : } else {
951 : // 下发post recv, 注:HCCP完成QP建链后需要两端握手确认QP状态OK后才能发起通信
952 0 : CHK_RET(InitTagRecvWqe());
953 0 : if (!(remoteIsHdc_ && (deviceLogicId_ == HOST_DEVICE_ID))) {
954 0 : CHK_RET(InitDataRecvWqe());
955 : }
956 : }
957 :
958 : // 为提高收发处理速度,提前准备post send需要的wr模板
959 0 : CHK_SAFETY_FUNC_RET(memset_s(&envelopeWr_, sizeof(struct ibv_send_wr), 0, sizeof(struct ibv_send_wr)));
960 0 : envelopeWr_.sg_list = &envelopeSge_;
961 0 : envelopeWr_.next = nullptr;
962 0 : envelopeWr_.num_sge = 1;
963 0 : envelopeWr_.opcode = IBV_WR_SEND;
964 0 : envelopeWr_.send_flags = IBV_SEND_SIGNALED | IBV_SEND_INLINE;
965 :
966 0 : CHK_SAFETY_FUNC_RET(memset_s(&dataReadWr_, sizeof(struct ibv_send_wr), 0, sizeof(struct ibv_send_wr)));
967 0 : dataReadWr_.sg_list = &dataReadSge_;
968 0 : dataReadWr_.next = nullptr;
969 0 : dataReadWr_.num_sge = 1;
970 0 : dataReadWr_.opcode = IBV_WR_RDMA_READ;
971 0 : dataReadWr_.send_flags = IBV_SEND_SIGNALED | IBV_SEND_FENCE;
972 :
973 0 : CHK_SAFETY_FUNC_RET(memset_s(&dataWriteWr_, sizeof(struct ibv_send_wr), 0, sizeof(struct ibv_send_wr)));
974 0 : dataWriteWr_.sg_list = &dataWriteSge_;
975 0 : dataWriteWr_.next = nullptr;
976 0 : dataWriteWr_.num_sge = 1;
977 0 : dataWriteWr_.opcode = IBV_WR_RDMA_WRITE;
978 0 : dataWriteWr_.send_flags = IBV_SEND_FENCE;
979 :
980 0 : CHK_SAFETY_FUNC_RET(memset_s(¬ifyWriteWr_, sizeof(struct ibv_send_wr), 0, sizeof(struct ibv_send_wr)));
981 0 : notifyWriteWr_.sg_list = ¬ifyWriteSge_;
982 0 : notifyWriteWr_.next = nullptr;
983 0 : notifyWriteWr_.num_sge = 1;
984 0 : notifyWriteWr_.opcode = IBV_WR_RDMA_WRITE;
985 0 : notifyWriteWr_.send_flags = IBV_SEND_SIGNALED | IBV_SEND_FENCE;
986 :
987 0 : CHK_SAFETY_FUNC_RET(memset_s(&dataAckWr_, sizeof(struct ibv_send_wr), 0, sizeof(struct ibv_send_wr)));
988 0 : dataAckWr_.sg_list = &dataAckSge_;
989 0 : dataAckWr_.next = nullptr;
990 0 : dataAckWr_.num_sge = 1;
991 0 : dataAckWr_.opcode = IBV_WR_SEND_WITH_IMM;
992 0 : dataAckWr_.send_flags = IBV_SEND_FENCE | IBV_SEND_INLINE;
993 :
994 0 : return HCCL_SUCCESS;
995 : }
996 :
997 0 : HcclResult TransportHeterogRoce::CreateCqAndQp()
998 : {
999 0 : HCCL_INFO("TransportHeterogRoce CreateCqAndQp");
1000 0 : CHK_RET(CreateQpWithCq(nicRdmaHandle_, -1, -1, nullptr, nullptr, tagQpInfo_, isHdcMode_, isESMode_));
1001 0 : CHK_RET(CreateQpWithCq(nicRdmaHandle_, -1, -1, nullptr, nullptr, dataQpInfo_, isHdcMode_, isESMode_));
1002 0 : return HCCL_SUCCESS;
1003 : }
1004 :
1005 0 : HcclResult TransportHeterogRoce::DestroyCqAndQp()
1006 : {
1007 0 : HCCL_INFO("TransportHeterogRoce DestroyCqAndQp");
1008 0 : CHK_RET(DestroyQpWithCq(tagQpInfo_, isHdcMode_));
1009 0 : tagQpInfo_ = QpInfo();
1010 0 : CHK_RET(DestroyQpWithCq(dataQpInfo_, isHdcMode_));
1011 0 : dataQpInfo_ = QpInfo();
1012 0 : return HCCL_SUCCESS;
1013 : }
1014 :
1015 0 : HcclResult TransportHeterogRoce::QpConnect(bool& completed)
1016 : {
1017 0 : CHK_RET(HrtRaQpNonBlockConnectAsync(tagQpInfo_.qpHandle, initSM_.locInitInfo.socketInfo[0].fdHandle));
1018 0 : CHK_RET(HrtRaQpNonBlockConnectAsync(dataQpInfo_.qpHandle, initSM_.locInitInfo.socketInfo[1].fdHandle));
1019 :
1020 0 : completed = true;
1021 0 : return HCCL_SUCCESS;
1022 : }
1023 :
1024 0 : HcclResult TransportHeterogRoce::RegMr(void* mem, u64 size, u32& lkey, bool isTagQpHandle)
1025 : {
1026 0 : HCCL_DEBUG("reg mr mem[%p] size[%llu Byte]", mem, size);
1027 0 : if (size == 0) {
1028 0 : lkey = 0;
1029 0 : return HCCL_SUCCESS;
1030 : }
1031 0 : CHK_PTR_NULL(mem);
1032 :
1033 0 : if (isTagQpHandle || IsRamdHandleLevelMr()) {
1034 0 : CHK_RET(mrManager_->GetKey(mem, size, lkey));
1035 : } else {
1036 0 : CHK_RET(dataQpMrManager_->GetKey(mem, size, lkey));
1037 : }
1038 0 : return HCCL_SUCCESS;
1039 : }
1040 :
1041 0 : HcclResult TransportHeterogRoce::DeregMr(void* mem, u64 size, bool isTagQpHandle)
1042 : {
1043 0 : HCCL_DEBUG("dereg mr mem[%p] size[%llu Byte]", mem, size);
1044 0 : if (size == 0) {
1045 0 : return HCCL_SUCCESS;
1046 : }
1047 :
1048 0 : if (isTagQpHandle || IsRamdHandleLevelMr()) {
1049 0 : CHK_RET(mrManager_->ReleaseKey(mem, size));
1050 : } else {
1051 0 : CHK_RET(dataQpMrManager_->ReleaseKey(mem, size));
1052 : }
1053 0 : return HCCL_SUCCESS;
1054 : }
1055 :
1056 0 : HcclResult TransportHeterogRoce::RoceConnectSocket(SocketConnectInfoT conn[], u32 num, bool& completed)
1057 : {
1058 0 : if (initSM_.locInitInfo.role == CLIENT_ROLE_SOCKET) {
1059 0 : return ConnectSocket(conn, num, completed);
1060 : } else {
1061 0 : completed = true;
1062 0 : return HCCL_SUCCESS;
1063 : }
1064 : }
1065 :
1066 0 : HcclResult TransportHeterogRoce::FlushSendQueue(bool& completed)
1067 : {
1068 0 : if (envelopeBacklogQueue_.size() > 0) {
1069 0 : HcclEnvelope tmpEnvelopeInfo;
1070 0 : while (!envelopeBacklogQueue_.empty()) {
1071 0 : tmpEnvelopeInfo = envelopeBacklogQueue_.front();
1072 0 : CHK_RET(SendEnvelope(tmpEnvelopeInfo));
1073 0 : envelopeBacklogQueue_.pop();
1074 : }
1075 : }
1076 0 : completed = true;
1077 0 : return HCCL_SUCCESS;
1078 : }
1079 :
1080 0 : HcclResult TransportHeterogRoce::EnterStateProcess(ConnState nextState)
1081 : {
1082 0 : switch (nextState) {
1083 0 : case ConnState::CONN_STATE_CONNECT_CHECK_SOCKET:
1084 0 : initSM_.socketNum = 1;
1085 0 : break;
1086 0 : case ConnState::CONN_STATE_GET_CHECK_SOCKET:
1087 0 : initSM_.socketNum = 1;
1088 0 : initSM_.completeNum = 0;
1089 0 : break;
1090 0 : case ConnState::CONN_STATE_SEND_CF:
1091 : case ConnState::CONN_STATE_RECV_CF:
1092 0 : initSM_.size = HETEROG_MAX_FRAME_LEN;
1093 0 : initSM_.completeSize = 0;
1094 0 : break;
1095 0 : case ConnState::CONN_STATE_CHECK_CF:
1096 0 : CHK_RET(CheckConsistentFrame());
1097 0 : CHK_RET(TryTransition(HCCL_SUCCESS, true, ConnState::CONN_STATE_CONNECT_ALL_SOCKET));
1098 0 : break;
1099 0 : case ConnState::CONN_STATE_CONNECT_ALL_SOCKET:
1100 0 : initSM_.socketNum = initSM_.locInitInfo.socketConnInfo.size() - 1;
1101 0 : break;
1102 0 : case ConnState::CONN_STATE_GET_ALL_SOCKET:
1103 0 : initSM_.socketNum = initSM_.locInitInfo.socketInfo.size() - 1;
1104 0 : initSM_.completeNum = 0;
1105 0 : break;
1106 0 : case ConnState::CONN_STATE_SEND_STATUS:
1107 0 : initSM_.size = sizeof(initSM_.locInitInfo.signal);
1108 0 : initSM_.completeSize = 0;
1109 0 : break;
1110 0 : case ConnState::CONN_STATE_RECV_STATUS:
1111 0 : initSM_.size = sizeof(initSM_.remInitInfo.signal);
1112 0 : initSM_.completeSize = 0;
1113 0 : break;
1114 0 : case ConnState::CONN_STATE_COMPLETE:
1115 0 : HCCL_INFO("link[%s]: connect complete", initSM_.locInitInfo.socketInfo[0].tag);
1116 0 : break;
1117 0 : default:
1118 0 : HCCL_INFO("link[%s]: state[%u] no need to do anything", initSM_.locInitInfo.socketInfo[0].tag, nextState);
1119 : }
1120 :
1121 0 : return HCCL_SUCCESS;
1122 : }
1123 : // 需要循环检查的状态
1124 0 : HcclResult TransportHeterogRoce::LoopStateProcess()
1125 : {
1126 0 : HcclResult testRet = HCCL_SUCCESS;
1127 0 : bool completed = false;
1128 0 : switch (GetState()) {
1129 0 : case ConnState::CONN_STATE_CONNECT_CHECK_SOCKET:
1130 0 : testRet = RoceConnectSocket(initSM_.locInitInfo.socketConnInfo.data(), initSM_.socketNum, completed);
1131 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_GET_CHECK_SOCKET));
1132 0 : break;
1133 0 : case ConnState::CONN_STATE_GET_CHECK_SOCKET:
1134 0 : testRet = GetSocket(
1135 0 : initSM_.locInitInfo.role, initSM_.locInitInfo.socketInfo.data(), initSM_.socketNum, initSM_.completeNum,
1136 : completed);
1137 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_SEND_CF));
1138 0 : break;
1139 0 : case ConnState::CONN_STATE_SEND_CF:
1140 0 : testRet = SocketSend(
1141 0 : initSM_.locInitInfo.socketInfo[0].fdHandle, initSM_.locInitInfo.checkFrame, initSM_.size,
1142 0 : initSM_.completeSize, completed);
1143 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_RECV_CF));
1144 0 : break;
1145 0 : case ConnState::CONN_STATE_RECV_CF:
1146 0 : testRet = SocketRecv(
1147 0 : initSM_.locInitInfo.socketInfo[0].fdHandle, initSM_.remInitInfo.checkFrame, initSM_.size,
1148 0 : initSM_.completeSize, completed);
1149 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_CHECK_CF));
1150 0 : break;
1151 0 : case ConnState::CONN_STATE_CONNECT_ALL_SOCKET:
1152 0 : testRet = RoceConnectSocket(
1153 0 : reinterpret_cast<SocketConnectInfoT*>(initSM_.locInitInfo.socketConnInfo.data()) + 1, initSM_.socketNum,
1154 : completed);
1155 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_GET_ALL_SOCKET));
1156 0 : break;
1157 0 : case ConnState::CONN_STATE_GET_ALL_SOCKET:
1158 0 : testRet = GetSocket(
1159 : initSM_.locInitInfo.role,
1160 0 : reinterpret_cast<struct SocketInfoT*>(initSM_.locInitInfo.socketInfo.data()) + 1, initSM_.socketNum,
1161 0 : initSM_.completeNum, completed);
1162 0 : testRet = ((testRet == HCCL_SUCCESS) && completed) ? CreatSignalMesg() : testRet;
1163 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_CONNECT_QP));
1164 0 : break;
1165 0 : case ConnState::CONN_STATE_CONNECT_QP:
1166 0 : testRet = QpConnect(completed);
1167 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_GET_QP));
1168 0 : break;
1169 0 : case ConnState::CONN_STATE_GET_QP:
1170 0 : testRet = GetQpStatus(completed);
1171 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_SEND_STATUS));
1172 0 : break;
1173 0 : case ConnState::CONN_STATE_SEND_STATUS:
1174 0 : testRet = SocketSend(
1175 0 : initSM_.locInitInfo.socketInfo[SOCKET_FOR_SENDRECV_QP].fdHandle, &(initSM_.locInitInfo.signal),
1176 0 : initSM_.size, initSM_.completeSize, completed);
1177 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_RECV_STATUS));
1178 0 : break;
1179 0 : case ConnState::CONN_STATE_RECV_STATUS:
1180 0 : testRet = SocketRecv(
1181 0 : initSM_.locInitInfo.socketInfo[SOCKET_FOR_SENDRECV_QP].fdHandle, &(initSM_.remInitInfo.signal),
1182 0 : initSM_.size, initSM_.completeSize, completed);
1183 0 : fdHandle_ = initSM_.locInitInfo.socketInfo[SOCKET_FOR_SENDRECV_QP].fdHandle;
1184 0 : testRet = ((testRet == HCCL_SUCCESS) && completed && !isRawConn_) ? ExchangeSignalMesg() : testRet;
1185 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_FLUSH_QUEUE));
1186 0 : break;
1187 0 : case ConnState::CONN_STATE_FLUSH_QUEUE: {
1188 : // 为防止Isend中积压信封入队和TestSome中flush积压信封队列并发问题,
1189 : // 该处flush积压信封队列并状态迁移完成后,再解锁。
1190 0 : std::unique_lock<std::mutex> lock(envelopeBacklogQueueLock_);
1191 0 : testRet = FlushSendQueue(completed);
1192 0 : CHK_RET(TryTransition(testRet, completed, ConnState::CONN_STATE_COMPLETE));
1193 0 : break;
1194 0 : }
1195 0 : default:
1196 0 : HCCL_ERROR(
1197 : "Establish communication connection failed[%s]: state[%u]", initSM_.locInitInfo.socketInfo[0].tag,
1198 : GetState());
1199 0 : return HCCL_E_INTERNAL;
1200 : }
1201 0 : return HCCL_SUCCESS;
1202 : }
1203 :
1204 0 : HcclResult TransportHeterogRoce::GetSocketInfos(std::vector<std::vector<HcclSocketInfo>>& socketInfos)
1205 : {
1206 0 : std::vector<HcclSocketInfo> hcclSocketInfo;
1207 0 : for (SocketInfoT raSocketInfo : initSM_.locInitInfo.socketInfo) {
1208 0 : hcclSocketInfo.push_back({raSocketInfo.socketHandle, raSocketInfo.fdHandle});
1209 : }
1210 0 : socketInfos.push_back(hcclSocketInfo);
1211 0 : return HCCL_SUCCESS;
1212 0 : }
1213 :
1214 0 : void TransportHeterogRoce::GetTransportResourceInfo(const TransportResourceInfo& transportResourceInfo)
1215 : {
1216 0 : tagQpInfo_.flag = transportResourceInfo.flag;
1217 0 : tagQpInfo_.qpMode = transportResourceInfo.qpMode;
1218 0 : dataQpInfo_.flag = transportResourceInfo.flag;
1219 0 : dataQpInfo_.qpMode = transportResourceInfo.qpMode;
1220 0 : isHdcMode_ = transportResourceInfo.isHdcMode;
1221 0 : deviceLogicId_ = transportResourceInfo.deviceLogicId;
1222 0 : memBlockNum_ = transportResourceInfo.memBlockNum;
1223 0 : remoteIsHdc_ = transportResourceInfo.remoteIsHdc;
1224 0 : isESMode_ = transportResourceInfo.isESMode;
1225 0 : isGlobalMrmanagerInit_ = transportResourceInfo.isGlobalMrmanagerInit;
1226 0 : hdcHostWqeBatchNum_ = transportResourceInfo.hdcHostWqeBatchNum;
1227 0 : HCCL_INFO(
1228 : "tagQpInfo_.flag[%d] tagQpInfo_.qpMode[%d] dataQpInfo_.flag[%d] dataQpInfo_.qpMode[%d] isHdcMode_[%d] "
1229 : "deviceLogicId_[%d] memBlockNum_[%u] remoteIsHdc_[%d] isESMode_[%d] isGlobalMrmanagerInit_[%d] "
1230 : "hdcHostWqeBatchNum_[%u]",
1231 : tagQpInfo_.flag, tagQpInfo_.qpMode, dataQpInfo_.flag, dataQpInfo_.qpMode, isHdcMode_, deviceLogicId_,
1232 : memBlockNum_, remoteIsHdc_, isESMode_, isGlobalMrmanagerInit_, hdcHostWqeBatchNum_);
1233 0 : }
1234 :
1235 0 : HcclResult TransportHeterogRoce::PollCq(QpInfo& qpInfo, bool isSend, s32& num, struct ibv_wc* wc)
1236 : {
1237 0 : if (!isHdcMode_ || tagQpInfo_.qpMode == NORMAL_QP_MODE) {
1238 0 : if (isSend) {
1239 0 : CHK_RET(hrtIbvPollCq(qpInfo.sendCq, HCCL_POLL_CQ_DEPTH, wc, num));
1240 : } else {
1241 0 : CHK_RET(hrtIbvPollCq(qpInfo.recvCq, HCCL_POLL_CQ_DEPTH, wc, num));
1242 : }
1243 0 : } else {
1244 0 : s32 ret = hrtRaPollCq(qpInfo.qpHandle, isSend, HCCL_POLL_CQ_ONETIME, wc);
1245 0 : if (ret >= 0 && static_cast<u32>(ret) <= HCCL_POLL_CQ_ONETIME) {
1246 0 : num = ret;
1247 : } else {
1248 0 : HCCL_ERROR("call trace: hcclRet -> %d", ret);
1249 0 : return HCCL_E_REMOTE;
1250 : }
1251 : }
1252 0 : return HCCL_SUCCESS;
1253 : }
1254 :
1255 0 : HcclResult TransportHeterogRoce::GetRemoteIsendDoneSignal(std::shared_ptr<LocalIpcNotify>& signal)
1256 : {
1257 0 : signal = remoteIsendDoneSignal_;
1258 0 : CHK_SMART_PTR_NULL(signal);
1259 0 : return HCCL_SUCCESS;
1260 : }
1261 :
1262 0 : HcclResult TransportHeterogRoce::GetRemoteImrecvDoneSignal(std::shared_ptr<LocalIpcNotify>& signal)
1263 : {
1264 0 : signal = remoteImrecvDoneSignal_;
1265 0 : CHK_SMART_PTR_NULL(signal);
1266 0 : return HCCL_SUCCESS;
1267 : }
1268 :
1269 0 : HcclResult TransportHeterogRoce::GetNotifySize()
1270 : {
1271 : DevType devType;
1272 0 : CHK_RET(hrtHalGetDeviceType(index_, devType));
1273 :
1274 0 : if (devType == DevType::DEV_TYPE_910) {
1275 0 : notifySize_ = 8; // 910A 每个notify占8个字节
1276 0 : } else if ((devType == DevType::DEV_TYPE_910B) || (devType == DevType::DEV_TYPE_910_93)) {
1277 0 : notifySize_ = 4; // 910B/910_93 每个notify占4个字节
1278 : } else {
1279 0 : notifySize_ = 8; // 其余芯片类型每个notify占8个字节
1280 : }
1281 0 : HCCL_INFO("devType[%d] notifySize[%d]", devType, notifySize_);
1282 0 : return HCCL_SUCCESS;
1283 : }
1284 :
1285 0 : HcclResult TransportHeterogRoce::CreateRdmaSignal(
1286 : std::shared_ptr<LocalIpcNotify>& localNotify, HcclRdmaSignalInfo& rdmaSignalInfo, MemType notifyType)
1287 : {
1288 0 : EXCEPTION_CATCH((localNotify = std::make_shared<LocalIpcNotify>()), return HCCL_E_PTR);
1289 0 : CHK_SMART_PTR_NULL(localNotify);
1290 0 : s32 pid = 0;
1291 0 : CHK_RET(SalGetBareTgid(&pid)); // 当前进程id
1292 0 : CHK_RET(localNotify->Init(deviceLogicId_, deviceLogicId_));
1293 0 : s64 recvId = 0xFFFFFFFF00000000 | (static_cast<s64>(pid) & 0xFFFFFFFF);
1294 0 : CHK_RET(localNotify->Grant(recvId));
1295 :
1296 0 : u64 notifyOffset = 0;
1297 0 : u64 notifyBaseVa = 0; // notify寄存器虚拟地址
1298 0 : u64 notifyTotalSize = 0;
1299 0 : CHK_RET(HrtRaGetNotifyBaseAddr(nicRdmaHandle_, ¬ifyBaseVa, ¬ifyTotalSize));
1300 0 : CHK_RET(localNotify->GetNotifyOffset(notifyOffset));
1301 0 : u64 notifyVa = notifyBaseVa + notifyOffset;
1302 :
1303 0 : rdmaSignalInfo.mrRegFlag = 0;
1304 0 : rdmaSignalInfo.notifyAddr = reinterpret_cast<void*>(notifyVa);
1305 0 : rdmaSignalInfo.len = notifySize_;
1306 0 : rdmaSignalInfo.type = notifyType;
1307 :
1308 0 : struct MrInfoT mrInfo = {};
1309 0 : mrInfo.addr = rdmaSignalInfo.notifyAddr;
1310 0 : mrInfo.size = rdmaSignalInfo.len;
1311 0 : mrInfo.access = access_;
1312 0 : CHK_RET(HrtRaMrReg(dataQpInfo_.qpHandle, &mrInfo));
1313 0 : rdmaSignalInfo.lkey = mrInfo.lkey;
1314 0 : return HCCL_SUCCESS;
1315 : }
1316 :
1317 0 : HcclResult TransportHeterogRoce::PsRdmaDbSend(uint32_t dbindex, uint64_t dbinfo, rtStream_t stream)
1318 : {
1319 0 : CHK_RET(hrtSetDevice(index_));
1320 0 : s32 ret = hrtRDMADBSend(dbindex, dbinfo, stream);
1321 0 : CHK_PRT_RET(
1322 : ret != RT_ERROR_NONE,
1323 : HCCL_ERROR(
1324 : "[rtRDMADBSend]errNo[0x%016llx] rt rdma send fail, "
1325 : "return[%d]. para: dbindex[%u]dbinfo[%llu].",
1326 : HCCL_ERROR_CODE(HCCL_E_RUNTIME), ret, dbindex, dbinfo),
1327 : HCCL_E_RUNTIME);
1328 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
1329 : // ps侧DbSend对当前线程SetDevice后,改变了原GE通信域初始化时setdevice 0
1330 : // 若后面使用本线程save会导致getctx失败获取不到通信域句柄,所以需要在此处重新set回默认
1331 0 : CHK_RET(hrtSetDevice(0));
1332 : }
1333 0 : return HCCL_SUCCESS;
1334 : }
1335 :
1336 0 : HcclResult TransportHeterogRoce::CreateDevMemForNotify(DeviceMem& devMem, u64 size, u32 value)
1337 : {
1338 0 : HCCL_INFO("Use dev mem for notify value");
1339 0 : void* devMemAddr{nullptr};
1340 0 : CHK_RET(hrtSetDevice(index_));
1341 0 : CHK_RET(HrtDevMalloc(&devMemAddr, size));
1342 :
1343 0 : devMemPtrs_.emplace_back(devMemAddr);
1344 :
1345 0 : devMem = DeviceMem::create(devMemAddr, size);
1346 0 : CHK_RET(hrtMemcpy(devMemAddr, size, &value, sizeof(u32), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
1347 :
1348 0 : return HCCL_SUCCESS;
1349 : }
1350 :
1351 0 : HcclResult TransportHeterogRoce::CreateHostMemForNotify(DeviceMem& devMem, u64 size, u32 value, bool needMap)
1352 : {
1353 0 : HCCL_INFO("PS use host mem for notify value");
1354 0 : u64 memLen = size + SMALL_PAGE_SIZE;
1355 0 : s8* ptr = new (std::nothrow) s8[memLen];
1356 0 : CHK_PTR_NULL(ptr);
1357 0 : hostMemPtr_.emplace_back(ptr);
1358 0 : u64 pageSizeNum = reinterpret_cast<u64>(ptr) / SMALL_PAGE_SIZE;
1359 0 : void* ptrVoid = reinterpret_cast<void*>((pageSizeNum + 1) * SMALL_PAGE_SIZE);
1360 0 : void* devVirAddr = ptrVoid;
1361 0 : if (needMap) {
1362 0 : s32 ret = dataQpMrManager_->MapMem(ptrVoid, size, devVirAddr);
1363 0 : if (ret != 0 || devVirAddr == nullptr) {
1364 0 : HCCL_ERROR("PS malloc device mem fail[%d]", HCCL_E_MEMORY);
1365 0 : return HCCL_E_MEMORY;
1366 : }
1367 :
1368 0 : devMem = DeviceMem::create(devVirAddr, size);
1369 0 : CHK_RET(hrtMemcpy(
1370 : ptrVoid, notifyMem_.size(), &value, sizeof(u32), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_HOST));
1371 0 : return HCCL_SUCCESS;
1372 : }
1373 :
1374 0 : devMem = DeviceMem::create(devVirAddr, size);
1375 0 : s32 ret = memcpy_s(ptrVoid, notifyMem_.size(), &value, sizeof(u32));
1376 0 : if (ret < 0) {
1377 0 : HCCL_ERROR("memcpy_s fail[%d]", ret);
1378 0 : return HCCL_E_MEMORY;
1379 : }
1380 :
1381 0 : return HCCL_SUCCESS;
1382 : }
1383 :
1384 0 : HcclResult TransportHeterogRoce::CreateNotifyValueBuffer()
1385 : {
1386 0 : if (notifyMem_.ptr() == nullptr) {
1387 0 : u32 notifyVaule = 1;
1388 :
1389 0 : if (deviceLogicId_ == HOST_DEVICE_ID && isHdcMode_ && (dataQpInfo_.qpMode != NORMAL_QP_MODE)) {
1390 : // ES多机AI server的PS,申请device内存
1391 0 : CHK_RET(CreateDevMemForNotify(notifyMem_, notifyValueSize_, notifyVaule));
1392 0 : } else {
1393 0 : CHK_RET(CreateHostMemForNotify(notifyMem_, notifyValueSize_, notifyVaule, isHdcMode_));
1394 : }
1395 :
1396 0 : CHK_PRT_RET(!notifyMem_.ptr(), HCCL_ERROR("CreateNotifyValueBuffer malloc failed."), HCCL_E_MEMORY);
1397 : }
1398 :
1399 0 : struct MrInfoT mrInfo = {};
1400 0 : mrInfo.addr = notifyMem_.ptr();
1401 0 : mrInfo.size = notifyValueSize_;
1402 0 : mrInfo.access = access_;
1403 0 : CHK_RET(HrtRaMrReg(dataQpInfo_.qpHandle, &mrInfo));
1404 :
1405 0 : notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].mrRegFlag = REG_VALID;
1406 0 : notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].addr = notifyMem_.ptr();
1407 0 : notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].len = notifyValueSize_;
1408 0 : notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].memType = MemType::NOTIFY_VALUE_MEM;
1409 0 : notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].lkey = mrInfo.lkey;
1410 0 : return HCCL_SUCCESS;
1411 : }
1412 :
1413 0 : HcclResult TransportHeterogRoce::DeleteNotifyValueBuffer()
1414 : {
1415 0 : for (u64 i = 0; i < devMemPtrs_.size(); i++) {
1416 0 : if (devMemPtrs_[i] != nullptr) {
1417 0 : CHK_RET(HrtDevFree(devMemPtrs_[i]));
1418 0 : devMemPtrs_[i] = nullptr;
1419 : }
1420 : }
1421 :
1422 0 : devMemPtrs_.clear();
1423 :
1424 0 : for (u64 i = 0; i < hostMemPtr_.size(); i++) {
1425 0 : if (hostMemPtr_[i] != nullptr) {
1426 0 : delete[] hostMemPtr_[i];
1427 0 : hostMemPtr_[i] = nullptr;
1428 : }
1429 : }
1430 :
1431 0 : hostMemPtr_.clear();
1432 :
1433 0 : return HCCL_SUCCESS;
1434 : }
1435 :
1436 0 : HcclResult TransportHeterogRoce::RecoverNotifyMsg(HcclRdmaSignalInfo* remoteRdmaSignal, u64 signalNum)
1437 : {
1438 0 : if (signalNum <= 0) {
1439 0 : return HCCL_E_NOT_FOUND;
1440 : }
1441 :
1442 0 : for (u64 i = 0; i < signalNum; i++) {
1443 0 : u32 tmpMemType = (remoteRdmaSignal + i)->type;
1444 0 : notifyMemMsg_[tmpMemType].mrRegFlag = (remoteRdmaSignal + i)->mrRegFlag;
1445 0 : notifyMemMsg_[tmpMemType].addr = (remoteRdmaSignal + i)->notifyAddr;
1446 0 : notifyMemMsg_[tmpMemType].len = (remoteRdmaSignal + i)->len;
1447 0 : notifyMemMsg_[tmpMemType].memType = (remoteRdmaSignal + i)->memType;
1448 0 : notifyMemMsg_[tmpMemType].rkey = (remoteRdmaSignal + i)->lkey;
1449 : }
1450 :
1451 0 : return HCCL_SUCCESS;
1452 : }
1453 :
1454 0 : HcclResult TransportHeterogRoce::CreatSignalMesg()
1455 : {
1456 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
1457 : // ps
1458 : // 310soc的ps不需要申请notify value,直接返回
1459 0 : if (!isHdcMode_ && !remoteIsHdc_) {
1460 0 : return HCCL_SUCCESS;
1461 : }
1462 : // Notify start
1463 0 : if (isHdcMode_) {
1464 0 : CHK_RET(hrtSetDevice(index_));
1465 : }
1466 0 : CHK_RET(CreateNotifyValueBuffer());
1467 : } else {
1468 : // worker
1469 : // 310soc的worker不需要在这里申请notify,直接返回
1470 0 : if (!isHdcMode_) {
1471 0 : return HCCL_SUCCESS;
1472 : }
1473 : // Notify start
1474 0 : CHK_RET(GetNotifySize());
1475 0 : CHK_RET(CreateRdmaSignal(remoteIsendDoneSignal_, rdmaSignal_[0], MemType::SEND_NOTIFY_MEM));
1476 0 : CHK_RET(CreateRdmaSignal(remoteImrecvDoneSignal_, rdmaSignal_[1], MemType::RECV_NOTIFY_MEM));
1477 : }
1478 :
1479 0 : return HCCL_SUCCESS;
1480 : }
1481 :
1482 0 : HcclResult TransportHeterogRoce::ExchangeSignalMesg()
1483 : {
1484 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
1485 : // ps
1486 : // Notify start
1487 0 : HcclRdmaSignalInfo remoteRdmaSignal[REMOTE_RDMA_SIGNAL_SIZE];
1488 0 : CHK_RET(
1489 : hrtRaSocketBlockRecv(fdHandle_, remoteRdmaSignal, sizeof(HcclRdmaSignalInfo) * REMOTE_RDMA_SIGNAL_SIZE));
1490 0 : CHK_RET(RecoverNotifyMsg(remoteRdmaSignal, REMOTE_RDMA_SIGNAL_SIZE));
1491 : } else {
1492 : // worker
1493 : // Notify start
1494 0 : CHK_RET(hrtRaSocketBlockSend(fdHandle_, rdmaSignal_, sizeof(HcclRdmaSignalInfo) * REMOTE_RDMA_SIGNAL_SIZE));
1495 : }
1496 0 : return HCCL_SUCCESS;
1497 : }
1498 :
1499 0 : HcclResult TransportHeterogRoce::RecordNotifyWithReq(Stream& stream, RdmaNotifyOp type, HcclRequestInfo*& request)
1500 : {
1501 0 : TransData sendData{};
1502 0 : TransportEndPointParam epParam{};
1503 :
1504 0 : CHK_RET(GenerateSendRequest(sendData, epParam, request));
1505 0 : request->transportRequest.requestType = HcclRequestType::HCCL_REQUEST_RECV;
1506 0 : u64 wrId = reinterpret_cast<uint64_t>(request);
1507 0 : CHK_RET(RecordNotify(stream, type, wrId));
1508 :
1509 0 : s32 notifyFlag = HCCL_TEST_INCOMPLETED;
1510 0 : TIME_PRINT(CHK_RET(this->Wait(*request, notifyFlag)));
1511 :
1512 0 : return HCCL_SUCCESS;
1513 : }
1514 :
1515 0 : HcclResult TransportHeterogRoce::RecordNotify(Stream& stream, RdmaNotifyOp type, u64 wrId)
1516 : {
1517 0 : HCCL_INFO(
1518 : "RecordNotify notifyType[%u], wrId[%llu] isHdcMode_[%d] qpMode[%d]", type, wrId, isHdcMode_,
1519 : dataQpInfo_.qpMode);
1520 0 : MemType opType = MemType::MEM_TYPE_RESERVED;
1521 0 : if (type == RdmaNotifyOp::SEND_NOTIFY) {
1522 0 : opType = MemType::SEND_NOTIFY_MEM;
1523 0 : } else if (type == RdmaNotifyOp::RECV_NOTIFY) {
1524 0 : opType = MemType::RECV_NOTIFY_MEM;
1525 : } else {
1526 0 : HCCL_ERROR("TransportHeterogRoce::TYPE is not supported.");
1527 0 : return HCCL_E_PARA;
1528 : }
1529 :
1530 0 : if (!isHdcMode_ || dataQpInfo_.qpMode == NORMAL_QP_MODE) {
1531 0 : notifyWriteSge_.addr = static_cast<u64>(reinterpret_cast<uintptr_t>(notifyMem_.ptr()));
1532 0 : notifyWriteSge_.length = notifyMemMsg_[static_cast<u32>(opType)].len;
1533 0 : notifyWriteSge_.lkey = notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].lkey;
1534 :
1535 0 : notifyWriteWr_.wr_id = wrId;
1536 : notifyWriteWr_.wr.rdma.remote_addr
1537 0 : = static_cast<u64>(reinterpret_cast<uintptr_t>(notifyMemMsg_[static_cast<u32>(opType)].addr));
1538 0 : notifyWriteWr_.wr.rdma.rkey = notifyMemMsg_[static_cast<u32>(opType)].rkey;
1539 :
1540 0 : struct ibv_send_wr* badWr = nullptr;
1541 0 : HCCL_INFO(
1542 : "notify write: remote addr[%llu] length[%d] wrId[%llu]", notifyWriteWr_.wr.rdma.remote_addr,
1543 : notifyWriteSge_.length, notifyWriteWr_.wr_id);
1544 0 : CHK_RET(hrtIbvPostSend(dataQpInfo_.qp, ¬ifyWriteWr_, &badWr));
1545 0 : } else {
1546 0 : struct SgList list = {};
1547 0 : list.addr = static_cast<u64>(reinterpret_cast<uintptr_t>(notifyMem_.ptr()));
1548 0 : list.len = notifyMemMsg_[static_cast<u32>(opType)].len;
1549 0 : list.lkey = notifyMemMsg_[static_cast<u32>(MemType::NOTIFY_VALUE_MEM)].lkey;
1550 :
1551 0 : struct SendWrV2 wr {};
1552 0 : wr.wrId = wrId;
1553 0 : wr.bufList = &list;
1554 0 : wr.bufNum = 1; /* 此处list只有一个,设置为1 */
1555 0 : wr.dstAddr = static_cast<u64>(reinterpret_cast<uintptr_t>(notifyMemMsg_[static_cast<u32>(opType)].addr));
1556 0 : wr.rkey = notifyMemMsg_[static_cast<u32>(opType)].rkey;
1557 0 : wr.op = static_cast<u32>(RdmaOp::OP_WRITE); /* RDMA_WRITE: 0 */
1558 0 : wr.sendFlag = RA_SEND_SIGNALED | RA_SEND_FENCE;
1559 0 : struct SendWrRsp opRsp = {};
1560 0 : CHK_RET(HrtRaSendWrV2(dataQpInfo_.qpHandle, &wr, &opRsp, GetWorkflowMode()));
1561 0 : CHK_RET(DoorBellSend(dataQpInfo_.qpMode, opRsp));
1562 : }
1563 :
1564 0 : return HCCL_SUCCESS;
1565 : }
1566 :
1567 0 : HcclResult TransportHeterogRoce::DoorBellSend(const s32 qpMode, const SendWrRsp& opRsp, void* stream)
1568 : {
1569 0 : if (qpMode == OPBASE_QP_MODE || qpMode == OPBASE_QP_MODE_EXT || qpMode == OFFLINE_QP_MODE_EXT) {
1570 0 : HCCL_DEBUG("entry PsRdmaDbSend");
1571 0 : u32 dbIndex = static_cast<u32>(opRsp.db.dbIndex);
1572 0 : u64 dbInfo = static_cast<u64>(opRsp.db.dbInfo);
1573 0 : CHK_RET(PsRdmaDbSend(dbIndex, dbInfo, stream));
1574 0 : } else {
1575 0 : HCCL_DEBUG("entry hrtRDMASend");
1576 0 : u32 qpn = opRsp.wqeTmp.sqIndex;
1577 0 : u32 wqe_index = opRsp.wqeTmp.wqeIndex;
1578 0 : CHK_RET(hrtRDMASend(qpn, wqe_index, stream));
1579 : }
1580 :
1581 0 : return HCCL_SUCCESS;
1582 : }
1583 :
1584 0 : HcclResult TransportHeterogRoce::MrManagerInit()
1585 : {
1586 : // mrManager_管理信封内存
1587 0 : if (IsRamdHandleLevelMr()) {
1588 : // 通信域初始化时外部还未传入全局内存,需要在这里面手动去初始化需要的全局内存
1589 0 : CHK_PTR_NULL(mrManager_);
1590 0 : std::map<MrMapKey, MrInfo> unRegMrMap = MrManager::GetInstance().GetUnregMap();
1591 0 : mrManager_->InitUnRegMrMap(unRegMrMap);
1592 : // 使用全局的MrManager时dataQp也使用全局的MrManager
1593 0 : dataQpMrManager_ = mrManager_;
1594 0 : return HCCL_SUCCESS;
1595 0 : }
1596 0 : mrManager_ = new (nothrow) MrManager();
1597 0 : CHK_PTR_NULL(mrManager_);
1598 0 : std::map<MrMapKey, MrInfo> unRegMrMap = MrManager::GetInstance().GetUnregMap();
1599 0 : CHK_PRT(mrManager_->Init(tagQpInfo_.qpHandle, index_, deviceLogicId_ == HOST_DEVICE_ID, unRegMrMap));
1600 :
1601 : // dataQpManager_管理数据收发内存
1602 0 : dataQpMrManager_ = new (nothrow) MrManager();
1603 0 : CHK_PTR_NULL(dataQpMrManager_);
1604 0 : unRegMrMap = MrManager::GetInstance().GetUnregMap();
1605 0 : CHK_PRT(dataQpMrManager_->Init(dataQpInfo_.qpHandle, index_, deviceLogicId_ == HOST_DEVICE_ID, unRegMrMap));
1606 :
1607 0 : return HCCL_SUCCESS;
1608 0 : }
1609 :
1610 0 : HcclResult TransportHeterogRoce::MrManagerDeInit()
1611 : {
1612 0 : if (IsRamdHandleLevelMr()) {
1613 : // 若mrManager是全局的,那么就在通信类外部释放
1614 0 : dataQpMrManager_ = nullptr;
1615 0 : return HCCL_SUCCESS;
1616 : }
1617 0 : HCCL_INFO("entry MrManagerDeInit");
1618 0 : CHK_PTR_NULL(mrManager_);
1619 0 : CHK_PRT(mrManager_->DeInit(tagQpInfo_.qpHandle));
1620 0 : delete mrManager_;
1621 :
1622 0 : CHK_PTR_NULL(dataQpMrManager_);
1623 0 : CHK_PRT(dataQpMrManager_->DeInit(dataQpInfo_.qpHandle));
1624 0 : delete dataQpMrManager_;
1625 :
1626 0 : return HCCL_SUCCESS;
1627 : }
1628 :
1629 0 : HcclResult TransportHeterogRoce::PreHdcResource()
1630 : {
1631 : // hdc模式下在通信类内部注册内存
1632 0 : CHK_PRT(MrManagerInit());
1633 : // worker侧不做信封内存注册
1634 0 : if (deviceLogicId_ == HOST_DEVICE_ID) {
1635 0 : if (!IsRamdHandleLevelMr()) {
1636 0 : CHK_PRT(MemBlocksManagerInit());
1637 0 : CHK_RET(mrManager_->GetKey(
1638 : tagMemBlocksManager_->GetMemAddr(), tagMemBlocksManager_->GetMemSize(), blockMemLkey_));
1639 : }
1640 :
1641 : const std::unique_ptr<HeterogMemBlocksManager>& memBlocksManagerPtr
1642 0 : = (IsRamdHandleLevelMr()) ? memBlocksManager_ : tagMemBlocksManager_;
1643 : // mrmanager是全局的时使用的信封内存管理类也是外部传入的全局的管理类
1644 0 : hostAddrBegin_ = (u64)memBlocksManagerPtr->GetMemAddr();
1645 :
1646 0 : devAddrBegin_ = MrManager::g_devAddr;
1647 0 : recvWqeBatchNum_ = hdcHostWqeBatchNum_;
1648 0 : recvWqeBatchThreshold_ = hdcHostWqeBatchNum_;
1649 0 : recvWqeBatchSupplement_ = RECV_WQE_HDC_BATCH_SUPPLEMENT;
1650 0 : HCCL_INFO(
1651 : "PreHdcResource IsRamdHandleLevelMr[%d] recvWqeBatchNum_[%u] recvWqeBatchThreshold_[%u]",
1652 : IsRamdHandleLevelMr(), recvWqeBatchNum_, recvWqeBatchThreshold_);
1653 0 : if (useDevMem_) {
1654 : #ifndef CCL_KERNEL
1655 : CHK_RET(hrtSetDevice(index_));
1656 : // device内存申请跟host内存一样大的内存
1657 : u64 memSize = memBlocksManagerPtr->GetMemSize();
1658 : s32 ret = HrtDevMalloc(&deviceEvePtr_, memSize);
1659 : if (ret != 0 || deviceEvePtr_ == nullptr) {
1660 : HCCL_ERROR("PS HrtDevMalloc device mem fail ret=[%d]", ret);
1661 : return HCCL_E_MEMORY;
1662 : }
1663 :
1664 : struct MrInfoT mrInfo = {nullptr};
1665 : mrInfo.addr = deviceEvePtr_;
1666 : mrInfo.size = memSize;
1667 : mrInfo.access = RA_ACCESS_LOCAL_WRITE | RA_ACCESS_REMOTE_WRITE | RA_ACCESS_REMOTE_READ;
1668 : CHK_RET(HrtRaMrReg(dataQpInfo_.qpHandle, &mrInfo));
1669 : deviceEveLkey_ = mrInfo.lkey;
1670 : HCCL_INFO("index_[%u] deviceEvePtr_[%p] memSize[%llu]", index_, deviceEvePtr_, memSize);
1671 : #endif
1672 : }
1673 0 : CHK_RET(InitTagRecvWqe());
1674 0 : return HCCL_SUCCESS;
1675 : }
1676 :
1677 0 : return HCCL_SUCCESS;
1678 : }
1679 :
1680 0 : HcclResult TransportHeterogRoce::MemBlocksManagerInit()
1681 : {
1682 : // 初始化信封内存
1683 0 : tagMemBlocksManager_.reset(new (std::nothrow) HeterogMemBlocksManager());
1684 0 : CHK_SMART_PTR_NULL(tagMemBlocksManager_);
1685 0 : CHK_RET(tagMemBlocksManager_->Init(memBlockNum_));
1686 :
1687 0 : return HCCL_SUCCESS;
1688 : }
1689 :
1690 0 : HcclResult TransportHeterogRoce::MemBlocksManagerDeInit()
1691 : {
1692 0 : while (wqeBlockLists_.size() > 0) {
1693 0 : CHK_RET(FreeMemBlock(wqeBlockLists_.front()));
1694 : }
1695 :
1696 0 : if (IsRamdHandleLevelMr()) {
1697 : // MrManager是全局的时候在通信类外部统一释放
1698 0 : return HCCL_SUCCESS;
1699 : }
1700 0 : CHK_SMART_PTR_NULL(tagMemBlocksManager_);
1701 0 : CHK_PTR_NULL(mrManager_);
1702 0 : CHK_RET(mrManager_->ReleaseKey(tagMemBlocksManager_->GetMemAddr(), tagMemBlocksManager_->GetMemSize()));
1703 0 : tagMemBlocksManager_ = nullptr;
1704 :
1705 0 : return HCCL_SUCCESS;
1706 : }
1707 :
1708 0 : void TransportHeterogRoce::GetLinkTag(std::string& tag)
1709 : {
1710 0 : tag = initSM_.locInitInfo.socketInfo[0].tag;
1711 0 : return;
1712 : }
1713 :
1714 0 : bool TransportHeterogRoce::IsRamdHandleLevelMr()
1715 : {
1716 : // 非hdc模式或者AI-Server910B场景ps下不分平面时以RdmaHandle粒度注册MR
1717 0 : return (!isHdcMode_ || (isHdcMode_ && isGlobalMrmanagerInit_));
1718 : }
1719 :
1720 : } // namespace hccl
|