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