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_device_roce_mem.h"
12 : #include "log.h"
13 : #include "dispatcher_pub.h"
14 : #include "adapter_verbs.h"
15 :
16 : namespace hccl {
17 : std::atomic<u64> TransportDeviceRoceMem::wrIdOffset_;
18 :
19 0 : TransportDeviceRoceMem::TransportDeviceRoceMem(const std::unique_ptr<NotifyPool> ¬ifyPool,
20 : const HcclNetDevCtx &netDevCtx, const HcclDispatcher &dispatcher, AttrInfo &attrInfo, bool aicpuUnfoldMode,
21 0 : const HcclQpInfoV2 &qpInfo)
22 : : TransportMem(notifyPool, netDevCtx, dispatcher, attrInfo, aicpuUnfoldMode),
23 0 : timeout_{std::chrono::microseconds((attrInfo.timeout == INVALID_UINT) ? 0 : attrInfo.timeout)}, qpInfo_{qpInfo}
24 : {
25 0 : }
26 :
27 0 : TransportDeviceRoceMem::~TransportDeviceRoceMem()
28 : {
29 0 : }
30 :
31 0 : HcclResult TransportDeviceRoceMem::ExchangeMemDesc(const RmaMemDescs &localMemDescs, RmaMemDescs &remoteMemDescs,
32 : u32 &actualNumOfRemote)
33 : {
34 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support ExchangeMemDesc");
35 0 : return HCCL_E_NOT_SUPPORT;
36 : }
37 :
38 0 : HcclResult TransportDeviceRoceMem::EnableMemAccess(const RmaMemDesc &remoteMemDesc, RmaMem &remoteMem)
39 : {
40 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support EnableMemAccess");
41 0 : return HCCL_E_NOT_SUPPORT;
42 : }
43 :
44 0 : HcclResult TransportDeviceRoceMem::DisableMemAccess(const RmaMemDesc &remoteMemDesc)
45 : {
46 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support DisableMemAccess");
47 0 : return HCCL_E_NOT_SUPPORT;
48 : }
49 :
50 0 : HcclResult TransportDeviceRoceMem::SetSocket(const std::shared_ptr<HcclSocket> &socket)
51 : {
52 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support SetSocket");
53 0 : return HCCL_E_NOT_SUPPORT;
54 : }
55 :
56 0 : HcclResult TransportDeviceRoceMem::Connect(s32 timeoutSec)
57 : {
58 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support Connect");
59 0 : return HCCL_E_NOT_SUPPORT;
60 : }
61 :
62 0 : HcclResult TransportDeviceRoceMem::Write(const HcclBuf &remoteMem, const HcclBuf &localMem, const rtStream_t &stream)
63 : {
64 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support HcclBuf Write");
65 0 : return HCCL_E_NOT_SUPPORT;
66 : }
67 :
68 0 : HcclResult TransportDeviceRoceMem::Read(const HcclBuf &localMem, const HcclBuf &remoteMem, const rtStream_t &stream)
69 : {
70 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support HcclBuf Read");
71 0 : return HCCL_E_NOT_SUPPORT;
72 : }
73 :
74 0 : HcclResult TransportDeviceRoceMem::Write(const RmaOpMem &remoteMem, const RmaOpMem &localMem, const rtStream_t &stream)
75 : {
76 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support RmaOpMem Write");
77 0 : return HCCL_E_NOT_SUPPORT;
78 : }
79 :
80 0 : HcclResult TransportDeviceRoceMem::Read(const RmaOpMem &localMem, const RmaOpMem &remoteMem, const rtStream_t &stream)
81 : {
82 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support RmaOpMem Read");
83 0 : return HCCL_E_NOT_SUPPORT;
84 : }
85 :
86 0 : HcclResult TransportDeviceRoceMem::AddOpFence(const rtStream_t &stream)
87 : {
88 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support HOST AddOpFence");
89 0 : return HCCL_E_NOT_SUPPORT;
90 : }
91 :
92 0 : HcclResult TransportDeviceRoceMem::GetTransInfo(HcclQpInfoV2 &qpInfo, u32 *lkey, u32 *rkey, HcclBuf *localMem,
93 : HcclBuf *remoteMem, u32 num)
94 : {
95 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support GetTransInfo");
96 0 : return HCCL_E_NOT_SUPPORT;
97 : }
98 :
99 0 : HcclResult TransportDeviceRoceMem::WaitOpFence(const rtStream_t &stream)
100 : {
101 0 : HCCL_ERROR("TransportDeviceRoceMem doesn't support WaitOpFence");
102 0 : return HCCL_E_NOT_SUPPORT;
103 : }
104 :
105 0 : HcclResult TransportDeviceRoceMem::TransportDeviceRoceMem::BatchWrite(const std::vector<MemDetails> &remoteMems,
106 : const std::vector<MemDetails> &localMems, Stream &stream)
107 : {
108 0 : return BatchOp(stream, localMems, remoteMems, false, false);
109 : }
110 :
111 0 : HcclResult TransportDeviceRoceMem::BatchRead(const std::vector<MemDetails> &localMems,
112 : const std::vector<MemDetails> &remoteMems, Stream &stream)
113 : {
114 0 : return BatchOp(stream, localMems, remoteMems, true, false);
115 : }
116 :
117 0 : HcclResult TransportDeviceRoceMem::AddOpFence(const MemDetails &localFenceMem, const MemDetails &remoteFenceMem,
118 : Stream &stream)
119 : {
120 0 : std::vector<MemDetails> localMems(1, localFenceMem);
121 0 : std::vector<MemDetails> remoteMems(1, remoteFenceMem);
122 0 : return BatchOp(stream, localMems, remoteMems, true, true);
123 0 : }
124 :
125 0 : HcclResult TransportDeviceRoceMem::DoorBellSend(Stream &stream, u64 dbInfo, u32 wrDataLen, bool fence)
126 : {
127 0 : HCCL_DEBUG("[DoorBellSend] dbIndex[%#x] dbInfo[%#llx] remoteRankId[%u]", qpInfo_.dbIndex, dbInfo, remoteRankId_);
128 0 : RdmaTaskInfo rdmaInfo;
129 0 : WrInformation wrInfo;
130 0 : wrInfo.notifyId = 0;
131 0 : wrInfo.wrData.memList.len = wrDataLen;
132 0 : rdmaInfo.wrInfos.emplace_back(wrInfo);
133 0 : rdmaInfo.rdmaType = fence ? RdmaType::RDMA_SEND_NOTIFY : RdmaType::RDMA_SEND_PAYLOAD;
134 0 : rdmaInfo.remoteRank = remoteRankId_;
135 0 : HcclResult ret = HCCL_SUCCESS;
136 0 : DispatcherPub *dispatcher = static_cast<DispatcherPub *>(dispatcher_);
137 0 : ret = dispatcher->RdmaSend(qpInfo_.dbIndex, dbInfo, stream, rdmaInfo);
138 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
139 : HCCL_ERROR("[DoorBellSend] RdmaSend failed, ret[%u]. dbIndex[%#x] dbInfo[%#llx] remoteRankId[%u]", ret,
140 : qpInfo_.dbIndex, dbInfo, rdmaInfo.remoteRank),
141 : ret);
142 0 : std::vector<Stream> subStreams;
143 0 : ret = dispatcher->LaunchTasksEx(stream, subStreams);
144 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
145 : HCCL_ERROR("[DoorBellSend] LaunchTask failed, ret[%u]. dbIndex[%#x] dbInfo[%#llx] remoteRankId[%u]", ret,
146 : qpInfo_.dbIndex, dbInfo, rdmaInfo.remoteRank),
147 : ret);
148 0 : return HCCL_SUCCESS;
149 0 : }
150 :
151 0 : HcclResult TransportDeviceRoceMem::FillMemDetails(std::vector<MemDetails> &localMemList,
152 : std::vector<MemDetails> &remoteMemList, MemDetails &localMem, MemDetails &remoteMem)
153 : {
154 0 : CHK_PRT_RET(localMem.size != remoteMem.size, HCCL_ERROR("[TransportDeviceRoceMem][FillMemDetails] "
155 : "local buffer size[%llu] is not equal to remote buffer size[%llu]", localMem.size, remoteMem.size),
156 : HCCL_E_PARA);
157 0 : u64 remainingBytes = localMem.size;
158 0 : while (remainingBytes > 0) {
159 0 : const u64 chunkBytes = (remainingBytes > MAX_RDMA_WQE_SIZE) ? MAX_RDMA_WQE_SIZE : remainingBytes;
160 0 : localMem.size = chunkBytes;
161 0 : remoteMem.size = chunkBytes;
162 0 : localMemList.emplace_back(localMem);
163 0 : remoteMemList.emplace_back(remoteMem);
164 0 : localMem.addr += chunkBytes;
165 0 : remoteMem.addr += chunkBytes;
166 0 : remainingBytes -= chunkBytes;
167 : }
168 0 : return HCCL_SUCCESS;
169 : }
170 :
171 0 : HcclResult TransportDeviceRoceMem::BatchOp(Stream &stream, const std::vector<MemDetails> &localMems,
172 : const std::vector<MemDetails> &remoteMems, bool isRead, bool fence)
173 : {
174 0 : constexpr u32 MAX_RDMA_WQE_NUM = 64; // related to qp depth
175 0 : CHK_PRT_RET(localMems.size() != remoteMems.size(), HCCL_ERROR("[TransportDeviceRoceMem][BatchOp] "
176 : "local buffer num[%llu] is not equal to remote buffer num[%llu]", localMems.size(), remoteMems.size()),
177 : HCCL_E_PARA);
178 0 : u64 dbInfo = 0;
179 0 : u32 wqeCount = 0;
180 0 : u64 wrDataLen = 0;
181 0 : const u32 memNum = localMems.size();
182 0 : for (u32 index = 0; index < memNum; index++) {
183 0 : MemDetails localMem = localMems[index];
184 0 : MemDetails remoteMem = remoteMems[index];
185 0 : wrDataLen += localMem.size;
186 0 : std::vector<MemDetails> localMemList;
187 0 : std::vector<MemDetails> remoteMemList;
188 0 : CHK_RET(FillMemDetails(localMemList, remoteMemList, localMem, remoteMem));
189 0 : CHK_RET(BatchPostSend(stream, dbInfo, localMemList, remoteMemList, isRead, fence, wqeCount, wrDataLen));
190 0 : if (wqeCount >= MAX_RDMA_WQE_NUM) {
191 0 : CHK_RET(DoorBellSend(stream, dbInfo, wrDataLen, fence));
192 0 : wqeCount = 0;
193 0 : wrDataLen = 0;
194 : }
195 0 : }
196 0 : if (wqeCount != 0) {
197 0 : CHK_RET(DoorBellSend(stream, dbInfo, wrDataLen, fence));
198 : }
199 0 : return HCCL_SUCCESS;
200 : }
201 :
202 0 : HcclResult TransportDeviceRoceMem::BatchPostSend(Stream &stream, u64 &dbInfo, std::vector<MemDetails> &localMemList,
203 : std::vector<MemDetails> &remoteMemList, bool isRead, bool fence, u32 &wqeCount, u64 &wrDataLen)
204 : {
205 0 : const u32 wrTotalCount = localMemList.size();
206 0 : u32 sendWrCount = 0;
207 0 : while (sendWrCount < wrTotalCount) {
208 0 : const u32 wrCount = std::min(wrTotalCount - sendWrCount, SEND_WR_LEN);
209 0 : CHK_RET(PostSend(stream, dbInfo, &(localMemList[sendWrCount]), &(remoteMemList[sendWrCount]), wrCount,
210 : isRead, fence, wqeCount, wrDataLen));
211 0 : sendWrCount += wrCount;
212 0 : wqeCount += wrCount;
213 : }
214 0 : return HCCL_SUCCESS;
215 : }
216 :
217 0 : HcclResult TransportDeviceRoceMem::PostSend(Stream &stream, u64 &dbInfo, struct MemDetails *localMems,
218 : struct MemDetails *remoteMems, u32 memNum, bool isRead, bool fence, u32 &wqeCount, u64 &wrDataLen)
219 : {
220 0 : constexpr u32 RETRY_DELAY_THRESH = 100;
221 0 : u32 retryCount = 0;
222 0 : auto startTime = std::chrono::steady_clock::now();
223 0 : HcclResult ret = HCCL_E_NETWORK;
224 0 : while (ret != HCCL_SUCCESS) {
225 0 : RdmaOp opCode = isRead ? RdmaOp::OP_READ : RdmaOp::OP_WRITE;
226 0 : ret = RdmaPostSend(dbInfo, localMems, remoteMems, memNum, opCode, fence);
227 0 : if (ret == HCCL_E_AGAIN) {
228 0 : if ((retryCount == 0) && (wqeCount != 0)) {
229 0 : HCCL_WARNING("[PostSend] retry with DoorBellSend, isRead[%u] remoteRankId[%u] wqeCount[%u]", isRead,
230 : remoteRankId_, wqeCount);
231 0 : CHK_RET(DoorBellSend(stream, dbInfo, wrDataLen, fence));
232 0 : wqeCount = 0;
233 0 : wrDataLen = 0;
234 0 : } else {
235 0 : CHK_PRT_RET(timeout_ == std::chrono::microseconds(0),
236 : HCCL_ERROR("[PostSend] failed without retry, isRead[%u] remoteRankId[%u]", isRead, remoteRankId_),
237 : ret);
238 0 : auto elapsedTime = std::chrono::duration_cast<std::chrono::microseconds>(
239 0 : std::chrono::steady_clock::now() - startTime);
240 0 : CHK_PRT_RET(elapsedTime >= timeout_,
241 : HCCL_ERROR("[PostSend] failed after timeout, elapsedTime[%lld us] isRead[%u] remoteRankId[%u]",
242 : elapsedTime.count(), isRead, remoteRankId_),
243 : ret);
244 0 : if (retryCount % RETRY_DELAY_THRESH == 0) {
245 0 : HCCL_WARNING("[PostSend] retryCount[%u] after failed, elapsedTime[%lld us] isRead[%u] "
246 : "remoteRankId[%u]", retryCount, elapsedTime.count(), isRead, remoteRankId_);
247 : }
248 0 : SaluSleep(ONE_MILLISECOND_OF_USLEEP *
249 0 : std::min(CeilDiv(retryCount, RETRY_DELAY_THRESH), RETRY_DELAY_THRESH));
250 : }
251 0 : ++retryCount;
252 0 : continue; // to retry
253 0 : }
254 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
255 : HCCL_ERROR("[PostSend] HnsPostSend failed[%u], isRead[%u] remoteRankId[%u]", ret, isRead, remoteRankId_),
256 : ret);
257 0 : if (retryCount != 0) {
258 0 : HCCL_INFO("[PostSend] retry success, isRead[%u] remoteRankId[%u]", isRead, remoteRankId_);
259 : }
260 : }
261 0 : return HCCL_SUCCESS;
262 : }
263 :
264 0 : HcclResult TransportDeviceRoceMem::RdmaPostSend(u64 &dbInfo, MemDetails *localMems, MemDetails *remoteMems, u32 memNum,
265 : RdmaOp opCode, bool fence)
266 : {
267 0 : CHK_PTR_NULL(localMems);
268 0 : CHK_PTR_NULL(remoteMems);
269 :
270 0 : CHK_PRT_RET(memNum > SEND_WR_LEN,
271 : HCCL_ERROR("[TransportDeviceRoceMem][RdmaPostSend] buffer size is:%u over SEND_WR_LEN: %u", memNum, SEND_WR_LEN),
272 : HCCL_E_PARA);
273 0 : const u32 last = memNum - 1;
274 0 : struct ibv_send_wr sendWr[SEND_WR_LEN] = {0};
275 0 : struct ibv_sge sge[SEND_WR_LEN] = {0};
276 0 : for (u32 index = 0; index < memNum; index++) {
277 : // 设置WR的SGE
278 0 : sge[index].addr = localMems[index].addr;
279 0 : sge[index].length = remoteMems[index].size;
280 0 : sge[index].lkey = localMems[index].key;
281 :
282 : // 设置WR属性
283 0 : sendWr[index].wr_id = wrIdOffset_.fetch_add(1, std::memory_order_relaxed);
284 0 : sendWr[index].num_sge = 1; // 只有一个SGE
285 0 : sendWr[index].sg_list = &sge[index];
286 0 : sendWr[index].wr.rdma.remote_addr = remoteMems[index].addr;
287 0 : sendWr[index].wr.rdma.rkey = remoteMems[index].key;
288 0 : sendWr[index].next = (index == last) ? nullptr : &sendWr[index + 1]; // 第一个WR指向第二个WR
289 0 : sendWr[index].send_flags = (index == last) ?
290 : (fence ? (IBV_SEND_SIGNALED | IBV_SEND_FENCE) : IBV_SEND_SIGNALED) : 0; // 最后一个WR才需要回复CQE
291 0 : sendWr[index].opcode = static_cast<enum ibv_wr_opcode>(opCode);
292 0 : HCCL_DEBUG("[TransportDeviceRoceMem][RdmaPostSend] Direct ibv_post_send[%llu], opcode=[0x%x], "
293 : "remote_addr=[0x%llx], size=[%u], fence[%u]", wrIdOffset_.load(), sendWr[index].opcode,
294 : sendWr[index].wr.rdma.remote_addr, sendWr[index].sg_list->length, fence);
295 : }
296 :
297 0 : struct ibv_send_wr *badWr = nullptr;
298 0 : struct WrExpRsp exp_rsp = {0};
299 0 : struct ibv_qp *qp = reinterpret_cast<struct ibv_qp *>(qpInfo_.qpPtr);
300 0 : CHK_PTR_NULL(qp);
301 0 : HCCL_DEBUG("[TransportDeviceRoceMem][RdmaPostSend] qp=%p, handle=%u, qp_num=%u, qp_type=%d, qp_stat=%d", qp,
302 : qp->handle, qp->qp_num, qp->qp_type, qp->state);
303 0 : HcclResult ret = HrtHnsIbvExpPostSend(qp, &sendWr[0], &badWr, &exp_rsp);
304 0 : CHK_PRT_RET(ret != HCCL_SUCCESS && ret != HCCL_E_AGAIN,
305 : HCCL_ERROR("[TransportDeviceRoceMem][RdmaPostSend] failed, qp=%p, handle=%u, qp_num=%u, qp_type=%d, qp_stat=%d",
306 : qp, qp->handle, qp->qp_num, qp->qp_type, qp->state),
307 : ret);
308 0 : if (ret == HCCL_SUCCESS) {
309 0 : dbInfo = exp_rsp.db_info;
310 : }
311 0 : return ret;
312 0 : }
313 : } // namespace hccl
|