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 175 : 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 34 : 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 0 : bool IsBothEndResSufficient() const
135 : {
136 0 : return locResStatus_ == CcuResStatus::RES_OK && rmtResStatus_ == CcuResStatus::RES_OK;
137 : }
138 : static HcclResult
139 : ConstructMsgOnlyTransport(Hccl::Socket* socket, std::unique_ptr<CcuTransport>& impl, CcuResStatus status);
140 :
141 : // 下面接口为平台层接口,不能在框架层使用
142 : uint32_t GetDieId() const;
143 : uint32_t GetChannelId() const;
144 : HcclResult GetLocCkeByIndex(const uint32_t index, uint32_t& locCkeId) const;
145 : HcclResult GetLocXnByIndex(const uint32_t index, uint32_t& locXnId) const;
146 : HcclResult GetRmtCkeByIndex(const uint32_t index, uint32_t& rmtCkeId) const;
147 : HcclResult GetRmtXnByIndex(const uint32_t index, uint32_t& rmtXnId) const;
148 : HcclResult GetRmtWishCntXnAddr(const std::string& resGroupTag, uint64_t& wishCntXnAddr) const;
149 : HcclResult GetLocBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
150 : HcclResult GetRmtBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
151 : HcclResult GetCkeNum(uint32_t& ckeNum) const;
152 : HcclResult GetRmtVarAddrByIndex(uint32_t index, uint64_t& rmtXnAddr) const;
153 : HcclResult GetRmtSignalAddrByIndex(uint32_t index, uint64_t& rmtCkeAddr) const;
154 : HcclResult GetRmtCcuBufferTokenInfo(uint32_t& rmtTokenId, uint32_t& rmtTokenValue) const;
155 : std::string Describe() const;
156 : HcclResult Describe(std::string& dfxMsg);
157 :
158 : public:
159 : struct Attribution {
160 : Hccl::OpMode opMode{Hccl::OpMode::OPBASE};
161 : u32 devicePhyId{0};
162 : std::vector<char> handshakeMsg{};
163 1 : std::string Describe() const
164 : {
165 : return Hccl::StringFormat(
166 1 : "CcuTransportAttribution[opMode=%s, devicePhyId=%u, handshakeMsg=%s]", opMode.Describe().c_str(),
167 2 : devicePhyId, Hccl::Bytes2hex(handshakeMsg.data(), handshakeMsg.size()).c_str());
168 : }
169 : };
170 :
171 2 : std::vector<char>& GetRmtHandshakeMsg() // 返回握手消息
172 : {
173 2 : return rmtHandshakeMsg_;
174 : }
175 :
176 1 : std::vector<char>& GetLocalHandshakeMsg() // 返回握手消息
177 : {
178 1 : return attr_.handshakeMsg;
179 : }
180 :
181 2 : void SetHandshakeMsg(const std::vector<char>& handshakeMsg) { attr_.handshakeMsg = handshakeMsg; }
182 :
183 : private:
184 : HcclResult StatusMachine();
185 : HcclResult AppendCkes(uint32_t ckesNum);
186 : HcclResult AppendXns(uint32_t xnsNum);
187 : HcclResult AppendCntXns();
188 : HcclResult SendFinish();
189 : HcclResult RecvFinish();
190 : HcclResult CheckFinish();
191 : HcclResult RecvDataProcess();
192 : HcclResult RecvTransInfoProcess();
193 : HcclResult ReleaseTransRes();
194 : HcclResult SendConnAndTransInfo();
195 : HcclResult RecvConnAndTransInfo();
196 : HcclResult SendDataSize();
197 : HcclResult RecvDataSize();
198 : HcclResult SendTransInfo();
199 : HcclResult RecvTransInfo();
200 : HcclResult HandshakeMsgPack(Hccl::BinaryStream& binaryStream);
201 : HcclResult ConnInfoPack(Hccl::BinaryStream& binaryStream) const;
202 : HcclResult TransResPack(Hccl::BinaryStream& binaryStream);
203 : HcclResult TransCntXnResPack(Hccl::BinaryStream& binaryStream);
204 : HcclResult BufferInfoPack(Hccl::BinaryStream& binaryStream, std::vector<CclBufferInfo>& bufferVec) const;
205 : HcclResult HandshakeMsgUnpack(Hccl::BinaryStream& binaryStream);
206 : HcclResult ConnInfoUnpackProc(Hccl::BinaryStream& binaryStream) const;
207 : HcclResult TransResUnpackProc(Hccl::BinaryStream& binaryStream);
208 : HcclResult TransCntXnResUnpackProc(Hccl::BinaryStream& binaryStream);
209 : HcclResult BufferInfoUnpack(Hccl::BinaryStream& binaryStream);
210 : HcclResult GetRmtVarAddrByXnId(const uint32_t rmtXnId, uint64_t& rmtXnAddr) const;
211 :
212 : HcclResult ReturnErrorStatus(const std::string& funcName);
213 :
214 : private:
215 : // 保存transport中需要使用的cke,xn等ccu资源
216 : struct TransRes {
217 : std::vector<uint32_t> ckes{};
218 : std::vector<uint32_t> xns{};
219 : std::map<std::string, uint32_t> cntXns{}; // {groupTag, wishCntXn}
220 : };
221 :
222 : uint32_t dieId_{0};
223 : int32_t devLogicId_{0};
224 : Attribution attr_{};
225 : std::vector<char> rmtHandshakeMsg_{0}; // 远端握手消息
226 : Hccl::Socket* socket_{nullptr};
227 : std::unique_ptr<CcuConnection> ccuConnection_;
228 : TransRes locRes_{};
229 : TransRes rmtRes_{};
230 : TransStatus transStatus_{TransStatus::INVALID};
231 : CcuResStatus locResStatus_{CcuResStatus::RES_UNKNOWN};
232 : CcuResStatus rmtResStatus_{CcuResStatus::RES_UNKNOWN};
233 : std::vector<std::vector<ResInfo>> ckesRes_{};
234 : std::vector<std::vector<ResInfo>> xnsRes_{};
235 : std::vector<CclBufferInfo> locBufferInfos_{};
236 : CclBufferInfo rmtHcclBufferInfo_{};
237 : std::vector<std::unique_ptr<Hccl::RemoteUbRmaBuffer>> rmtBufferVec_{};
238 : uint32_t exchangeDataSize_{0};
239 : std::vector<char> recvData_{};
240 : std::vector<char> recvTrans_{};
241 : std::vector<char> sendData_{};
242 : std::vector<char> sendTrans_{};
243 : std::vector<char> recvFinishMsg_{};
244 : std::vector<char> sendFinishMsg_{};
245 : bool cacheValid_ = false; // GetUserRemoteMem 的缓存标识
246 : std::mutex remoteMemsMutex_; // 远端内存列表互斥锁
247 : std::vector<CommMem> remoteUserMems_; // 内存基本信息缓存
248 : std::vector<std::string> memInfoCopies_; // 储存 Tag 字符串副本
249 : std::vector<char*> memInfoPointers_; // Tag 缓存
250 : };
251 :
252 : HcclResult BuildCcuConnection(
253 : const CcuTransport::CcuConnectionInfo& ccuConnectionInfo, std::unique_ptr<CcuConnection>& ccuConnection);
254 :
255 : HcclResult CcuCreateTransport(
256 : Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
257 : const CcuTransport::CclBufferInfo& cclBufferInfo, std::unique_ptr<CcuTransport>& ccuTransport);
258 :
259 : HcclResult CcuCreateTransport(
260 : Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
261 : const std::vector<CcuTransport::CclBufferInfo>& bufferInfos, std::unique_ptr<CcuTransport>& ccuTransport);
262 : } // namespace hcomm
263 : #endif // HCOMM_CCU_TRANSPORT_H
|