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 <numeric>
12 : #include "sal_pub.h"
13 : #include "hccl_one_sided_conn.h"
14 : #include "p2p_mgmt_pub.h"
15 :
16 : namespace hccl {
17 : using namespace std;
18 :
19 0 : HcclOneSidedConn::HcclOneSidedConn(
20 : const HcclNetDevCtx& netDevCtx, const HcclRankLinkInfo& localRankInfo, const HcclRankLinkInfo& remoteRankInfo,
21 : std::unique_ptr<HcclSocketManager>& socketManager, std::unique_ptr<NotifyPool>& notifyPool,
22 : const HcclDispatcher& dispatcher, const bool& useRdma, u32 sdid, u32 serverId, u32 trafficClass, u32 serviceLevel,
23 0 : bool aicpuUnfoldMode, bool isStandardCard, bool isNeedEnableP2P)
24 0 : : localRankInfo_(localRankInfo),
25 0 : socketManager_(socketManager),
26 0 : notifyPool_(notifyPool),
27 0 : aicpuUnfoldMode_(aicpuUnfoldMode),
28 0 : isStandardCard_(isStandardCard),
29 0 : isNeedEnableP2P_(isNeedEnableP2P)
30 : {
31 0 : netDevCtx_ = netDevCtx;
32 0 : remoteRankInfo_ = remoteRankInfo;
33 0 : useRdma_ = useRdma;
34 0 : TransportMem::AttrInfo attrInfo{};
35 0 : attrInfo.localRankId = localRankInfo.userRank;
36 0 : attrInfo.remoteRankId = remoteRankInfo.userRank;
37 0 : attrInfo.sdid = sdid;
38 0 : attrInfo.serverId = serverId;
39 0 : attrInfo.trafficClass = trafficClass;
40 0 : attrInfo.serviceLevel = serviceLevel;
41 0 : if (useRdma) {
42 0 : transportMemPtr_ = TransportMem::Create(
43 0 : TransportMem::TpType::ROCE, notifyPool, netDevCtx, dispatcher, attrInfo, aicpuUnfoldMode_);
44 : } else {
45 0 : transportMemPtr_ = TransportMem::Create(
46 0 : TransportMem::TpType::IPC, notifyPool, netDevCtx, dispatcher, attrInfo, aicpuUnfoldMode_);
47 : }
48 0 : CHK_SMART_PTR_RET_NULL(transportMemPtr_);
49 0 : }
50 :
51 0 : HcclOneSidedConn::~HcclOneSidedConn()
52 : {
53 0 : if ((isStandardCard_ && !useRdma_) && isNeedEnableP2P_) {
54 0 : if (!enableP2PDevices_.empty()) {
55 0 : P2PMgmtPub::DisableP2P(enableP2PDevices_);
56 0 : enableP2PDevices_.clear();
57 : }
58 : }
59 0 : }
60 :
61 0 : HcclResult HcclOneSidedConn::Connect(const std::string& commIdentifier, s32 timeoutSec)
62 : {
63 0 : const auto startTime = TIME_NOW();
64 0 : if (aicpuUnfoldMode_) {
65 0 : CHK_RET(DeviceMem::alloc(transportDataDevice_, sizeof(TransportDeviceNormalData)));
66 : }
67 : // 创建socket用于交换数据
68 0 : std::string newTag;
69 0 : if (localRankInfo_.userRank < remoteRankInfo_.userRank) {
70 : // 本端为SERVER,对端为CLIENT
71 0 : newTag = string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
72 0 : + string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
73 0 : + commIdentifier;
74 : } else {
75 0 : newTag = string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
76 0 : + string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
77 0 : + commIdentifier;
78 : }
79 0 : HCCL_DEBUG("[HcclOneSidedConn][Connect]socket tag:%s", newTag.c_str());
80 :
81 : // 1、通信域初始化时会做非标卡且非310P场景的EnableP2P操作
82 : // 2、此处补全标卡且不使用RDMA场景下的EnableP2P操作
83 0 : HCCL_INFO(
84 : "[HcclOneSidedConn][Connect]localRankId[%u]-localDevicePhyId[%u], remoteRankId[%u]-remoteDevicePhyId[%u], "
85 : "isStandardCard[%s], useRdma[%s], isNeedEnableP2P[%s]",
86 : localRankInfo_.userRank, localRankInfo_.devicePhyId, remoteRankInfo_.userRank, remoteRankInfo_.devicePhyId,
87 : isStandardCard_ ? "true" : "false", useRdma_ ? "true" : "false", isNeedEnableP2P_ ? "true" : "false");
88 :
89 0 : if ((isStandardCard_ && !useRdma_) && isNeedEnableP2P_) {
90 0 : std::vector<u32> enableP2PDevices;
91 0 : enableP2PDevices.push_back(remoteRankInfo_.devicePhyId);
92 0 : HCCL_INFO(
93 : "[HcclOneSidedConn][Connect]localDevicePhyId[%u] enable p2p with remoteDevicePhyId[%u]",
94 : localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId);
95 0 : HcclResult ret = P2PMgmtPub::EnableP2P(enableP2PDevices);
96 0 : CHK_PRT_RET(
97 : ret != HCCL_SUCCESS,
98 : HCCL_ERROR(
99 : "[HcclOneSidedConn][Connect]Enable P2P Failed, localPhyId[%u], remotephyId[%u], ret[%u]",
100 : localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId, ret),
101 : ret);
102 0 : enableP2PDevices_.push_back(remoteRankInfo_.devicePhyId);
103 0 : }
104 :
105 : // EnableP2P需要和WaitP2PEnabled匹配使用,此处需要对1、2两处的EnableP2P做WaitP2PEnabled处理
106 0 : if ((!isStandardCard_ || !useRdma_) && isNeedEnableP2P_) {
107 0 : std::vector<u32> waitP2PEnabledDevices;
108 0 : waitP2PEnabledDevices.push_back(remoteRankInfo_.devicePhyId);
109 0 : HCCL_INFO(
110 : "[HcclOneSidedConn][Connect]localDevicePhyId[%u] wait p2p enable with remoteDevicePhyId[%u]",
111 : localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId);
112 0 : HcclResult ret = P2PMgmtPub::WaitP2PEnabled(waitP2PEnabledDevices, [this]() -> bool {
113 0 : return socketManager_->GetStopFlag();
114 : });
115 0 : CHK_PRT_RET(
116 : ret != HCCL_SUCCESS,
117 : HCCL_ERROR(
118 : "[HcclOneSidedConn][Connect]Wait Enable P2P Failed, src devicePhyId[%u], dst devicePhyId[%u], ret[%u]",
119 : localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId, ret),
120 : ret);
121 0 : }
122 :
123 0 : std::vector<std::shared_ptr<HcclSocket>> connectSockets;
124 0 : CHK_RET(socketManager_->CreateSingleLinkSocket(
125 : newTag, netDevCtx_, remoteRankInfo_, connectSockets, true, true, timeoutSec));
126 0 : CHK_RET(transportMemPtr_->SetDataSocket(connectSockets[0]));
127 0 : socket_ = connectSockets[0];
128 :
129 0 : if (useRdma_) {
130 : // 创建socket用于QP建链
131 0 : newTag += "_QP";
132 0 : auto timeCostSec = std::chrono::duration_cast<std::chrono::seconds>(TIME_NOW() - startTime).count();
133 0 : auto timeLeft = timeoutSec - timeCostSec;
134 0 : CHK_RET(socketManager_->CreateSingleLinkSocket(
135 : newTag, netDevCtx_, remoteRankInfo_, connectSockets, true, true, timeLeft));
136 0 : CHK_RET(transportMemPtr_->SetSocket(connectSockets[0]));
137 0 : rdmaSocket_ = connectSockets[0];
138 :
139 0 : if (timeoutSec == -1) {
140 : // timeout为-1,超时时间设为最大值
141 0 : CHK_RET(transportMemPtr_->Connect(INT_MAX));
142 : } else {
143 : // 超时时间减去已消耗的时间,避免接口整体耗时超过入参的秒数
144 0 : timeCostSec = std::chrono::duration_cast<std::chrono::seconds>(TIME_NOW() - startTime).count();
145 0 : timeLeft = timeoutSec - timeCostSec;
146 0 : CHK_PRT_RET(
147 : timeLeft <= 0,
148 : HCCL_ERROR(
149 : "[HcclOneSidedConn][Connect] Connect timeout. comm[%s], timeoutSec[%d s]", commIdentifier.c_str(),
150 : timeoutSec),
151 : HCCL_E_TIMEOUT);
152 :
153 : // Transport建链:notify资源创建+QP建链
154 0 : CHK_RET(transportMemPtr_->Connect(timeLeft));
155 : }
156 : }
157 :
158 0 : return HCCL_SUCCESS;
159 0 : }
160 :
161 0 : void HcclOneSidedConn::CleanSocketResource(const std::string& commIdentifier)
162 : {
163 : HcclSocketRole role;
164 0 : std::string newTag;
165 0 : if (localRankInfo_.userRank < remoteRankInfo_.userRank) {
166 : // 本端为SERVER,对端为CLIENT
167 0 : role = HcclSocketRole::SOCKET_ROLE_SERVER;
168 0 : newTag = string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
169 0 : + string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
170 0 : + commIdentifier;
171 : } else {
172 0 : role = HcclSocketRole::SOCKET_ROLE_CLIENT;
173 0 : newTag = string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
174 0 : + string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
175 0 : + commIdentifier;
176 : }
177 0 : if (socket_ != nullptr) {
178 0 : HCCL_INFO(
179 : "[HcclOneSidedConn][%s]abort and delete socket with remote[%u] tag[%s]", __func__, remoteRankInfo_.userRank,
180 : newTag.c_str());
181 0 : std::map<u32, std::vector<std::shared_ptr<HcclSocket>>> socketsMap;
182 0 : std::vector<std::shared_ptr<HcclSocket>> rankSockets{socket_};
183 0 : socketsMap.insert(std::make_pair(remoteRankInfo_.userRank, rankSockets));
184 0 : socketManager_->AbortAndDeleteSocket(newTag, role, socketsMap);
185 0 : }
186 :
187 0 : if (rdmaSocket_ != nullptr) {
188 0 : newTag += "_QP";
189 0 : HCCL_INFO(
190 : "[HcclOneSidedConn][%s]abort and delete rdmaSocket with remote[%u] tag[%s]", __func__,
191 : remoteRankInfo_.userRank, newTag.c_str());
192 0 : std::map<u32, std::vector<std::shared_ptr<HcclSocket>>> socketsMap;
193 0 : std::vector<std::shared_ptr<HcclSocket>> rankSockets{rdmaSocket_};
194 0 : socketsMap.insert(std::make_pair(remoteRankInfo_.userRank, rankSockets));
195 0 : socketManager_->AbortAndDeleteSocket(newTag, role, socketsMap);
196 0 : }
197 0 : return;
198 0 : }
199 :
200 0 : HcclResult HcclOneSidedConn::ExchangeIpcProcessInfo(const ProcessInfo& localProcess, ProcessInfo& remoteProcess)
201 : {
202 0 : HCCL_DEBUG(
203 : "[HcclOneSidedConn][ExchangeIpcProcessInfo] localRank[%u] exchange process info", localRankInfo_.userRank);
204 0 : if (socket_->GetLocalRole() == HcclSocketRole::SOCKET_ROLE_CLIENT) {
205 0 : CHK_RET(socket_->Recv(&remoteProcess, sizeof(ProcessInfo)));
206 0 : CHK_RET(socket_->Send(&localProcess, sizeof(ProcessInfo)));
207 : } else {
208 0 : CHK_RET(socket_->Send(&localProcess, sizeof(ProcessInfo)));
209 0 : CHK_RET(socket_->Recv(&remoteProcess, sizeof(ProcessInfo)));
210 : }
211 0 : return HCCL_SUCCESS;
212 : }
213 :
214 0 : HcclResult HcclOneSidedConn::ExchangeMemDesc(
215 : const HcclMemDescs& localMemDescs, HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
216 : {
217 0 : TransportMem::RmaMemDesc* localMemDescArray
218 : = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(localMemDescs.array));
219 0 : TransportMem::RmaMemDescs localRmaMemDescs = {localMemDescArray, localMemDescs.arrayLength};
220 0 : TransportMem::RmaMemDesc* remoteMemDescArray
221 : = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(remoteMemDescs.array));
222 0 : TransportMem::RmaMemDescs remoteRmaMemDescs = {remoteMemDescArray, remoteMemDescs.arrayLength};
223 :
224 0 : return transportMemPtr_->ExchangeMemDesc(localRmaMemDescs, remoteRmaMemDescs, actualNumOfRemote);
225 : }
226 :
227 0 : HcclResult HcclOneSidedConn::GetMemType(const char* description, RmaMemType& memType)
228 : {
229 0 : std::string tempDesc = std::string(description, TRANSPORT_EMD_ESC_SIZE);
230 0 : std::istringstream iss(tempDesc);
231 : // 定义需要跳过的变量的大小
232 : const std::vector<size_t> skip_sizes = {
233 : sizeof(u8), // type
234 : sizeof(void*), // addr
235 : sizeof(u64), // size
236 : sizeof(void*) // devAddr
237 0 : };
238 : // 计算偏移量
239 0 : size_t offset = std::accumulate(skip_sizes.begin(), skip_sizes.end(), 0);
240 : // 定位到 memType 的位置
241 0 : iss.seekg(offset);
242 0 : iss.read(reinterpret_cast<char_t*>(&memType), sizeof(memType));
243 0 : CHK_PRT_RET(
244 : memType >= RmaMemType::TYPE_NUM,
245 : HCCL_ERROR("[HcclOneSidedConn][GetMemType] get memType failed memType[%d]", static_cast<int>(memType)),
246 : HCCL_E_INTERNAL);
247 0 : return HCCL_SUCCESS;
248 0 : }
249 :
250 0 : void HcclOneSidedConn::EnableMemAccess(const HcclMemDesc& remoteMemDesc, HcclMem& remoteMem)
251 : {
252 : // 数据第一次转换
253 0 : HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access.");
254 0 : const RmaMemDesc* remoteRmaMemDesc = static_cast<const RmaMemDesc*>(static_cast<const void*>(remoteMemDesc.desc));
255 0 : string tempStr = RmaMemDescCopyToStr(*remoteRmaMemDesc);
256 : RmaMemType memType;
257 0 : EXCEPTION_THROW_IF_ERR(
258 : GetMemType(remoteRmaMemDesc->memDesc, memType), "[HcclOneSidedConn][EnableMemAccess] get memType failed");
259 0 : auto iter = memDescMap_.find(tempStr);
260 0 : if (iter != memDescMap_.end()) {
261 0 : HcclBuf& outBuf = iter->second;
262 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf.addr), outBuf.len);
263 0 : auto resultPair = remoteRmaBufferMgr_.Add(tempKey, outBuf.handle);
264 0 : EXCEPTION_THROW_IF_COND_ERR(
265 : resultPair.first == remoteRmaBufferMgr_.End(),
266 : "[HcclOneSidedConn][EnableMemAccess]The memory that is expected to enable"
267 : " overlaps with the memory that has been enabled, please check params");
268 0 : remoteMem.type = static_cast<HcclMemType>(memType); // GE会检查这个字段先从字符串中获取
269 0 : remoteMem.addr = outBuf.addr;
270 0 : remoteMem.size = outBuf.len;
271 0 : return;
272 : }
273 :
274 : HcclBuf outBuf;
275 0 : EXCEPTION_THROW_IF_ERR(
276 : HcclMemImport(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE, true, &outBuf, netDevCtx_),
277 : "[HcclOneSidedConn][EnableMemAccess] Enable memory access failed.");
278 0 : remoteMem.type = static_cast<HcclMemType>(memType); // GE会检查这个字段先从字符串中获取
279 0 : remoteMem.addr = outBuf.addr;
280 0 : remoteMem.size = outBuf.len;
281 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf.addr), outBuf.len);
282 0 : auto resultPair = remoteRmaBufferMgr_.Add(tempKey, outBuf.handle);
283 :
284 0 : EXCEPTION_THROW_IF_COND_ERR(
285 : resultPair.first == remoteRmaBufferMgr_.End(),
286 : "[HcclOneSidedConn][EnableMemAccess]The memory that is expected to enable"
287 : " overlaps with the memory that has been enabled, please check params");
288 0 : HCCL_INFO(
289 : "[HcclOneSidedConn][EnableMemAccess] after insert remoteRmaBufferMgr_ size[%d]", remoteRmaBufferMgr_.size());
290 :
291 0 : HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] before insert memDescMap_ size[%d]", memDescMap_.size());
292 :
293 0 : memDescMap_.emplace(tempStr, outBuf);
294 0 : HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] after insert memDescMap_ size[%d]", memDescMap_.size());
295 0 : HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access success.");
296 0 : }
297 :
298 0 : void HcclOneSidedConn::DisableMemAccess(const HcclMemDesc& remoteMemDesc)
299 : {
300 : // 数据第一次转换
301 0 : const RmaMemDesc* remoteRmaMemDesc = static_cast<const RmaMemDesc*>(static_cast<const void*>(remoteMemDesc.desc));
302 0 : string tempStr = RmaMemDescCopyToStr(*remoteRmaMemDesc);
303 0 : auto it = memDescMap_.find(tempStr);
304 0 : EXCEPTION_THROW_IF_COND_ERR(it == memDescMap_.end(), "Can't find hcclmem by key");
305 :
306 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(it->second.addr), it->second.len);
307 0 : HcclBuf& buf = it->second;
308 : try {
309 0 : if (remoteRmaBufferMgr_.Del(tempKey)) {
310 0 : EXCEPTION_THROW_IF_COND_ERR(HcclMemClose(&buf) != HCCL_SUCCESS, "Close remote memory failed.");
311 0 : HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] before erase memDescMap_ size[%d]", memDescMap_.size());
312 0 : memDescMap_.erase(remoteRmaMemDesc->memDesc);
313 0 : HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] after erase memDescMap_ size[%d]", memDescMap_.size());
314 : // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
315 0 : HCCL_INFO("[TransportIpcMem][DisableMemAccess]Memory reference count is 0, disable memory access.");
316 : } else {
317 : // 删除失败:输入key是表中某一最相近key的全集,计数不为0(存在其他remoteRank使用),返回false
318 0 : HCCL_INFO("[TransportIpcMem][DisableMemAccess]Memory reference count is larger than 0"
319 : "(used by other RemoteRank), do not disable memory.");
320 : }
321 0 : } catch (std::out_of_range& e) {
322 0 : HCCL_ERROR("[TransportIpcMem][DisableMemAccess] catch RmaBufferMgr Del exception: %s", e.what());
323 0 : EXCEPTION_THROW_IF_COND_ERR(true, "[TransportIpcMem][DisableMemAccess] catch RmaBufferMgr Del exception");
324 0 : }
325 0 : HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] Disable memory access success.");
326 0 : }
327 :
328 0 : void HcclOneSidedConn::BatchWrite(const HcclOneSideOpDesc* oneSideDescs, u32 descNum, const rtStream_t& stream)
329 : {
330 0 : for (u32 i = 0; i < descNum; i++) {
331 0 : if (oneSideDescs[i].count == 0) {
332 0 : HCCL_WARNING("[HcclOneSidedConn][BatchWrite] Desc item[%u] count is 0.", i);
333 : }
334 : u32 unitSize;
335 0 : EXCEPTION_THROW_IF_ERR(
336 : SalGetDataTypeSize(oneSideDescs[i].dataType, unitSize),
337 : "[HcclOneSidedConn][BatchWrite] Get dataType size failed!");
338 0 : u64 byteSize = oneSideDescs[i].count * unitSize;
339 0 : HCCL_DEBUG(
340 : "[HcclOneSidedConn][BatchWrite] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i,
341 : oneSideDescs[i].localAddr, oneSideDescs[i].remoteAddr, byteSize);
342 :
343 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(oneSideDescs[i].remoteAddr), byteSize);
344 :
345 0 : auto rmaBuffer = remoteRmaBufferMgr_.Find(tempKey);
346 0 : EXCEPTION_THROW_IF_COND_ERR(!rmaBuffer.first, "Can't find remoteBuffer by key");
347 :
348 0 : HcclBuf localMem = {oneSideDescs[i].localAddr, byteSize, nullptr};
349 0 : HcclBuf remoteMem = {oneSideDescs[i].remoteAddr, byteSize, rmaBuffer.second};
350 0 : EXCEPTION_THROW_IF_ERR(
351 : transportMemPtr_->Write(remoteMem, localMem, stream),
352 : "[HcclOneSidedConn][BatchWrite] transportMem WriteAsync failed.");
353 : }
354 0 : EXCEPTION_THROW_IF_ERR(transportMemPtr_->AddOpFence(stream), "[HcclOneSidedConn][BatchWrite] AddOpFence failed.");
355 0 : }
356 :
357 0 : void HcclOneSidedConn::BatchRead(const HcclOneSideOpDesc* oneSideDescs, u32 descNum, const rtStream_t& stream)
358 : {
359 0 : for (u32 i = 0; i < descNum; i++) {
360 0 : if (oneSideDescs[i].count == 0) {
361 0 : HCCL_WARNING("[HcclOneSidedConn][BatchRead] Desc item[%u] count is 0.", i);
362 : }
363 : u32 unitSize;
364 0 : EXCEPTION_THROW_IF_ERR(
365 : SalGetDataTypeSize(oneSideDescs[i].dataType, unitSize),
366 : "[HcclOneSidedConn][BatchRead] Get dataType size failed!");
367 0 : u64 byteSize = oneSideDescs[i].count * unitSize;
368 0 : HCCL_DEBUG(
369 : "[HcclOneSidedConn][BatchRead] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i,
370 : oneSideDescs[i].localAddr, oneSideDescs[i].remoteAddr, byteSize);
371 :
372 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(oneSideDescs[i].remoteAddr), byteSize);
373 :
374 0 : auto rmaBuffer = remoteRmaBufferMgr_.Find(tempKey);
375 0 : EXCEPTION_THROW_IF_COND_ERR(!rmaBuffer.first, "Can't find remoteBuffer by key");
376 :
377 0 : HcclBuf localMem = {oneSideDescs[i].localAddr, byteSize, nullptr};
378 0 : HcclBuf remoteMem = {oneSideDescs[i].remoteAddr, byteSize, rmaBuffer.second};
379 0 : EXCEPTION_THROW_IF_ERR(
380 : transportMemPtr_->Read(localMem, remoteMem, stream),
381 : "[HcclOneSidedConn][BatchRead] transportMem ReadAsync failed.");
382 : }
383 0 : EXCEPTION_THROW_IF_ERR(transportMemPtr_->AddOpFence(stream), "[HcclOneSidedConn][BatchRead] AddOpFence failed.");
384 0 : }
385 :
386 0 : HcclResult HcclOneSidedConn::GetTransInfo(
387 : HcclOneSideOpDescParam* descParam, const HcclOneSideOpDesc* desc, u32 descNum, u64& transportDataAddr,
388 : u64& transportDataSize)
389 : {
390 0 : std::vector<u32> lkeys(descNum);
391 0 : std::vector<u32> rkeys(descNum);
392 0 : std::vector<HcclBuf> localMems(descNum);
393 0 : std::vector<HcclBuf> remoteMems(descNum);
394 0 : for (u32 i = 0; i < descNum - 1; ++i) { // last element is signal
395 0 : u32 unitSize = 0;
396 0 : CHK_RET(SalGetDataTypeSize(desc[i].dataType, unitSize));
397 0 : u64 bufSize = desc[i].count * unitSize;
398 0 : HCCL_DEBUG(
399 : "[HcclOneSidedConn][GetTransInfo] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i, desc[i].localAddr,
400 : desc[i].remoteAddr, bufSize);
401 :
402 0 : BufferKey<uintptr_t, u64> bufKey(reinterpret_cast<uintptr_t>(desc[i].remoteAddr), bufSize);
403 0 : auto rmaBuffer = remoteRmaBufferMgr_.Find(bufKey);
404 0 : CHK_PRT_RET(
405 : !rmaBuffer.first,
406 : HCCL_ERROR("[GetTransInfo] Can't find remoteBuffer by key[%p][%llu]", desc[i].remoteAddr, bufSize),
407 : HCCL_E_PARA);
408 :
409 0 : localMems[i] = {desc[i].localAddr, bufSize, nullptr};
410 0 : remoteMems[i] = {desc[i].remoteAddr, bufSize, rmaBuffer.second};
411 : }
412 0 : CHK_RET(transportMemPtr_->GetTransInfo(
413 : transportData_.qpInfo, lkeys.data(), rkeys.data(), localMems.data(), remoteMems.data(), descNum));
414 0 : CHK_RET(hrtMemSyncCopy(
415 : transportDataDevice_.ptr(), transportDataDevice_.size(), reinterpret_cast<void*>(&transportData_),
416 : sizeof(transportData_), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
417 0 : transportDataAddr = reinterpret_cast<u64>(transportDataDevice_.ptr());
418 0 : transportDataSize = transportDataDevice_.size();
419 0 : for (u32 i = 0; i < descNum - 1; ++i) {
420 0 : u32 unitSize = 0;
421 0 : CHK_RET(SalGetDataTypeSize(desc[i].dataType, unitSize));
422 0 : descParam[i].dataType = static_cast<u8>(desc[i].dataType);
423 0 : descParam[i].count = localMems[i].len / unitSize;
424 0 : descParam[i].localAddr = reinterpret_cast<u64>(localMems[i].addr);
425 0 : descParam[i].remoteAddr = reinterpret_cast<u64>(remoteMems[i].addr);
426 0 : descParam[i].lkey = lkeys[i];
427 0 : descParam[i].rkey = rkeys[i];
428 : }
429 0 : descParam[descNum - 1].dataType = static_cast<u8>(HcclDataType::HCCL_DATA_TYPE_UINT8);
430 0 : descParam[descNum - 1].count = localMems[descNum - 1].len; // 因为dataType是UINT8,所以count等于len
431 0 : descParam[descNum - 1].localAddr = reinterpret_cast<u64>(localMems[descNum - 1].addr);
432 0 : descParam[descNum - 1].remoteAddr = reinterpret_cast<u64>(remoteMems[descNum - 1].addr);
433 0 : descParam[descNum - 1].lkey = lkeys[descNum - 1];
434 0 : descParam[descNum - 1].rkey = rkeys[descNum - 1];
435 0 : return HCCL_SUCCESS;
436 0 : }
437 :
438 0 : HcclResult HcclOneSidedConn::WaitOpFence(const rtStream_t& stream)
439 : {
440 0 : CHK_RET(transportMemPtr_->WaitOpFence(stream));
441 0 : return HCCL_SUCCESS;
442 : }
443 :
444 : HcclResult
445 0 : HcclOneSidedConn::ConnectWithRemote(const std::string& commIdentifier, ProcessInfo localProcess, s32 timeoutSec)
446 : {
447 0 : CHK_RET(Connect(commIdentifier, timeoutSec));
448 0 : if (!useRdma_) {
449 0 : CHK_RET(ExchangeIpcProcessInfo(localProcess, remoteProcess_));
450 : }
451 0 : return HCCL_SUCCESS;
452 : }
453 :
454 0 : HcclResult HcclOneSidedConn::GetRemoteProcessInfo(ProcessInfo& remoteProcess)
455 : {
456 0 : remoteProcess = remoteProcess_;
457 0 : return HCCL_SUCCESS;
458 : }
459 :
460 0 : HcclResult HcclOneSidedConn::ExchangeMemDesc(const HcclMemDescs& localMemDescs)
461 : {
462 0 : constexpr u32 exchangeCntPerLoop = MAX_REMOTE_MEM_NUM;
463 0 : u32 localMemOffset = 0;
464 0 : u32 localMemCnt = localMemDescs.arrayLength;
465 0 : remoteMemDescsVec_.resize(exchangeCntPerLoop);
466 0 : actualNumOfRemote_ = 0;
467 :
468 : while (true) {
469 : // 每轮循环最多交换 exchangeCntPerLoop 个 memDesc
470 0 : u32 sendLocalCnt = localMemCnt > exchangeCntPerLoop ? exchangeCntPerLoop : localMemCnt;
471 0 : TransportMem::RmaMemDesc* localMemDescArray
472 0 : = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(localMemDescs.array)) + localMemOffset;
473 0 : TransportMem::RmaMemDescs localRmaMemDescs = {localMemDescArray, sendLocalCnt};
474 :
475 0 : if (remoteMemDescsVec_.size() - actualNumOfRemote_ < exchangeCntPerLoop) {
476 0 : remoteMemDescsVec_.resize(remoteMemDescsVec_.size() + exchangeCntPerLoop);
477 : }
478 : TransportMem::RmaMemDesc* remoteMemDescArray
479 0 : = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(&remoteMemDescsVec_[actualNumOfRemote_]));
480 0 : TransportMem::RmaMemDescs remoteRmaMemDescs = {remoteMemDescArray, exchangeCntPerLoop};
481 :
482 0 : u32 actualNumOfRemote = 0;
483 0 : CHK_RET(transportMemPtr_->ExchangeMemDesc(localRmaMemDescs, remoteRmaMemDescs, actualNumOfRemote));
484 0 : localMemOffset += sendLocalCnt;
485 0 : localMemCnt -= sendLocalCnt;
486 0 : actualNumOfRemote_ += actualNumOfRemote;
487 :
488 0 : if (actualNumOfRemote < exchangeCntPerLoop && sendLocalCnt < exchangeCntPerLoop) {
489 : // 循环结束条件,下轮没有memDesc要发 且 对端也没有memDesc要发
490 0 : break;
491 : }
492 0 : }
493 0 : return HCCL_SUCCESS;
494 : }
495 :
496 0 : HcclResult HcclOneSidedConn::EnableMemAccess()
497 : {
498 0 : CHK_PRT_RET(
499 : actualNumOfRemote_ > remoteMemDescsVec_.size(),
500 : HCCL_ERROR(
501 : "[HcclOneSidedConn][EnableMemAccess] actualNumOfRemote[%u] is larger than remoteMemDescsVec.size[%zu]",
502 : actualNumOfRemote_, remoteMemDescsVec_.size()),
503 : HCCL_E_INTERNAL);
504 :
505 : HcclMem remoteMem;
506 0 : for (u32 i = 0; i < actualNumOfRemote_; i++) {
507 : // 创建HcclMemDesc对象
508 0 : HcclMemDesc* remoteMemDesc = static_cast<HcclMemDesc*>(static_cast<void*>(&remoteMemDescsVec_.at(i)));
509 0 : this->EnableMemAccess(*remoteMemDesc, remoteMem);
510 : }
511 0 : return HCCL_SUCCESS;
512 : }
513 :
514 0 : HcclResult HcclOneSidedConn::DisableMemAccess()
515 : {
516 0 : CHK_PRT_RET(
517 : actualNumOfRemote_ > remoteMemDescsVec_.size(),
518 : HCCL_ERROR(
519 : "[HcclOneSidedConn][DisableMemAccess] actualNumOfRemote[%u] is larger than remoteMemDescsVec.size[%zu]",
520 : actualNumOfRemote_, remoteMemDescsVec_.size()),
521 : HCCL_E_INTERNAL);
522 :
523 0 : for (u32 i = 0; i < actualNumOfRemote_; i++) {
524 : // 创建HcclMemDesc对象
525 0 : HcclMemDesc* remoteMemDesc = static_cast<HcclMemDesc*>(static_cast<void*>(&remoteMemDescsVec_.at(i)));
526 0 : this->DisableMemAccess(*remoteMemDesc);
527 : }
528 :
529 0 : return HCCL_SUCCESS;
530 : }
531 : } // namespace hccl
|