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