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 "remote_access_impl.h"
12 : #include <algorithm>
13 :
14 : #include "transport_remote_access.h"
15 :
16 : namespace hccl {
17 : using namespace std;
18 :
19 0 : RemoteAccessImpl::RemoteAccessImpl() : userRank_(0), userRankNum_(0), serverNum_(0), rankNumPerServer_(0) {}
20 :
21 0 : RemoteAccessImpl::~RemoteAccessImpl() {}
22 :
23 0 : HcclResult RemoteAccessImpl::Init(u32 rank, const vector<MemRegisterAddr>& addrInfos, const RmaRankTable& rankTable)
24 : {
25 0 : HCCL_INFO("RemoteAccessImpl init start");
26 :
27 0 : userRank_ = rank;
28 0 : userRankNum_ = rankTable.rankNum;
29 0 : serverNum_ = rankTable.serverNum;
30 0 : if (serverNum_ == 0) {
31 0 : HCCL_ERROR("[RemoteAccessImpl][Init]errNo[0x%016llx] server num is zero", HCOM_ERROR_CODE(HCCL_E_PARA));
32 0 : return HCCL_E_PARA;
33 : }
34 :
35 0 : rankNumPerServer_ = userRankNum_ / serverNum_;
36 0 : HCCL_INFO(
37 : "RemoteAccessImpl Init userRank_[%u] userRankNum_[%u] serverNum_[%u] rankNumPerServer_[%u]", userRank_,
38 : userRankNum_, serverNum_, rankNumPerServer_);
39 :
40 0 : u32 rankInComm = userRank_ / rankNumPerServer_;
41 0 : CHK_PRT_RET(
42 : rankTable.deviceIps.empty(), HCCL_ERROR("[Init][RemoteAccessImpl]rankTable.rankList is empty"), HCCL_E_PARA);
43 0 : u32 devicePhyId = rankTable.devicePhyId;
44 0 : std::map<u32, std::vector<HcclIpAddress>> rankInfo; // rankIdInComm - deviceIp
45 0 : for (u32 rankIndex = 0; rankIndex < userRankNum_; rankIndex++) {
46 0 : if ((rankIndex % rankNumPerServer_) == (userRank_ % rankNumPerServer_)) { // 在同一平面
47 0 : u32 curRank = rankIndex / rankNumPerServer_; // 通信域内的第几个rank
48 0 : rankInfo.insert(std::pair<u32, std::vector<HcclIpAddress>>(curRank, rankTable.deviceIps[rankIndex]));
49 : }
50 : }
51 0 : comm_.reset(new (std::nothrow) CommRemoteAccess(rankInComm, devicePhyId, rankInfo, addrInfos));
52 0 : CHK_SMART_PTR_NULL(comm_);
53 0 : CHK_RET(comm_->Init());
54 0 : return HCCL_SUCCESS;
55 0 : }
56 :
57 0 : void RemoteAccessImpl::ParseRemoteAccessAddrInfo(
58 : const vector<HcomRemoteAccessAddrInfo>& addrInfos, map<u32, vector<HcomRemoteAccessAddrInfo>>& addrInfoMap)
59 : {
60 0 : for (u32 i = 0; i < addrInfos.size(); i++) {
61 0 : u32 remoteRankInComm = addrInfos[i].remotetRankID / rankNumPerServer_;
62 0 : addrInfoMap[remoteRankInComm].push_back(addrInfos[i]);
63 0 : HCCL_DEBUG(
64 : "ParseRemoteAccessAddrInfo localAddr[0x%016lx] remoteAddr[0x%016lx] length[%llu] "
65 : "remoteRankInComm[%u]",
66 : addrInfos[i].localAddr, addrInfos[i].remoteAddr, addrInfos[i].length, remoteRankInComm);
67 : }
68 0 : }
69 :
70 0 : HcclResult RemoteAccessImpl::IsInSamePlane(const u32 userRank, const vector<HcomRemoteAccessAddrInfo>& addrInfos)
71 : {
72 0 : for (u32 i = 0; i < addrInfos.size(); i++) {
73 0 : CHK_PRT_RET(
74 : (userRank % rankNumPerServer_) != (addrInfos[i].remotetRankID % rankNumPerServer_),
75 : HCCL_ERROR(
76 : "[Is][InSamePlane]The userrank[%u] and remoterank[%u] must be in the same plane", userRank,
77 : addrInfos[i].remotetRankID),
78 : HCCL_E_PARA);
79 : }
80 0 : return HCCL_SUCCESS;
81 : }
82 :
83 0 : HcclResult RemoteAccessImpl::RemoteWrite(const vector<HcomRemoteAccessAddrInfo>& addrInfos, HcclRtStream stream)
84 : {
85 0 : size_t infoSize = addrInfos.size();
86 0 : CHK_PRT_RET(addrInfos.empty(), HCCL_ERROR("[Remote][Write]addrInfos is empty!"), HCCL_E_PARA);
87 0 : CHK_RET(IsInSamePlane(userRank_, addrInfos));
88 :
89 0 : Stream streamObj(stream);
90 : // GE 保证传入的addrInfos按照remotetRankID排序,如果目标是同一个remotetRank,优化性能
91 0 : if (infoSize > 1 && addrInfos[0].remotetRankID == addrInfos[infoSize - 1].remotetRankID) {
92 0 : u32 remoteRankInComm = addrInfos[0].remotetRankID / rankNumPerServer_;
93 0 : CHK_PRT_RET(
94 : remoteRankInComm > (serverNum_ - 1),
95 : HCCL_ERROR(
96 : "[Remote][Write]remote write invalid rank id [%u] should be in [0, %u]!", remoteRankInComm,
97 : (serverNum_ - 1)),
98 : HCCL_E_PARA);
99 0 : std::shared_ptr<TransportRemoteAccess> transportPtr = comm_->GetTransportByRank(remoteRankInComm);
100 0 : CHK_SMART_PTR_NULL(transportPtr);
101 0 : CHK_RET(transportPtr->RemoteWrite(addrInfos, streamObj));
102 0 : } else {
103 0 : map<u32, vector<HcomRemoteAccessAddrInfo>> addrInfoMap;
104 0 : ParseRemoteAccessAddrInfo(addrInfos, addrInfoMap);
105 0 : for (auto it = addrInfoMap.begin(); it != addrInfoMap.end(); it++) {
106 0 : CHK_PRT_RET(
107 : it->first > (serverNum_ - 1),
108 : HCCL_ERROR(
109 : "[Remote][Write]remote write invalid rank id [%u] should be in [0, %u]!", it->first,
110 : (serverNum_ - 1)),
111 : HCCL_E_PARA);
112 0 : std::shared_ptr<TransportRemoteAccess> transportPtr = comm_->GetTransportByRank(it->first);
113 0 : CHK_SMART_PTR_NULL(transportPtr);
114 0 : CHK_RET(transportPtr->RemoteWrite(it->second, streamObj));
115 0 : }
116 0 : }
117 0 : return HCCL_SUCCESS;
118 0 : }
119 :
120 0 : HcclResult RemoteAccessImpl::RemoteRead(const vector<HcomRemoteAccessAddrInfo>& addrInfos, HcclRtStream stream)
121 : {
122 0 : HCCL_INFO("RemoteAccessImpl::RemoteRead");
123 0 : size_t infoSize = addrInfos.size();
124 0 : CHK_PRT_RET(addrInfos.empty(), HCCL_ERROR("[Remote][Read]addrInfos is empty!"), HCCL_E_PARA);
125 :
126 0 : CHK_RET(IsInSamePlane(userRank_, addrInfos));
127 :
128 0 : Stream streamObj(stream);
129 : // GE 保证传入的addrInfos按照remotetRankID排序,如果目标是同一个remotetRank,优化性能
130 0 : if (infoSize > 1 && addrInfos[0].remotetRankID == addrInfos[infoSize - 1].remotetRankID) {
131 0 : u32 remoteRankInComm = addrInfos[0].remotetRankID / rankNumPerServer_;
132 0 : CHK_PRT_RET(
133 : remoteRankInComm > (serverNum_ - 1),
134 : HCCL_ERROR(
135 : "[remote][Read]remote read invalid rank id [%u] should be in [0, %u]!", remoteRankInComm,
136 : (serverNum_ - 1)),
137 : HCCL_E_PARA);
138 :
139 0 : std::shared_ptr<TransportRemoteAccess> transportPtr = comm_->GetTransportByRank(remoteRankInComm);
140 0 : CHK_SMART_PTR_NULL(transportPtr);
141 0 : CHK_RET(transportPtr->RemoteRead(addrInfos, streamObj));
142 0 : } else {
143 0 : map<u32, vector<HcomRemoteAccessAddrInfo>> addrInfoMap;
144 0 : ParseRemoteAccessAddrInfo(addrInfos, addrInfoMap);
145 0 : for (auto it = addrInfoMap.begin(); it != addrInfoMap.end(); it++) {
146 0 : CHK_PRT_RET(
147 : it->first > (serverNum_ - 1),
148 : HCCL_ERROR(
149 : "[remote][Read]remote read invalid rank id [%u] should be in [0, %u]!", it->first,
150 : (serverNum_ - 1)),
151 : HCCL_E_PARA);
152 :
153 0 : std::shared_ptr<TransportRemoteAccess> transportPtr = comm_->GetTransportByRank(it->first);
154 0 : CHK_SMART_PTR_NULL(transportPtr);
155 :
156 0 : CHK_RET(transportPtr->RemoteRead(it->second, streamObj));
157 0 : }
158 0 : }
159 0 : return HCCL_SUCCESS;
160 0 : }
161 : } // namespace hccl
|