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