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