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