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