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