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 "symmetric_memory_agent.h"
12 : #include <chrono>
13 :
14 : namespace hccl {
15 : using namespace std;
16 :
17 : const string STR_IPC_MEM_EXCHANGE = "Exchange_Info";
18 : constexpr u32 USLEEP_ONE_THOUSAND = 1000;
19 : constexpr u32 RING_RANK_SIZE_MIN = 2;
20 :
21 64 : SymmetricMemoryAgent::SymmetricMemoryAgent(
22 : const std::unique_ptr<HcclSocketManager>& socketManager, u32 devicePhyId, s32 deviceLogicId,
23 : const HcclIpAddress& localVnicIp, const std::vector<RankInfo>& rankInfoList, u32 userRank, bool useSuperPodMode,
24 64 : const std::string& identifier)
25 64 : : socketManager_(socketManager),
26 64 : devicePhyId_(devicePhyId),
27 64 : deviceLogicId_(deviceLogicId),
28 64 : localVnicIp_(localVnicIp),
29 64 : rankInfoList_(rankInfoList),
30 64 : userRank_(userRank),
31 128 : rankSize_(rankInfoList.size()),
32 64 : useSuperPodMode_(useSuperPodMode),
33 64 : identifier_(identifier)
34 : {
35 64 : if (rankSize_ >= RING_RANK_SIZE_MIN) { // 当前数据交换算法使用超节点内大平面ring算法,需要和“左右”两边的rank建链
36 61 : leftRank_ = (userRank_ - 1 + rankSize_) % rankSize_;
37 61 : rightRank_ = (userRank_ + 1) % rankSize_;
38 : }
39 64 : }
40 :
41 64 : SymmetricMemoryAgent::~SymmetricMemoryAgent()
42 : {
43 64 : threadRun_ = false;
44 63 : if (recvThread_ && recvThread_->joinable()) {
45 9 : recvThread_->join();
46 9 : recvThread_ = nullptr;
47 : }
48 62 : if (vnicPortCtx_ != nullptr) {
49 8 : HcclNetCloseDev(vnicPortCtx_);
50 8 : vnicPortCtx_ = nullptr;
51 : }
52 62 : }
53 :
54 10 : HcclResult SymmetricMemoryAgent::Init()
55 : {
56 10 : CHK_PRT_RET(
57 : rankSize_ < RING_RANK_SIZE_MIN, HCCL_ERROR("[SymmetricMemoryAgent][Init] single rank communicator"),
58 : HCCL_E_PARA);
59 9 : CHK_RET(EstablishSockets());
60 9 : CHK_RET(InitRecvThread());
61 9 : return HCCL_SUCCESS;
62 : }
63 :
64 9 : HcclResult SymmetricMemoryAgent::InitRecvThread()
65 : {
66 9 : threadRun_ = true;
67 9 : recvThread_.reset(new (std::nothrow) std::thread(&SymmetricMemoryAgent::DealWithRequest, std::ref(*this)));
68 9 : CHK_SMART_PTR_NULL(recvThread_);
69 9 : return HCCL_SUCCESS;
70 : }
71 :
72 8 : HcclResult SymmetricMemoryAgent::EstablishSockets()
73 : {
74 8 : CHK_PRT_RET((vnicPortCtx_ != nullptr), HCCL_ERROR("[SymmetricMemoryAgent][Init] already initd"), HCCL_E_PARA);
75 8 : CHK_RET(HcclNetOpenDev(&vnicPortCtx_, NicType::VNIC_TYPE, devicePhyId_, deviceLogicId_, localVnicIp_));
76 8 : CHK_PTR_NULL(vnicPortCtx_);
77 :
78 8 : HCCL_INFO(
79 : "[SymmetricMemoryAgent][EstablishSockets] userRank[%u], leftRank_[%u], rightRank_[%u], rankSize_[%u]",
80 : userRank_, leftRank_, rightRank_, rankSize_);
81 29 : for (size_t i = 0; i < rankInfoList_.size(); i++) {
82 21 : if (rankInfoList_[i].userRank == leftRank_ || rankInfoList_[i].userRank == rightRank_) {
83 10 : HcclRankLinkInfo remoteLinkInfo;
84 10 : RankInfo dstRankInfo = rankInfoList_[i];
85 10 : remoteLinkInfo.userRank = dstRankInfo.userRank;
86 10 : remoteLinkInfo.devicePhyId = dstRankInfo.devicePhyId;
87 10 : remoteLinkInfo.ip = HcclIpAddress(dstRankInfo.devicePhyId);
88 10 : if (useSuperPodMode_) {
89 10 : CHK_RET(hrtRaGetSingleSocketVnicIpInfo(
90 : devicePhyId_, DeviceIdType::DEVICE_ID_TYPE_SDID, dstRankInfo.superDeviceId, remoteLinkInfo.ip));
91 : } else {
92 0 : CHK_RET(hrtRaGetSingleSocketVnicIpInfo(
93 : devicePhyId_, DeviceIdType::DEVICE_ID_TYPE_PHY_ID, dstRankInfo.devicePhyId, remoteLinkInfo.ip));
94 : }
95 : // 通信域未分配端口则使用默认端口
96 : remoteLinkInfo.port
97 10 : = dstRankInfo.deviceVnicPort == HCCL_INVALID_PORT ? HETEROG_CCL_PORT : dstRankInfo.deviceVnicPort;
98 10 : remoteLinkInfo.socketsPerLink = 1;
99 10 : string newTag = GenerateSocketTag(devicePhyId_, rankInfoList_[i].devicePhyId);
100 10 : std::vector<std::shared_ptr<HcclSocket>> tmpSockets;
101 : HcclResult ret
102 10 : = socketManager_->CreateSingleLinkSocket(newTag, vnicPortCtx_, remoteLinkInfo, tmpSockets, false, true);
103 10 : CHK_PRT_RET(
104 : ret != HCCL_SUCCESS,
105 : HCCL_ERROR(
106 : "[Create][DestSockets]Create single link sockets failed, "
107 : "local rank[%u], remote rank[%u]",
108 : userRank_, rankInfoList_[i].userRank),
109 : ret);
110 10 : if (tmpSockets.size() != 1) {
111 0 : HCCL_ERROR(
112 : "[SymmetricMemoryAgent][CreateVnic] socket number[%llu] is not 1 as expected!", tmpSockets.size());
113 0 : return HCCL_E_INTERNAL;
114 : }
115 : // 设置强制断链为关闭,避免进程退出时recv失败
116 10 : tmpSockets[0]->SetForceClose(false);
117 10 : mapRankIdconnectedSockets_[remoteLinkInfo.userRank] = (tmpSockets[0]);
118 10 : mapRankId2DevPhyId_[remoteLinkInfo.userRank] = remoteLinkInfo.devicePhyId;
119 10 : }
120 : }
121 :
122 18 : for (const auto& kv : mapRankIdconnectedSockets_) {
123 10 : CHK_PRT_RET(
124 : socketManager_->WaitLinkEstablish(kv.second) != HCCL_SUCCESS,
125 : HCCL_ERROR(
126 : "[SymmetricMemoryAgent][EstablishSockets] tag[%s] socket establish failed",
127 : kv.second->GetTag().c_str()),
128 : HCCL_E_INTERNAL);
129 : }
130 8 : return HCCL_SUCCESS;
131 : }
132 :
133 10 : std::string SymmetricMemoryAgent::GenerateSocketTag(u32 localRank, u32 remoteRank)
134 : {
135 10 : u32 small = localRank;
136 10 : u32 large = remoteRank;
137 :
138 10 : if (localRank > remoteRank) {
139 0 : small = remoteRank;
140 0 : large = localRank;
141 : }
142 :
143 : // Socket构造规则:前缀 + identifier + small + large
144 : std::string tag
145 10 : = STR_IPC_MEM_EXCHANGE + "_" + identifier_ + "_" + std::to_string(small) + ":" + std::to_string(large);
146 10 : return tag;
147 : }
148 :
149 10 : HcclResult SymmetricMemoryAgent::ExchangeInfo(void* inputPtr, void* outputPtr, u64 inputSize)
150 : {
151 10 : CHK_PTR_NULL(inputPtr);
152 9 : CHK_PTR_NULL(outputPtr);
153 8 : CHK_PRT_RET(inputSize == 0, HCCL_ERROR("Input size is 0"), HCCL_E_PARA);
154 : // 校验 inputSize 是否超过协议载荷上限
155 7 : CHK_PRT_RET(
156 : inputSize > PACKET_DATA_MAX_LEN,
157 : HCCL_ERROR("Input size %lu exceeds max payload %u", inputSize, PACKET_DATA_MAX_LEN), HCCL_E_PARA);
158 : // 校验是否建链成功
159 6 : CHK_PRT_RET(
160 : mapRankIdconnectedSockets_.find(rightRank_) == mapRankIdconnectedSockets_.end(),
161 : HCCL_ERROR("[ExchangeInfo] rightRank_%u socket not found in map", rightRank_), HCCL_E_INTERNAL);
162 4 : CHK_PRT_RET(
163 : mapRankIdconnectedSockets_.find(leftRank_) == mapRankIdconnectedSockets_.end(),
164 : HCCL_ERROR("[ExchangeInfo] leftRank_%u socket not found in map", leftRank_), HCCL_E_INTERNAL);
165 :
166 4 : HCCL_INFO(
167 : "[SymmetricMemoryAgent] start to ExchangeInfo, inputPtr[%p], outputPtr[%p], inputSize[%llu]", inputPtr,
168 : outputPtr, inputSize);
169 :
170 : // 重置本轮状态
171 4 : outputDataPtr_ = static_cast<u8*>(outputPtr);
172 4 : currentInputSize_ = inputSize; // 记录实际有效长度
173 4 : collectedCount_ = 0;
174 : // 本地数据处理:先把自己的一份拷到 Output 对应位置
175 4 : u8* selfDstPtr = outputDataPtr_ + (userRank_ * inputSize);
176 4 : CHK_SAFETY_FUNC_RET(memcpy_s(selfDstPtr, inputSize, inputPtr, inputSize));
177 4 : collectedCount_++;
178 :
179 : Packet dataPkt;
180 4 : dataPkt.type = MsgType::MSG_TYPE_DATA;
181 4 : dataPkt.rankId = userRank_;
182 4 : CHK_SAFETY_FUNC_RET(memset_s(dataPkt.data, PACKET_DATA_MAX_LEN, 0, PACKET_DATA_MAX_LEN));
183 4 : CHK_SAFETY_FUNC_RET(memcpy_s(dataPkt.data, PACKET_DATA_MAX_LEN, inputPtr, inputSize));
184 : {
185 4 : std::lock_guard<std::mutex> lock(queueMutex_);
186 4 : requestQueue_.push(dataPkt);
187 4 : }
188 4 : isProcessingTask_ = true;
189 :
190 4 : CHK_RET(WaitForCollectionComplete());
191 1 : HCCL_INFO("[SymmetricMemoryAgent] ExchangeInfo end");
192 1 : return HCCL_SUCCESS;
193 : }
194 :
195 4 : HcclResult SymmetricMemoryAgent::WaitForCollectionComplete()
196 : {
197 4 : std::unique_lock<std::mutex> lock(completionMutex_);
198 4 : auto timeout = std::chrono::seconds(GetExternalInputHcclLinkTimeOut());
199 4 : auto status = completionCv_.wait_for(lock, timeout);
200 4 : if (status == std::cv_status::timeout) {
201 6 : HCCL_ERROR("[SymmetricMemoryAgent] ExchangeInfo Timeout! Collected: %u/%u", collectedCount_.load(), rankSize_);
202 3 : return HCCL_E_TCP_TRANSFER;
203 : }
204 1 : return HCCL_SUCCESS;
205 4 : }
206 :
207 9 : void SymmetricMemoryAgent::DealWithRequest()
208 : {
209 9 : if (hrtSetDevice(deviceLogicId_) != HCCL_SUCCESS) {
210 0 : return;
211 : }
212 :
213 9 : std::vector<u8> leftRecvBuf(PACKET_TOTAL_LEN, 0);
214 9 : u32 leftRecvLen = 0;
215 :
216 1899 : while (threadRun_) {
217 1890 : if (isProcessingTask_) {
218 1883 : if (collectedCount_ < rankSize_) {
219 944 : u64 received = 0;
220 944 : std::unique_lock<std::mutex> lock(socketMutex_);
221 944 : HcclResult ret = mapRankIdconnectedSockets_[leftRank_]->IRecv(
222 944 : leftRecvBuf.data() + leftRecvLen, PACKET_TOTAL_LEN - leftRecvLen, received);
223 :
224 944 : CHK_PRT_CONT(
225 : (ret != HCCL_SUCCESS) && (ret != HCCL_E_AGAIN),
226 : HCCL_ERROR(
227 : "[SymmetricMemoryAgent][DealWithRequest] IRecv failed, ret[%d] remoteRank[%u] "
228 : "receivedSize[%llu]",
229 : ret, leftRank_, leftRecvLen));
230 :
231 944 : leftRecvLen += received;
232 944 : if (leftRecvLen == PACKET_TOTAL_LEN) {
233 4 : Packet* pkt = reinterpret_cast<Packet*>(leftRecvBuf.data());
234 4 : ProcessReceivedPacket(*pkt);
235 4 : leftRecvLen = 0;
236 : }
237 944 : }
238 1883 : std::lock_guard<std::mutex> lock(queueMutex_);
239 1883 : if (!requestQueue_.empty()) {
240 943 : Packet pkt = requestQueue_.front();
241 943 : std::unique_lock<std::mutex> sockLock(socketMutex_);
242 : HcclResult ret
243 943 : = mapRankIdconnectedSockets_[rightRank_]->Send(static_cast<void*>(&pkt), PACKET_TOTAL_LEN);
244 943 : if (ret == HCCL_SUCCESS) {
245 2 : requestQueue_.pop();
246 : } else {
247 941 : HCCL_ERROR(
248 : "[SymmetricMemoryAgent][DealWithRequest] Data(from rank[%u]) Send to rank[%u] failed.",
249 : pkt.rankId, rightRank_);
250 : }
251 943 : }
252 : // 检查是否完全结束, 退出条件: 数据全齐 && 队列空闲
253 1883 : if (requestQueue_.empty() && collectedCount_ == rankSize_) {
254 1 : std::unique_lock<std::mutex> lock(completionMutex_);
255 1 : HCCL_INFO("[SymmetricMemoryAgent] ExchangeInfo Complete.");
256 1 : isProcessingTask_ = false;
257 1 : completionCv_.notify_all();
258 1 : }
259 1883 : }
260 1890 : SaluSleep(USLEEP_ONE_THOUSAND);
261 : }
262 :
263 9 : hrtResetDevice(deviceLogicId_);
264 9 : }
265 :
266 4 : HcclResult SymmetricMemoryAgent::ProcessReceivedPacket(Packet& pkt)
267 : {
268 4 : if (pkt.rankId < rankSize_ && pkt.rankId != userRank_) {
269 4 : u8* dest = outputDataPtr_ + (pkt.rankId * currentInputSize_);
270 4 : CHK_SAFETY_FUNC_RET(memcpy_s(dest, currentInputSize_, pkt.data, currentInputSize_));
271 4 : collectedCount_++;
272 : }
273 8 : HCCL_INFO(
274 : "[SymmetricMemoryAgent][ProcessReceivedPacket] Data Recv from rank[%u]. Collected[%u / %u].", pkt.rankId,
275 : collectedCount_.load(), rankSize_);
276 : // Ring 转发逻辑:如果数据不是自己的,也不是右边Rank发出的(转了一圈),则转发给右边
277 4 : if (pkt.rankId != userRank_ && pkt.rankId != rightRank_) {
278 0 : std::lock_guard<std::mutex> lock(queueMutex_);
279 0 : requestQueue_.push(pkt);
280 0 : }
281 4 : return HCCL_SUCCESS;
282 : }
283 : } // namespace hccl
|