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 "template_utils.h"
12 : #include "log.h"
13 : #include "buffer.h"
14 :
15 : namespace Hccl {
16 0 : HcclResult GetUnitAllignSize(const AllignInfo& allignInfo, u64& unitAllignSize)
17 : {
18 0 : u32 dataSizePerVolume = DataTypeSizeGet(allignInfo.dataType);
19 :
20 0 : if (allignInfo.enableAllign) {
21 0 : CHK_PRT_RET(
22 : allignInfo.allignSize < dataSizePerVolume,
23 : HCCL_ERROR("[CollAlgFactory] Invalid input allignSize [%u].", allignInfo.allignSize),
24 : HcclResult::HCCL_E_PARA);
25 0 : unitAllignSize = (allignInfo.allignSize % dataSizePerVolume == 0) ? allignInfo.allignSize :
26 0 : allignInfo.allignSize * dataSizePerVolume;
27 : } else {
28 0 : unitAllignSize = dataSizePerVolume;
29 : }
30 0 : return HcclResult::HCCL_SUCCESS;
31 : }
32 :
33 6 : HcclResult GetAlgRank(const RankId virtRank, const std::vector<RankId>& tempVTopo, u32& algRank)
34 : {
35 6 : std::vector<RankId>::const_iterator topoVecIter = std::find(tempVTopo.begin(), tempVTopo.end(), virtRank);
36 6 : CHK_PRT_RET(
37 : topoVecIter == tempVTopo.end(), HCCL_ERROR("[CollAlgFactory] Invalid virtual Rank!"), HcclResult::HCCL_E_PARA);
38 6 : algRank = distance(tempVTopo.begin(), topoVecIter);
39 :
40 6 : return HcclResult::HCCL_SUCCESS;
41 : }
42 :
43 0 : HcclResult CalcRsAgSliceInfoConcurrMesh(
44 : const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const AllignInfo& allignInfo,
45 : const u64 dataSize, RankSliceInfo& sliceInfoVec)
46 : {
47 : // multi-dimensional mesh
48 : u64 unitAllignSize;
49 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
50 :
51 0 : u32 dimSize0 = tempVTopo[0].size();
52 0 : u32 dimSize1 = tempVTopo[1].size();
53 : u64 sliceSize0
54 0 : = min(dataSize, RoundUp(dataSize, ((dimSize0 + dimSize1) * unitAllignSize)) * dimSize0 * unitAllignSize);
55 0 : u64 sliceSize1 = dataSize - sliceSize0;
56 0 : u64 accumOff = 0;
57 0 : u32 tempRankSize = dimSize0 * dimSize1;
58 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize; rankIdx++) {
59 0 : SliceInfo slice0 = {accumOff, sliceSize0};
60 0 : sliceInfoVec[rankIdx][0] = slice0;
61 0 : accumOff += sliceSize0;
62 :
63 0 : SliceInfo slice1 = {accumOff, sliceSize1};
64 0 : sliceInfoVec[rankIdx][1] = slice1;
65 0 : accumOff += sliceSize1;
66 : }
67 :
68 0 : CHK_PRT_RET(
69 : (sliceInfoVec[tempRankSize - 1][1].offset + sliceInfoVec[tempRankSize - 1][1].size != dataSize * tempRankSize),
70 : HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
71 :
72 0 : return HcclResult::HCCL_SUCCESS;
73 : }
74 :
75 0 : HcclResult CalcRsAgSliceInfoMesh(
76 : const RankId myRank, const u32 tempRankSize, const AllignInfo& allignInfo, const u64 dataSize,
77 : RankSliceInfo& sliceInfoVec)
78 : {
79 : (void)allignInfo;
80 0 : u64 accumOff = 0;
81 0 : for (u32 rankIdx = 0; rankIdx < sliceInfoVec.size(); rankIdx++) {
82 0 : SliceInfo slice = {accumOff, dataSize};
83 0 : sliceInfoVec[rankIdx][0] = slice;
84 0 : accumOff += dataSize;
85 : }
86 0 : CHK_PRT_RET(
87 : (sliceInfoVec[tempRankSize - 1][0].offset + sliceInfoVec[tempRankSize - 1][0].size != dataSize * tempRankSize),
88 : HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
89 :
90 0 : return HcclResult::HCCL_SUCCESS;
91 : }
92 :
93 0 : HcclResult CalcRsAgSliceInfoRing(
94 : const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const AllignInfo& allignInfo,
95 : const u64 dataSize, RankSliceInfo& sliceInfoVec)
96 : {
97 0 : u32 queNum = tempVTopo.size();
98 0 : u32 tempRankSize = tempVTopo[0].size();
99 : u64 unitAllignSize;
100 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
101 :
102 0 : u64 queSliceSize = RoundUp(dataSize, (queNum * unitAllignSize)) * unitAllignSize;
103 :
104 0 : u64 resChunkSize = dataSize;
105 0 : std::vector<u64> queSlice;
106 0 : for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
107 : // split data on queues
108 0 : u64 currQueSliceSize = (resChunkSize > queSliceSize) ? queSliceSize : resChunkSize;
109 0 : queSlice.push_back(currQueSliceSize);
110 0 : resChunkSize -= currQueSliceSize;
111 : }
112 0 : CHK_PRT_RET(
113 : resChunkSize != 0, HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank),
114 : HcclResult::HCCL_E_INTERNAL);
115 :
116 0 : u64 accumOff = 0;
117 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize; rankIdx++) {
118 0 : for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
119 0 : u64 currSliceSize = queSlice[queIdx];
120 0 : SliceInfo currSlice = {accumOff, currSliceSize};
121 0 : accumOff += currSliceSize;
122 0 : sliceInfoVec[rankIdx][queIdx] = currSlice;
123 : }
124 :
125 0 : CHK_PRT_RET(
126 : (accumOff != dataSize * (rankIdx + 1)),
127 : HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank),
128 : HcclResult::HCCL_E_INTERNAL);
129 : }
130 :
131 0 : CHK_PRT_RET(
132 : (sliceInfoVec[tempRankSize - 1][queNum - 1].offset + sliceInfoVec[tempRankSize - 1][queNum - 1].size
133 : != dataSize * tempRankSize),
134 : HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
135 :
136 0 : return HcclResult::HCCL_SUCCESS;
137 0 : }
138 :
139 0 : HcclResult CalcRsAgSliceInfoNHR(
140 : const RankId myRank, const u32 tempRankSize, const AllignInfo& allignInfo, const u64 dataSize,
141 : RankSliceInfo& sliceInfoVec)
142 : {
143 : (void)allignInfo;
144 0 : u64 accumOff = 0;
145 0 : for (u32 rankIdx = 0; rankIdx < sliceInfoVec.size(); rankIdx++) {
146 0 : SliceInfo slice = {accumOff, dataSize};
147 0 : sliceInfoVec[rankIdx][0] = slice;
148 0 : accumOff += dataSize;
149 : }
150 :
151 0 : CHK_PRT_RET(
152 : (sliceInfoVec[tempRankSize - 1][0].offset + sliceInfoVec[tempRankSize - 1][0].size != dataSize * tempRankSize),
153 : HCCL_ERROR("[CollAlgFactory] Rank [%d], SliceInfo calculation error!", myRank), HcclResult::HCCL_E_INTERNAL);
154 :
155 0 : return HcclResult::HCCL_SUCCESS;
156 : }
157 :
158 0 : HcclResult CalcResLinksMesh(
159 : const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
160 : const u32 linkNumBtwPeers, AlgTempResReq& tempResReq)
161 : {
162 : u32 myAlgRank;
163 0 : CHK_RET(GetAlgRank(myRank, tempVTopo[0], myAlgRank));
164 :
165 0 : for (u32 queIdx = 0; queIdx < tempVTopo[0].size() - 1; queIdx++) {
166 : // find neighbors : virtualRank
167 0 : RankId neighborRank = tempVTopo[0][(myAlgRank + 1 + queIdx) % tempRankSize];
168 :
169 : // LinkNum
170 0 : tempResReq.links[neighborRank] = linkNumBtwPeers;
171 : }
172 :
173 0 : return HcclResult::HCCL_SUCCESS;
174 : }
175 :
176 0 : HcclResult CalcResLinksMesh2D(
177 : const RankId myRank, const std::vector<std::vector<RankId>>& tempVTopo, const u32 linkNumBtwPeers,
178 : AlgTempResReq& tempResReq)
179 : {
180 : u32 myAlgRank;
181 0 : for (u32 dim = 0; dim < tempVTopo.size(); dim++) {
182 0 : CHK_RET(GetAlgRank(myRank, tempVTopo[dim], myAlgRank));
183 0 : for (u32 queIdx = 0; queIdx < tempVTopo[dim].size() - 1; queIdx++) {
184 0 : u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo[dim].size());
185 0 : CHK_PRT_RET(neighborAlgRank > (tempVTopo[dim].size() - 1),
186 : HCCL_ERROR(
187 : "[CalcResLinksMesh2D] neighborAlgRank[%u] is invalid,"
188 : "the Max rank[%zu].",
189 : neighborAlgRank, tempVTopo[dim].size() - 1);
190 : , HcclResult::HCCL_E_INTERNAL);
191 0 : RankId neighborRank = tempVTopo[dim][neighborAlgRank];
192 0 : tempResReq.links[neighborRank] = linkNumBtwPeers;
193 : }
194 : }
195 :
196 0 : return HcclResult::HCCL_SUCCESS;
197 : }
198 :
199 0 : HcclResult GetDetourSendRecvLinksIn4P(
200 : const RankId myRank, const RankId neighborRank, const ResLinks& tempLinks,
201 : std::vector<std::vector<LinkDataIterator>>& sendRecvLinks)
202 : {
203 0 : HCCL_DEBUG("[CollAlgFactory] [GetDetourSendRecvLinksIn4P] Rank [%d], NeighborRank [%d].", myRank, neighborRank);
204 0 : LinkDataIterator neighborLinkDataIter = tempLinks.at(neighborRank).begin();
205 0 : while (neighborLinkDataIter != tempLinks.at(neighborRank).end()) {
206 0 : if ((*neighborLinkDataIter).GetDirection() == LinkDirection::BOTH) {
207 0 : sendRecvLinks[0][0] = (neighborLinkDataIter);
208 0 : sendRecvLinks[0][1] = (neighborLinkDataIter);
209 0 : } else if ((*neighborLinkDataIter).GetDirection() == LinkDirection::RECV_ONLY) {
210 : // 当前算法是根据Linkdata属性来判断哪条绕路链路负责收数据,哪条负责发数据,后续方案会改进
211 0 : sendRecvLinks[1][1] = (neighborLinkDataIter);
212 : } else {
213 0 : sendRecvLinks[1][0] = (neighborLinkDataIter);
214 : }
215 0 : neighborLinkDataIter++;
216 : }
217 0 : return HcclResult::HCCL_SUCCESS;
218 : }
219 :
220 0 : HcclResult CalcResLinksRing(
221 : const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
222 : AlgTempResReq& tempResReq)
223 : {
224 0 : std::vector<std::vector<RankId>>::const_iterator tempVTopoIter;
225 0 : for (tempVTopoIter = tempVTopo.begin(); tempVTopoIter != tempVTopo.end(); tempVTopoIter++) {
226 : // locate myRank in tempVTopo -> algRank
227 : u32 myAlgRank;
228 0 : CHK_RET(GetAlgRank(myRank, (*tempVTopoIter), myAlgRank));
229 :
230 : // find neighbors -> virtualRank
231 0 : RankId sendToRank = tempVTopoIter->at((myAlgRank + 1) % tempRankSize);
232 0 : RankId recvFromRank = tempVTopoIter->at((myAlgRank - 1 + tempRankSize) % tempRankSize); // virtualRank
233 :
234 : // LinkNum
235 0 : tempResReq.links[sendToRank] = 1;
236 0 : tempResReq.links[recvFromRank] = 1;
237 : }
238 0 : return HcclResult::HCCL_SUCCESS;
239 : }
240 :
241 0 : u32 GetLinkNum(const RankGraph* rankGraph, RankId srcRank, RankId dstRank)
242 : {
243 0 : std::set<u32> levelSet = rankGraph->GetLevels(srcRank);
244 0 : u32 linkNum = 0;
245 0 : for (u32 levelIdx : levelSet) {
246 0 : std::vector<NetInstance::Path> paths = rankGraph->GetPaths(levelIdx, srcRank, dstRank);
247 0 : linkNum += paths.size();
248 0 : }
249 0 : return linkNum;
250 0 : }
251 :
252 : // NHR的算法步数 = Ceil(log2(N))
253 0 : u32 GetNHRStepNum(u32 rankSize)
254 : {
255 0 : u32 nSteps = 0;
256 0 : for (u32 tmp = rankSize - 1; tmp != 0; tmp >>= 1, nSteps++) {
257 : }
258 0 : HCCL_DEBUG("[NHRBase][GetStepNumInterServer] rankSize[%u] nSteps[%u]", rankSize, nSteps);
259 :
260 0 : return nSteps;
261 : }
262 :
263 0 : HcclResult CalcResLinksNHR(
264 : const RankId myRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
265 : AlgTempResReq& tempResReq)
266 : {
267 0 : CHK_PRT_RET(
268 : tempVTopo.size() != 1,
269 : HCCL_ERROR("[CollAlgFactory][CalcResLinksNHR] invalid tempVTopo size[%zu]", tempVTopo.size()),
270 : HcclResult::HCCL_E_PARA);
271 0 : const std::vector<RankId>& tree = tempVTopo[0];
272 0 : CHK_PRT_RET(
273 : tree.size() != tempRankSize,
274 : HCCL_ERROR("[CollAlgFactory][CalcResLinksNHR] tempRankSize[%u] != tree.size[%zu]", tempRankSize, tree.size()),
275 : HcclResult::HCCL_E_PARA);
276 0 : u32 nSteps = GetNHRStepNum(tempRankSize);
277 :
278 : RankId sendToRank;
279 : RankId recvFromRank;
280 : // locate myRank in tempVTopo -> algRank
281 : u32 myAlgRank;
282 0 : CHK_RET(GetAlgRank(myRank, tree, myAlgRank));
283 :
284 0 : for (u32 currentStep = 0; currentStep < nSteps; currentStep++) {
285 0 : u32 deltaRank = nSteps - 1 - currentStep;
286 : // send info
287 0 : sendToRank = tree[(myAlgRank + (1 << deltaRank)) % tempRankSize];
288 : // receive Info
289 0 : recvFromRank = tree[(myAlgRank + tempRankSize - (1 << deltaRank)) % tempRankSize];
290 0 : tempResReq.links[sendToRank] = 1;
291 0 : tempResReq.links[recvFromRank] = 1;
292 : }
293 0 : return HcclResult::HCCL_SUCCESS;
294 : }
295 :
296 1 : HcclResult GetLocalSendRecvInfoforAlltoall(
297 : const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
298 : {
299 1 : u64 curSendDispls = 0;
300 1 : u64 curSendOffset = 0;
301 1 : u64 curRecvDispls = 0;
302 1 : u64 curRecvOffset = 0;
303 5 : for (u32 j = 0; j < userRankSize; j++) {
304 4 : u64 curSendCounts = opParam.all2AllDataDes.sendCount;
305 4 : u64 curSendLength = curSendCounts * DataTypeSizeGet(opParam.all2AllDataDes.sendType);
306 4 : localSendRecvInfo.sendCounts[j] = curSendCounts;
307 4 : localSendRecvInfo.sendDispls[j] = curSendDispls;
308 4 : localSendRecvInfo.sendLength[j] = curSendLength;
309 4 : localSendRecvInfo.sendOffset[j] = curSendOffset;
310 4 : curSendDispls += curSendCounts;
311 4 : curSendOffset += curSendLength;
312 :
313 4 : u64 curRecvCounts = opParam.all2AllDataDes.sendCount;
314 4 : u64 curRecvLength = curRecvCounts * DataTypeSizeGet(opParam.all2AllDataDes.recvType);
315 4 : localSendRecvInfo.recvCounts[j] = curRecvCounts;
316 4 : localSendRecvInfo.recvDispls[j] = curRecvDispls;
317 4 : localSendRecvInfo.recvLength[j] = curRecvLength;
318 4 : localSendRecvInfo.recvOffset[j] = curRecvOffset;
319 4 : curRecvDispls += curRecvCounts;
320 4 : curRecvOffset += curRecvLength;
321 12 : HCCL_DEBUG(
322 : "[GetLocalSendRecvInfoforAlltoall] rank[%u], sendCounts[%llu], sendDispls[%llu] "
323 : "recvCounts[%llu], recvDispls[%llu], sendLength[%llu], recvLength[%llu]",
324 : userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
325 : localSendRecvInfo.recvDispls[j], localSendRecvInfo.sendLength[j], localSendRecvInfo.recvLength[j]);
326 : }
327 1 : return HcclResult::HCCL_SUCCESS;
328 : }
329 :
330 0 : HcclResult GetLocalSendRecvInfoforAlltoallV(
331 : const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
332 : {
333 0 : CHK_PTR_NULL(opParam.all2AllVDataDes.sendCounts);
334 0 : CHK_PTR_NULL(opParam.all2AllVDataDes.sdispls);
335 0 : CHK_PTR_NULL(opParam.all2AllVDataDes.recvCounts);
336 0 : CHK_PTR_NULL(opParam.all2AllVDataDes.rdispls);
337 0 : for (u32 j = 0; j < userRankSize; j++) {
338 0 : u64 curSendCounts = *(static_cast<const u64*>(opParam.all2AllVDataDes.sendCounts) + j);
339 0 : u64 curSendDispls = *(static_cast<const u64*>(opParam.all2AllVDataDes.sdispls) + j);
340 0 : localSendRecvInfo.sendCounts[j] = curSendCounts;
341 0 : localSendRecvInfo.sendDispls[j] = curSendDispls;
342 0 : localSendRecvInfo.sendLength[j] = curSendCounts * DataTypeSizeGet(opParam.all2AllVDataDes.sendType);
343 0 : localSendRecvInfo.sendOffset[j] = curSendDispls * DataTypeSizeGet(opParam.all2AllVDataDes.sendType);
344 :
345 0 : u64 curRecvCounts = *(static_cast<const u64*>(opParam.all2AllVDataDes.recvCounts) + j);
346 0 : u64 curRecvDispls = *(static_cast<const u64*>(opParam.all2AllVDataDes.rdispls) + j);
347 0 : localSendRecvInfo.recvCounts[j] = curRecvCounts;
348 0 : localSendRecvInfo.recvDispls[j] = curRecvDispls;
349 0 : localSendRecvInfo.recvLength[j] = curRecvCounts * DataTypeSizeGet(opParam.all2AllVDataDes.recvType);
350 0 : localSendRecvInfo.recvOffset[j] = curRecvDispls * DataTypeSizeGet(opParam.all2AllVDataDes.recvType);
351 :
352 0 : HCCL_DEBUG(
353 : "[GetLocalSendRecvInfoforAlltoallV] rank[%u], sendCounts[%llu], sendDispls[%llu] "
354 : "recvCounts[%llu], recvDispls[%llu], sendLength[%llu], recvLength[%llu]",
355 : userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
356 : localSendRecvInfo.recvDispls[j], localSendRecvInfo.sendLength[j], localSendRecvInfo.recvLength[j]);
357 : }
358 0 : return HcclResult::HCCL_SUCCESS;
359 : }
360 :
361 0 : HcclResult GetLocalSendRecvInfoforAlltoallVC(
362 : const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
363 : {
364 0 : u64 curSendDispls = 0;
365 0 : u64 curSendOffset = 0;
366 0 : u64 curRecvDispls = 0;
367 0 : u64 curRecvOffset = 0;
368 0 : for (u32 j = 0; j < userRankSize; j++) {
369 0 : u64 curSendCounts
370 0 : = *(static_cast<const u64*>(opParam.all2AllVCDataDes.sendCountMatrix) + userRank * userRankSize + j);
371 0 : u64 curSendLength = curSendCounts * DataTypeSizeGet(opParam.all2AllVCDataDes.sendType);
372 0 : localSendRecvInfo.sendCounts[j] = curSendCounts;
373 0 : localSendRecvInfo.sendDispls[j] = curSendDispls;
374 0 : localSendRecvInfo.sendLength[j] = curSendLength;
375 0 : localSendRecvInfo.sendOffset[j] = curSendOffset;
376 0 : curSendDispls += curSendCounts;
377 0 : curSendOffset += curSendLength;
378 :
379 0 : u64 curRecvCounts
380 0 : = *(static_cast<const u64*>(opParam.all2AllVCDataDes.sendCountMatrix) + userRank + userRankSize * j);
381 0 : u64 curRecvLength = curRecvCounts * DataTypeSizeGet(opParam.all2AllVCDataDes.recvType);
382 0 : localSendRecvInfo.recvCounts[j] = curRecvCounts;
383 0 : localSendRecvInfo.recvDispls[j] = curRecvDispls;
384 0 : localSendRecvInfo.recvLength[j] = curRecvLength;
385 0 : localSendRecvInfo.recvOffset[j] = curRecvOffset;
386 0 : curRecvDispls += curRecvCounts;
387 0 : curRecvOffset += curRecvLength;
388 0 : HCCL_DEBUG(
389 : "[GetLocalSendRecvInfoforAlltoallVC] rank[%u], sendCounts[%llu], sendDispls[%llu] "
390 : "recvCounts[%llu], recvDispls[%llu]",
391 : userRank, localSendRecvInfo.sendCounts[j], localSendRecvInfo.sendDispls[j], localSendRecvInfo.recvCounts[j],
392 : localSendRecvInfo.recvDispls[j]);
393 : }
394 0 : return HcclResult::HCCL_SUCCESS;
395 : }
396 :
397 1 : HcclResult GetAlltoAllLocalSendRecvInfo(
398 : const CollAlgOperator& opParam, const u32 userRank, const u32 userRankSize, A2ASendRecvInfo& localSendRecvInfo)
399 : {
400 3 : HCCL_DEBUG("[GetAlltoAllLocalSendRecvInfo] rank[%u], userRankSize[%u]", userRank, userRankSize);
401 1 : localSendRecvInfo.sendCounts.resize(userRankSize, 0);
402 1 : localSendRecvInfo.sendDispls.resize(userRankSize, 0);
403 1 : localSendRecvInfo.sendLength.resize(userRankSize, 0);
404 1 : localSendRecvInfo.sendOffset.resize(userRankSize, 0);
405 :
406 1 : localSendRecvInfo.recvCounts.resize(userRankSize, 0);
407 1 : localSendRecvInfo.recvDispls.resize(userRankSize, 0);
408 1 : localSendRecvInfo.recvLength.resize(userRankSize, 0);
409 1 : localSendRecvInfo.recvOffset.resize(userRankSize, 0);
410 1 : if (opParam.opType == OpType::ALLTOALLV) {
411 0 : CHK_RET(GetLocalSendRecvInfoforAlltoallV(opParam, userRank, userRankSize, localSendRecvInfo));
412 1 : } else if (opParam.opType == OpType::ALLTOALL) {
413 1 : CHK_RET(GetLocalSendRecvInfoforAlltoall(opParam, userRank, userRankSize, localSendRecvInfo));
414 0 : } else if (opParam.opType == OpType::ALLTOALLVC) {
415 0 : CHK_RET(GetLocalSendRecvInfoforAlltoallVC(opParam, userRank, userRankSize, localSendRecvInfo));
416 0 : } else if (opParam.opType != OpType::HALFALLTOALLV) {
417 0 : HCCL_ERROR("Only support optype alltoall , alltoallv, halfalltoallv and alltoallvc !");
418 : }
419 3 : HCCL_DEBUG("[GetAlltoAllLocalSendRecvInfo] GetAlltoAllLocalSendRecvInfo success");
420 1 : return HcclResult::HCCL_SUCCESS;
421 : }
422 :
423 : /*
424 : * 一个基本的 Allreduce 数据切分函数,用于ReduceScatter + Allgather组合成的 Allreduce 算。
425 : * 输入的 dataSize 是一张卡上完整的数据量
426 : * 函数会将 dataSize 切分成 rankSize 份,最后一份尾块可能会比其他的切分出来的子块大。
427 : */
428 0 : HcclResult CalcSliceInfoAllReduce(
429 : const AllignInfo& allignInfo, const u32 rankSize, const u64 dataSize, RankSliceInfo& sliceInfoVec)
430 : {
431 0 : sliceInfoVec.clear();
432 0 : sliceInfoVec.resize(rankSize);
433 :
434 0 : u32 dataSizePerVolume = DataTypeSizeGet(allignInfo.dataType);
435 : u64 unitAllignSize;
436 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
437 0 : u64 unitPerSlice = dataSize / unitAllignSize / rankSize;
438 0 : HCCL_DEBUG("unitAllignSize[%llu] unitPerSlice[%llu]", unitAllignSize, unitPerSlice);
439 :
440 0 : u64 accumOff = 0;
441 : SliceInfo currSlice;
442 0 : for (u32 rankIdx = 0; rankIdx < rankSize; rankIdx++) {
443 0 : if (rankIdx == rankSize - 1) {
444 0 : currSlice.offset = accumOff;
445 0 : currSlice.size = dataSize - accumOff;
446 : } else {
447 0 : currSlice.offset = accumOff;
448 0 : currSlice.size = unitPerSlice * unitAllignSize;
449 : }
450 0 : CHK_PRT_RET(
451 : currSlice.size % dataSizePerVolume != 0,
452 : HCCL_ERROR(
453 : "[Calc][SliceInfo]rank[%u] slice size[%llu] is invalid, dataSizePerVolume[%llu]", rankIdx,
454 : currSlice.size, dataSizePerVolume),
455 : HcclResult::HCCL_E_INTERNAL);
456 0 : sliceInfoVec[rankIdx].push_back(currSlice);
457 0 : accumOff += currSlice.size;
458 : }
459 :
460 0 : CHK_PRT_RET(
461 : (sliceInfoVec[rankSize - 1][0].offset + sliceInfoVec[rankSize - 1][0].size != dataSize),
462 : HCCL_ERROR(
463 : "[CalcSliceInfoAllReduce] SliceInfo calculation error! DataSize[%llu], "
464 : "lastoffset[%llu], lastsize[%llu]",
465 : dataSize, sliceInfoVec[rankSize - 1][0].offset, sliceInfoVec[rankSize - 1][0].size),
466 : HcclResult::HCCL_E_INTERNAL);
467 :
468 0 : return HcclResult::HCCL_SUCCESS;
469 : }
470 :
471 0 : HcclResult BufferTypeToAddr(const BufferType& bufferType, CollAlgOperator& op, uint64_t& addr)
472 : {
473 0 : Buffer* buffer = op.GetBuffer(bufferType);
474 0 : CHK_PTR_NULL(buffer);
475 0 : addr = buffer->GetAddr();
476 0 : return HcclResult::HCCL_SUCCESS;
477 : }
478 :
479 0 : HcclResult CalcDataSplitRateForLinks(const std::vector<LinkData>& links, std::vector<float>& dataSplitRate)
480 : {
481 : // 取到第一个对端的link数量来作为数据切分的依据
482 0 : std::vector<u8> linkPortGroupSizes;
483 0 : linkPortGroupSizes.resize(links.size());
484 0 : for (u32 linkIdx = 0; linkIdx < links.size(); linkIdx++) {
485 0 : const LinkData& linkData = links[linkIdx];
486 0 : linkPortGroupSizes[linkIdx] = linkData.GetPortGroupSize();
487 : }
488 0 : u32 totalPortNum = accumulate(linkPortGroupSizes.begin(), linkPortGroupSizes.end(), 0);
489 0 : if (totalPortNum == 0) {
490 0 : HCCL_ERROR("totalPortNum is zero");
491 0 : return HcclResult::HCCL_E_INTERNAL;
492 : }
493 0 : for (u32 linkIdx = 0; linkIdx < linkPortGroupSizes.size(); linkIdx++) {
494 0 : dataSplitRate[linkIdx] = static_cast<float>(linkPortGroupSizes[linkIdx]) / totalPortNum;
495 : }
496 0 : return HcclResult::HCCL_SUCCESS;
497 0 : }
498 :
499 0 : DataSlice CalcDataSliceForLinks(
500 : const DataSlice& recvSrcSliceAllLinks, std::vector<float> dataSplitRate, u32 j, DataType dataType_)
501 : {
502 0 : BufferType type = recvSrcSliceAllLinks.GetType();
503 0 : u64 offset = recvSrcSliceAllLinks.GetOffset();
504 0 : u64 size = recvSrcSliceAllLinks.GetSize();
505 0 : u64 AccSize = 0;
506 0 : u64 typeSize = DataTypeSizeGet(dataType_);
507 0 : u64 dataCnt = size / typeSize;
508 0 : u64 linkNum = dataSplitRate.size();
509 0 : std::vector<DataSlice> dataSliceForLinks(linkNum);
510 0 : HCCL_INFO("[InsTempAllGatherNHR] Slice data for links");
511 0 : for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
512 0 : if (linkIdx != linkNum - 1) {
513 0 : dataSliceForLinks[linkIdx].SetSize(
514 0 : static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize);
515 : } else {
516 0 : dataSliceForLinks[linkIdx].SetSize(size - AccSize);
517 : }
518 0 : dataSliceForLinks[linkIdx].SetOffset(offset + AccSize);
519 0 : AccSize += dataSliceForLinks[linkIdx].GetSize();
520 0 : dataSliceForLinks[linkIdx].SetBufferType(type);
521 : }
522 0 : return dataSliceForLinks[j];
523 0 : }
524 :
525 : } // namespace Hccl
|