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