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_v2.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 1 : HcclResult HcclRegisterMemV2(HcclComm comm, u32 remoteRank, int type, void* addr, u64 size, HcclMemDesc* desc)
37 : {
38 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
39 1 : std::string commIdentifier = hcclCommunicator->GetId();
40 3 : HCCL_RUN_INFO(
41 : "Entry-%s:comm[%s], remoteRank[%u], memType[%d], memAddr[%p], memSize[%llu], memDescPtr[%p]", __func__,
42 : commIdentifier.c_str(), remoteRank, type, addr, size, desc);
43 :
44 1 : auto it = HCCL_MEM_TYPE_V2.find(type);
45 1 : CHK_PRT_RET(
46 : it == HCCL_MEM_TYPE_V2.end(),
47 : HCCL_ERROR("[HcclRegisterMemV2] HcclMemType[%d] is invalid, please check memory type", type), HCCL_E_PARA);
48 1 : HcclMemType memType = it->second;
49 1 : u32 localRank = INVALID_VALUE_RANKID;
50 1 : CHK_RET(hcclCommunicator->GetRankId(localRank));
51 :
52 1 : CHK_PRT_RET(
53 : remoteRank == localRank,
54 : HCCL_WARNING(
55 : "remoteRank[%u] is equal to localRank[%u], no need to "
56 : "register memory, return HcclRegisterMem success",
57 : remoteRank, localRank),
58 : HCCL_SUCCESS);
59 :
60 1 : CHK_PRT_RET(
61 : memType != HcclMemType::HCCL_MEM_TYPE_DEVICE && memType != HcclMemType::HCCL_MEM_TYPE_HOST,
62 : HCCL_ERROR("[HcclRegisterMem]memoryType[%d] must be device or host, please check type", type), HCCL_E_PARA);
63 1 : CHK_PRT_RET(
64 : size <= ONE_SIDE_HOST_MEM_ZERO,
65 : HCCL_ERROR(
66 : "[HcclRegisterMem]memory size[%llu] is invalid, "
67 : "please check memory size",
68 : size),
69 : HCCL_E_PARA);
70 1 : CHK_PRT_RET(
71 : memType == HcclMemType::HCCL_MEM_TYPE_DEVICE && size > ONE_SIDE_DEVICE_MEM_MAX_SIZE,
72 : HCCL_ERROR("[HcclRegisterMem]memory size[%llu] is too large, please check memory size", size), HCCL_E_PARA);
73 1 : CHK_PRT_RET(
74 : memType == HcclMemType::HCCL_MEM_TYPE_HOST,
75 : HCCL_ERROR("[HcclRegisterMem] HCCL_MEM_TYPE_HOST is not support, please check memory type"),
76 : HCCL_E_NOT_SUPPORT);
77 :
78 3 : HCCL_INFO("HcclRegisterMemV2 GetLocalRankID Success: localRank[%u]", localRank);
79 :
80 1 : Hccl::HcclOneSidedService* service = nullptr;
81 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
82 1 : CHK_PTR_NULL(service);
83 :
84 3 : HCCL_INFO("HcclRegisterMemV2 RegMem Begin");
85 :
86 : // HcclResult HcclOneSidedService::RegMem(void *addr, u64 size, HcclMemType type, RankId remoteRankId, HcclMemDesc
87 : // &localMemDesc)
88 1 : CHK_RET(service->RegMem(addr, size, memType, remoteRank, *desc));
89 :
90 3 : HCCL_INFO("HcclRegisterMemV2 RegMem End");
91 :
92 3 : HCCL_RUN_INFO(
93 : "%s success:commPtr[%p], remoteRank[%u], memType[%d], memAddr[%p], memSize[%llu], memDescPtr[%p]", __func__,
94 : comm, remoteRank, type, addr, size, desc);
95 1 : return HCCL_SUCCESS;
96 1 : }
97 :
98 1 : HcclResult HcclDeregisterMemV2(HcclComm comm, HcclMemDesc* desc)
99 : {
100 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
101 1 : std::string commIdentifier = hcclCommunicator->GetId();
102 3 : HCCL_RUN_INFO("Entry-%s:comm[%s], memDescPtr[%p]", __func__, commIdentifier.c_str(), desc);
103 :
104 1 : Hccl::HcclOneSidedService* service = nullptr;
105 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
106 1 : CHK_PTR_NULL(service);
107 :
108 3 : HCCL_INFO("HcclRegisterMemV2 DeregMem Begin");
109 4 : CHK_RET(service->DeregMem(*desc));
110 0 : HCCL_INFO("HcclRegisterMemV2 DeregMem End");
111 :
112 0 : HCCL_RUN_INFO("%s success:commPtr[%p], memDescPtr[%p]", __func__, comm, desc);
113 0 : return HCCL_SUCCESS;
114 1 : }
115 :
116 1 : HcclResult HcclExchangeMemDescV2(
117 : HcclComm comm, u32 remoteRank, HcclMemDescs* local, int timeout, HcclMemDescs* remote, u32* actualNum)
118 : {
119 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
120 1 : std::string commIdentifier = hcclCommunicator->GetId();
121 3 : HCCL_RUN_INFO(
122 : "Entry-%s:comm[%s], remoteRank[%u], localMemDescPtr[%p], timeout[%d s], remoteMemDescPtr[%p], "
123 : "actualNum[%u]",
124 : __func__, commIdentifier.c_str(), remoteRank, local, timeout, remote, *actualNum);
125 :
126 1 : u32 localRank = INVALID_VALUE_RANKID;
127 1 : CHK_RET(hcclCommunicator->GetRankId(localRank));
128 1 : CHK_PRT_RET(
129 : remoteRank == localRank,
130 : HCCL_WARNING(
131 : "remoteRank[%u] is equal to localRank[%u], no need to "
132 : "register memory, return HcclRegisterMem success",
133 : remoteRank, localRank),
134 : HCCL_SUCCESS);
135 :
136 1 : Hccl::HcclOneSidedService* service = nullptr;
137 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
138 1 : CHK_PTR_NULL(service);
139 :
140 3 : HCCL_INFO("HcclRegisterMemV2 ExchangeMemDesc Begin");
141 1 : CHK_RET(service->ExchangeMemDesc(remoteRank, *local, *remote, *actualNum));
142 3 : HCCL_INFO("HcclRegisterMemV2 ExchangeMemDesc end");
143 :
144 3 : HCCL_RUN_INFO(
145 : "%s success:commPtr[%p], remoteRank[%u], localMemDescPtr[%p], timeout[%d s], remoteMemDescPtr[%p], "
146 : "actualNum[%u]",
147 : __func__, comm, remoteRank, local, timeout, remote, *actualNum);
148 1 : return HCCL_SUCCESS;
149 1 : }
150 :
151 1 : HcclResult HcclEnableMemAccessV2(HcclComm comm, HcclMemDesc* remoteMemDesc, HcclMem* remoteMem)
152 : {
153 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
154 1 : std::string commIdentifier = hcclCommunicator->GetId();
155 3 : HCCL_RUN_INFO(
156 : "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& paraData, 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(
199 : paraData.descNum > MAX_DESC_NUM,
200 : HCCL_WARNING("[%s] the count of HcclOneSideOpDesc exceed specification.", batchString.c_str()), HCCL_E_PARA);
201 :
202 2 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
203 : // 同算子复用tag
204 2 : u32 localRank = INVALID_VALUE_RANKID;
205 2 : CHK_RET(hcclCommunicator->GetRankId(localRank));
206 :
207 4 : const std::string tag = batchString + "_" + std::to_string(localRank) + "_" + std::to_string(paraData.remoteRank)
208 4 : + "_" + hcclCommunicator->GetId();
209 2 : getTag = tag;
210 :
211 2 : u32 rankSize = INVALID_VALUE_RANKSIZE;
212 2 : CHK_RET_AND_PRINT_IDE(hcclCommunicator->GetRankSize(&rankSize), tag.c_str());
213 2 : CHK_RET(HcomCheckUserRankV2(rankSize, paraData.remoteRank));
214 2 : CHK_PRT_RET(
215 : paraData.remoteRank == localRank,
216 : HCCL_ERROR("[%s] the remoteRank can't be equal to localRank, please check.", batchString.c_str()), HCCL_E_PARA);
217 :
218 2 : s32 streamId = 0;
219 :
220 6 : HCCL_RUN_INFO(
221 : "Entry-%s::tag[%s], descNum[%u], streamId[%d], localRank[%u], remoteRank[%u]", __func__, tag.c_str(),
222 : paraData.descNum, streamId, localRank, paraData.remoteRank);
223 :
224 6 : HCCL_INFO("HcclBatchParaCheckV2 End");
225 2 : return HCCL_SUCCESS;
226 2 : }
227 :
228 4 : HcclResult HcclBatchPutV2(HcclComm comm, u32 remoteRank, HcclOneSideOpDesc* desc, u32 descNum, const rtStream_t stream)
229 : {
230 12 : HCCL_INFO("HcclBatchPutV2 Begin");
231 7 : CHK_PTR_NULL(comm);
232 6 : CHK_PTR_NULL(desc);
233 5 : CHK_PTR_NULL(stream);
234 1 : std::string getTag;
235 1 : CHK_PRT_RET(descNum == 0, HCCL_WARNING("[%s] the count of HcclOneSideOpDesc is zero.", __func__), HCCL_SUCCESS);
236 1 : HcclBatchData paraData = {comm, HcclCMDType::HCCL_CMD_BATCH_PUT, remoteRank, desc, descNum, stream};
237 1 : CHK_RET(HcclBatchParaCheckV2(comm, paraData, getTag));
238 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
239 1 : Hccl::HcclOneSidedService* service = nullptr;
240 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
241 1 : CHK_PTR_NULL(service);
242 :
243 3 : HCCL_INFO("HcclBatchPutV2 BatchPut Begin");
244 1 : CHK_RET(service->BatchPut(remoteRank, desc, descNum, stream));
245 3 : HCCL_INFO("HcclBatchPutV2 End");
246 1 : return HCCL_SUCCESS;
247 1 : }
248 :
249 4 : HcclResult HcclBatchGetV2(HcclComm comm, u32 remoteRank, HcclOneSideOpDesc* desc, u32 descNum, const rtStream_t stream)
250 : {
251 12 : HCCL_INFO("HcclBatchGetV2 Begin");
252 7 : CHK_PTR_NULL(comm);
253 6 : CHK_PTR_NULL(desc);
254 5 : CHK_PTR_NULL(stream);
255 1 : std::string getTag;
256 1 : CHK_PRT_RET(descNum == 0, HCCL_WARNING("[%s] the count of HcclOneSideOpDesc is zero.", __func__), HCCL_SUCCESS);
257 1 : HcclBatchData paraData = {comm, HcclCMDType::HCCL_CMD_BATCH_PUT, remoteRank, desc, descNum, stream};
258 1 : CHK_RET(HcclBatchParaCheckV2(comm, paraData, getTag));
259 1 : Hccl::HcclCommunicator* hcclCommunicator = static_cast<Hccl::HcclCommunicator*>(comm);
260 1 : Hccl::HcclOneSidedService* service = nullptr;
261 1 : CHK_RET(hcclCommunicator->GetOneSidedService(&service));
262 1 : CHK_PTR_NULL(service);
263 :
264 3 : HCCL_INFO("HcclBatchGetV2 BatchGet Begin");
265 1 : CHK_RET(service->BatchGet(remoteRank, desc, descNum, stream));
266 3 : HCCL_INFO("HcclBatchGetV2 End");
267 1 : return HCCL_SUCCESS;
268 1 : }
|