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 "ub_mem_transport.h"
12 : #include "serializable.h"
13 : #include "exchange_ub_buffer_dto.h"
14 : #include "exchange_ub_conn_dto.h"
15 : #include "local_ub_rma_buffer.h"
16 : #include "dev_capability.h"
17 : #include "dev_buffer.h"
18 : #include "../../common/dlprof_func_v2.h"
19 : #include "user_remote_mem_getter.h"
20 : #include "exception_util.h"
21 : #include "env_config/env_config_v2.h"
22 :
23 : namespace Hccl {
24 : constexpr u32 FINISH_MSG_SIZE = 128;
25 : constexpr char_t FINISH_MSG[FINISH_MSG_SIZE] = "Ub Comm Pipe ready!";
26 : constexpr u32 ONE_MILLISECOND_OF_USLEEP = 1000;
27 :
28 110 : UbMemTransport::UbMemTransport(
29 : CommonLocRes& commonLocRes, Attribution& attr, const LinkData& linkData, const Socket& socket,
30 110 : RdmaHandle rdmaHandle1, LocCntNotifyRes& locCntNotifyRes1, bool isRecvFirst)
31 : : BaseMemTransport(commonLocRes, attr, linkData, socket, TransportType::UB),
32 110 : rdmaHandle(rdmaHandle1),
33 110 : locCntNotifyRes(locCntNotifyRes1),
34 110 : isRecvFirst_(isRecvFirst)
35 : {
36 330 : HCCL_INFO("source: %s", locCntNotifyRes.Describe().c_str());
37 110 : }
38 :
39 8 : UbMemTransport::UbMemTransport(
40 : CommonLocRes& commonLocRes, Attribution& attr, const LinkData& linkData, const Socket& socket,
41 : RdmaHandle rdmaHandle1, LocCntNotifyRes& locCntNotifyRes1,
42 8 : std::function<void(u32 streamId, u32 taskId, const TaskParam& taskParam)> callback)
43 : : BaseMemTransport(commonLocRes, attr, linkData, socket, TransportType::UB, callback),
44 8 : rdmaHandle(rdmaHandle1),
45 8 : locCntNotifyRes(locCntNotifyRes1)
46 : {
47 24 : HCCL_INFO("source: %s", locCntNotifyRes.Describe().c_str());
48 8 : }
49 :
50 10 : std::string UbMemTransport::Describe() const
51 : {
52 : string msg = StringFormat(
53 20 : "UbMemTransport=[commonLocRes=%s, locCntNotifyRes=%s, ubStatus=%s, ", commonLocRes.Describe().c_str(),
54 30 : locCntNotifyRes.Describe().c_str(), ubStatus.Describe().c_str());
55 10 : msg += StringFormat("exchangeDataSize=%u, ", exchangeDataSize);
56 10 : msg += StringFormat("rmtNotifyNum=%zu, rmtCntNotifyVecNum=%zu]", rmtNotifyVec.size(), rmtCntNotifyVec.size());
57 10 : return msg;
58 0 : }
59 :
60 1 : HcclResult UbMemTransport::BuildDrainResource()
61 : {
62 : // notify作为read的落点
63 1 : if (drainNotify_ == nullptr) {
64 1 : bool devUsed = true;
65 1 : EXCEPTION_CATCH(drainNotify_ = std::make_unique<Hccl::UbLocalNotify>(rdmaHandle, devUsed), return HCCL_E_PTR);
66 3 : HCCL_INFO("[UbMemTransport][%s] drain notify created: %s", __func__, drainNotify_->Describe().c_str());
67 : }
68 :
69 : // 常量1内存供远端读取
70 1 : if (drainBuffer_ == nullptr) {
71 1 : u32 notifySize = Hccl::DevCapability::GetInstance().GetNotifySize();
72 :
73 1 : std::shared_ptr<Hccl::DevBuffer> constMem;
74 1 : EXCEPTION_CATCH(constMem = std::make_shared<Hccl::DevBuffer>(notifySize), return HCCL_E_PTR);
75 :
76 1 : Hccl::HrtMemcpy(
77 2 : reinterpret_cast<void*>(constMem->GetAddr()), constMem->GetSize(), &NORMAL_NOTIFY_VAL,
78 : sizeof(NORMAL_NOTIFY_VAL), RT_MEMCPY_HOST_TO_DEVICE);
79 :
80 1 : EXCEPTION_CATCH(
81 : drainBuffer_ = std::make_unique<Hccl::LocalUbRmaBuffer>(constMem, rdmaHandle), return HCCL_E_PTR);
82 3 : HCCL_INFO(
83 : "[UbMemTransport][%s] drain buffer created: addr[0x%llx], size[%zu]", __func__,
84 : static_cast<unsigned long long>(drainBuffer_->GetAddr()), drainBuffer_->GetSize());
85 1 : }
86 :
87 1 : return HCCL_SUCCESS;
88 : }
89 :
90 0 : HcclResult UbMemTransport::Describe(std::string& dfxMsg)
91 : {
92 0 : HCCL_INFO("UbMemTransport Describe connNum[%u]", connNum);
93 0 : for (u32 i = 0; i < connNum; i++) {
94 0 : CHK_RET(commonLocRes.connVec[i]->Describe(dfxMsg));
95 : }
96 0 : return HCCL_SUCCESS;
97 : }
98 :
99 8 : MemoryBuffer UbMemTransport::GetLocMemBuffer(const RmaBufferSlice& locSlice) const
100 : {
101 8 : return MemoryBuffer(locSlice.addr, locSlice.size, locSlice.buf->GetMemHandle());
102 : }
103 :
104 8 : MemoryBuffer UbMemTransport::GetRmtMemBuffer(const RmtRmaBufferSlice& rmtSlice) const
105 : {
106 8 : return MemoryBuffer(rmtSlice.addr, rmtSlice.size, rmtSlice.buf->GetMemHandle());
107 : }
108 :
109 5 : MemoryBuffer UbMemTransport::GetRmtNotifyMemBuffer(u32 index)
110 : {
111 : return MemoryBuffer(
112 5 : rmtNotifyVec[index]->GetAddr(), rmtNotifyVec[index]->GetSize(), rmtNotifyVec[index]->GetMemHandle());
113 : }
114 :
115 4 : MemoryBuffer UbMemTransport::GetRmtCntNotifyMemBuffer(const WithNotifyIn& withNotify)
116 : {
117 4 : auto index = withNotify.index_;
118 : return MemoryBuffer(
119 4 : rmtCntNotifyVec[index]->GetAddr(), rmtCntNotifyVec[index]->GetSize(), rmtCntNotifyVec[index]->GetMemHandle());
120 : }
121 :
122 5 : static void SubmitTask(const TaskUbDbSend& ubSend, const Stream& stream)
123 : {
124 15 : HCCL_INFO("SubmitTask UbDbSend ");
125 : HrtUbDbInfo info;
126 5 : info.dbNum = 1;
127 5 : info.wrCqe = 0; // 默认值是0 不会cqe 如果传1,驱动分发,会给hccl cqe,用于维护ci指针。
128 5 : info.info[0].functionId = ubSend.GetFuncId();
129 5 : info.info[0].dieId = ubSend.GetDieId();
130 5 : info.info[0].jettyId = ubSend.GetJettyId();
131 5 : info.info[0].piValue = ubSend.GetPiVal();
132 5 : HrtUbDbSend(info, stream.GetPtr());
133 0 : }
134 :
135 1 : static void SubmitTask(const TaskUbDirectSend& ubDirectSend, const Stream& stream)
136 : {
137 3 : HCCL_INFO("SubmitTask UbDirectSend");
138 1 : if (ubDirectSend.GetDwqeSize() != DWQE_SIZE_64 && ubDirectSend.GetDwqeSize() != DWQE_SIZE_128) {
139 : std::string msg
140 0 : = StringFormat("dwqe size is not valid, cannot submit task, dwqeSize=%u", ubDirectSend.GetDwqeSize());
141 0 : THROW<InternalException>(msg);
142 0 : }
143 : HrtUbWqeInfo info;
144 1 : info.wrCqe = 0;
145 1 : info.functionId = ubDirectSend.GetFuncId();
146 1 : info.dieId = ubDirectSend.GetDieId();
147 1 : info.jettyId = ubDirectSend.GetJettyId();
148 1 : info.wqe = const_cast<u8*>(ubDirectSend.GetDwqePtr());
149 1 : info.wqePtrLen = ubDirectSend.GetDwqeSize();
150 1 : info.wqeSize = info.wqePtrLen == DWQE_SIZE_64 ? 0 : 1;
151 1 : HrtUbDirectSend(info, stream.GetPtr());
152 0 : }
153 :
154 5 : static void SubmitTask(const TaskWriteValue& taskWriteValue, const Stream& stream)
155 : {
156 15 : HCCL_INFO("begin HrtWriteValue");
157 5 : HrtWriteValue(taskWriteValue.GetDbAddr(), taskWriteValue.GetPiVal(), stream.GetPtr());
158 0 : HCCL_INFO("finished HrtWriteValue");
159 0 : }
160 :
161 : template <typename TaskType>
162 3 : std::function<void(const BaseTask&, const Stream&)> GetSubmitUbTaskFunction()
163 : {
164 14 : return [](const BaseTask& task, const Stream& stream) {
165 11 : SubmitTask(static_cast<const TaskType&>(task), stream);
166 3 : };
167 : }
168 :
169 : std::map<TaskType, std::function<void(const BaseTask&, const Stream&)>> g_ubTaskSubmitRuleMap
170 : = {{TaskType::UB_SEND, GetSubmitUbTaskFunction<TaskUbDbSend>()},
171 : {TaskType::UB_DIRECT_SEND, GetSubmitUbTaskFunction<TaskUbDirectSend>()},
172 : {TaskType::WRITE_VALUE, GetSubmitUbTaskFunction<TaskWriteValue>()}};
173 :
174 13 : static void SubmitUbTask(unique_ptr<BaseTask> task, const Stream& stream)
175 : {
176 13 : if (task != nullptr) {
177 11 : g_ubTaskSubmitRuleMap.at(task->GetType())(*task.get(), stream);
178 : }
179 2 : }
180 :
181 5 : void UbMemTransport::SubmitNotify(const MemoryBuffer& rmtNotify, u64 data, const Stream& stream)
182 : {
183 5 : SqeConfig config;
184 10 : SubmitUbTask(commonLocRes.connVec[0]->PrepareInlineWrite(rmtNotify, data, config), stream);
185 0 : }
186 :
187 1 : void UbMemTransport::Post(u32 index, const Stream& stream)
188 : {
189 1 : TaskParam taskParam{};
190 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
191 :
192 1 : SubmitNotify(GetRmtNotifyMemBuffer(index), NORMAL_NOTIFY_VAL, stream);
193 :
194 0 : taskParam.taskType = TaskParamType::TASK_NOTIFY_RECORD;
195 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
196 : ;
197 0 : taskParam.taskPara.Notify.notifyID = rmtNotifyVec[index]->GetAddr();
198 0 : taskParam.taskPara.Notify.value = NORMAL_NOTIFY_VAL;
199 :
200 0 : SaveDfxTaskInfo(taskParam);
201 1 : }
202 :
203 1 : void UbMemTransport::Wait(u32 index, const Stream& stream, u32 timeout)
204 : {
205 1 : TaskParam taskParam{};
206 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
207 :
208 1 : commonLocRes.notifyVec[index]->Wait(stream, timeout);
209 :
210 1 : taskParam.taskType = TaskParamType::TASK_NOTIFY_WAIT;
211 1 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
212 1 : taskParam.taskPara.Notify.notifyID = commonLocRes.notifyVec[index]->GetNotify()->GetId();
213 1 : taskParam.taskPara.Notify.value = NORMAL_NOTIFY_VAL;
214 :
215 1 : SaveDfxTaskInfo(taskParam);
216 1 : }
217 :
218 1 : void UbMemTransport::Read(const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const Stream& stream)
219 : {
220 1 : TaskParam taskParam{};
221 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
222 :
223 1 : SqeConfig config;
224 1 : config.wqeMode = WqeMode::DWQE;
225 1 : SubmitUbTask(
226 3 : commonLocRes.connVec[0]->PrepareRead(GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), config), stream);
227 :
228 0 : taskParam.taskType = TaskParamType::TASK_RDMA;
229 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
230 0 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(locSlice.addr);
231 0 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmtSlice.addr);
232 0 : taskParam.taskPara.DMA.size = rmtSlice.size;
233 0 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
234 0 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
235 0 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_READ;
236 0 : SaveDfxTaskInfo(taskParam);
237 1 : }
238 :
239 1 : void UbMemTransport::ReadReduce(
240 : const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const ReduceIn& reduceIn, const Stream& stream)
241 : {
242 1 : TaskParam taskParam{};
243 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
244 :
245 1 : SqeConfig config;
246 1 : config.wqeMode = WqeMode::DWQE;
247 1 : SubmitUbTask(
248 4 : commonLocRes.connVec[0]->PrepareReadReduce(
249 2 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), reduceIn.dataType, reduceIn.reduceOp, config),
250 : stream);
251 :
252 0 : taskParam.taskType = TaskParamType::TASK_UB_REDUCE_INLINE;
253 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
254 0 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(locSlice.addr);
255 0 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmtSlice.addr);
256 0 : taskParam.taskPara.DMA.size = rmtSlice.size;
257 0 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
258 0 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
259 0 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_READ;
260 :
261 0 : SaveDfxTaskInfo(taskParam);
262 1 : }
263 :
264 1 : void UbMemTransport::Write(const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const Stream& stream)
265 : {
266 1 : TaskParam taskParam{};
267 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
268 :
269 1 : SqeConfig config;
270 1 : config.wqeMode = WqeMode::DWQE;
271 1 : SubmitUbTask(
272 3 : commonLocRes.connVec[0]->PrepareWrite(GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), config), stream);
273 0 : taskParam.taskType = TaskParamType::TASK_RDMA;
274 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
275 0 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(locSlice.addr);
276 0 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmtSlice.addr);
277 0 : taskParam.taskPara.DMA.size = locSlice.size;
278 0 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
279 0 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
280 0 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_WRITE;
281 :
282 0 : SaveDfxTaskInfo(taskParam);
283 1 : }
284 :
285 1 : void UbMemTransport::WriteReduce(
286 : const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const ReduceIn& reduceIn, const Stream& stream)
287 : {
288 1 : TaskParam taskParam{};
289 1 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
290 :
291 1 : SqeConfig config;
292 1 : config.wqeMode = WqeMode::DWQE;
293 1 : SubmitUbTask(
294 4 : commonLocRes.connVec[0]->PrepareWriteReduce(
295 2 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), reduceIn.dataType, reduceIn.reduceOp, config),
296 : stream);
297 :
298 0 : taskParam.taskType = TaskParamType::TASK_UB_REDUCE_INLINE;
299 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
300 0 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(locSlice.addr);
301 0 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmtSlice.addr);
302 0 : taskParam.taskPara.DMA.size = locSlice.size;
303 0 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
304 0 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
305 0 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_WRITE;
306 :
307 0 : SaveDfxTaskInfo(taskParam);
308 1 : }
309 :
310 5 : void UbMemTransport::WriteWithNotify(
311 : const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const WithNotifyIn& withNotify,
312 : const Stream& stream)
313 : {
314 5 : if (locSlice.size == 0) {
315 2 : return SubmitWriteEmptyWithNotify(withNotify, stream);
316 : }
317 :
318 3 : if (withNotify.notifyType_ == TransportNotifyType::NORMAL) {
319 1 : return SubmitWriteWithNotify(
320 2 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), NORMAL_NOTIFY_VAL,
321 2 : GetRmtNotifyMemBuffer(withNotify.index_), stream);
322 2 : } else if (withNotify.notifyType_ == TransportNotifyType::COUNT) {
323 1 : return SubmitWriteWithNotify(
324 2 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), withNotify.userData_,
325 2 : GetRmtCntNotifyMemBuffer(withNotify), stream);
326 : } else {
327 1 : std::string msg = StringFormat("%s error", withNotify.Describe().c_str());
328 1 : THROW<InternalException>(msg);
329 1 : }
330 : }
331 :
332 5 : void UbMemTransport::WriteReduceWithNotify(
333 : const RmaBufferSlice& locSlice, const RmtRmaBufferSlice& rmtSlice, const ReduceIn& reduceIn,
334 : const WithNotifyIn& withNotify, const Stream& stream)
335 : {
336 5 : if (locSlice.size == 0) {
337 2 : return SubmitWriteEmptyWithNotify(withNotify, stream);
338 : }
339 :
340 3 : if (withNotify.notifyType_ == TransportNotifyType::NORMAL) {
341 1 : SubmitWriteReduceWithNotify(
342 1 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), reduceIn, NORMAL_NOTIFY_VAL,
343 2 : GetRmtNotifyMemBuffer(withNotify.index_), stream);
344 2 : } else if (withNotify.notifyType_ == TransportNotifyType::COUNT) {
345 1 : SubmitWriteReduceWithNotify(
346 1 : GetRmtMemBuffer(rmtSlice), GetLocMemBuffer(locSlice), reduceIn, withNotify.userData_,
347 2 : GetRmtCntNotifyMemBuffer(withNotify), stream);
348 : } else {
349 1 : std::string msg = StringFormat("%s error", withNotify.Describe().c_str());
350 1 : THROW<InternalException>(msg);
351 1 : }
352 : }
353 :
354 4 : void UbMemTransport::SubmitWriteEmptyWithNotify(const WithNotifyIn& withNotify, const Stream& stream)
355 : {
356 4 : TaskParam taskParam{};
357 4 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
358 4 : u32 value = NORMAL_NOTIFY_VAL;
359 :
360 4 : if (withNotify.notifyType_ == TransportNotifyType::NORMAL) {
361 2 : SubmitNotify(GetRmtNotifyMemBuffer(withNotify.index_), NORMAL_NOTIFY_VAL, stream);
362 2 : } else if (withNotify.notifyType_ == TransportNotifyType::COUNT) {
363 2 : SubmitNotify(GetRmtCntNotifyMemBuffer(withNotify), withNotify.userData_, stream);
364 0 : value = withNotify.userData_;
365 : } else {
366 0 : std::string msg = StringFormat("%s error", withNotify.Describe().c_str());
367 0 : THROW<InternalException>(msg);
368 0 : }
369 :
370 0 : taskParam.taskType = TaskParamType::TASK_NOTIFY_WAIT;
371 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
372 0 : taskParam.taskPara.Notify.notifyID = INVALID_VALUE_NOTIFYID;
373 0 : taskParam.taskPara.Notify.value = value;
374 :
375 0 : SaveDfxTaskInfo(taskParam);
376 4 : }
377 :
378 2 : void UbMemTransport::SubmitWriteWithNotify(
379 : const MemoryBuffer& rmt, const MemoryBuffer& loc, u64 data, const MemoryBuffer& rmtNotify, const Stream& stream)
380 : {
381 2 : TaskParam taskParam{};
382 2 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
383 :
384 2 : SqeConfig config;
385 2 : config.wqeMode = WqeMode::DWQE;
386 4 : SubmitUbTask(commonLocRes.connVec[0]->PrepareWriteWithNotify(rmt, loc, data, rmtNotify, config), stream);
387 :
388 0 : taskParam.taskType = TaskParamType::TASK_WRITE_WITH_NOTIFY;
389 0 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
390 0 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(loc.addr);
391 0 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmt.addr);
392 0 : taskParam.taskPara.DMA.size = loc.size;
393 0 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
394 0 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
395 0 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_WRITE;
396 :
397 0 : SaveDfxTaskInfo(taskParam);
398 2 : }
399 :
400 2 : void UbMemTransport::SubmitWriteReduceWithNotify(
401 : const MemoryBuffer& rmt, const MemoryBuffer& loc, const ReduceIn& reduceIn, u64 data, const MemoryBuffer& rmtNotify,
402 : const Stream& stream)
403 : {
404 2 : TaskParam taskParam{};
405 2 : taskParam.beginTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
406 :
407 2 : SqeConfig config;
408 2 : config.wqeMode = WqeMode::DWQE;
409 2 : SubmitUbTask(
410 4 : commonLocRes.connVec[0]->PrepareWriteReduceWithNotify(
411 : rmt, loc, reduceIn.dataType, reduceIn.reduceOp, data, rmtNotify, config),
412 : stream);
413 :
414 2 : taskParam.taskType = TaskParamType::TASK_WRITE_REDUCE_WITH_NOTIFY;
415 2 : taskParam.endTime = DlProfFunc::GetInstance().dlMsprofSysCycleTime();
416 2 : taskParam.taskPara.DMA.src = reinterpret_cast<const void*>(loc.addr);
417 2 : taskParam.taskPara.DMA.dst = reinterpret_cast<const void*>(rmt.addr);
418 2 : taskParam.taskPara.DMA.size = loc.size;
419 2 : taskParam.taskPara.DMA.notifyID = INVALID_VALUE_NOTIFYID;
420 2 : taskParam.taskPara.DMA.linkType = DfxLinkType::UB;
421 2 : taskParam.taskPara.DMA.dmaOp = DmaOp::HCCL_DMA_WRITE;
422 :
423 2 : SaveDfxTaskInfo(taskParam);
424 2 : }
425 :
426 5 : bool UbMemTransport::IsResReady()
427 : {
428 7 : for (auto& it : commonLocRes.connVec) {
429 3 : CHECK_NULLPTR(it, StringFormat("[UbMemTransport::%s] failed, connection pointer is nullptr", __func__));
430 :
431 3 : RmaConnType connType = it->GetRmaConnType();
432 3 : if (connType != RmaConnType::UB) {
433 0 : THROW<InternalException>(
434 0 : "[UbMemTransport::%s] connection type[%s] is not ub", __func__, connType.Describe().c_str());
435 : }
436 :
437 3 : auto status = it->GetStatus();
438 3 : if (status != RmaConnStatus::EXCHANGEABLE && status != RmaConnStatus::READY) {
439 1 : return false;
440 : }
441 : }
442 :
443 12 : HCCL_INFO("[UbMemTransport::IsResReady] all resources ready.");
444 4 : return true;
445 : }
446 :
447 4 : bool UbMemTransport::IsConnsReady()
448 : {
449 4 : for (u32 i = 0; i < connNum; i++) {
450 1 : if (commonLocRes.connVec[i]->GetStatus() != RmaConnStatus::READY) {
451 1 : return false;
452 : }
453 : }
454 9 : HCCL_INFO("conns are ready.");
455 3 : return true;
456 : }
457 :
458 15 : HcclResult UbMemTransport::StatusMachine()
459 : {
460 18 : TRY_CATCH_RETURN(
461 : if (socket == nullptr) {
462 : HCCL_ERROR("[UbMemTransport][StatusMachine]socket is nullptr, please check");
463 : return HcclResult::HCCL_E_INTERNAL;
464 : } SocketStatus socketStatus
465 : = isHost_ ? socket->GetStatus() : socket->GetAsyncStatus();
466 : if (socketStatus == Hccl::SocketStatus::INIT || socketStatus == Hccl::SocketStatus::TIMEOUT) {
467 : HCCL_ERROR("[UbMemTransport][StatusMachine]socket timeout or no link, please check");
468 : return HcclResult::HCCL_E_INTERNAL;
469 : }
470 :
471 : if (socketStatus != Hccl::SocketStatus::OK) {
472 : SaluSleep(ONE_MILLISECOND_OF_USLEEP); // 防止get sockets冲高CtrlCPU
473 : return HcclResult::HCCL_SUCCESS; // 操作成功,保持当前状态
474 : } switch (ubStatus) {
475 : case UbStatus::INIT:
476 : CHK_RET(HandleInitStatus());
477 : break;
478 : case UbStatus::SEND_DATA:
479 : CHK_RET(HandleSendAllStatus());
480 : break;
481 : case UbStatus::RECV_SIZE:
482 : CHK_RET(HandleRecvSizeStatus());
483 : break;
484 : case UbStatus::RECV_DATA:
485 : CHK_RET(HandleRecvDataStatus());
486 : break;
487 : case UbStatus::PROCESS_DATA:
488 : CHK_RET(HandleProcessDataStatus());
489 : break;
490 : case UbStatus::SEND_FIN:
491 : CHK_RET(HandleSendFinStatus());
492 : break;
493 : case UbStatus::RECV_FIN:
494 : CHK_RET(HandleRecvFinStatus());
495 : break;
496 : case UbStatus::SET_READY:
497 : CHK_RET(HandleSetReadyStatus());
498 : break;
499 : default:
500 : break;
501 : });
502 14 : return HCCL_SUCCESS;
503 : }
504 :
505 4 : HcclResult UbMemTransport::HandleInitStatus()
506 : {
507 4 : ubStatus = isRecvFirst_ ? UbStatus::RECV_SIZE : UbStatus::SEND_DATA;
508 4 : baseStatus = TransportStatus::SOCKET_OK;
509 4 : return HCCL_SUCCESS;
510 : }
511 :
512 4 : HcclResult UbMemTransport::HandleSendAllStatus()
513 : {
514 4 : if (IsResReady()) {
515 3 : CHK_RET(SendAll());
516 3 : ubStatus = isRecvFirst_ ? UbStatus::PROCESS_DATA : UbStatus::RECV_SIZE;
517 : }
518 4 : return HCCL_SUCCESS;
519 : }
520 :
521 5 : HcclResult UbMemTransport::HandleRecvSizeStatus()
522 : {
523 8 : CHK_RET(RecvDataSize());
524 4 : ubStatus = UbStatus::RECV_DATA;
525 4 : return HCCL_SUCCESS;
526 : }
527 :
528 5 : HcclResult UbMemTransport::HandleRecvDataStatus()
529 : {
530 8 : CHK_RET(RecvExchangeData());
531 4 : ubStatus = isRecvFirst_ ? UbStatus::SEND_DATA : UbStatus::PROCESS_DATA;
532 4 : return HCCL_SUCCESS;
533 : }
534 :
535 5 : HcclResult UbMemTransport::HandleProcessDataStatus()
536 : {
537 5 : bool needSendFinish = false;
538 8 : CHK_RET(RecvDataProcess(needSendFinish));
539 4 : if (needSendFinish) {
540 2 : ubStatus = UbStatus::SEND_FIN;
541 : } else {
542 2 : SetBaseStatusReady();
543 2 : ubStatus = UbStatus::READY;
544 : }
545 4 : return HCCL_SUCCESS;
546 : }
547 :
548 4 : HcclResult UbMemTransport::HandleSendFinStatus()
549 : {
550 4 : if (IsConnsReady()) {
551 6 : CHK_RET(SendFinish());
552 2 : ubStatus = UbStatus::RECV_FIN;
553 : }
554 3 : return HCCL_SUCCESS;
555 : }
556 :
557 3 : HcclResult UbMemTransport::HandleRecvFinStatus()
558 : {
559 6 : CHK_RET(RecvFinish());
560 2 : ubStatus = UbStatus::SET_READY;
561 2 : return HCCL_SUCCESS;
562 : }
563 :
564 2 : HcclResult UbMemTransport::HandleSetReadyStatus()
565 : {
566 2 : SetBaseStatusReady();
567 2 : ubStatus = UbStatus::READY;
568 2 : return HCCL_SUCCESS;
569 : }
570 :
571 13 : TransportStatus UbMemTransport::GetStatus()
572 : {
573 26 : if (baseStatus == TransportStatus::READY || baseStatus == TransportStatus::CONNECT_FAILED
574 26 : || baseStatus == TransportStatus::SOCKET_TIMEOUT) {
575 0 : return baseStatus;
576 : }
577 :
578 13 : HcclResult ret = StatusMachine();
579 13 : if (ret != HCCL_SUCCESS) {
580 0 : HCCL_ERROR("[UbMemTransport::GetStatus] StatusMachine failed, ret=%d", ret);
581 0 : baseStatus = TransportStatus::CONNECT_FAILED;
582 : }
583 13 : return baseStatus;
584 : }
585 :
586 1 : HcclResult UbMemTransport::SendAll()
587 : {
588 1 : notifyNum = commonLocRes.notifyVec.size();
589 1 : bufferNum = commonLocRes.bufferVec.size();
590 1 : connNum = commonLocRes.connVec.size();
591 1 : cntNotifyNum = locCntNotifyRes.vec.size();
592 :
593 1 : cntNotifyDescSize = locCntNotifyRes.desc.size();
594 :
595 3 : HCCL_INFO(
596 : "notifyNum=%u, bufferNum=%u, connNum=%u, cntNotifyNum=%u, cntNotifyDescSize=%u, isHost_[%d]", notifyNum,
597 : bufferNum, connNum, cntNotifyNum, cntNotifyDescSize, isHost_);
598 :
599 1 : BinaryStream binaryStream;
600 1 : HandshakeMsgPack(binaryStream);
601 1 : NotifyVecPack(binaryStream);
602 1 : BufferVecPack(binaryStream, commonLocRes.bufferVec);
603 1 : CntNotifyVecPack(binaryStream);
604 1 : CntNotifyDescPack(binaryStream);
605 1 : CHK_RET(DrainBufPack(binaryStream));
606 1 : ConnVecPack(binaryStream);
607 :
608 1 : sendDataPack_.resize(sizeof(u32));
609 1 : binaryStream.Dump(sendDataPack_);
610 1 : u32 dataSize = sendDataPack_.size() - sizeof(u32);
611 1 : CHK_SAFETY_FUNC_RET(memcpy_s(sendDataPack_.data(), sizeof(u32), &dataSize, sizeof(u32)));
612 :
613 1 : bool ret = false;
614 1 : if (isHost_) {
615 0 : ret = socket->Send(sendDataPack_.data(), sendDataPack_.size());
616 : } else {
617 1 : socket->SendAsync(sendDataPack_.data(), sendDataPack_.size());
618 1 : ret = true;
619 : }
620 1 : if (!ret) {
621 0 : HCCL_ERROR("[UbMemTransport::SendAll] Send failed");
622 0 : return HCCL_E_INTERNAL;
623 : }
624 3 : HCCL_INFO("[UbMemTransport::%s] Send size[%zu] of data success.", __func__, sendDataPack_.size());
625 1 : return HCCL_SUCCESS;
626 1 : }
627 :
628 1 : HcclResult UbMemTransport::RecvDataSize()
629 : {
630 : // 接收数据包尺寸
631 1 : bool ret = false;
632 3 : HCCL_DEBUG("Starting to recv message size[%zu] bytes, isHost_[%d]", sizeof(exchangeDataSize), isHost_);
633 1 : if (isHost_) {
634 0 : ret = socket->Recv(&exchangeDataSize, sizeof(exchangeDataSize));
635 : } else {
636 1 : socket->RecvAsync(reinterpret_cast<u8*>(&exchangeDataSize), sizeof(exchangeDataSize));
637 1 : ret = true;
638 : }
639 1 : if (!ret) {
640 0 : HCCL_ERROR("[UbMemTransport::RecvDataSize] Recv size failed");
641 0 : return HCCL_E_INTERNAL;
642 : }
643 3 : HCCL_INFO(
644 : "[UbMemTransport::%s] Receive size[%u] of data success. [%zu] bytes received.", __func__, exchangeDataSize,
645 : sizeof(exchangeDataSize));
646 1 : return HCCL_SUCCESS;
647 : }
648 :
649 1 : HcclResult UbMemTransport::SendExchangeData()
650 : {
651 1 : bool ret = false;
652 3 : HCCL_DEBUG("Starting to send message size[%zu] bytes, isHost_[%d]", sendData.size(), isHost_);
653 1 : if (isHost_) {
654 0 : ret = socket->Send(sendData.data(), sendData.size());
655 : } else {
656 1 : socket->SendAsync(sendData.data(), sendData.size());
657 1 : ret = true;
658 : }
659 1 : if (!ret) {
660 0 : HCCL_ERROR("[UbMemTransport::SendExchangeData] Send data failed");
661 0 : return HCCL_E_INTERNAL;
662 : }
663 3 : HCCL_INFO("send data %s, size=%zu", GetLinkDescInfo().c_str(), sendData.size());
664 1 : return HCCL_SUCCESS;
665 : }
666 :
667 0 : HcclResult UbMemTransport::RecvExchangeData()
668 : {
669 0 : recvData.resize(exchangeDataSize);
670 0 : bool ret = false;
671 0 : HCCL_DEBUG("Starting to recv message size[%zu] bytes, isHost_[%d]", recvData.size(), isHost_);
672 0 : if (isHost_) {
673 0 : ret = socket->Recv(recvData.data(), recvData.size());
674 : } else {
675 0 : socket->RecvAsync(reinterpret_cast<u8*>(recvData.data()), recvData.size());
676 0 : ret = true;
677 : }
678 0 : if (!ret) {
679 0 : HCCL_ERROR("[UbMemTransport::RecvExchangeData] Recv data failed");
680 0 : return HCCL_E_INTERNAL;
681 : }
682 :
683 0 : HCCL_INFO("recv data %s, size=%zu", GetLinkDescInfo().c_str(), recvData.size());
684 0 : return HCCL_SUCCESS;
685 : }
686 :
687 0 : HcclResult UbMemTransport::RecvDataProcess(bool& needSendFinish)
688 : {
689 0 : HCCL_INFO(
690 : "RecvDataProcess: link=%s, size=%zu, exchangeDataSize=%u", GetLinkDescInfo().c_str(), recvData.size(),
691 : exchangeDataSize);
692 0 : BinaryStream binaryStream(recvData);
693 0 : HcclResult ret = HandshakeMsgUnpack(binaryStream);
694 0 : if (ret != HCCL_SUCCESS) {
695 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] HandshakeMsgUnpack failed, ret=%d", ret);
696 0 : return ret;
697 : }
698 :
699 0 : ret = RmtBufferVecUnpackProc(notifyNum, binaryStream, rmtNotifyVec, UbRmtBufType::NOTIFY);
700 0 : if (ret != HCCL_SUCCESS) {
701 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] RmtBufferVecUnpackProc notify failed, ret=%d", ret);
702 0 : return ret;
703 : }
704 :
705 0 : ret = RmtBufferVecUnpackProc(bufferNum, binaryStream, rmtBufferVec, UbRmtBufType::BUFFER);
706 0 : if (ret != HCCL_SUCCESS) {
707 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] RmtBufferVecUnpackProc buffer failed, ret=%d", ret);
708 0 : return ret;
709 : }
710 :
711 0 : ret = RmtBufferVecUnpackProc(cntNotifyNum, binaryStream, rmtCntNotifyVec, UbRmtBufType::CNT_NOTIFY);
712 0 : if (ret != HCCL_SUCCESS) {
713 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] RmtBufferVecUnpackProc cntNotify failed, ret=%d", ret);
714 0 : return ret;
715 : }
716 :
717 0 : ret = CntNotifyDescUnpack(binaryStream);
718 0 : if (ret != HCCL_SUCCESS) {
719 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] CntNotifyDescUnpack failed, ret=%d", ret);
720 0 : return ret;
721 : }
722 :
723 0 : ret = DrainBufUnpack(binaryStream);
724 0 : if (ret != HCCL_SUCCESS) {
725 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] DrainBufUnpack failed, ret=%d", ret);
726 0 : return ret;
727 : }
728 :
729 0 : ret = ConnVecUnpackProc(binaryStream, needSendFinish);
730 0 : if (ret != HCCL_SUCCESS) {
731 0 : HCCL_ERROR("[UbMemTransport::RecvDataProcess] ConnVecUnpackProc failed, ret=%d", ret);
732 0 : return ret;
733 : }
734 :
735 0 : return HCCL_SUCCESS;
736 0 : }
737 :
738 6 : void UbMemTransport::BufferVecPack(BinaryStream& binaryStream, std::vector<LocalRmaBuffer*>& bufferVec)
739 : {
740 6 : binaryStream << static_cast<u32>(bufferVec.size());
741 18 : HCCL_INFO("start pack %s bufferVec", transportType.Describe().c_str());
742 6 : u32 pos = 0;
743 13 : for (auto& it : bufferVec) {
744 7 : binaryStream << pos;
745 7 : if (it != nullptr) { // 非空的buffer,从buffer中获取 dto
746 7 : std::unique_ptr<Serializable> dto = it->GetExchangeDto();
747 7 : dto->Serialize(binaryStream);
748 21 : HCCL_INFO("pack buffer pos=%u dto %s", pos, dto->Describe().c_str());
749 7 : } else { // 空的buffer,dto所有字段为0(size=0)
750 0 : ExchangeUbBufferDto exchangeDto;
751 0 : exchangeDto.Serialize(binaryStream);
752 0 : HCCL_INFO("pack buffer pos=%u, dto is null %s", pos, exchangeDto.Describe().c_str());
753 0 : }
754 7 : pos++;
755 : }
756 6 : }
757 :
758 1 : void UbMemTransport::CntNotifyVecPack(BinaryStream& binaryStream)
759 : {
760 1 : binaryStream << cntNotifyNum;
761 3 : HCCL_INFO("pack UB cntNotify num=%u, %s", cntNotifyNum, GetLinkDescInfo().c_str());
762 1 : u32 pos = 0;
763 2 : for (auto& it : locCntNotifyRes.vec) {
764 1 : binaryStream << pos;
765 1 : std::unique_ptr<Serializable> dto = it->GetExchangeDto();
766 1 : dto->Serialize(binaryStream);
767 3 : HCCL_INFO("pack cntNotify pos=%u, dto %s", pos, dto->Describe().c_str());
768 1 : pos++;
769 1 : }
770 1 : }
771 :
772 1 : void UbMemTransport::CntNotifyDescPack(BinaryStream& binaryStream)
773 : {
774 1 : binaryStream << cntNotifyDescSize;
775 3 : HCCL_INFO("pack cntNotify desc size=%u %s", cntNotifyDescSize, GetLinkDescInfo().c_str());
776 3 : HCCL_INFO("pack cntNotify desc =%s", Bytes2hex(locCntNotifyRes.desc.data(), locCntNotifyRes.desc.size()).c_str());
777 3 : for (auto& it : locCntNotifyRes.desc) {
778 2 : binaryStream << it;
779 : }
780 1 : }
781 :
782 1 : HcclResult UbMemTransport::DrainBufPack(BinaryStream& binaryStream)
783 : {
784 : // 打包交换信息前,进行资源创建
785 1 : CHK_RET(BuildDrainResource());
786 :
787 : // 只需交换 常量buffer信息 供远端读
788 3 : HCCL_INFO("start pack drain buffer");
789 1 : if (drainBuffer_ != nullptr) { // 非空的buffer,从buffer中获取 dto
790 1 : std::unique_ptr<Serializable> dto = drainBuffer_->GetExchangeDto();
791 1 : dto->Serialize(binaryStream);
792 3 : HCCL_INFO("pack drain buffer dto %s", dto->Describe().c_str());
793 1 : } else { // 空的buffer,dto所有字段为0(size=0)
794 0 : ExchangeUbBufferDto exchangeDto;
795 0 : exchangeDto.Serialize(binaryStream);
796 0 : HCCL_INFO("pack drain buffer, dto is null %s", exchangeDto.Describe().c_str());
797 0 : }
798 :
799 1 : return HCCL_SUCCESS;
800 : }
801 :
802 0 : HcclResult UbMemTransport::DrainBufUnpack(BinaryStream& binaryStream)
803 : {
804 0 : HCCL_INFO("start unpack drain buffer");
805 0 : ExchangeUbBufferDto dto;
806 0 : dto.Deserialize(binaryStream);
807 :
808 0 : if (dto.size == 0) {
809 0 : rmtDrainBuffer_ = nullptr;
810 0 : HCCL_WARNING("unpack drain buffer dto is null");
811 : } else {
812 0 : rmtDrainBuffer_ = std::make_unique<RemoteUbRmaBuffer>(rdmaHandle, dto);
813 0 : HCCL_INFO("unpack drain buffer rmtDrainBuffer=%s", rmtDrainBuffer_->Describe().c_str());
814 : }
815 :
816 0 : return HCCL_SUCCESS;
817 0 : }
818 :
819 0 : HcclResult UbMemTransport::CntNotifyDescUnpack(BinaryStream& binaryStream)
820 : {
821 : u32 descSize;
822 0 : binaryStream >> descSize;
823 0 : if (descSize != cntNotifyDescSize) {
824 0 : HCCL_ERROR(
825 : "[UbMemTransport::CntNotifyDescUnpack] size=%u is not equal to rmtNum=%u", descSize, cntNotifyDescSize);
826 0 : return HCCL_E_PARA;
827 : }
828 0 : rmtCntNotifyDesc.clear();
829 0 : u32 pos = 0;
830 0 : for (pos = 0; pos < descSize; pos++) {
831 : char c;
832 0 : binaryStream >> c;
833 0 : rmtCntNotifyDesc.push_back(c);
834 : }
835 0 : HCCL_INFO("unpack cntNotify Desc=%s", Bytes2hex(rmtCntNotifyDesc.data(), rmtCntNotifyDesc.size()).c_str());
836 0 : return HCCL_SUCCESS;
837 : }
838 :
839 3 : HcclResult UbMemTransport::RmtBufferVecUnpackProc(
840 : u32 locNum, BinaryStream& binaryStream, RemoteBufferVec& bufferVec, UbRmtBufType type)
841 : {
842 : u32 rmtNum;
843 3 : binaryStream >> rmtNum;
844 3 : if (UNLIKELY(type == UbRmtBufType::BUFFER && rmtNum > MAX_BUFFER_NUM)) {
845 0 : HCCL_ERROR("[UbMemTransport][RmtBufferVecUnpackProc] rmtNum[%u] exceeds limit[%u]", rmtNum, MAX_BUFFER_NUM);
846 0 : return HCCL_E_PARA;
847 : }
848 :
849 : // 允许本端和远端交换内存数量不一致
850 9 : HCCL_INFO("unpack %s %s, locNum=%u, rmtNum=%u", type.Describe().c_str(), GetLinkDescInfo().c_str(), locNum, rmtNum);
851 :
852 7 : for (u32 i = 0; i < rmtNum; i++) {
853 : u32 pos;
854 4 : binaryStream >> pos;
855 4 : ExchangeUbBufferDto dto;
856 4 : dto.Deserialize(binaryStream);
857 4 : if (bufferVec.size() > pos) {
858 : // 对于之前已经加过的资源,无需追加
859 0 : continue;
860 : }
861 :
862 12 : HCCL_INFO("unpack %s pos=%u, dto %s", type.Describe().c_str(), pos, dto.Describe().c_str());
863 4 : if (dto.size == 0) { // size为0,则为 remote 空buffer
864 0 : HCCL_INFO("unpack nullptr, pos=%u", pos);
865 0 : bufferVec.push_back(nullptr);
866 0 : FillRmtRmaBufferVec(nullptr, type);
867 : } else { // size非0,则构造一个remote buffer
868 4 : bufferVec.push_back(make_unique<RemoteUbRmaBuffer>(rdmaHandle, dto));
869 4 : FillRmtRmaBufferVec(bufferVec.back().get(), type);
870 12 : HCCL_INFO("unpack buffer pos=%u, rmtRmaBuffer=%s", pos, bufferVec.back()->Describe().c_str());
871 : }
872 4 : }
873 :
874 3 : return HCCL_SUCCESS;
875 : }
876 :
877 1 : HcclResult UbMemTransport::ConnVecUnpackProc(BinaryStream& binaryStream, bool& needSendFinish)
878 : {
879 : u32 rmtConnNum;
880 1 : binaryStream >> rmtConnNum;
881 3 : HCCL_INFO("start unpack conn %s connNum=%u, rmtConnNum=%u", GetLinkDescInfo().c_str(), connNum, rmtConnNum);
882 1 : if (connNum != rmtConnNum) {
883 0 : HCCL_ERROR("[UbMemTransport::ConnVecUnpackProc] connNum=%u is not equal to rmtConnNum=%u", connNum, rmtConnNum);
884 0 : return HCCL_E_PARA;
885 : }
886 :
887 1 : needSendFinish = false; // 不需要发送 finish
888 2 : for (u32 i = 0; i < rmtConnNum; i++) {
889 : u32 pos;
890 1 : binaryStream >> pos;
891 1 : ExchangeUbConnDto rmtDto;
892 1 : rmtDto.Deserialize(binaryStream);
893 3 : HCCL_INFO("unpack connection pos=%u dto %s", pos, rmtDto.Describe().c_str());
894 1 : if (commonLocRes.connVec[i]->GetStatus() != RmaConnStatus::READY) {
895 0 : HCCL_INFO(
896 : "parse and import pos=%u, rmt dto to connection[%s]", pos, commonLocRes.connVec[i]->Describe().c_str());
897 0 : commonLocRes.connVec[i]->ParseRmtExchangeDto(rmtDto);
898 0 : commonLocRes.connVec[i]->ImportRmtDto();
899 0 : needSendFinish = true; // connection 建链,需要发送finish
900 : }
901 1 : }
902 1 : return HCCL_SUCCESS;
903 : }
904 :
905 4 : void UbMemTransport::FillRmtRmaBufferVec(RemoteRmaBuffer* rmaBuffer, UbRmtBufType type)
906 : {
907 4 : if (type == UbRmtBufType::BUFFER) {
908 4 : rmtRmaBufferVec.push_back(rmaBuffer);
909 : }
910 4 : }
911 :
912 2 : HcclResult UbMemTransport::SendFinish()
913 : {
914 6 : HCCL_INFO("start send Finish Msg %s [%s], isHost_[%d]", GetLinkDescInfo().c_str(), FINISH_MSG, isHost_);
915 2 : sendFinishMsg = std::vector<char>(FINISH_MSG, FINISH_MSG + FINISH_MSG_SIZE);
916 2 : bool ret = false;
917 2 : if (isHost_) {
918 0 : ret = socket->Send(sendFinishMsg.data(), FINISH_MSG_SIZE);
919 : } else {
920 2 : socket->SendAsync(sendFinishMsg.data(), FINISH_MSG_SIZE);
921 1 : ret = true;
922 : }
923 1 : if (!ret) {
924 0 : HCCL_ERROR("[UbMemTransport::SendFinish] Send finish msg failed");
925 0 : return HCCL_E_INTERNAL;
926 : }
927 3 : HCCL_INFO("end send Finish Msg %s [%s]", GetLinkDescInfo().c_str(), FINISH_MSG);
928 1 : return HCCL_SUCCESS;
929 : }
930 :
931 1 : HcclResult UbMemTransport::RecvFinish()
932 : {
933 1 : recvFinishMsg.resize(FINISH_MSG_SIZE);
934 3 : HCCL_INFO("start recv Finish Msg %s [%s], isHost_[%d]", GetLinkDescInfo().c_str(), FINISH_MSG, isHost_);
935 1 : bool ret = false;
936 1 : if (isHost_) {
937 0 : ret = socket->Recv(recvFinishMsg.data(), FINISH_MSG_SIZE);
938 : } else {
939 1 : socket->RecvAsync(reinterpret_cast<u8*>(recvFinishMsg.data()), FINISH_MSG_SIZE);
940 1 : ret = true;
941 : }
942 1 : if (!ret) {
943 0 : HCCL_ERROR("[UbMemTransport::RecvFinish] Recv finish msg failed");
944 0 : return HCCL_E_INTERNAL;
945 : }
946 3 : HCCL_INFO("end recv Finish Msg %s [%s]", GetLinkDescInfo().c_str(), FINISH_MSG);
947 1 : return HCCL_SUCCESS;
948 : }
949 :
950 49 : std::vector<char> UbMemTransport::GetUniqueId()
951 : {
952 49 : if (baseStatus != TransportStatus::READY) {
953 4 : MACRO_THROW(InternalException, StringFormat("transport status is not ready, please check"));
954 : }
955 48 : u32 type = static_cast<u32>(transportType);
956 48 : BinaryStream binaryStream;
957 48 : binaryStream << type;
958 48 : binaryStream << notifyNum;
959 48 : binaryStream << bufferNum;
960 48 : binaryStream << static_cast<u32>(rmtBufferVec.size());
961 48 : binaryStream << connNum;
962 :
963 : // [header...][notifyUniqueId...][rmtNotifyUniqueId...][rmtBufferUniqueIds...]
964 48 : auto notifyUniqueIds = GetNotifyUniqueIds();
965 48 : binaryStream << notifyUniqueIds;
966 :
967 48 : auto rmtNotifyUniqueIds = GetRmtBufferUniqueIds(rmtNotifyVec, UbRmtBufType::NOTIFY);
968 48 : binaryStream << rmtNotifyUniqueIds;
969 :
970 48 : auto rmtBufferUniqueIds = GetRmtBufferUniqueIds(rmtBufferVec, UbRmtBufType::BUFFER);
971 48 : binaryStream << rmtBufferUniqueIds;
972 :
973 48 : auto connUniqueIds = GetConnUniqueIds();
974 48 : binaryStream << connUniqueIds;
975 :
976 48 : std::vector<char> result;
977 48 : binaryStream.Dump(result);
978 48 : return result;
979 48 : }
980 :
981 0 : std::vector<char> UbMemTransport::GetUniqueIdV2()
982 : {
983 0 : if (baseStatus != TransportStatus::READY) {
984 0 : MACRO_THROW(
985 : InternalException,
986 : StringFormat("transport status[%d] is not ready[%d], please check.", baseStatus, TransportStatus::READY));
987 : }
988 0 : u32 type = static_cast<u32>(transportType);
989 0 : BinaryStream binaryStream;
990 0 : binaryStream << type;
991 0 : binaryStream << notifyNum;
992 0 : binaryStream << bufferNum;
993 0 : binaryStream << static_cast<u32>(rmtBufferVec.size());
994 0 : binaryStream << connNum;
995 :
996 0 : auto notifyUniqueIds = GetNotifyUniqueIds();
997 0 : binaryStream << notifyUniqueIds;
998 :
999 0 : auto rmtNotifyUniqueIds = GetRmtBufferUniqueIds(rmtNotifyVec, UbRmtBufType::NOTIFY);
1000 0 : binaryStream << rmtNotifyUniqueIds;
1001 :
1002 0 : for (auto& it : commonLocRes.bufferVec) {
1003 0 : locBufferVec.emplace_back(reinterpret_cast<LocalUbRmaBuffer*>(it));
1004 : }
1005 :
1006 0 : auto locBufferUniqueIds = GetLocBufferUniqueIds(locBufferVec, UbRmtBufType::BUFFER);
1007 0 : binaryStream << locBufferUniqueIds;
1008 :
1009 0 : auto rmtBufferUniqueIds = GetRmtBufferUniqueIds(rmtBufferVec, UbRmtBufType::BUFFER);
1010 0 : binaryStream << rmtBufferUniqueIds;
1011 :
1012 0 : auto drainUniqueIds = GetDrainUniqueIds();
1013 0 : binaryStream << drainUniqueIds;
1014 :
1015 0 : auto connUniqueIds = GetConnUniqueIds();
1016 0 : binaryStream << connUniqueIds;
1017 :
1018 0 : std::vector<char> result;
1019 0 : binaryStream.Dump(result);
1020 0 : return result;
1021 0 : }
1022 :
1023 0 : std::vector<char> UbMemTransport::PackConnData()
1024 : {
1025 0 : if (baseStatus != TransportStatus::READY) {
1026 0 : MACRO_THROW(
1027 : InternalException,
1028 : StringFormat("transport status[%d] is not ready[%d], please check.", baseStatus, TransportStatus::READY));
1029 : }
1030 0 : u32 type = static_cast<u32>(transportType);
1031 0 : BinaryStream binaryStream;
1032 0 : binaryStream << type;
1033 0 : binaryStream << connNum;
1034 :
1035 0 : auto connUniqueIds = GetConnUniqueIds();
1036 0 : binaryStream << connUniqueIds;
1037 :
1038 0 : std::vector<char> result;
1039 0 : binaryStream.Dump(result);
1040 0 : return result;
1041 0 : }
1042 :
1043 : std::vector<char>
1044 1 : UbMemTransport::GetSingleRmtBufferUniqueId(u64 addr, u64 size, u32 tokenId, u32 tokenValue, u32 notifyId) const
1045 : {
1046 1 : BinaryStream binaryStream;
1047 1 : binaryStream << addr;
1048 1 : binaryStream << size;
1049 1 : binaryStream << tokenId;
1050 1 : binaryStream << tokenValue;
1051 1 : binaryStream << notifyId;
1052 3 : HCCL_INFO("UbMemTransport RmtBuffer[addr=0x%llx, size=0x%llx, notifyId=%u]", addr, size, notifyId);
1053 1 : std::vector<char> result;
1054 1 : binaryStream.Dump(result);
1055 1 : return result;
1056 1 : }
1057 :
1058 1 : std::vector<char> UbMemTransport::GetSingleLocBufferUniqueId(u64 addr, u64 size, u32 tokenId, u32 tokenValue) const
1059 : {
1060 1 : BinaryStream binaryStream;
1061 1 : binaryStream << addr;
1062 1 : binaryStream << size;
1063 1 : binaryStream << tokenId;
1064 1 : binaryStream << tokenValue;
1065 3 : HCCL_INFO("UbMemTransport LocBuffer[addr=0x%llx, size=0x%llx]", addr, size);
1066 1 : std::vector<char> result;
1067 1 : binaryStream.Dump(result);
1068 1 : return result;
1069 1 : }
1070 :
1071 48 : std::vector<char> UbMemTransport::GetNotifyUniqueIds()
1072 : {
1073 144 : HCCL_INFO("start packing all notify uniqueIds");
1074 48 : std::vector<char> result(0);
1075 94 : for (auto& it : commonLocRes.notifyVec) {
1076 138 : HCCL_INFO("ubMemTransport Notify %s", it->Describe().c_str());
1077 46 : auto uniqueId = it->GetUniqueId();
1078 46 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1079 46 : }
1080 48 : return result;
1081 0 : }
1082 :
1083 0 : std::vector<char> UbMemTransport::GetDrainUniqueIds() const
1084 : {
1085 0 : HCCL_INFO("start packing drain resources uniqueIds");
1086 0 : std::vector<char> result(0);
1087 0 : std::vector<char> uniqueId;
1088 :
1089 : // pack drain notify
1090 0 : if (drainNotify_ != nullptr) {
1091 0 : auto dto = drainNotify_->GetExchangeDto();
1092 0 : ExchangeUbBufferDto* rawDto = static_cast<ExchangeUbBufferDto*>(dto.get());
1093 0 : uniqueId = GetSingleRmtBufferUniqueId(
1094 0 : rawDto->addr, rawDto->size, rawDto->tokenId, rawDto->tokenValue, rawDto->notifyId);
1095 0 : HCCL_INFO("UbMemTransport::GetDrainUniqueIds, %s", drainNotify_->Describe().c_str());
1096 0 : } else {
1097 0 : uniqueId = GetSingleRmtBufferUniqueId(0, 0, 0, 0, UINT32_MAX); // 填充一个空的buffer
1098 0 : HCCL_INFO("UbMemTransport::GetDrainUniqueIds, drainNotify_ null buffer");
1099 : }
1100 0 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1101 :
1102 : // pack drain rmtMem
1103 0 : if (rmtDrainBuffer_ != nullptr) {
1104 0 : uniqueId = GetSingleRmtBufferUniqueId(
1105 0 : rmtDrainBuffer_->GetAddr(), rmtDrainBuffer_->GetSize(), rmtDrainBuffer_->GetTokenId(),
1106 0 : rmtDrainBuffer_->GetTokenValue(), rmtDrainBuffer_->GetNotifyId());
1107 0 : HCCL_INFO("UbMemTransport::GetDrainUniqueIds, %s", rmtDrainBuffer_->Describe().c_str());
1108 : } else {
1109 0 : uniqueId = GetSingleRmtBufferUniqueId(0, 0, 0, 0, UINT32_MAX); // 填充一个空的buffer
1110 0 : HCCL_INFO("UbMemTransport::GetDrainUniqueIds, rmtDrainBuffer_ null buffer");
1111 : }
1112 0 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1113 :
1114 0 : return result;
1115 0 : }
1116 :
1117 96 : std::vector<char> UbMemTransport::GetRmtBufferUniqueIds(RemoteBufferVec& bufferVec, UbRmtBufType type) const
1118 : {
1119 288 : HCCL_INFO("start packing all remote buffer %s uniqueIds", type.Describe().c_str());
1120 96 : std::vector<char> result(0);
1121 96 : for (auto& it : bufferVec) {
1122 0 : std::vector<char> uniqueId;
1123 0 : if (it != nullptr) {
1124 0 : uniqueId = GetSingleRmtBufferUniqueId(
1125 0 : it->GetAddr(), it->GetSize(), it->GetTokenId(), it->GetTokenValue(), it->GetNotifyId());
1126 0 : HCCL_INFO("UbMemTransport::GetRmtBufferUniqueIds, %s", it->Describe().c_str());
1127 : } else {
1128 0 : uniqueId = GetSingleRmtBufferUniqueId(0, 0, 0, 0, UINT32_MAX); // 填充一个空的buffer
1129 0 : HCCL_INFO("UbMemTransport::GetRmtBufferUniqueIds, null buffer");
1130 : }
1131 0 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1132 0 : }
1133 96 : return result;
1134 0 : }
1135 :
1136 0 : std::vector<char> UbMemTransport::GetLocBufferUniqueIds(LocalBufferVec& bufferVec, UbRmtBufType type) const
1137 : {
1138 0 : HCCL_INFO("start packing all local buffer %s uniqueIds", type.Describe().c_str());
1139 0 : std::vector<char> result(0);
1140 0 : for (auto& it : bufferVec) {
1141 0 : std::vector<char> uniqueId;
1142 0 : if (it != nullptr) {
1143 0 : uniqueId = GetSingleLocBufferUniqueId(it->GetAddr(), it->GetSize(), it->GetTokenId(), it->GetTokenValue());
1144 0 : HCCL_INFO("UbMemTransport::GetLocBufferUniqueIds, %s", it->Describe().c_str());
1145 : } else {
1146 0 : uniqueId = GetSingleLocBufferUniqueId(0, 0, 0, 0); // 填充一个空的buffer
1147 0 : HCCL_INFO("UbMemTransport::GetLocBufferUniqueIds, null buffer");
1148 : }
1149 0 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1150 0 : }
1151 0 : return result;
1152 0 : }
1153 :
1154 48 : std::vector<char> UbMemTransport::GetConnUniqueIds()
1155 : {
1156 144 : HCCL_INFO("start packing all conn uniqueIds");
1157 48 : std::vector<char> result(0);
1158 94 : for (auto& it : commonLocRes.connVec) {
1159 138 : HCCL_INFO("[UbMemTransport::%s] conn[%s]", __func__, it->Describe().c_str());
1160 46 : auto uniqueId = it->GetUniqueId();
1161 46 : result.insert(result.end(), uniqueId.begin(), uniqueId.end());
1162 46 : }
1163 48 : return result;
1164 0 : }
1165 :
1166 3 : void UbMemTransport::SaveDfxTaskInfo(const TaskParam& taskParam)
1167 : {
1168 : u32 taskId;
1169 : u32 streamId;
1170 3 : HrtGetTaskIdAndStreamID(taskId, streamId);
1171 :
1172 3 : callback(streamId, taskId, taskParam);
1173 3 : }
1174 :
1175 2 : HcclResult UbMemTransport::GetRemoteMems(uint32_t* memNum, CommMem** remoteMem, char*** memInfos)
1176 : {
1177 2 : std::lock_guard<std::mutex> lock(remoteMemsMutex_);
1178 : Hccl::RemoteMemCtx<std::unique_ptr<RemoteUbRmaBuffer>> remoteMemCtx{
1179 2 : cacheValid_, rmtBufferVec, remoteUserMems_, memInfoCopies_, memInfoPointers_, remoteMem, memInfos, memNum};
1180 2 : CHK_RET(GetRemoteUserMems(remoteMemCtx));
1181 2 : return HCCL_SUCCESS;
1182 2 : }
1183 :
1184 5 : HcclResult UbMemTransport::CheckSocketStatus(std::string socketOpreator)
1185 : {
1186 5 : CHK_PTR_NULL(socket);
1187 5 : auto timeout = std::chrono::seconds(Hccl::EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
1188 5 : auto startTime = std::chrono::steady_clock::now();
1189 5 : uint32_t retryCount = 0;
1190 : while (true) {
1191 5 : SocketStatus socketStatus = socket->GetAsyncStatus();
1192 5 : if (socketStatus == SocketStatus::OK) {
1193 : auto elapsed
1194 4 : = std::chrono::duration_cast<std::chrono::milliseconds>(std::chrono::steady_clock::now() - startTime)
1195 4 : .count();
1196 12 : HCCL_INFO(
1197 : "[UbMemTransport][%s] socket transport operation[%s] success, elapsed[%lld]ms, retryCount[%u]",
1198 : __func__, socketOpreator.c_str(), elapsed, retryCount);
1199 4 : break;
1200 : }
1201 1 : if ((std::chrono::steady_clock::now() - startTime) >= timeout || socketStatus == Hccl::SocketStatus::TIMEOUT) {
1202 : auto elapsed
1203 1 : = std::chrono::duration_cast<std::chrono::milliseconds>(std::chrono::steady_clock::now() - startTime)
1204 1 : .count();
1205 3 : HCCL_ERROR(
1206 : "[UbMemTransport][%s] socket transport operation[%s] timeout after %lld sec, elapsed[%lld]ms, "
1207 : "retryCount[%u]",
1208 : __func__, socketOpreator.c_str(), timeout, elapsed, retryCount);
1209 1 : return HCCL_E_TIMEOUT;
1210 : }
1211 0 : retryCount++;
1212 0 : }
1213 4 : return HCCL_SUCCESS;
1214 : }
1215 :
1216 3 : HcclResult UbMemTransport::UpdateMemInfo(std::vector<LocalRmaBuffer*>& bufferVecTemp)
1217 : {
1218 3 : if (bufferVecTemp.size() == 0) {
1219 3 : HCCL_WARNING("[UbMemTransport][UpdateMemInfo] bufferNum is 0.");
1220 1 : return HCCL_SUCCESS;
1221 : }
1222 6 : HCCL_INFO("[UbMemTransport][UpdateMemInfo] bufferNum[%zu]", bufferVecTemp.size());
1223 2 : sendData.clear();
1224 2 : BinaryStream sendStream;
1225 2 : std::vector<std::unique_ptr<RemoteUbRmaBuffer>> rmtBufferTemp{};
1226 84 : TRY_CATCH_RETURN([&]() -> void {
1227 : BufferVecPack(sendStream, bufferVecTemp);
1228 : sendStream.Dump(sendData);
1229 : u32 sendSize = sendData.size();
1230 : socket->SendAsync(&sendSize, sizeof(sendSize));
1231 : HCCL_INFO(
1232 : "[UbMemTransport][UpdateMemInfo] Send size[%u] of data success. [%zu] bytes sent.", sendSize,
1233 : sizeof(sendSize));
1234 : HcclResult result = CheckSocketStatus("SendDataSize");
1235 : CHK_RET_THROW(
1236 : InternalException, StringFormat("[UbMemTransport][UpdateMemInfo] failed to send dataSize."), result);
1237 : RecvDataSize();
1238 : result = CheckSocketStatus("RecvDataSize");
1239 : CHK_RET_THROW(
1240 : InternalException, StringFormat("[UbMemTransport][UpdateMemInfo] failed to receive dataSize."), result);
1241 : SendExchangeData();
1242 : result = CheckSocketStatus("SendExchangeData");
1243 : CHK_RET_THROW(InternalException, StringFormat("[UbMemTransport][UpdateMemInfo] failed to send data."), result);
1244 : RecvExchangeData();
1245 : result = CheckSocketStatus("RecvExchangeData");
1246 : CHK_RET_THROW(
1247 : InternalException, StringFormat("[UbMemTransport][UpdateMemInfo] failed to receive data."), result);
1248 : BinaryStream recvStream(recvData);
1249 : RmtBufferVecUnpackProc(bufferNum, recvStream, rmtBufferTemp, UbRmtBufType::BUFFER);
1250 : }());
1251 2 : rmtBufferVec.insert(
1252 1 : rmtBufferVec.end(), std::make_move_iterator(rmtBufferTemp.begin()),
1253 : std::make_move_iterator(rmtBufferTemp.end()));
1254 1 : commonLocRes.bufferVec.insert(commonLocRes.bufferVec.end(), bufferVecTemp.begin(), bufferVecTemp.end());
1255 1 : cacheValid_ = false;
1256 1 : return HCCL_SUCCESS;
1257 2 : }
1258 :
1259 0 : HcclResult UbMemTransport::Init()
1260 : {
1261 0 : for (auto& ubConn : commonLocRes.connVec) {
1262 0 : TRY_CATCH_RETURN(ubConn->Connect());
1263 : }
1264 :
1265 0 : return HCCL_SUCCESS;
1266 : }
1267 :
1268 0 : HcclResult UbMemTransport::DeInit() const
1269 : {
1270 0 : socket->Destroy();
1271 0 : return HCCL_SUCCESS;
1272 : }
1273 :
1274 0 : HcclResult UbMemTransport::GetRemoteSeg(const void* addr, u64 len, u64* seg)
1275 : {
1276 0 : if (rmtBufferVec.empty()) {
1277 0 : HCCL_ERROR("[UbMemTransport::%s] rmtBufferVec is empty.", __func__);
1278 0 : return HCCL_E_INTERNAL;
1279 : }
1280 :
1281 0 : bool isAddrInRange = false;
1282 0 : for (auto& it : rmtBufferVec) {
1283 0 : Buffer iterBuf(it->GetAddr(), it->GetSize());
1284 0 : if (iterBuf.Contains(reinterpret_cast<uintptr_t>(addr), len)) {
1285 0 : *seg = it->GetSegVa();
1286 0 : isAddrInRange = true;
1287 0 : break;
1288 : }
1289 0 : }
1290 :
1291 0 : if (!isAddrInRange) {
1292 0 : return HCCL_E_INTERNAL;
1293 : }
1294 0 : return HCCL_SUCCESS;
1295 : }
1296 :
1297 : } // namespace Hccl
|