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 "hccl_one_sided_conn_v2.h"
12 : #include "connections_builder.h"
13 : #include "hccl_net_dev.h"
14 : #include "hccl_mem.h"
15 : #include "communicator_impl.h"
16 : #include "transport_urma_mem.h"
17 : namespace Hccl {
18 : using namespace std;
19 :
20 5 : HcclOneSidedConn::HcclOneSidedConn(CommunicatorImpl* comm, LinkData linkData) : comm_(comm), linkData_(linkData) {}
21 :
22 5 : HcclOneSidedConn::~HcclOneSidedConn()
23 : {
24 6 : for (const auto& pair : desc2netDevMap_) {
25 1 : const HcclNetDev& hcclNetDev = pair.second;
26 1 : HcclResult ret = HcclNetDevClose(hcclNetDev);
27 1 : if (ret != HCCL_SUCCESS) {
28 0 : HCCL_ERROR(
29 : "[HcclOneSidedConn][~HcclOneSidedConn]HcclNetDevClose failed, descStr[%s], ret[%d].",
30 : pair.first.c_str(), ret);
31 : }
32 : }
33 5 : }
34 1 : HcclResult HcclOneSidedConn::Connect([[maybe_unused]] const std::string& commId)
35 : {
36 3 : HCCL_INFO("[HcclOneSidedConn]Connect start");
37 :
38 : // Socket/RmaConnection建链
39 1 : vector<LinkData> links;
40 1 : links.push_back(linkData_);
41 1 : comm_->GetSocketManager().BatchCreateSockets(links);
42 1 : make_unique<ConnectionsBuilder>(*comm_)->BatchBuild(comm_->GetId(), links);
43 1 : comm_->GetMemTransportManager()->BatchBuildOneSidedTransports(links);
44 :
45 : // Transport粒度申请notify,aicpu76行那个
46 2 : for (auto& link : links) {
47 1 : comm_->GetConnLocalNotifyManager().ApplyFor(link.GetRemoteRankId(), link);
48 : }
49 :
50 : // 推动式建链
51 1 : WaitOneSidedTransportReady();
52 :
53 : // 保存socket
54 1 : SocketConfig socketConfig(linkData_.GetRemoteRankId(), linkData_, comm_->GetEstablishLinkSocketTag());
55 1 : socket_ = comm_->GetSocketManager().GetConnectedSocket(socketConfig);
56 1 : if (socket_ == nullptr) {
57 0 : HCCL_ERROR("[HcclOneSidedConn]socket_ is nullptr");
58 0 : return HCCL_E_PTR;
59 : }
60 :
61 : // 创建TransportUrmaMem并保存
62 1 : transportMemPtr_ = make_shared<TransportUrmaMem>(
63 1 : comm_->GetMemTransportManager()->GetOneSidedTransport(linkData_), remoteHcclBufMgr_);
64 1 : return HCCL_SUCCESS;
65 1 : }
66 :
67 1 : void HcclOneSidedConn::WaitOneSidedTransportReady()
68 : {
69 1 : auto timeout = std::chrono::seconds(EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
70 1 : HcclUs startTime = std::chrono::steady_clock::now();
71 : while (true) {
72 1 : if (comm_->GetMemTransportManager()->IsAllOneSidedTransportReady()) {
73 1 : break;
74 : }
75 0 : if ((std::chrono::steady_clock::now() - startTime) >= timeout) {
76 0 : RPT_INPUT_ERR(
77 : true, "EI0006", std::vector<std::string>({"reason"}),
78 : std::vector<std::string>({"WaitOneSidedTransportReady timeout, SOCKET_TIMEOUT."}));
79 0 : THROW<InternalException>("WaitOneSidedTransportReady timeout.");
80 : }
81 : }
82 1 : }
83 :
84 1 : HcclResult HcclOneSidedConn::SendLocalMemDesc(const HcclMemDescs& localMemDescs)
85 : {
86 3 : HCCL_INFO("[HcclOneSidedConn]SendLocalMemDesc start");
87 1 : socket_->Send((u8*)(&localMemDescs.arrayLength), sizeof(u32));
88 3 : HCCL_INFO("send localMemDescs.arrayLength:%u", localMemDescs.arrayLength);
89 1 : if (localMemDescs.arrayLength == 0) {
90 0 : HCCL_INFO("localMemDescs.arrayLength[%u], no need to send data", localMemDescs.arrayLength);
91 : } else {
92 3 : HCCL_INFO("send descSize:%u", localMemDescs.arrayLength * sizeof(HcclMemDesc));
93 1 : if (static_cast<u64>(localMemDescs.arrayLength) > static_cast<u64>(UINT32_MAX) / sizeof(HcclMemDesc)) {
94 0 : THROW<InternalException>("integer overflow occurs");
95 : }
96 1 : socket_->Send((u8*)(localMemDescs.array), localMemDescs.arrayLength * sizeof(HcclMemDesc));
97 : }
98 1 : return HCCL_SUCCESS;
99 : }
100 :
101 2 : HcclResult HcclOneSidedConn::ReceiveRemoteMemDesc(HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
102 : {
103 6 : HCCL_INFO("[HcclOneSidedConn]ReceiveRemoteMemDesc start");
104 2 : socket_->Recv((u8*)(&actualNumOfRemote), sizeof(u32));
105 2 : remoteMemDescs.arrayLength = actualNumOfRemote;
106 6 : HCCL_INFO("receive actualNumOfRemote:%u", actualNumOfRemote);
107 2 : if (remoteMemDescs.arrayLength == 0) {
108 3 : HCCL_INFO("actualNumOfRemote[%u], no need to receive data", remoteMemDescs.arrayLength);
109 : } else {
110 1 : if (remoteMemDescs.array == nullptr) {
111 3 : HCCL_ERROR(
112 : "[HcclOneSidedConn]remoteMemDescs.array is nullptr but actualNumOfRemote[%u] > 0", actualNumOfRemote);
113 1 : return HCCL_E_PTR;
114 : }
115 0 : HCCL_INFO("receive descSize:%u", actualNumOfRemote * sizeof(HcclMemDesc));
116 0 : socket_->Recv((u8*)remoteMemDescs.array, actualNumOfRemote * sizeof(HcclMemDesc));
117 : }
118 1 : return HCCL_SUCCESS;
119 : }
120 :
121 1 : HcclResult HcclOneSidedConn::ExchangeMemDesc(
122 : const HcclMemDescs& localMemDescs, HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
123 : {
124 3 : HCCL_INFO("[HcclOneSidedConn]ExchangeMemDesc start");
125 1 : CHK_PRT_RET(
126 : (localMemDescs.array == nullptr) && (remoteMemDescs.array == nullptr),
127 : HCCL_ERROR(
128 : "[HcclOneSidedConn]localMemDesc array and remoteMemDesc array are both nullptr, do not need to exchange"),
129 : HCCL_E_PARA);
130 1 : CHK_PRT_RET(
131 : (localMemDescs.arrayLength == 0) && (remoteMemDescs.arrayLength == 0),
132 : HCCL_ERROR(
133 : "[HcclOneSidedConn]localMemDesc arrayLength = %u and remoteMemDescs arrayLength = %u , do "
134 : "not need to exchange",
135 : localMemDescs.arrayLength, remoteMemDescs.arrayLength),
136 : HCCL_E_PARA);
137 :
138 1 : if (socket_ == nullptr) {
139 1 : CHK_RET(Connect(comm_->GetId()));
140 : }
141 :
142 1 : if (socket_->GetRole() == SocketRole::CLIENT) {
143 : // 先收后发
144 0 : CHK_RET(ReceiveRemoteMemDesc(remoteMemDescs, actualNumOfRemote));
145 0 : CHK_RET(SendLocalMemDesc(localMemDescs));
146 : } else {
147 : // 先发后收
148 1 : CHK_RET(SendLocalMemDesc(localMemDescs));
149 1 : CHK_RET(ReceiveRemoteMemDesc(remoteMemDescs, actualNumOfRemote));
150 : }
151 :
152 : // 校验remoteDescs中的remoteRankId和conn对象中保存的localRankId是否一样
153 1 : for (u32 i = 0; i < actualNumOfRemote; i++) {
154 0 : CHK_PTR_NULL(remoteMemDescs.array);
155 0 : const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDescs.array[i].desc);
156 0 : CHK_PTR_NULL(remoteRmaMemDesc);
157 0 : RankId tempRankId = remoteRmaMemDesc->remoteRankId;
158 0 : HCCL_INFO("[TransportMem][ExchangeMemDesc]tempRankId:%u, userRank:%u", tempRankId, comm_->GetMyRank());
159 0 : if (tempRankId != comm_->GetMyRank()) {
160 0 : HCCL_ERROR(
161 : "[TransportMem][ExchangeMemDesc]localRank[%u] receive remoteMemDesc from wrong localRank[%u], "
162 : "connection is for localRank[%u]",
163 : comm_->GetMyRank(), tempRankId, comm_->GetMyRank());
164 0 : return HCCL_E_INTERNAL;
165 : }
166 : }
167 :
168 1 : return HCCL_SUCCESS;
169 : }
170 :
171 2 : HcclResult HcclOneSidedConn::EnableMemAccess(const HcclMemDesc& remoteMemDesc, HcclMem& remoteMem)
172 : {
173 6 : HCCL_INFO("[HcclOneSidedConn]EnableMemAccess start");
174 : // 反序列化remoteMemDesc
175 2 : const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDesc.desc);
176 2 : std::vector<char> tempDesc(TRANSPORT_EMD_ESC_SIZE);
177 2 : tempDesc.assign(remoteRmaMemDesc->memDesc, remoteRmaMemDesc->memDesc + TRANSPORT_EMD_ESC_SIZE);
178 2 : ExchangeUbBufferDto dto;
179 2 : BinaryStream remoteRdmaRmaBufferStream(tempDesc);
180 2 : dto.Deserialize(remoteRdmaRmaBufferStream);
181 :
182 : // 导入内存描述符
183 2 : shared_ptr<HcclBuf> outBuf = make_shared<HcclBuf>();
184 2 : string tempStr = string(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE);
185 2 : auto iter = desc2HcclBufMapRemoteUb_.find(tempStr);
186 2 : if (iter != desc2HcclBufMapRemoteUb_.end()) {
187 0 : outBuf = iter->second;
188 : } else {
189 2 : HcclNetDevInfos info;
190 2 : info.addr.protoType = HcclNetDevice::ConvertHcclProtoToLinkProto(linkData_.GetLocalPort().GetProto());
191 2 : info.addr.type = HCCL_ADDR_TYPE_IP_V4;
192 2 : info.netdevDeployment = HcclNetDevice::ConvertDeploymentType(linkData_.GetLocalPort().GetType());
193 2 : info.devicePhyId = comm_->GetDevicePhyId();
194 2 : info.addr.addr = linkData_.GetLocalPort().GetAddr().GetBinaryAddress().addr;
195 : HcclNetDev netDev;
196 2 : HcclResult ret = HcclNetDevOpen(&info, &netDev);
197 2 : if (ret != HCCL_SUCCESS) {
198 3 : HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]HcclNetDevOpen failed, ret[%d]", ret);
199 1 : return ret;
200 : }
201 1 : desc2netDevMap_.emplace(tempStr, netDev);
202 1 : ret = HcclMemImport(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE, true, outBuf.get(), netDev);
203 1 : if (ret != HCCL_SUCCESS) {
204 0 : HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]EnableMemAccess failed, ret [%d]", ret);
205 0 : return ret;
206 : }
207 : }
208 :
209 : // 填充remoteMem
210 1 : remoteMem.type = static_cast<HcclMemType>(dto.memType);
211 1 : remoteMem.addr = outBuf->addr;
212 1 : remoteMem.size = outBuf->len;
213 :
214 : // 添加计数器
215 1 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf->addr), outBuf->len);
216 1 : auto resultPair = remoteHcclBufMgr_.Add(tempKey, outBuf);
217 1 : if (resultPair.first == remoteHcclBufMgr_.End()) {
218 0 : HCCL_ERROR("[HcclOneSidedConn][EnableMemAccess]The memory overlaps with the memory has been enabled");
219 0 : return HCCL_E_INTERNAL;
220 : }
221 :
222 : // 存储remoteHcclBuf
223 1 : desc2HcclBufMapRemoteUb_.emplace(tempStr, outBuf);
224 3 : HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access success.");
225 1 : return HCCL_SUCCESS;
226 2 : }
227 :
228 3 : HcclResult HcclOneSidedConn::DisableMemAccess(const HcclMemDesc& remoteMemDesc)
229 : {
230 9 : HCCL_INFO("[HcclOneSidedConn]DisableMemAccess start");
231 : // 将HcclMemDesc转化为RmaMemDesc
232 3 : const RmaMemDesc* remoteRmaMemDesc = reinterpret_cast<const RmaMemDesc*>(remoteMemDesc.desc);
233 3 : string tempStr = string(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE);
234 3 : auto it = desc2HcclBufMapRemoteUb_.find(tempStr);
235 3 : if (it == desc2HcclBufMapRemoteUb_.end()) {
236 6 : HCCL_ERROR("[HcclOneSidedConn][DisableMemAccess]Can't find hcclmem by key.");
237 2 : return HCCL_E_INTERNAL;
238 : }
239 :
240 : // 计数器删除HcclBuf
241 1 : HcclBuf* buf = it->second.get();
242 1 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(it->second->addr), it->second->len);
243 : // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
244 : // 删除失败:输入key是表中某一最相近key的全集,计数-1后不为0(说明存在其他remoteRank使用),返回false
245 1 : auto resultPair = remoteHcclBufMgr_.Del(tempKey);
246 1 : if (resultPair) {
247 1 : HcclResult ret = HcclMemClose(buf);
248 1 : if (ret != HCCL_SUCCESS) {
249 0 : HCCL_ERROR("[HcclOneSidedConn][DisableMemAccess]Close remote memory failed. ret[%d]", ret);
250 0 : return HCCL_E_INTERNAL;
251 : }
252 3 : desc2HcclBufMapRemoteUb_.erase(remoteRmaMemDesc->memDesc);
253 : }
254 3 : HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] Disable memory access success.");
255 1 : return HCCL_SUCCESS;
256 3 : }
257 :
258 2 : HcclResult HcclOneSidedConn::BatchBufferSlice(
259 : const HcclOneSideOpDesc* oneSideDescs, u32 descNum,
260 : vector<HcclAicpuLocBufLite>& hostBatchPutGetLocalBufferSliceBufs,
261 : vector<HcclAicpuLocBufLite>& hostBatchPutGetRemoteBufferSliceBufs)
262 : {
263 4 : RmaBufferSlice localRmaBufferSlice[descNum] = {};
264 4 : RmtRmaBufferSlice remoteRmaBufferSlice[descNum] = {};
265 :
266 2 : if (transportMemPtr_ != nullptr) {
267 2 : CHK_RET(transportMemPtr_->BatchBufferSlice(oneSideDescs, descNum, localRmaBufferSlice, remoteRmaBufferSlice));
268 : } else {
269 0 : THROW<InternalException>("transportMemPtr is nullptr");
270 : }
271 :
272 4 : for (u32 i = 0; i < descNum; i++) {
273 2 : hostBatchPutGetLocalBufferSliceBufs[i].addr = localRmaBufferSlice[i].addr;
274 2 : hostBatchPutGetLocalBufferSliceBufs[i].size = localRmaBufferSlice[i].size;
275 2 : hostBatchPutGetLocalBufferSliceBufs[i].tokenId
276 2 : = static_cast<LocalUbRmaBuffer*>(localRmaBufferSlice[i].buf)->GetTokenId();
277 2 : hostBatchPutGetLocalBufferSliceBufs[i].tokenValue
278 2 : = static_cast<LocalUbRmaBuffer*>(localRmaBufferSlice[i].buf)->GetTokenValue();
279 6 : HCCL_INFO(
280 : "hostBatchPutGetLocalBufferSliceBufs, addr=0x%llx, size=0x%llx",
281 : hostBatchPutGetLocalBufferSliceBufs[i].addr, hostBatchPutGetLocalBufferSliceBufs[i].size);
282 :
283 2 : hostBatchPutGetRemoteBufferSliceBufs[i].addr = remoteRmaBufferSlice[i].addr;
284 2 : hostBatchPutGetRemoteBufferSliceBufs[i].size = remoteRmaBufferSlice[i].size;
285 2 : hostBatchPutGetRemoteBufferSliceBufs[i].tokenId
286 2 : = static_cast<RemoteUbRmaBuffer*>(remoteRmaBufferSlice[i].buf)->GetTokenId();
287 2 : hostBatchPutGetRemoteBufferSliceBufs[i].tokenValue
288 2 : = static_cast<RemoteUbRmaBuffer*>(remoteRmaBufferSlice[i].buf)->GetTokenValue();
289 6 : HCCL_INFO(
290 : "hostBatchPutGetRemoteBufferSliceBufs, addr=0x%llx, size=0x%llx",
291 : hostBatchPutGetRemoteBufferSliceBufs[i].addr, hostBatchPutGetRemoteBufferSliceBufs[i].size);
292 : }
293 2 : return HCCL_SUCCESS;
294 2 : }
295 : } // namespace Hccl
|