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 TRANSPORT_HETEROG_PUB_H
12 : #define TRANSPORT_HETEROG_PUB_H
13 : #include <unordered_map>
14 : #include <atomic>
15 : #include <functional>
16 : #include <vector>
17 : #include <queue>
18 : #include <stack>
19 : #include "memory_alloc_ring.h"
20 : #include "private_types.h"
21 : #include "local_ipc_notify.h"
22 :
23 : namespace hccl {
24 : enum class ConnState {
25 : CONN_STATE_IDLE,
26 : CONN_STATE_CONNECT_CHECK_SOCKET,
27 : CONN_STATE_GET_CHECK_SOCKET,
28 : CONN_STATE_SEND_CF,
29 : CONN_STATE_RECV_CF,
30 : CONN_STATE_CHECK_CF,
31 : CONN_STATE_CONNECT_ALL_SOCKET,
32 : CONN_STATE_GET_ALL_SOCKET,
33 : CONN_STATE_CONNECT_QP,
34 : CONN_STATE_GET_QP,
35 : CONN_STATE_SEND_STATUS,
36 : CONN_STATE_RECV_STATUS,
37 : CONN_STATE_FLUSH_QUEUE,
38 : CONN_STATE_GET_TAG_QP_ATTR,
39 : CONN_STATE_SEND_TAG_QP_INFO,
40 : CONN_STATE_RECV_TAG_QP_INFO,
41 : CONN_STATE_MODIFY_TAG_QP,
42 : CONN_STATE_GET_DATA_QP_ATTR,
43 : CONN_STATE_SEND_DATA_QP_INFO,
44 : CONN_STATE_RECV_DATA_QP_INFO,
45 : CONN_STATE_MODIFY_DATA_QP,
46 : CONN_STATE_COMPLETE
47 : };
48 :
49 : enum class RdmaNotifyOp { SEND_NOTIFY, RECV_NOTIFY, NUM };
50 : constexpr u32 HETEROG_MAX_FRAME_LEN = 128;
51 : struct InitInfo {
52 : s32 protocolType = 0; // 0:ROCE; 1:TCP
53 : u32 role = 0;
54 : u32 signal = 0;
55 : u8 checkFrame[HETEROG_MAX_FRAME_LEN] = {0};
56 : std::vector<SocketInfoT> socketInfo;
57 : std::vector<SocketConnectInfoT> socketConnInfo;
58 : };
59 :
60 : struct InitStateMachine {
61 : InitInfo locInitInfo;
62 : InitInfo remInitInfo;
63 : u64 size = 0;
64 : u64 completeSize = 0;
65 : u32 socketNum = 0;
66 : u32 completeNum = 0;
67 : };
68 :
69 : static constexpr u32 SYNC_SIGNAL = 0xFFFFFFFF;
70 : constexpr u32 HCCL_POLL_CQ_DEPTH = 32;
71 :
72 : struct TransportEndPointInfoHash {
73 0 : std::size_t operator()(const TransportEndPointInfo& t) const
74 : {
75 0 : return std::hash<u32>()(t.commId) ^ std::hash<u32>()(t.rank) ^ std::hash<u32>()(t.tag);
76 : }
77 : };
78 :
79 : using HcclReceivedEnvelope
80 : = std::unordered_map<TransportEndPointInfo, std::queue<HcclEnvelopeSummary>, TransportEndPointInfoHash>;
81 :
82 : class TransportHeterog {
83 : public:
84 : explicit TransportHeterog(
85 : const std::string& tag, HcclIpAddress& selfIp, HcclIpAddress& peerIp, u32 peerPort, u32 selfPort,
86 : const TransportResourceInfo& transportResourceInfo);
87 : explicit TransportHeterog(const TransportResourceInfo& transportResourceInfo);
88 : virtual ~TransportHeterog();
89 : virtual HcclResult Init() = 0;
90 : virtual HcclResult Init(u32 localUserRank, u32 remoteUserRank);
91 : virtual HcclResult Init(SocketInfoT& socketInfo, RdmaHandle rdmaHandle, MrHandle mrHandle);
92 : virtual HcclResult Deinit() = 0;
93 : virtual HcclResult
94 : Isend(const TransData& sendData, const TransportEndPointParam& epParam, HcclRequestInfo*& request)
95 : = 0;
96 : virtual HcclResult Send(const TransData& sendData, const TransportEndPointParam& epParam) = 0;
97 : virtual HcclResult
98 : Improbe(const TransportEndPointParam& epParam, s32& matched, HcclMessageInfo*& msg, HcclStatus& status)
99 : = 0;
100 : virtual HcclResult Imrecv(const TransData& recvData, HcclMessageInfo& msg, HcclRequestInfo*& request) = 0;
101 : virtual HcclResult Test(HcclRequestInfo& request, s32& flag, HcclStatus& compState) = 0;
102 : virtual HcclResult
103 : Improbe(const TransportEndPointParam& epParam, s32& matched, HcclMessageInfo*& msg, HcclStatus& status, bool& flag);
104 : virtual HcclResult
105 : Imrecv(const TransData& recvData, HcclMessageInfo& msg, HcclRequestInfo*& request, bool flag, bool needRecordFlag);
106 : virtual HcclResult ImrecvScatter(
107 : void* buf[], int count[], int bufCount, HcclDataType datatype, HcclMessageInfo& msg, HcclRequestInfo*& request);
108 : HcclResult SetDeviceIndex(s32 index);
109 : u32 GetRecvEnvelopNum();
110 : void AddRecvEnvelopNum();
111 : void SubRecvEnvelopNum();
112 : virtual HcclResult BlockSend(
113 : const TransData& sendData, const TransportEndPointParam& epParam, HcclRequestInfo*& request, s32 waitTimeOut);
114 : virtual HcclResult BlockRecv(
115 : const TransData& recvData, bool matched, TransportHeterog*& transport, s32 waitTimeOut, s32 waitPayloadTimeOut);
116 : HcclResult CheckAndPushBuildLink();
117 : HcclResult WaitBuildLinkComplete();
118 : virtual HcclResult Iwrite(const TransData& sendData, const HcclEnvelope& envelope, HcclRequestInfo*& request);
119 : virtual HcclResult GetRemoteIsendDoneSignal(std::shared_ptr<LocalIpcNotify>& signal);
120 : virtual HcclResult GetRemoteImrecvDoneSignal(std::shared_ptr<LocalIpcNotify>& signal);
121 :
122 : ConnState GetState();
123 : virtual void GetLinkTag(std::string& tag);
124 : void SetForceClose();
125 : HcclResult SocketSend(const FdHandle fdHandle, void* data, u64 size, u64& sentSize, bool& completed);
126 : HcclResult SocketRecv(const FdHandle fdHandle, void* data, u64 size, u64& recvSize, bool& completed);
127 : static void RecordRankTableCrc(const u32 crcValue);
128 :
129 : protected:
130 : HcclResult CheckRecvMsgAndRequestBuffer();
131 : HcclResult
132 : GenerateSendRequest(const TransData& sendData, const TransportEndPointParam& epParam, HcclRequestInfo*& request);
133 : HcclResult GenerateRecvRequest(const TransData& recvData, const HcclMessageInfo& msg, HcclRequestInfo*& request);
134 : HcclResult GenerateRecvScatterRequest(const HcclMessageInfo& msg, HcclRequestInfo*& request);
135 : HcclResult FreeRequest(HcclRequestInfo& request) const;
136 : HcclResult
137 : CheckTransportEndPointInfo(const TransportEndPointInfo& epInfo, const TransportEndPointInfo& epInfoCheck) const;
138 : HcclResult CheckRecvEnvelope(const TransData& recvDataCheck, const HcclEnvelopeSummary& envelope);
139 : HcclResult CheckRecvScatterEnvelope(
140 : void* buf[], int count[], int bufCount, HcclDataType datatype, const HcclEnvelopeSummary& envelope);
141 : HcclResult GenerateRecvMessage(HcclEnvelopeSummary& recvEnvelope, HcclMessageInfo*& msg, HcclStatus& status);
142 : HcclResult FreeRecvMessage(HcclMessageInfo& msg) const;
143 : HcclResult ProbeNothing(s32& flag, HcclMessageInfo*& msg, HcclStatus& status) const;
144 : HcclResult ConnectSocket(SocketConnectInfoT conn[], u32 num, bool& completed);
145 : HcclResult GetSocket(u32 role, struct SocketInfoT info[], u32 num, u32& connectedNum, bool& completed);
146 : HcclResult SocketClose();
147 : HcclResult CheckConsistentFrame();
148 : HcclResult ConnectAsync();
149 : HcclResult PrepareSocketInfo(s32 type, s32 linkNum, const std::string& clientTag, const std::string& serverTag);
150 : HcclResult InitTransportConnect(s32 type, s32 linkNum);
151 : HcclResult InitTransportConnect(s32 type, u32 role, s32 linkNum, u32 tag);
152 : HcclResult AddSocketWhiteList(std::string& tag);
153 : HcclResult TryTransition(HcclResult ret, bool completed, ConnState nextState);
154 : virtual HcclResult EnterStateProcess(ConnState nextState) = 0;
155 : virtual HcclResult LoopStateProcess() = 0;
156 :
157 : const std::string transTag_;
158 : SocketHandle nicSocketHandle_;
159 : HcclIpAddress selfIp_;
160 : HcclIpAddress peerIp_;
161 : u32 peerPort_;
162 : u32 selfPort_;
163 : struct InitStateMachine initSM_ {};
164 : std::atomic<ConnState> connState_{ConnState::CONN_STATE_IDLE};
165 : const std::unique_ptr<LocklessRingMemoryAllocate<HcclMessageInfo>>& pMsgInfosMem_;
166 : const std::unique_ptr<LocklessRingMemoryAllocate<HcclRequestInfo>>& pReqInfosMem_;
167 : s32 index_ = 0;
168 : u32 recvEnvelopNum_;
169 : bool isHdcMode_ = false;
170 : u32 localRank_ = 0;
171 : u32 remoteRank_ = 0;
172 : bool remoteIsHdc_ = false; // 连接对端为310时,为false
173 : bool isESMode_ = false;
174 : bool forceClose_ = false; // 设置socket batch close时是否为强制关闭,而非超时“优雅”关闭
175 :
176 : static std::atomic<u32> rankTableCrc_;
177 : };
178 : } // namespace hccl
179 : #endif
|