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 : #include "ub_memory_transport.h"
11 : #include "exchange_ipc_buffer_dto.h"
12 : #include "coll_operator_check.h"
13 :
14 : namespace Hccl {
15 : constexpr u32 ONE_HUNDRED_MICROSECOND_OF_USLEEP = 100;
16 6 : UbMemoryTransport::UbMemoryTransport(const std::shared_ptr<Buffer> cclBuffer, const std::shared_ptr<Buffer> aivTagBuffer, const std::shared_ptr<Buffer> aivOffloadTagBuffer, Socket *socket, int32_t deviceLogicId)
17 42 : : cclBuffer(cclBuffer), aivTagBuffer(aivTagBuffer), aivOffloadTagBuffer(aivOffloadTagBuffer), socket(socket), deviceLogicId(deviceLogicId)
18 : {
19 6 : }
20 :
21 1 : HcclResult UbMemoryTransport::Init()
22 : {
23 1 : localBufferVec.push_back(make_unique<LocalIpcRmaBuffer>(cclBuffer));
24 1 : localBufferVec.push_back(make_unique<LocalIpcRmaBuffer>(aivTagBuffer));
25 1 : localBufferVec.push_back(make_unique<LocalIpcRmaBuffer>(aivOffloadTagBuffer));
26 1 : return HCCL_SUCCESS;
27 : }
28 :
29 2 : LocalIpcRmaBuffer *UbMemoryTransport::GetLocMemBuffer(const u32 bufIndex) const
30 : {
31 6 : HCCL_INFO("[%s] start", __func__);
32 2 : if (bufIndex >= localBufferVec.size()) {
33 3 : HCCL_ERROR("[%s] bufIndex[%u] is invalid, size[%u]", __func__, bufIndex, localBufferVec.size());
34 1 : THROW<InternalException>(
35 2 : StringFormat("[%s] bufIndex[%u] is invalid, size[%u]", __func__, bufIndex, localBufferVec.size()));
36 : }
37 1 : return localBufferVec[bufIndex].get();
38 : }
39 :
40 2 : RemoteIpcRmaBuffer *UbMemoryTransport::GetRmtMemBuffer(const u32 bufIndex) const
41 : {
42 6 : HCCL_INFO("[%s] start", __func__);
43 2 : if (bufIndex >= rmtBufferVec.size()) {
44 3 : HCCL_ERROR("[%s] bufIndex[%u] is invalid, size[%u]", __func__, bufIndex, rmtBufferVec.size());
45 1 : THROW<InternalException>(
46 2 : StringFormat("[%s] bufIndex[%u] is invalid, size[%u]", __func__, bufIndex, rmtBufferVec.size()));
47 : }
48 1 : return rmtBufferVec[bufIndex].get();
49 : }
50 :
51 :
52 10 : UbMemoryTransport::UBTransportStatus UbMemoryTransport::GetStatus()
53 : {
54 10 : UbMemoryTransport::UBTransportStatus status = UbMemoryTransport::UBTransportStatus::CONNECT_FAILED;
55 : try {
56 10 : status = StateMachine();
57 1 : } catch (HcclException &e) {
58 3 : HCCL_ERROR(e.what());
59 1 : return UbMemoryTransport::UBTransportStatus::CONNECT_FAILED;
60 1 : } catch (exception &e) {
61 0 : HCCL_ERROR(e.what());
62 0 : return UbMemoryTransport::UBTransportStatus::CONNECT_FAILED;
63 0 : } catch (...) {
64 0 : HCCL_ERROR("Unknown error occured when StateMachine!");
65 0 : return UbMemoryTransport::UBTransportStatus::CONNECT_FAILED;
66 0 : }
67 9 : return status;
68 : }
69 :
70 10 : UbMemoryTransport::UBTransportStatus UbMemoryTransport::StateMachine()
71 : {
72 10 : if (ubStatus == UBTransportStatus::READY) {
73 1 : return ubStatus;
74 : }
75 9 : SocketStatus socketStatus = socket->GetAsyncStatus();
76 9 : if (socketStatus == SocketStatus::INIT) {
77 1 : THROW<InternalException>("[UbMemoryTransport][GetStatus]socket timeout or no link, please check");
78 8 : } else if (socketStatus == SocketStatus::TIMEOUT) {
79 1 : return UBTransportStatus::SOCKET_TIMEOUT;
80 7 : } else if (socketStatus != SocketStatus::OK) {
81 0 : return ubStatus;
82 : }
83 7 : switch (ubStatus) {
84 1 : case UBTransportStatus::INIT:
85 1 : ubStatus = UBTransportStatus::SOCKET_OK;
86 1 : break;
87 1 : case UBTransportStatus::SOCKET_OK:
88 1 : SendMemInfo();
89 1 : ubStatus = UBTransportStatus::SEND_MEM_INFO;
90 1 : break;
91 1 : case UBTransportStatus::SEND_MEM_INFO:
92 1 : RecvMemInfo();
93 1 : ubStatus = UBTransportStatus::RECV_MEM_INFO;
94 1 : break;
95 1 : case UBTransportStatus::RECV_MEM_INFO:
96 1 : RecvMemProcess();
97 1 : ubStatus = UBTransportStatus::RECV_MEM_INFO_PROCESS;
98 1 : break;
99 : // 预留状态机,当前方案未使用
100 1 : case UBTransportStatus::RECV_MEM_INFO_PROCESS:
101 1 : ubStatus = UBTransportStatus::SEND_NAME;
102 1 : break;
103 1 : case UBTransportStatus::SEND_NAME:
104 1 : ubStatus = UBTransportStatus::RECV_NAME;
105 1 : break;
106 1 : case UBTransportStatus::RECV_NAME:
107 1 : ubStatus = UBTransportStatus::READY;
108 1 : break;
109 0 : default:
110 0 : break;
111 : }
112 7 : return ubStatus;
113 : }
114 :
115 0 : void UbMemoryTransport::SendMemInfo()
116 : {
117 0 : HCCL_INFO("[%s] start", __func__);
118 :
119 0 : BinaryStream binaryStream;
120 :
121 0 : HandshakeMsgPack(binaryStream);
122 0 : BufferPack(binaryStream);
123 :
124 0 : std::vector<char> data;
125 0 : binaryStream.Dump(data);
126 0 : socket->SendAsync(&data[0], data.size());
127 0 : exchangeDataSize = data.size();
128 0 : HCCL_INFO("[%s] finished", __func__);
129 0 : }
130 :
131 0 : void UbMemoryTransport::HandshakeMsgPack(BinaryStream &binaryStream)
132 : {
133 0 : binaryStream << static_cast<u32>(locOpAcceState);
134 0 : binaryStream << localHandshakeMsg;
135 0 : HCCL_INFO("[UbMemoryTransport][%s] start pack handshakeMsg", __func__);
136 0 : }
137 :
138 0 : void UbMemoryTransport::RecvMemInfo()
139 : {
140 0 : recvDataMsg.resize(exchangeDataSize);
141 0 : socket->RecvAsync(reinterpret_cast<u8 *>(&recvDataMsg[0]), recvDataMsg.size());
142 0 : HCCL_INFO("recv data, size=%llu, data=%s", recvDataMsg.size(), Bytes2hex(recvDataMsg.data(), recvDataMsg.size()).c_str());
143 0 : }
144 :
145 0 : void UbMemoryTransport::RecvMemProcess()
146 : {
147 0 : BinaryStream binaryStream(recvDataMsg);
148 0 : HandshakeMsgUnpack(binaryStream);
149 0 : RmtBufferUnpackProc(binaryStream);
150 0 : }
151 :
152 0 : void UbMemoryTransport::HandshakeMsgUnpack(BinaryStream &binaryStream)
153 : {
154 0 : u32 rmtAccelerator{0};
155 0 : binaryStream >> rmtAccelerator;
156 0 : HCCL_INFO("[UbMemoryTransport::HandshakeMsgUnpack], rmtAccelerator[%u]", rmtAccelerator);
157 0 : rmtOpAcceState = static_cast<AcceleratorState::Value>(rmtAccelerator);
158 :
159 0 : if (rmtOpAcceState != locOpAcceState) {
160 0 : THROW<InvalidParamsException>(
161 0 : StringFormat("[UbMemoryTransport::HandshakeMsgUnpack] Accelerator information check fail. "
162 : "locOpAccelerator[%s], rmtOpAccelerator[%s]",
163 0 : locOpAcceState.Describe().c_str(), rmtOpAcceState.Describe().c_str()));
164 : }
165 :
166 0 : rmtHandshakeMsg.clear();
167 0 : binaryStream >> rmtHandshakeMsg;
168 :
169 0 : if (localHandshakeMsg.size() != rmtHandshakeMsg.size()) {
170 0 : THROW<InvalidParamsException>(StringFormat("handshakeMsg size=%u is not equal to rmt=%u",
171 : localHandshakeMsg.size(), rmtHandshakeMsg.size()));
172 : }
173 :
174 0 : auto localCollOperator = CollOperator::GetPackedData(localHandshakeMsg);
175 0 : auto remoteCollOperator = CollOperator::GetPackedData(rmtHandshakeMsg);
176 0 : CheckCollOperator(localCollOperator, remoteCollOperator); // 两端算子参数一致性校验
177 :
178 0 : HCCL_INFO("[UbMemoryTransport][%s] start unpack handshakeMsg", __func__);
179 0 : }
180 :
181 0 : void UbMemoryTransport::BufferPack(BinaryStream &binaryStream)
182 : {
183 0 : u32 vecSize = localBufferVec.size();
184 0 : binaryStream << vecSize;
185 0 : for (auto &it : localBufferVec) {
186 0 : if (it != nullptr) { // 非空的buffer,从buffer中获取 dto
187 0 : std::unique_ptr<Serializable> dto = it->GetExchangeDto();
188 0 : HCCL_INFO("[%s] dto[%s]", __func__, dto->Describe().c_str());
189 0 : dto->Serialize(binaryStream);
190 0 : } else { // 空的buffer,dto所有字段为0(size=0)
191 0 : ExchangeIpcBufferDto exchangeDto;
192 0 : exchangeDto.Serialize(binaryStream);
193 0 : }
194 : }
195 0 : }
196 :
197 0 : void UbMemoryTransport::RmtBufferUnpackProc(BinaryStream &binaryStream)
198 : {
199 0 : rmtBufferVec.clear();
200 0 : rmtRmaBufferVec.clear();
201 0 : u32 vecSize{0};
202 0 : binaryStream >> vecSize;
203 0 : HCCL_INFO("vecSize=%u", vecSize);
204 0 : for (u32 pos = 0; pos < vecSize; ++pos) {
205 0 : ExchangeIpcBufferDto dto;
206 0 : dto.Deserialize(binaryStream);
207 0 : HCCL_INFO("[%s] dto[%s]", __func__, dto.Describe().c_str());
208 :
209 0 : if (dto.size == 0) { // size为0,则为 remote 空buffer
210 0 : HCCL_INFO("unpack nullptr, pos=%u", pos);
211 0 : rmtBufferVec.push_back(nullptr);
212 0 : rmtRmaBufferVec.push_back(nullptr);
213 : } else { // size非0,则构造一个remote buffer
214 0 : rmtBufferVec.push_back(make_unique<RemoteIpcRmaBuffer>(dto));
215 0 : rmtRmaBufferVec.push_back(rmtBufferVec.back().get());
216 : }
217 0 : }
218 0 : }
219 :
220 1 : std::string UbMemoryTransport::Describe() const
221 : {
222 1 : std::string description = "";
223 :
224 1 : description = StringFormat("deviceLogicId:%d,", deviceLogicId);
225 1 : description += StringFormat("UBTransportStatus:%u", static_cast<u32>(ubStatus));
226 1 : return description;
227 0 : }
228 :
229 : } // namespace Hccl
|