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 "one_sided_service_adapt_v2.h"
12 : #include "task_param.h"
13 : #include "hccl_one_sided_data.h"
14 : #include "hccl_one_sided_service.h"
15 : #include "hccl_communicator.h"
16 : #include "hccl_common_v2.h"
17 : #include "log.h"
18 : #include "param_check_v2.h"
19 :
20 : using namespace std;
21 : using namespace Hccl;
22 :
23 : constexpr u64 ONE_SIDE_DEVICE_MEM_MAX_SIZE = 64llu * 1024 * 1024 * 1024; // device侧支持内存注册大小上限为64GB
24 : constexpr u64 ONE_SIDE_HOST_MEM_MAX_SIZE = 1024llu * 1024 * 1024 * 1024; // host侧支持内存注册大小上限为1TB
25 : constexpr u64 ONE_SIDE_HOST_MEM_ZERO = 0;
26 : constexpr u64 MAX_DESC_NUM = 64; // 批量操作描述符个数上限
27 : constexpr u64 MEM_TYPE_DEVICE = 0;
28 : constexpr u64 MEM_TYPE_HOST = 0;
29 : constexpr u64 MEM_TYPE_NUM = 0;
30 :
31 : const std::map<int, HcclMemType> HCCL_MEM_TYPE_V2 {
32 : {MEM_TYPE_DEVICE, HcclMemType::HCCL_MEM_TYPE_DEVICE},
33 : {MEM_TYPE_HOST, HcclMemType::HCCL_MEM_TYPE_HOST},
34 : {MEM_TYPE_NUM, HcclMemType::HCCL_MEM_TYPE_NUM}
35 : };
36 :
37 1 : HcclResult HcclRegisterMemV2(HcclComm comm, u32 remoteRank, int type, void *addr, u64 size, HcclMemDesc *desc)
38 : {
39 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
40 1 : std::string commIdentifier = hcclCommunicator->GetId();
41 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], remoteRank[%u], memType[%d], memAddr[%p], memSize[%llu], memDescPtr[%p]",
42 : __func__, commIdentifier.c_str(), remoteRank, type, addr, size, desc);
43 :
44 1 : auto it = HCCL_MEM_TYPE_V2.find(type);
45 1 : CHK_PRT_RET(it == HCCL_MEM_TYPE_V2.end(),
46 : HCCL_ERROR("[HcclRegisterMemV2] HcclMemType[%d] is invalid, please check memory type", type), HCCL_E_PARA);
47 1 : HcclMemType memType = it->second;
48 1 : u32 localRank = INVALID_VALUE_RANKID;
49 1 : CHK_RET(hcclCommunicator->GetRankId(localRank));
50 :
51 1 : CHK_PRT_RET(remoteRank == localRank,
52 : HCCL_WARNING("remoteRank[%u] is equal to localRank[%u], no need to "
53 : "register memory, return HcclRegisterMem success",
54 : remoteRank,
55 : localRank),
56 : HCCL_SUCCESS);
57 :
58 1 : CHK_PRT_RET(memType != HcclMemType::HCCL_MEM_TYPE_DEVICE && memType != HcclMemType::HCCL_MEM_TYPE_HOST,
59 : HCCL_ERROR("[HcclRegisterMem]memoryType[%d] must be device or host, please check type", type),
60 : HCCL_E_PARA);
61 1 : CHK_PRT_RET(size <= ONE_SIDE_HOST_MEM_ZERO,
62 : HCCL_ERROR("[HcclRegisterMem]memory size[%llu] is invalid, "
63 : "please check memory size",
64 : size),
65 : HCCL_E_PARA);
66 1 : CHK_PRT_RET(memType == HcclMemType::HCCL_MEM_TYPE_DEVICE && size > ONE_SIDE_DEVICE_MEM_MAX_SIZE,
67 : HCCL_ERROR("[HcclRegisterMem]memory size[%llu] is too large, please check memory size", size),
68 : HCCL_E_PARA);
69 1 : CHK_PRT_RET(memType == HcclMemType::HCCL_MEM_TYPE_HOST ,
70 : HCCL_ERROR("[HcclRegisterMem] HCCL_MEM_TYPE_HOST is not support, please check memory type"),
71 : HCCL_E_NOT_SUPPORT);
72 :
73 3 : HCCL_INFO("HcclRegisterMemV2 GetLocalRankID Success: localRank[%u]", localRank);
74 :
75 1 : Hccl::HcclOneSidedService *service = nullptr;
76 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
77 1 : CHK_PTR_NULL(service);
78 :
79 3 : HCCL_INFO("HcclRegisterMemV2 RegMem Begin");
80 :
81 : //HcclResult HcclOneSidedService::RegMem(void *addr, u64 size, HcclMemType type, RankId remoteRankId, HcclMemDesc &localMemDesc)
82 1 : CHK_RET(service->RegMem(addr, size, memType, remoteRank, *desc));
83 :
84 3 : HCCL_INFO("HcclRegisterMemV2 RegMem End");
85 :
86 3 : HCCL_RUN_INFO("%s success:commPtr[%p], remoteRank[%u], memType[%d], memAddr[%p], memSize[%llu], memDescPtr[%p]",
87 : __func__,
88 : comm,
89 : remoteRank,
90 : type,
91 : addr,
92 : size,
93 : desc);
94 1 : return HCCL_SUCCESS;
95 1 : }
96 :
97 1 : HcclResult HcclDeregisterMemV2(HcclComm comm, HcclMemDesc *desc)
98 : {
99 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
100 1 : std::string commIdentifier = hcclCommunicator->GetId();
101 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], memDescPtr[%p]", __func__, commIdentifier.c_str(), desc);
102 :
103 1 : Hccl::HcclOneSidedService *service = nullptr;
104 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
105 1 : CHK_PTR_NULL(service);
106 :
107 3 : HCCL_INFO("HcclRegisterMemV2 DeregMem Begin");
108 4 : CHK_RET(service->DeregMem(*desc));
109 0 : HCCL_INFO("HcclRegisterMemV2 DeregMem End");
110 :
111 0 : HCCL_RUN_INFO("%s success:commPtr[%p], memDescPtr[%p]", __func__, comm, desc);
112 0 : return HCCL_SUCCESS;
113 1 : }
114 :
115 1 : HcclResult HcclExchangeMemDescV2(
116 : HcclComm comm, u32 remoteRank, HcclMemDescs *local, int timeout, HcclMemDescs *remote, u32 *actualNum)
117 : {
118 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
119 1 : std::string commIdentifier = hcclCommunicator->GetId();
120 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], remoteRank[%u], localMemDescPtr[%p], timeout[%d s], remoteMemDescPtr[%p], "
121 : "actualNum[%u]", __func__, commIdentifier.c_str(), remoteRank, local, timeout, remote, *actualNum);
122 :
123 1 : u32 localRank = INVALID_VALUE_RANKID;
124 1 : CHK_RET(hcclCommunicator->GetRankId(localRank));
125 1 : CHK_PRT_RET(remoteRank == localRank,
126 : HCCL_WARNING("remoteRank[%u] is equal to localRank[%u], no need to "
127 : "register memory, return HcclRegisterMem success",
128 : remoteRank,
129 : localRank),
130 : HCCL_SUCCESS);
131 :
132 1 : Hccl::HcclOneSidedService *service = nullptr;
133 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
134 1 : CHK_PTR_NULL(service);
135 :
136 3 : HCCL_INFO("HcclRegisterMemV2 ExchangeMemDesc Begin");
137 1 : CHK_RET(service->ExchangeMemDesc(remoteRank, *local, *remote, *actualNum));
138 3 : HCCL_INFO("HcclRegisterMemV2 ExchangeMemDesc end");
139 :
140 3 : HCCL_RUN_INFO("%s success:commPtr[%p], remoteRank[%u], localMemDescPtr[%p], timeout[%d], remoteMemDescPtr[%p], "
141 : "actualNum[%u]",
142 : __func__,
143 : comm,
144 : remoteRank,
145 : local,
146 : timeout,
147 : remote,
148 : *actualNum);
149 1 : return HCCL_SUCCESS;
150 1 : }
151 :
152 1 : HcclResult HcclEnableMemAccessV2(HcclComm comm, HcclMemDesc *remoteMemDesc, HcclMem *remoteMem)
153 : {
154 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
155 1 : std::string commIdentifier = hcclCommunicator->GetId();
156 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], remoteMemDescPtr[%p], remoteMemPtr[%p]", __func__, commIdentifier.c_str(), remoteMemDesc,
157 : remoteMem);
158 :
159 1 : Hccl::HcclOneSidedService *service = nullptr;
160 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
161 1 : CHK_PTR_NULL(service);
162 :
163 3 : HCCL_INFO("HcclRegisterMemV2 EnableMemAccess Begin");
164 1 : CHK_RET(service->EnableMemAccess(*remoteMemDesc, *remoteMem));
165 3 : HCCL_INFO("HcclRegisterMemV2 EnableMemAccess End");
166 :
167 3 : HCCL_RUN_INFO(
168 : "%s success:commPtr[%p], remoteMemDescPtr[%p], remoteMemPtr[%p]", __func__, comm, remoteMemDesc, remoteMem);
169 1 : return HCCL_SUCCESS;
170 1 : }
171 :
172 1 : HcclResult HcclDisableMemAccessV2(HcclComm comm, HcclMemDesc *remoteMemDesc)
173 : {
174 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
175 1 : std::string commIdentifier = hcclCommunicator->GetId();
176 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], remoteMemDescPtr[%p]", __func__, commIdentifier.c_str(), remoteMemDesc);
177 :
178 1 : Hccl::HcclOneSidedService *service = nullptr;
179 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
180 1 : CHK_PTR_NULL(service);
181 :
182 3 : HCCL_INFO("HcclRegisterMemV2 DisableMemAccess Begin");
183 1 : CHK_RET(service->DisableMemAccess(*remoteMemDesc));
184 3 : HCCL_INFO("HcclRegisterMemV2 DisableMemAccess End");
185 :
186 3 : HCCL_RUN_INFO("%s success:commPtr[%p], remoteMemDescPtr[%p]", __func__, comm, remoteMemDesc);
187 1 : return HCCL_SUCCESS;
188 1 : }
189 :
190 2 : inline static HcclResult HcclBatchParaCheckV2(HcclComm comm, HcclBatchData ¶Data, std::string &getTag)
191 : {
192 6 : HCCL_INFO("HcclBatchParaCheckV2 Begin");
193 : // 参数校验和适配
194 2 : CHK_PTR_NULL(paraData.comm);
195 2 : CHK_PTR_NULL(paraData.stream);
196 2 : CHK_PTR_NULL(paraData.desc);
197 2 : std::string batchString = (paraData.cmdType == HcclCMDType::HCCL_CMD_BATCH_GET) ? "BatchGet" : "BatchPut";
198 2 : CHK_PRT_RET(paraData.descNum > MAX_DESC_NUM, HCCL_WARNING("[%s] the count of HcclOneSideOpDesc exceed specification.",
199 : batchString.c_str()), HCCL_E_PARA);
200 :
201 2 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
202 : // 同算子复用tag
203 2 : u32 localRank = INVALID_VALUE_RANKID;
204 2 : CHK_RET(hcclCommunicator->GetRankId(localRank));
205 :
206 4 : const std::string tag = batchString + "_" + std::to_string(localRank) + "_" + std::to_string(paraData.remoteRank)
207 4 : + "_" + hcclCommunicator->GetId();
208 2 : getTag = tag;
209 :
210 2 : u32 rankSize = INVALID_VALUE_RANKSIZE;
211 2 : CHK_RET_AND_PRINT_IDE(hcclCommunicator->GetRankSize(&rankSize), tag.c_str());
212 2 : CHK_RET(HcomCheckUserRankV2(rankSize, paraData.remoteRank));
213 2 : CHK_PRT_RET(paraData.remoteRank == localRank,
214 : HCCL_ERROR("[%s] the remoteRank can't be equal to localRank, please check.", batchString.c_str()), HCCL_E_PARA);
215 :
216 2 : s32 streamId = 0;
217 :
218 6 : HCCL_RUN_INFO("Entry-%s::tag[%s], descNum[%u], streamId[%d], localRank[%u], remoteRank[%u]", __func__,
219 : tag.c_str(), paraData.descNum, streamId, localRank, paraData.remoteRank);
220 :
221 6 : HCCL_INFO("HcclBatchParaCheckV2 End");
222 2 : return HCCL_SUCCESS;
223 2 : }
224 :
225 4 : HcclResult HcclBatchPutV2(HcclComm comm, u32 remoteRank, HcclOneSideOpDesc* desc, u32 descNum, const rtStream_t stream)
226 : {
227 12 : HCCL_INFO("HcclBatchPutV2 Begin");
228 7 : CHK_PTR_NULL(comm);
229 6 : CHK_PTR_NULL(desc);
230 5 : CHK_PTR_NULL(stream);
231 1 : std::string getTag;
232 1 : CHK_PRT_RET(descNum == 0, HCCL_WARNING("[%s] the count of HcclOneSideOpDesc is zero.",
233 : __func__), HCCL_SUCCESS);
234 1 : HcclBatchData paraData = {comm, HcclCMDType::HCCL_CMD_BATCH_PUT, remoteRank, desc, descNum, stream};
235 1 : CHK_RET(HcclBatchParaCheckV2(comm, paraData, getTag));
236 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
237 1 : Hccl::HcclOneSidedService *service = nullptr;
238 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
239 1 : CHK_PTR_NULL(service);
240 :
241 3 : HCCL_INFO("HcclBatchPutV2 BatchPut Begin");
242 1 : CHK_RET(service->BatchPut(remoteRank, desc, descNum, stream));
243 3 : HCCL_INFO("HcclBatchPutV2 End");
244 1 : return HCCL_SUCCESS;
245 1 : }
246 :
247 4 : HcclResult HcclBatchGetV2(HcclComm comm, u32 remoteRank, HcclOneSideOpDesc* desc, u32 descNum, const rtStream_t stream)
248 : {
249 12 : HCCL_INFO("HcclBatchGetV2 Begin");
250 7 : CHK_PTR_NULL(comm);
251 6 : CHK_PTR_NULL(desc);
252 5 : CHK_PTR_NULL(stream);
253 1 : std::string getTag;
254 1 : CHK_PRT_RET(descNum == 0, HCCL_WARNING("[%s] the count of HcclOneSideOpDesc is zero.",
255 : __func__), HCCL_SUCCESS);
256 1 : HcclBatchData paraData = {comm, HcclCMDType::HCCL_CMD_BATCH_PUT, remoteRank, desc, descNum, stream};
257 1 : CHK_RET(HcclBatchParaCheckV2(comm, paraData, getTag));
258 1 : Hccl::HcclCommunicator *hcclCommunicator = static_cast<Hccl::HcclCommunicator *>(comm);
259 1 : Hccl::HcclOneSidedService *service = nullptr;
260 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
261 1 : CHK_PTR_NULL(service);
262 :
263 3 : HCCL_INFO("HcclBatchGetV2 BatchGet Begin");
264 1 : CHK_RET(service->BatchGet(remoteRank, desc, descNum, stream));
265 3 : HCCL_INFO("HcclBatchGetV2 End");
266 1 : return HCCL_SUCCESS;
267 1 : }
|