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