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