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 : #ifndef HCOMM_CCU_TRANSPORT_H
12 : #define HCOMM_CCU_TRANSPORT_H
13 :
14 : #include <memory>
15 : #include <vector>
16 : #include <shared_mutex>
17 :
18 : #include "../../../../../legacy/ascend950/unified_platform/resource/socket/socket.h"
19 : #include "env_config.h"
20 : #include "op_mode.h"
21 : #include "binary_stream.h"
22 : #include "ccu_conn.h"
23 : #include "remote_rma_buffer.h"
24 :
25 : namespace hcomm {
26 :
27 : class CcuTransport {
28 : public:
29 : // 缩减channel预留的cke、xn量,避免与ccu instance需求量冲突
30 : static constexpr uint32_t INIT_CKE_NUM = 4;
31 : static constexpr uint32_t INIT_XN_NUM = 4;
32 166 : MAKE_ENUM(
33 : TransStatus, INIT, SEND_DATA_SIZE, RECV_DATA_SIZE, SEND_ALL_INFO, RECV_ALL_INFO, SEND_TRANS_RES, RECV_TRANS_RES,
34 : SEND_FIN, RECV_FIN, RECVING_FIN, RECVING_TRANS_RES, READY, CONNECT_FAILED, SOCKET_TIMEOUT)
35 :
36 : enum class CcuResStatus : uint8_t {
37 : RES_UNKNOWN = 0, // 待确认(初始值)
38 : RES_OK = 1, // 资源充足
39 : RES_UNAVAIL = 2, // 资源不足
40 : RES_FAILED = 3, // 资源创建硬失败
41 : };
42 :
43 : struct CclBufferInfo {
44 : uint64_t addr{0};
45 : uint32_t size{0};
46 : uint32_t tokenId{0};
47 : uint32_t tokenValue{0};
48 : CommMemType type{COMM_MEM_TYPE_INVALID};
49 : std::array<char, HCCL_RES_TAG_MAX_LEN> memInfo{};
50 :
51 62 : explicit CclBufferInfo() = default;
52 49 : CclBufferInfo(const uint64_t addr, const uint32_t size, const uint32_t tokenId, const uint32_t tokenValue)
53 49 : : addr(addr),
54 49 : size(size),
55 49 : tokenId(tokenId),
56 49 : tokenValue(tokenValue)
57 49 : {}
58 :
59 16 : CclBufferInfo(
60 : const uint64_t addr, const uint32_t size, const uint32_t tokenId, const uint32_t tokenValue,
61 : const CommMemType type, const std::array<char, HCCL_RES_TAG_MAX_LEN>& memInfo)
62 16 : : addr(addr),
63 16 : size(size),
64 16 : tokenId(tokenId),
65 16 : tokenValue(tokenValue),
66 16 : type(type),
67 16 : memInfo(memInfo)
68 16 : {}
69 :
70 2 : void Pack(Hccl::BinaryStream& binaryStream) const
71 : {
72 2 : binaryStream << addr << size << tokenId << tokenValue << type;
73 : // 逐个字节传输
74 512 : for (uint32_t i = 0; i < HCCL_RES_TAG_MAX_LEN; ++i) {
75 510 : binaryStream << static_cast<u8>(memInfo[i]);
76 : }
77 2 : HCCL_INFO("Pack Ccl Buffer Info: addr[%llu] size[%u] memInfo[%s]", addr, size, memInfo.data());
78 2 : }
79 :
80 2 : void Unpack(Hccl::BinaryStream& binaryStream)
81 : {
82 2 : binaryStream >> addr >> size >> tokenId >> tokenValue >> type;
83 512 : for (uint32_t i = 0; i < HCCL_RES_TAG_MAX_LEN; ++i) {
84 : u8 byte;
85 510 : binaryStream >> byte;
86 510 : memInfo[i] = static_cast<char>(byte);
87 : }
88 4 : HCCL_INFO("Unpack Ccl Buffer Info: addr[%llu] size[%u] memInfo[%s]", addr, size, memInfo.data());
89 2 : }
90 : };
91 :
92 30 : MAKE_ENUM(CcuConnectionType, UBC_TP, UB_CTP, UBC_CTP = UB_CTP);
93 : struct CcuConnectionInfo {
94 : CcuConnectionType type{CcuConnectionType::UBC_TP};
95 : CommAddr locAddr{};
96 : CommAddr rmtAddr{};
97 : CcuChannelInfo channelInfo{};
98 : std::vector<CcuJetty*> ccuJettys{};
99 : uint32_t qos{EnvConfig::UB_QOS_DEFAULT};
100 :
101 : explicit CcuConnectionInfo() = default;
102 15 : CcuConnectionInfo(
103 : const CcuConnectionType type, const CommAddr& locAddr, const CommAddr& rmtAddr,
104 : const CcuChannelInfo& channelInfo, const std::vector<CcuJetty*>& ccuJettys,
105 : uint32_t qos = EnvConfig::UB_QOS_DEFAULT)
106 15 : : type(type),
107 15 : locAddr(locAddr),
108 15 : rmtAddr(rmtAddr),
109 15 : channelInfo(channelInfo),
110 15 : ccuJettys(ccuJettys),
111 15 : qos(qos)
112 15 : {}
113 : };
114 :
115 : CcuTransport(Hccl::Socket* socket, std::unique_ptr<CcuConnection>&& connection, const CclBufferInfo& locCclBufInfo);
116 : CcuTransport(
117 : Hccl::Socket* socket, std::unique_ptr<CcuConnection>&& connection,
118 : const std::vector<CclBufferInfo>& bufferInfos);
119 : CcuTransport(const CcuTransport& that) = delete;
120 : CcuTransport& operator=(const CcuTransport& other) = delete;
121 : ~CcuTransport();
122 : HcclResult Init();
123 : TransStatus GetStatus();
124 : void Clean();
125 : HcclResult GetRemoteMems(uint32_t* memNum, CommMem** remoteMem, char*** memInfos);
126 : HcclResult CheckSocketStatus();
127 : HcclResult UpdateMemInfo(std::vector<CcuTransport::CclBufferInfo>& bufferVecTemp);
128 : HcclResult ResUpdate(std::vector<std::string>& resGroupTags);
129 :
130 11 : CcuResStatus GetLocResStatus() const { return locResStatus_; }
131 33 : void SetLocResStatus(CcuResStatus status) { locResStatus_ = status; }
132 8 : bool IsLocResUnavailable() const { return locResStatus_ == CcuResStatus::RES_UNAVAIL; }
133 8 : bool IsRmtResUnavailable() const { return rmtResStatus_ == CcuResStatus::RES_UNAVAIL; }
134 : static HcclResult
135 : ConstructMsgOnlyTransport(Hccl::Socket* socket, std::unique_ptr<CcuTransport>& impl, CcuResStatus status);
136 :
137 : // 下面接口为平台层接口,不能在框架层使用
138 : uint32_t GetDieId() const;
139 : uint32_t GetChannelId() const;
140 : HcclResult GetLocCkeByIndex(const uint32_t index, uint32_t& locCkeId) const;
141 : HcclResult GetLocXnByIndex(const uint32_t index, uint32_t& locXnId) const;
142 : HcclResult GetRmtCkeByIndex(const uint32_t index, uint32_t& rmtCkeId) const;
143 : HcclResult GetRmtXnByIndex(const uint32_t index, uint32_t& rmtXnId) const;
144 : HcclResult GetRmtWishCntXnAddr(const std::string& resGroupTag, uint64_t& wishCntXnAddr) const;
145 : HcclResult GetLocBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
146 : HcclResult GetRmtBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
147 : HcclResult GetCkeNum(uint32_t& ckeNum) const;
148 : HcclResult GetRmtVarAddrByIndex(uint32_t index, uint64_t& rmtXnAddr) const;
149 : HcclResult GetRmtSignalAddrByIndex(uint32_t index, uint64_t& rmtCkeAddr) const;
150 : HcclResult GetRmtCcuBufferTokenInfo(uint32_t& rmtTokenId, uint32_t& rmtTokenValue) const;
151 : std::string Describe() const;
152 : HcclResult Describe(std::string& dfxMsg);
153 :
154 : public:
155 : struct Attribution {
156 : Hccl::OpMode opMode{Hccl::OpMode::OPBASE};
157 : u32 devicePhyId{0};
158 : std::vector<char> handshakeMsg{};
159 1 : std::string Describe() const
160 : {
161 : return Hccl::StringFormat(
162 1 : "CcuTransportAttribution[opMode=%s, devicePhyId=%u, handshakeMsg=%s]", opMode.Describe().c_str(),
163 2 : devicePhyId, Hccl::Bytes2hex(handshakeMsg.data(), handshakeMsg.size()).c_str());
164 : }
165 : };
166 :
167 2 : std::vector<char>& GetRmtHandshakeMsg() // 返回握手消息
168 : {
169 2 : return rmtHandshakeMsg_;
170 : }
171 :
172 1 : std::vector<char>& GetLocalHandshakeMsg() // 返回握手消息
173 : {
174 1 : return attr_.handshakeMsg;
175 : }
176 :
177 2 : void SetHandshakeMsg(const std::vector<char>& handshakeMsg) { attr_.handshakeMsg = handshakeMsg; }
178 :
179 : private:
180 : HcclResult StatusMachine();
181 : HcclResult AppendCkes(uint32_t ckesNum);
182 : HcclResult AppendXns(uint32_t xnsNum);
183 : HcclResult AppendCntXns();
184 : HcclResult SendFinish();
185 : HcclResult RecvFinish();
186 : HcclResult CheckFinish();
187 : HcclResult RecvDataProcess();
188 : HcclResult RecvTransInfoProcess();
189 : HcclResult ReleaseTransRes();
190 : HcclResult SendConnAndTransInfo();
191 : HcclResult RecvConnAndTransInfo();
192 : HcclResult SendDataSize();
193 : HcclResult RecvDataSize();
194 : HcclResult SendTransInfo();
195 : HcclResult RecvTransInfo();
196 : HcclResult HandshakeMsgPack(Hccl::BinaryStream& binaryStream);
197 : HcclResult ConnInfoPack(Hccl::BinaryStream& binaryStream) const;
198 : HcclResult TransResPack(Hccl::BinaryStream& binaryStream);
199 : HcclResult TransCntXnResPack(Hccl::BinaryStream& binaryStream);
200 : HcclResult BufferInfoPack(Hccl::BinaryStream& binaryStream, std::vector<CclBufferInfo>& bufferVec) const;
201 : HcclResult HandshakeMsgUnpack(Hccl::BinaryStream& binaryStream);
202 : HcclResult ConnInfoUnpackProc(Hccl::BinaryStream& binaryStream) const;
203 : HcclResult TransResUnpackProc(Hccl::BinaryStream& binaryStream);
204 : HcclResult TransCntXnResUnpackProc(Hccl::BinaryStream& binaryStream);
205 : HcclResult BufferInfoUnpack(Hccl::BinaryStream& binaryStream);
206 : HcclResult GetRmtVarAddrByXnId(const uint32_t rmtXnId, uint64_t& rmtXnAddr) const;
207 :
208 : HcclResult ReturnErrorStatus(const std::string& funcName);
209 :
210 : private:
211 : // 保存transport中需要使用的cke,xn等ccu资源
212 : struct TransRes {
213 : std::vector<uint32_t> ckes{};
214 : std::vector<uint32_t> xns{};
215 : std::map<std::string, uint32_t> cntXns{}; // {groupTag, wishCntXn}
216 : };
217 :
218 : uint32_t dieId_{0};
219 : int32_t devLogicId_{0};
220 : Attribution attr_{};
221 : std::vector<char> rmtHandshakeMsg_{0}; // 远端握手消息
222 : Hccl::Socket* socket_{nullptr};
223 : std::unique_ptr<CcuConnection> ccuConnection_;
224 : TransRes locRes_{};
225 : TransRes rmtRes_{};
226 : TransStatus transStatus_{TransStatus::INVALID};
227 : CcuResStatus locResStatus_{CcuResStatus::RES_UNKNOWN};
228 : CcuResStatus rmtResStatus_{CcuResStatus::RES_UNKNOWN};
229 : std::vector<std::vector<ResInfo>> ckesRes_{};
230 : std::vector<std::vector<ResInfo>> xnsRes_{};
231 : std::vector<CclBufferInfo> locBufferInfos_{};
232 : CclBufferInfo rmtHcclBufferInfo_{};
233 : std::vector<std::unique_ptr<Hccl::RemoteUbRmaBuffer>> rmtBufferVec_{};
234 : uint32_t exchangeDataSize_{0};
235 : std::vector<char> recvData_{};
236 : std::vector<char> recvTrans_{};
237 : std::vector<char> sendData_{};
238 : std::vector<char> sendTrans_{};
239 : std::vector<char> recvFinishMsg_{};
240 : std::vector<char> sendFinishMsg_{};
241 : bool cacheValid_ = false; // GetUserRemoteMem 的缓存标识
242 : std::mutex remoteMemsMutex_; // 远端内存列表互斥锁
243 : std::vector<CommMem> remoteUserMems_; // 内存基本信息缓存
244 : std::vector<std::string> memInfoCopies_; // 储存 Tag 字符串副本
245 : std::vector<char*> memInfoPointers_; // Tag 缓存
246 : };
247 :
248 : HcclResult BuildCcuConnection(
249 : const CcuTransport::CcuConnectionInfo& ccuConnectionInfo, std::unique_ptr<CcuConnection>& ccuConnection);
250 :
251 : HcclResult CcuCreateTransport(
252 : Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
253 : const CcuTransport::CclBufferInfo& cclBufferInfo, std::unique_ptr<CcuTransport>& ccuTransport);
254 :
255 : HcclResult CcuCreateTransport(
256 : Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
257 : const std::vector<CcuTransport::CclBufferInfo>& bufferInfos, std::unique_ptr<CcuTransport>& ccuTransport);
258 : } // namespace hcomm
259 : #endif // HCOMM_CCU_TRANSPORT_H
|