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 "comm_ahc_base_pub.h"
12 : #include "alg_template_register.h"
13 : #include "calc_ahc_template_register.h"
14 :
15 : #include <iostream>
16 : #include <fstream>
17 :
18 : namespace hccl {
19 :
20 : //AHC 通信关系注册
21 43 : AHCCommCalcFuncRegistry::AHCCommCalcFuncRegistry()
22 : {
23 43 : commCalcFuncCreators_.resize(static_cast<u32>(AHCTemplateType::AHC_TEMPLATE_RESERVED), nullptr);
24 43 : }
25 :
26 129 : AHCCommCalcFuncRegistry &AHCCommCalcFuncRegistry::Instance()
27 : {
28 129 : static AHCCommCalcFuncRegistry globalAlgTemplateRegistry;
29 129 : return globalAlgTemplateRegistry;
30 : }
31 :
32 129 : HcclResult AHCCommCalcFuncRegistry::Register(AHCTemplateType type, AHCCommCalcFuncPtr funPtr)
33 : {
34 129 : if (type >= AHCTemplateType::AHC_TEMPLATE_RESERVED) {
35 0 : HCCL_ERROR("[AHCCommCalcFuncRegistry]template type[%d] out of range.", type);
36 0 : return HcclResult::HCCL_E_INTERNAL;
37 : }
38 :
39 129 : const std::lock_guard<std::mutex> lock(mu_);
40 129 : if (commCalcFuncCreators_[static_cast<u32>(type)] != nullptr) {
41 0 : HCCL_ERROR("[AHCCommCalcFuncRegistry]template type[%d] already registered.", type);
42 0 : return HcclResult::HCCL_E_INTERNAL;
43 : }
44 129 : commCalcFuncCreators_[static_cast<u32>(type)] = funPtr;
45 129 : return HcclResult::HCCL_SUCCESS;
46 129 : }
47 :
48 0 : AHCCommCalcFuncPtr AHCCommCalcFuncRegistry::GetCommCalcFunction(AHCTemplateType type)
49 : {
50 0 : if ( type >= AHCTemplateType::AHC_TEMPLATE_RESERVED) {
51 0 : HCCL_ERROR("[AHCCommCalcFuncRegistry]template type[%d] out of range.", type);
52 0 : return nullptr;
53 : }
54 :
55 0 : if (commCalcFuncCreators_[static_cast<u32>(type)] == nullptr) {
56 0 : HCCL_DEBUG("[AHCCommCalcFuncRegistry]Creator for template type[%d] has not registered.", type);
57 0 : return nullptr;
58 : }
59 0 : HCCL_DEBUG("[AHCCommCalcFuncRegistry][GetCommCalcFunction]get template by type[%d]", type);
60 0 : return commCalcFuncCreators_[static_cast<u32>(type)];
61 : }
62 :
63 : //AHC 核心算法逻辑
64 0 : CommAHCBaseInfo::CommAHCBaseInfo(const std::vector<std::vector<u32>> &subGroups)
65 0 : : minSubGroupIdx_(0), maxSubGroupIdx_(0), rankSize_(0), isAlignBound_(true), isContinusSlice_(true), subGroups_(subGroups)
66 : {
67 : //rank 到 group index 的map 初始化以及最大最小分组下标的初始化
68 0 : u32 minSubGroupSize = subGroups_[0].size();
69 0 : u32 maxSubGroupSize = subGroups_[0].size();
70 0 : u32 curIdx = 0;
71 0 : u32 curOffset = 0;
72 :
73 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
74 0 : rankSize_ = rankSize_ + subGroups[i].size();
75 0 : if (subGroups_[i].size() < minSubGroupSize) {
76 0 : minSubGroupSize = subGroups_[i].size();
77 0 : minSubGroupIdx_ = i;
78 : }
79 0 : if (subGroups_[i].size() > maxSubGroupSize) {
80 0 : maxSubGroupSize = subGroups_[i].size();
81 0 : maxSubGroupIdx_ = i;
82 : }
83 0 : for (u32 j = 0; j < subGroups_[i].size(); ++j) {
84 0 : rankGroupMap_.insert(std::make_pair(subGroups_[i][j], i));
85 0 : rankCommMap_.insert(std::make_pair(subGroups_[i][j], curIdx));
86 0 : curIdx++;
87 : }
88 0 : groupOriginOffset_.insert(std::make_pair(i, curOffset));
89 0 : curOffset = curOffset + subGroups_[i].size();
90 : }
91 :
92 0 : HCCL_DEBUG("[CommAHCBaseInfo] minSubGroupSize[%u] maxSubGroupSize[%u]", minSubGroupSize, maxSubGroupSize);
93 0 : }
94 :
95 0 : CommAHCBaseInfo::~CommAHCBaseInfo()
96 : {
97 0 : }
98 :
99 0 : HcclResult CommAHCBaseInfo::Init(AHCOpType opType, std::map<AHCConcOpType, TemplateType> &ahcAlgOption)
100 : {
101 : (void) opType;
102 0 : return HCCL_SUCCESS;
103 : }
104 :
105 0 : HcclResult CommAHCBaseInfo::DisposeSubGroups(const u32 rank, const std::vector<std::vector<std::vector<u32>>> &globalSubGroups,
106 : std::vector<std::vector<u32>> &level0SubGroups, std::vector<std::vector<u32>> &level1SubGroups)
107 : {
108 0 : bool isRankLevel0SubGroup = false;
109 0 : for (u32 i = 0; i < globalSubGroups.size(); i++) {
110 0 : std::vector<u32> level1SubGroup;
111 0 : for (u32 j = 0; j < globalSubGroups[i].size(); j++) {
112 0 : std::vector<u32> curSubGroup = globalSubGroups[i][j];
113 0 : for (u32 k = 0; k < curSubGroup.size(); k++) {
114 0 : if (curSubGroup[k] == rank) {
115 0 : isRankLevel0SubGroup = true;
116 : }
117 0 : level1SubGroup.push_back(curSubGroup[k]);
118 : }
119 0 : }
120 0 : if(isRankLevel0SubGroup) {
121 0 : level0SubGroups = globalSubGroups[i];
122 0 : isRankLevel0SubGroup = false;
123 : }
124 0 : level1SubGroups.push_back(level1SubGroup);
125 0 : }
126 0 : return HCCL_SUCCESS;
127 : }
128 :
129 0 : HcclResult CommAHCBaseInfo::DisposeSubGroups(const u32 rank, const std::vector<std::vector<std::vector<u32>>> &globalSubGroups,
130 : std::vector<std::vector<u32>> &level0SubGroups, std::vector<std::vector<u32>> &level1SubGroups,
131 : u64 &globalTotalSliceSegment, u32 &rankSizeLevel0)
132 : {
133 0 : u64 tmpTotalSliceSegmentLevel1 = 1;
134 0 : bool isRankLevel0SubGroup = false;
135 0 : for (u32 i = 0; i < globalSubGroups.size(); i++) {
136 0 : std::vector<u32> level1SubGroup;
137 0 : u64 tmpTotalSliceSegmentLevel0 = 1;
138 0 : for (u32 j = 0; j < globalSubGroups[i].size(); j++) {
139 0 : std::vector<u32> curSubGroup = globalSubGroups[i][j];
140 0 : u64 level0GroupSize = static_cast<u64>(curSubGroup.size());
141 0 : tmpTotalSliceSegmentLevel0 = tmpTotalSliceSegmentLevel0 * level0GroupSize / std::__gcd(tmpTotalSliceSegmentLevel0, level0GroupSize);
142 0 : for (u32 k = 0; k < curSubGroup.size(); k++) {
143 0 : if (curSubGroup[k] == rank) {
144 0 : isRankLevel0SubGroup = true;
145 : }
146 0 : level1SubGroup.push_back(curSubGroup[k]);
147 : }
148 0 : }
149 0 : if(isRankLevel0SubGroup) {
150 0 : level0SubGroups = globalSubGroups[i];
151 0 : rankSizeLevel0 = level1SubGroup.size();
152 0 : isRankLevel0SubGroup = false;
153 : }
154 0 : tmpTotalSliceSegmentLevel0 = tmpTotalSliceSegmentLevel0 * static_cast<u64>(level1SubGroup.size());
155 0 : globalTotalSliceSegment = globalTotalSliceSegment * tmpTotalSliceSegmentLevel0 / std::__gcd(globalTotalSliceSegment, tmpTotalSliceSegmentLevel0);
156 0 : u64 level1GroupSize = static_cast<u64>(level1SubGroup.size());
157 0 : tmpTotalSliceSegmentLevel1 = tmpTotalSliceSegmentLevel1 * level1GroupSize / std::__gcd(tmpTotalSliceSegmentLevel1, level1GroupSize);
158 0 : level1SubGroups.push_back(level1SubGroup);
159 0 : }
160 0 : tmpTotalSliceSegmentLevel1 = tmpTotalSliceSegmentLevel1 * static_cast<u64>(level1SubGroups.size());
161 0 : globalTotalSliceSegment = globalTotalSliceSegment * tmpTotalSliceSegmentLevel1 / std::__gcd(globalTotalSliceSegment, tmpTotalSliceSegmentLevel1);
162 0 : return HCCL_SUCCESS;
163 : }
164 :
165 522 : HcclResult CommAHCBaseInfo::InitConcAlgOption(std::map<AHCConcOpType, TemplateType> &ahcAlgOption)
166 : {
167 : //初始化设置拼接算法,intra NHR,inter RING ; 每个 level+conc 类型对应的算子类型约束一致
168 : std::map<AHCConcOpType, TemplateType> ahcAlgOptionInstance= {
169 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER}, TemplateType::TEMPLATE_REDUCESCATTER_NHR},
170 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_ALLREDUCE}, TemplateType::TEMPLATE_ALL_REDUCE_NHR},
171 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_ALLGATHER}, TemplateType::TEMPLATE_ALL_GATHER_NHR},
172 :
173 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER}, TemplateType::TEMPLATE_REDUCESCATTER_RING},
174 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_ALLREDUCE}, TemplateType::TEMPLATE_ALL_REDUCE_RING},
175 0 : {{AHCLevel::AHC_LEVEL_0, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_ALLGATHER}, TemplateType::TEMPLATE_ALL_GATHER_RING},
176 :
177 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER}, TemplateType::TEMPLATE_REDUCESCATTER_NHR},
178 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_ALLREDUCE}, TemplateType::TEMPLATE_ALL_REDUCE_NHR},
179 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTRA, AHCOpType::AHC_OP_TYPE_ALLGATHER}, TemplateType::TEMPLATE_ALL_GATHER_NHR},
180 :
181 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER}, TemplateType::TEMPLATE_REDUCESCATTER_RING},
182 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_ALLREDUCE}, TemplateType::TEMPLATE_ALL_REDUCE_RING},
183 0 : {{AHCLevel::AHC_LEVEL_1, ConcType::CONC_INTER, AHCOpType::AHC_OP_TYPE_ALLGATHER}, TemplateType::TEMPLATE_ALL_GATHER_RING}
184 1044 : };
185 522 : ahcAlgOption = ahcAlgOptionInstance;
186 522 : return HCCL_SUCCESS;
187 522 : }
188 :
189 0 : HcclResult CommAHCBaseInfo::SetIsAlignBound(bool isAlignBound)
190 : {
191 0 : isAlignBound_ = isAlignBound;
192 0 : HCCL_DEBUG("[CommAHCBaseInfo][SetIsAlignBound] isAlignBound_ set [%d].", isAlignBound_);
193 0 : return HCCL_SUCCESS;
194 : }
195 :
196 0 : HcclResult CommAHCBaseInfo::SetGlobalTotalSliceSegment(u64 globalTotalSliceSegment)
197 : {
198 : (void) globalTotalSliceSegment;
199 0 : return HCCL_SUCCESS;
200 : }
201 :
202 0 : HcclResult CommAHCBaseInfo::ParseInputSlice(const std::vector<Slice> &physicalSlices)
203 : {
204 0 : totalSize_ = 0;
205 :
206 0 : for (u32 i = 0; i < physicalSlices.size(); i++) {
207 0 : HCCL_DEBUG("[CommAHCBaseInfo][ParseInputSlice] physicalSlices index[%u] offset[%llu] size[%llu]",
208 : i, physicalSlices[i].offset, physicalSlices[i].size);
209 :
210 0 : if ( i >=1 && physicalSlices[i].size != 0) {
211 0 : if (physicalSlices[i-1].offset + physicalSlices[i-1].size != physicalSlices[i].offset) {
212 0 : isContinusSlice_ = false;
213 : }
214 : }
215 0 : totalSize_ = totalSize_ + physicalSlices[i].size;
216 : }
217 0 : HCCL_DEBUG("[CommAHCBaseInfo][ParseInputSlice] totalSize_[%llu] isContinusSlice_[%u]", totalSize_, isContinusSlice_);
218 0 : return HCCL_SUCCESS;
219 : }
220 :
221 0 : HcclResult CommAHCBaseInfo::TrasLogicSliceToPhysical(std::vector<Slice> &slices, const std::vector<Slice> &physicalSlices)
222 : {
223 0 : for (u32 i = 0; i < slices.size(); i++) {
224 0 : u64 logicOffset = 0;
225 0 : bool translateSuccess = false;
226 :
227 0 : HCCL_DEBUG("[CommAHCBaseInfo][TrasLogicSliceToPhysical] translate slice offset[%llu] size[%llu]", slices[i].offset, slices[i].size);
228 :
229 0 : for (u32 j = 0; j < physicalSlices.size(); j++) {
230 0 : bool startOffsetInRange = ((slices[i].offset >= logicOffset) && (slices[i].offset <= (logicOffset + physicalSlices[j].size - 1)));
231 0 : bool endOffsetInRange = (((slices[i].offset + slices[i].size - 1) >= logicOffset) &&
232 0 : ((slices[i].offset + slices[i].size - 1) <= (logicOffset + physicalSlices[j].size - 1)));
233 :
234 0 : if (isContinusSlice_ && startOffsetInRange) { //连续silie,检查逻辑slice起始边界在物理slice范围,则正常翻译
235 0 : translateSuccess = true;
236 0 : HCCL_DEBUG("[CommAHCBaseInfo][TrasLogicSliceToPhysical] translate to continuous slice offset[%llu] size[%llu] in physical slice offset[%llu] size[%llu]",
237 : slices[i].offset, slices[i].size, physicalSlices[j].offset, physicalSlices[j].size);
238 0 : break;
239 0 : } else if (startOffsetInRange && endOffsetInRange ) { //非连续slice,检查逻辑slice起始和结束边界在物理slice范围,则正常翻译
240 0 : slices[i].offset = physicalSlices[j].offset + slices[i].offset - logicOffset;
241 0 : translateSuccess = true;
242 0 : HCCL_DEBUG("[CommAHCBaseInfo][TrasLogicSliceToPhysical] translate to slice offset[%llu] size[%llu] in physical slice offset[%llu] size[%llu]",
243 : slices[i].offset, slices[i].size, physicalSlices[j].offset, physicalSlices[j].size);
244 0 : break;
245 0 : } else if (!isContinusSlice_ && startOffsetInRange && !endOffsetInRange && slices[i].size != 0) {//逻辑slice跨越非连续物理slice边界,异常退出
246 0 : HCCL_ERROR("[CommAHCBaseInfo][TrasLogicSliceToPhysical] logic slice index[%u] offset[%llu] size[%llu],\
247 : physical index[%u] start offset[%llu] end offset[%llu]", i, slices[i].offset, slices[i].size, j, logicOffset,
248 : (logicOffset + physicalSlices[j].size));
249 0 : return HCCL_E_PARA;
250 : }
251 :
252 0 : logicOffset = logicOffset + physicalSlices[j].size;
253 : }
254 :
255 : //检查翻译结果
256 0 : if (slices[i].size == 0) {
257 : // 0 切片特殊处理
258 0 : slices[i].offset = logicOffset;
259 0 : } else if (!translateSuccess) {
260 0 : HCCL_ERROR("[CommAHCBaseInfo][TrasLogicSliceToPhysical] slice index[%u] offset[%llu] size[%llu] translate ERROR",
261 : i, slices[i].offset, slices[i].size);
262 0 : return HCCL_E_PARA;
263 : }
264 : }
265 :
266 0 : return HCCL_SUCCESS;
267 : }
268 :
269 1381 : HcclResult CommAHCBaseInfo::CheckGlobalGroups(std::vector<std::vector<std::vector<u32>>> &globalSubGroups)
270 : {
271 1381 : if (globalSubGroups.size() == 0) {
272 0 : HCCL_ERROR("[CommAHCBaseInfo][globalSubGroups] globalSubGroups.size() == 0, globalSubGroups init ERROR");
273 0 : return HCCL_E_PARA;
274 : }
275 :
276 2762 : for (u32 i = 0; i < globalSubGroups.size(); i++) {
277 1381 : CHK_RET(CheckSubGroups(globalSubGroups[i]));
278 : }
279 1381 : return HCCL_SUCCESS;
280 : }
281 :
282 1381 : HcclResult CommAHCBaseInfo::CheckSubGroups(std::vector<std::vector<u32>> &subGroups)
283 : {
284 1381 : if (subGroups.size() == 0) {
285 0 : HCCL_ERROR("[CommAHCBaseInfo][CheckSubGroups] subGroups.size() == 0, subGroups init ERROR");
286 0 : return HCCL_E_PARA;
287 : }
288 :
289 3320 : for (u32 i = 0; i < subGroups.size(); i++) {
290 1939 : if (subGroups[i].size() == 0) {
291 0 : HCCL_ERROR("[CommAHCBaseInfo][CheckSubGroups] subGroups[%u].size() == 0, subGroups[] init ERROR", i);
292 0 : return HCCL_E_PARA;
293 : }
294 5852 : for (u32 j = 0; j < subGroups[i].size(); j++){
295 3913 : HCCL_DEBUG("[CommAHCBaseInfo][CheckSubGroups] subGroups[%u][%u] = %u", i, j, subGroups[i][j]);
296 : }
297 : }
298 1381 : return HCCL_SUCCESS;
299 : }
300 :
301 0 : void CommAHCBaseInfo::GetIntraCommGroup(u32 rank, std::vector<u32> &intraCommGroup)
302 : {
303 0 : u32 groupIndex = rankGroupMap_[rank];
304 0 : intraCommGroup = subGroups_[groupIndex];
305 0 : }
306 :
307 0 : void CommAHCBaseInfo::GetInterCommGroupIdxList(u32 rank, std::vector<u32> &interCommGroupIdxList)
308 : {
309 : //broke 方式的合法vetor大小为0或1,AHC 方式的合法vetor大小大于等于1
310 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
311 0 : for (u32 j = 0; j < logicCardCommGroups_[i].size(); ++j) {
312 0 : if (rank == logicCardCommGroups_[i][j]) {
313 0 : interCommGroupIdxList.push_back(i);
314 : }
315 : }
316 : }
317 0 : }
318 :
319 0 : void CommAHCBaseInfo::GetInterCommGroupList(u32 rank, std::vector<std::vector<u32>> &interCommGroupList)
320 : {
321 : //broke 方式的合法vetor大小为0或1,AHC 方式的合法vetor大小大于等于1
322 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
323 0 : for (u32 j = 0; j < logicCardCommGroups_[i].size(); ++j) {
324 0 : if (rank == logicCardCommGroups_[i][j]) {
325 0 : interCommGroupList.push_back(logicCardCommGroups_[i]);
326 : }
327 : }
328 : }
329 0 : }
330 :
331 0 : HcclResult CommAHCBaseInfo::CalcDstRanks(u32 rank, std::set<u32> &dstRanks, AHCLevel ahcLevel)
332 : {
333 : //组内和组间通信域计算
334 0 : std::vector<u32> intraCommGroup;
335 0 : std::vector<u32> interCommGroupIdxList;
336 :
337 0 : GetIntraCommGroup(rank, intraCommGroup);
338 0 : GetInterCommGroupIdxList(rank, interCommGroupIdxList);
339 0 : for (u32 i = 0; i < intraCommGroup.size(); i++) {
340 0 : HCCL_DEBUG("[CommAHCBaseInfo][CalcDstRanks] intraCommGroup[%u] = [%u]", i, intraCommGroup[i]);
341 : }
342 0 : for (u32 i = 0; i < interCommGroupIdxList.size(); i++) {
343 0 : for (u32 j = 0; j < logicCardCommGroups_[interCommGroupIdxList[i]].size(); j++) {
344 0 : HCCL_DEBUG("[CommAHCBaseInfo][CalcDstRanks] Rank[%u] logicCardCommGroups_[%u][%u] = [%u]",
345 : rank, interCommGroupIdxList[i], j, logicCardCommGroups_[interCommGroupIdxList[i]][j]);
346 : }
347 : }
348 :
349 : //组内通信关系计算
350 0 : AHCConcOpType concOpType;
351 0 : concOpType.ahcLevel = ahcLevel;
352 0 : concOpType.concType = ConcType::CONC_INTRA;
353 0 : concOpType.ahcOpType = AHCOpType::AHC_OP_TYPE_ALLREDUCE;
354 :
355 0 : TemplateType algType = ahcAlgOption_[concOpType];
356 0 : HCCL_DEBUG("[CommAHCBaseInfo][CalcDstRanks] Level[%u] ConcType[%u] choose algType[%u]",
357 : ahcLevel, ConcType::CONC_INTRA, algType);
358 :
359 0 : auto iterAHCCaclTemplateType = templateToAHCCalcTemplateMap.find(algType);
360 0 : if (iterAHCCaclTemplateType == templateToAHCCalcTemplateMap.end()) {
361 0 : HCCL_ERROR("[CommAHCBaseInfo][CalcDstRanks] intra algo type[%u] is invalid, is not register.", algType);
362 0 : return HCCL_E_PARA;
363 : }
364 :
365 0 : AHCCommCalcFuncPtr intraFunctionPtr = AHCCommCalcFuncRegistry::Instance().GetCommCalcFunction(iterAHCCaclTemplateType->second);
366 0 : CHK_PTR_NULL(intraFunctionPtr);
367 0 : intraFunctionPtr(GetIntraRank(rank), intraCommGroup, dstRanks);
368 :
369 : //组间通信关系计算
370 0 : concOpType.concType = ConcType::CONC_INTER;
371 0 : algType = ahcAlgOption_[concOpType];
372 0 : HCCL_DEBUG("[CommAHCBaseInfo][CalcDstRanks] Level[%u] ConcType[%u] choose algType[%u]",
373 : ahcLevel, ConcType::CONC_INTER, algType);
374 :
375 0 : iterAHCCaclTemplateType = templateToAHCCalcTemplateMap.find(algType);
376 0 : if (iterAHCCaclTemplateType == templateToAHCCalcTemplateMap.end()) {
377 0 : HCCL_ERROR("[CommAHCBaseInfo][CalcDstRanks] inter algo type[%u] is invalid, is not register.", algType);
378 0 : return HCCL_E_PARA;
379 : }
380 :
381 0 : AHCCommCalcFuncPtr interFunctionPtr = AHCCommCalcFuncRegistry::Instance().GetCommCalcFunction(iterAHCCaclTemplateType->second);
382 0 : CHK_PTR_NULL(interFunctionPtr);
383 0 : for (u32 i = 0; i < interCommGroupIdxList.size(); ++i) {
384 0 : interFunctionPtr(GetInterRank(interCommGroupIdxList[i], rank), logicCardCommGroups_[interCommGroupIdxList[i]], dstRanks);
385 : }
386 :
387 0 : return HCCL_SUCCESS;
388 0 : }
389 :
390 0 : HcclResult CommAHCBaseInfo::GetNslbDstRanks(u32 rank, std::vector<u32> &dstRanks)
391 : {
392 0 : HCCL_DEBUG("[NSLB-AHC] entry GetNslbDstRanks rank[%u]", rank);
393 0 : std::vector<u32> intraCommGroup;
394 0 : std::vector<u32> interCommGroupIdxList;
395 :
396 0 : GetIntraCommGroup(rank, intraCommGroup);
397 0 : GetInterCommGroupIdxList(rank, interCommGroupIdxList);
398 :
399 : //组间通信关系计算
400 0 : AHCConcOpType concOpType;
401 0 : concOpType.ahcLevel = AHCLevel::AHC_LEVEL_0;
402 0 : concOpType.concType = ConcType::CONC_INTER;
403 0 : concOpType.ahcOpType = AHCOpType::AHC_OP_TYPE_ALLREDUCE;
404 :
405 0 : TemplateType algType = ahcAlgOption_[concOpType];
406 0 : auto iterAHCCaclTemplateType = templateToAHCCalcTemplateMap.find(algType);
407 0 : if (iterAHCCaclTemplateType == templateToAHCCalcTemplateMap.end()) {
408 0 : HCCL_ERROR("[CommAHCBaseInfo][CalcDstRanks] inter algo type[%u] is invalid, is not register.", algType);
409 0 : return HCCL_E_PARA;
410 : }
411 :
412 0 : AHCCommCalcFuncPtr interFunctionPtr = AHCCommCalcFuncRegistry::Instance().GetCommCalcFunction(iterAHCCaclTemplateType->second);
413 0 : CHK_PTR_NULL(interFunctionPtr);
414 0 : for (u32 i = 0; i < interCommGroupIdxList.size(); ++i) {
415 0 : AHCTemplateType type = iterAHCCaclTemplateType->second;
416 0 : GetDstRanksByType(type, GetInterRank(interCommGroupIdxList[i], rank), logicCardCommGroups_[interCommGroupIdxList[i]], dstRanks);
417 : }
418 :
419 0 : return HCCL_SUCCESS;
420 0 : }
421 :
422 0 : u32 CommAHCBaseInfo::GetIntraRank(const u32 rank)
423 : {
424 0 : u32 intraRank = 0;
425 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); i++) {
426 0 : if (subGroups_[rankGroupMap_[rank]][i] == rank) {
427 0 : intraRank = i;
428 0 : return intraRank;
429 : }
430 : }
431 0 : HCCL_DEBUG("[CommAHCBaseInfo][GetIntraRank] rank[%u] not found", rank);
432 0 : return intraRank;
433 : }
434 :
435 0 : u32 CommAHCBaseInfo::GetInterRank(const u32 groupIdx, const u32 rank)
436 : {
437 0 : u32 subGroupsIdx = rankGroupMap_[rank];
438 0 : u32 interRank = interRankList_[groupIdx][subGroupsIdx];
439 0 : HCCL_DEBUG("[CommAHCBaseInfo][GetInterRank] rank[%u] group[%u] interRank[%u]", rank, groupIdx, interRank);
440 0 : return interRank;
441 : }
442 :
443 0 : u32 CommAHCBaseInfo::GetCommRank(const u32 rank)
444 : {
445 0 : return rankCommMap_[rank];
446 : }
447 :
448 0 : HcclResult CommAHCBaseInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
449 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &intraLinksVector,
450 : std::vector<std::vector<Slice>> &intraSlicesVector)
451 : {
452 : (void)rank;
453 : (void)dataUnitSize;
454 : (void)count;
455 : (void)links;
456 : (void)intraLinksVector;
457 : (void)intraSlicesVector;
458 0 : return HCCL_SUCCESS;
459 : }
460 :
461 0 : HcclResult CommAHCBaseInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
462 : const std::vector<LINK> &links, std::vector<LINK> &intraLinks, std::vector<Slice> &intraSlices)
463 : {
464 : (void)rank;
465 : (void)dataUnitSize;
466 : (void)count;
467 : (void)links;
468 : (void)intraLinks;
469 : (void)intraSlices;
470 0 : return HCCL_SUCCESS;
471 : }
472 :
473 0 : HcclResult CommAHCBaseInfo::CalcInterSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
474 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
475 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
476 : {
477 : (void)rank;
478 : (void)dataUnitSize;
479 : (void)count;
480 : (void)links;
481 : (void)interLinksVector;
482 : (void)interSlicesVector;
483 : (void)logicCardList;
484 0 : return HCCL_SUCCESS;
485 : }
486 :
487 0 : HcclResult CommAHCBaseInfo::GetIntraAlgTemplateOpInstance(const AHCOpType opType, std::unique_ptr<AlgTemplateBase> &tempAlg,
488 : const HcclDispatcher &dispatcher, const u64 reduceAttr,
489 : bool extendFlag, AHCExtendPreparePara extendPara, AHCLevel ahcLevel)
490 : {
491 0 : return GetAlgTemplateOpInstance(opType, tempAlg, dispatcher, reduceAttr, extendFlag, extendPara, ahcLevel, ConcType::CONC_INTRA);
492 : }
493 :
494 0 : HcclResult CommAHCBaseInfo::GetInterAlgTemplateOpInstance(const AHCOpType opType, std::unique_ptr<AlgTemplateBase> &tempAlg,
495 : const HcclDispatcher &dispatcher, const u64 reduceAttr,
496 : bool extendFlag, AHCExtendPreparePara extendPara, AHCLevel ahcLevel)
497 : {
498 0 : return GetAlgTemplateOpInstance(opType, tempAlg, dispatcher, reduceAttr, extendFlag, extendPara, ahcLevel, ConcType::CONC_INTER);
499 : }
500 :
501 0 : HcclResult CommAHCBaseInfo::GetAlgTemplateOpInstance(const AHCOpType opType, std::unique_ptr<AlgTemplateBase> &tempAlg,
502 : const HcclDispatcher &dispatcher, const u64 reduceAttr,
503 : bool extendFlag, AHCExtendPreparePara extendPara, AHCLevel ahcLevel, ConcType concType)
504 : {
505 0 : AHCConcOpType ahcConcOpType;
506 0 : ahcConcOpType.ahcLevel = ahcLevel;
507 0 : ahcConcOpType.concType = concType;
508 0 : ahcConcOpType.ahcOpType = opType;
509 :
510 0 : TemplateType algType = ahcAlgOption_[ahcConcOpType];
511 :
512 0 : HCCL_DEBUG("[CommAHCBaseInfo][GetAlgTemplateOpInstance] Level[%u] ConcType[%u] choose algType[%u]",
513 : ahcLevel, concType, algType);
514 :
515 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(algType, dispatcher);
516 0 : CHK_SMART_PTR_NULL(tempAlg);
517 :
518 : /*特殊属性传递*/
519 : //reduceAttr 传递
520 0 : if(opType == AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER || opType == AHCOpType::AHC_OP_TYPE_ALLREDUCE) {
521 0 : if (algType == TemplateType::TEMPLATE_REDUCESCATTER_NHR) {
522 0 : CHK_RET(tempAlg->Prepare(reduceAttr, false));
523 : } else {
524 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
525 : }
526 : }
527 :
528 : //AHC 扩展属性传递
529 0 : if (extendFlag) {
530 0 : CHK_RET(tempAlg->Prepare(extendPara));
531 : }
532 :
533 0 : return HCCL_SUCCESS;
534 : }
535 :
536 0 : bool CommAHCBaseInfo::IsNeedInterProc(const u32 rank)
537 : {
538 : (void)rank;
539 0 : return true;
540 : }
541 :
542 0 : CommBrokeAlignInfo::CommBrokeAlignInfo(const std::vector<std::vector<u32>> &subGroups)
543 0 : : CommAHCBaseInfo(subGroups)
544 : {
545 0 : }
546 :
547 0 : CommBrokeAlignInfo::~CommBrokeAlignInfo()
548 : {
549 0 : }
550 :
551 0 : HcclResult CommBrokeAlignInfo::Init(AHCOpType opType, std::map<AHCConcOpType, TemplateType> &ahcAlgOption)
552 : {
553 0 : ahcAlgOption_= ahcAlgOption;
554 :
555 : // 参数检查
556 0 : opType_ = opType;
557 0 : CHK_RET(CheckSubGroups(subGroups_));
558 :
559 : //初始化broke 对齐的组间通信域相关信息
560 0 : for (u32 i = 0; i < subGroups_[minSubGroupIdx_].size(); ++i) {
561 0 : std::map<u32, u32> interRankOrder;
562 0 : std::vector<u32> logicGroup;
563 0 : for (u32 j = 0; j < subGroups_.size(); ++j) {
564 0 : logicGroup.push_back(subGroups_[j][i]);
565 0 : interRankOrder.insert(std::make_pair(j, j));
566 : }
567 0 : interRankList_.push_back(interRankOrder);
568 0 : logicCardCommGroups_.push_back(logicGroup);
569 0 : }
570 :
571 : // Reduce-Scatter 及 All-Gather 增加建链信息
572 0 : if (opType_ != AHCOpType::AHC_OP_TYPE_ALLREDUCE) {
573 : // 生成 broke中Reduce-scatter的执行顺序及通信关系分组
574 0 : for (u32 i = subGroups_[minSubGroupIdx_].size(); i < subGroups_[maxSubGroupIdx_].size(); ++i) {
575 0 : std::vector<u32> logicGroup;
576 0 : std::vector<u32> tmpCompleteGroupOrder;
577 0 : std::vector<u32> tmpEmptyGroupOrder;
578 0 : std::map<u32, u32> interRankOrder;
579 0 : u32 curCompleteIdx = 0;
580 0 : for (u32 j = 0; j < subGroups_.size(); ++j) {
581 : // 填充需要得到数据的分组信息
582 0 : if (subGroups_[j].size() > i) {
583 0 : tmpCompleteGroupOrder.push_back(j);
584 0 : interRankOrder.insert(std::make_pair(j, curCompleteIdx));
585 0 : logicGroup.insert(logicGroup.begin() + curCompleteIdx, subGroups_[j][i % subGroups_[j].size()]);
586 0 : curCompleteIdx++;
587 : } else {
588 0 : tmpEmptyGroupOrder.push_back(j);
589 0 : logicGroup.push_back(subGroups_[j][i % subGroups_[j].size()]);
590 : }
591 : }
592 : // 填充用空片参与运算的分组信息
593 0 : for (u32 j = 0; j < tmpEmptyGroupOrder.size(); ++j) {
594 0 : interRankOrder.insert(std::make_pair(tmpEmptyGroupOrder[j], curCompleteIdx));
595 0 : curCompleteIdx++;
596 : }
597 0 : interRankList_.push_back(interRankOrder);
598 0 : logicCardCommGroups_.push_back(logicGroup);
599 0 : completeGroupOrder_.insert(std::make_pair(i, tmpCompleteGroupOrder));
600 0 : emptyGroupOrder_.insert(std::make_pair(i, tmpEmptyGroupOrder));
601 0 : }
602 : }
603 :
604 0 : return HCCL_SUCCESS;
605 : }
606 :
607 0 : bool CommBrokeAlignInfo::IsNeedInterProc(const u32 rank)
608 : {
609 0 : u32 intraRank = GetIntraRank(rank);
610 0 : HCCL_DEBUG("[CommBrokeAlignInfo][IsNeedInterProc] rank[%u] intraRank[%u] minSize[%u]",
611 : rank, intraRank, subGroups_[minSubGroupIdx_].size());
612 0 : if ( intraRank > (subGroups_[minSubGroupIdx_].size() - 1)) {
613 0 : return false;
614 : }
615 0 : return true;
616 : }
617 :
618 : // Reduce-Scatter 及 All-Gather 组内切片逻辑
619 0 : HcclResult CommBrokeAlignInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
620 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &intraLinksVector,
621 : std::vector<std::vector<Slice>> &intraSlicesVector)
622 : {
623 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
624 :
625 0 : u64 sliceSizeAligned = totalSize_ / rankSize_;
626 0 : u64 curoffset = 0;
627 :
628 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] calculate sliceSizeAligned[%llu]", sliceSizeAligned);
629 :
630 0 : for (u32 k = 0; k < subGroups_.size(); ++k) {
631 : // 满片分组处理过程
632 0 : for (u32 j = 0; j < subGroups_[k].size() / subGroups_[rankGroupMap_[rank]].size(); ++j) {
633 0 : std::vector<Slice> intraSlices;
634 0 : std::vector<LINK> intraLinks;
635 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
636 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
637 0 : intraLinks.push_back(links[curRank]);
638 0 : Slice slice;
639 0 : slice.size = sliceSizeAligned;
640 0 : slice.offset = curoffset;
641 0 : curoffset = curoffset + slice.size;
642 0 : intraSlices.push_back(slice);
643 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], link[%u] slices[%u].offset=%llu, slices[%u].size=%llu",
644 : rank, curRank, i, slice.offset, i, slice.size);
645 : }
646 0 : intraLinksVector.push_back(intraLinks);
647 0 : intraSlicesVector.push_back(intraSlices);
648 0 : }
649 0 : std::vector<Slice> intraSlices;
650 0 : std::vector<LINK> intraLinks;
651 : // 涉及空片分组非零切片处理过程
652 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
653 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
654 0 : intraLinks.push_back(links[curRank]);
655 0 : Slice slice;
656 0 : slice.size = i < subGroups_[k].size() % subGroups_[rankGroupMap_[rank]].size() ? sliceSizeAligned : 0;
657 0 : slice.offset = curoffset;
658 0 : curoffset = curoffset + slice.size;
659 0 : intraSlices.push_back(slice);
660 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], link[%u] slices[%u].offset=%llu, slices[%u].size=%llu",
661 : rank, curRank, i, slice.offset, i, slice.size);
662 : }
663 0 : intraLinksVector.push_back(intraLinks);
664 0 : intraSlicesVector.push_back(intraSlices);
665 0 : }
666 :
667 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
668 0 : return HCCL_SUCCESS;
669 : }
670 :
671 : // All-Reduce 组内切片逻辑
672 0 : HcclResult CommBrokeAlignInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
673 : const std::vector<LINK> &links, std::vector<LINK> &intraLinks, std::vector<Slice> &intraSlices)
674 : {
675 : // 计算组内每个rank结果上的offset和size
676 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
677 :
678 0 : u64 sliceSizeCalculated = (count + (static_cast<u32>(subGroups_[minSubGroupIdx_].size()) - 1)) / subGroups_[minSubGroupIdx_].size() * dataUnitSize;
679 0 : u64 totalSize = count * dataUnitSize;
680 0 : u64 residueSize = totalSize;
681 : u64 sliceSizeAligned;
682 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
683 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
684 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
685 : } else {
686 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
687 : }
688 :
689 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
690 0 : intraLinks.push_back(links[subGroups_[rankGroupMap_[rank]][i]]);
691 0 : Slice slice;
692 0 : if (i < subGroups_[minSubGroupIdx_].size()) {
693 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
694 0 : slice.offset = totalSize - residueSize;
695 0 : residueSize -= slice.size;
696 : } else {
697 0 : slice.size = 0;
698 0 : slice.offset = totalSize - residueSize;
699 : }
700 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
701 : rank, i, slice.offset, i, slice.size);
702 0 : intraSlices.push_back(slice);
703 : }
704 :
705 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
706 :
707 0 : return HCCL_SUCCESS;
708 : }
709 :
710 : // 组间切片逻辑统一对外接口
711 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
712 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
713 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
714 : {
715 0 : HcclResult ret = HCCL_SUCCESS;
716 0 : switch (opType_) {
717 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE:
718 0 : ret = CalcInterSlicesAndLinksForAR(rank, dataUnitSize, count, links, interLinksVector, interSlicesVector);
719 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
720 : HCCL_ERROR("[CommBrokeAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks step",
721 : rank, count), ret);
722 0 : break;
723 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER:
724 : case AHCOpType::AHC_OP_TYPE_ALLGATHER:
725 0 : ret = CalcInterSlicesAndLinksForRS(rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
726 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
727 : HCCL_ERROR("[CommBrokeAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks step",
728 : rank, count), ret);
729 0 : break;
730 0 : default:
731 0 : ret = HCCL_SUCCESS;
732 : }
733 0 : return ret;
734 : }
735 :
736 0 : HcclResult CommBrokeAlignInfo::PrepareIntraSlices(const u32 rank, const u32 dataUnitSize, const u64 count,
737 : std::vector<Slice> &intraSlices) const
738 : {
739 : (void)dataUnitSize;
740 : (void)count;
741 :
742 : // 计算组内每个rank结果上的offset和size
743 0 : HCCL_DEBUG("[CommBrokeAlignInfo][PrepareIntraSlices] begin calc intra slices and links rank[%u] ranksize[%u]", rank, rankSize_);
744 :
745 0 : u64 sliceSizeAligned = totalSize_ / rankSize_;
746 0 : u64 curoffset = 0;
747 :
748 0 : for (u32 i = 0; i < rankSize_; ++i) {
749 0 : Slice slice;
750 0 : slice.size = sliceSizeAligned;
751 0 : slice.offset = curoffset;
752 0 : curoffset = curoffset + slice.size;
753 0 : intraSlices.push_back(slice);
754 0 : HCCL_DEBUG("[CommBrokeAlignInfo][PrepareIntraSlices] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
755 : rank, i, slice.offset, i, slice.size);
756 : }
757 0 : HCCL_DEBUG("[CommBrokeAlignInfo][PrepareIntraSlices] end calc intra slices and links rank[%u]", rank);
758 0 : return HCCL_SUCCESS;
759 : }
760 :
761 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinksForRS(const u32 rank, const u32 dataUnitSize, const u64 count,
762 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
763 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
764 : {
765 0 : std::vector<Slice> intraSlices;
766 :
767 0 : CHK_RET(PrepareIntraSlices(rank, dataUnitSize, count, intraSlices));
768 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u] begin inter", rank);
769 0 : u32 intraRank = GetIntraRank(rank);
770 0 : u32 groupCountForRank = subGroups_[maxSubGroupIdx_].size() / subGroups_[rankGroupMap_[rank]].size();
771 0 : if (subGroups_[maxSubGroupIdx_].size() % subGroups_[rankGroupMap_[rank]].size() > intraRank) {
772 0 : groupCountForRank++;
773 : }
774 :
775 0 : for (u32 k = 0; k < groupCountForRank; ++k) {
776 0 : std::vector<Slice> interSlices;
777 0 : std::vector<LINK> interLinks;
778 0 : u32 curGroupIdx = intraRank + k * subGroups_[rankGroupMap_[rank]].size();
779 0 : if (curGroupIdx < subGroups_[minSubGroupIdx_].size()) { // 参与运算的所有 slice 都是有数据的
780 0 : logicCardList.push_back(rankGroupMap_[rank]);
781 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
782 0 : Slice curSlice = intraSlices[groupOriginOffset_[i] + curGroupIdx];
783 0 : interLinks.push_back(links[subGroups_[i][intraRank]]);
784 0 : interSlices.push_back(curSlice);
785 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], curIdx[%u], subGroup[%u], groupIdx[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
786 : rank, subGroups_[i][intraRank], groupOriginOffset_[i] + curGroupIdx, i, curGroupIdx, groupOriginOffset_[i], curSlice.offset, groupOriginOffset_[i], curSlice.size);
787 : }
788 : } else { // 部分空片参与运算
789 0 : Slice emptySlice;
790 0 : emptySlice.size = 0;
791 0 : emptySlice.offset = 0;
792 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
793 0 : u32 curSubgroupsIdx = i < completeGroupOrder_[curGroupIdx].size() ?
794 0 : completeGroupOrder_[curGroupIdx][i] : emptyGroupOrder_[curGroupIdx][i - completeGroupOrder_[curGroupIdx].size()];
795 0 : if (curSubgroupsIdx == rankGroupMap_[rank]) {
796 0 : logicCardList.push_back(i);
797 : }
798 0 : Slice curSlice = i < completeGroupOrder_[curGroupIdx].size() ?
799 0 : intraSlices[groupOriginOffset_[curSubgroupsIdx] + curGroupIdx] : emptySlice;
800 0 : interLinks.push_back(links[subGroups_[curSubgroupsIdx][curGroupIdx % subGroups_[curSubgroupsIdx].size()]]);
801 0 : interSlices.push_back(curSlice);
802 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], curIdx[%u], subGroup[%u], groupIdx[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
803 : rank, subGroups_[curSubgroupsIdx][curGroupIdx % subGroups_[curSubgroupsIdx].size()],
804 : groupOriginOffset_[curSubgroupsIdx] + curGroupIdx, curSubgroupsIdx, curGroupIdx,
805 : groupOriginOffset_[curSubgroupsIdx], curSlice.offset, groupOriginOffset_[curSubgroupsIdx], curSlice.size);
806 : }
807 : }
808 0 : interLinksVector.push_back(interLinks);
809 0 : interSlicesVector.push_back(interSlices);
810 0 : }
811 :
812 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u] end inter", rank);
813 0 : return HCCL_SUCCESS;
814 0 : }
815 :
816 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinksForAR(const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK> &links,
817 : std::vector<std::vector<LINK>> &interLinksVector, std::vector<std::vector<Slice>> &interSlicesVector)
818 : {
819 : // 查找自己位于组内的第几个rank
820 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] begin calc inter slices and links rank[%u]", rank);
821 :
822 0 : u32 intraRank = GetIntraRank(rank);
823 :
824 0 : std::vector<Slice> intraSlices;
825 0 : std::vector<LINK> intraLinks;
826 0 : CHK_RET(CalcIntraSlicesAndLinks(rank, dataUnitSize, count, links, intraLinks, intraSlices));
827 :
828 : // 计算组间每个rank结果上的offset和size
829 : u64 sliceSizeCalculated =
830 0 : (intraSlices[intraRank].size / dataUnitSize + (static_cast<u32>(subGroups_.size()) - 1)) / subGroups_.size() * dataUnitSize;
831 0 : u64 totalSize = intraSlices[intraRank].size;
832 0 : u64 residueSize = totalSize;
833 : u64 sliceSizeAligned;
834 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
835 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
836 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
837 : } else {
838 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
839 : }
840 :
841 0 : std::vector<LINK> interLinks;
842 0 : std::vector<Slice> interSlices;
843 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
844 0 : interLinks.push_back(links[subGroups_[i][intraRank]]);
845 0 : Slice slice;
846 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
847 0 : slice.offset = intraSlices[intraRank].offset + totalSize - residueSize;
848 0 : residueSize -= slice.size;
849 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
850 : rank, i, slice.offset, i, slice.size);
851 0 : interSlices.push_back(slice);
852 : }
853 0 : interLinksVector.push_back(interLinks);
854 0 : interSlicesVector.push_back(interSlices);
855 :
856 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] end calc inter slices and links rank[%u]", rank);
857 0 : return HCCL_SUCCESS;
858 0 : }
859 :
860 0 : CommAHCAlignInfo::CommAHCAlignInfo(const std::vector<std::vector<u32>> &subGroups)
861 0 : : CommAHCBaseInfo(subGroups)
862 : {
863 0 : }
864 :
865 0 : CommAHCAlignInfo::~CommAHCAlignInfo()
866 : {
867 0 : }
868 :
869 0 : HcclResult CommAHCAlignInfo::Init(AHCOpType opType, std::map<AHCConcOpType, TemplateType> &ahcAlgOption)
870 : {
871 0 : ahcAlgOption_= ahcAlgOption;
872 :
873 : // 参数检查
874 0 : opType_ = opType;
875 0 : CHK_RET(CheckSubGroups(subGroups_));
876 :
877 : //初始化slice相关信息
878 0 : CHK_RET(InitSliceInfo());
879 :
880 : //计算 logicCard 相关信息;
881 0 : InitLogicCardInfo();
882 :
883 : //初始化相关Map信息
884 0 : CHK_RET(InitMapInfo());
885 :
886 0 : return HCCL_SUCCESS;
887 : }
888 :
889 0 : HcclResult CommAHCAlignInfo::InitSliceInfo()
890 : {
891 : // 计算 totalSliceSegment_ ,即所有分组大小的最小公倍数, 以及 interRankOrder
892 0 : totalSliceSegment_ = subGroups_[0].size();
893 : //u32 groupSizeGcd;
894 0 : for (u32 i = 1; i < subGroups_.size(); ++i) {
895 0 : u32 groupSize = static_cast<u32>(subGroups_[i].size());
896 0 : totalSliceSegment_ = totalSliceSegment_ * groupSize / std::__gcd(totalSliceSegment_, groupSize);
897 : }
898 0 : globalTotalSliceSegment_ = rankSize_ * totalSliceSegment_;
899 0 : HCCL_DEBUG("[CommAHCAlignInfo][InitSliceInfo] totalSliceSegment [%u]", totalSliceSegment_);
900 :
901 : //计算 logicCardSliceSize_ ;
902 0 : std::set<u32> sliceOffset;
903 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
904 0 : for (u32 j = 0; j < subGroups_[i].size(); ++j) {
905 0 : u32 rankSliceSize = (totalSliceSegment_ / subGroups_[i].size()) * (j + 1);
906 0 : sliceOffset.insert(rankSliceSize);
907 0 : HCCL_DEBUG("[CommAHCAlignInfo][InitSliceInfo] sliceOffset [%u]",rankSliceSize);
908 : }
909 : }
910 0 : sliceOffset.insert(static_cast<u32>(0));
911 0 : logicCardSliceOffset_.resize(sliceOffset.size());
912 0 : std::copy(sliceOffset.begin(), sliceOffset.end(), logicCardSliceOffset_.begin());
913 :
914 0 : std::vector<u32>::iterator itPre = logicCardSliceOffset_.begin();
915 0 : std::vector<u32>::iterator itNext = logicCardSliceOffset_.begin();
916 0 : itNext++;
917 0 : while(itNext != logicCardSliceOffset_.end()) {
918 0 : auto boundDiff = (*itNext) - (*itPre);
919 0 : logicCardSliceSize_.push_back(boundDiff);
920 0 : itPre++;
921 0 : itNext++;
922 : }
923 :
924 0 : CHK_PRT_RET(logicCardSliceSize_.size() !=(logicCardSliceOffset_.size() - 1),
925 : HCCL_ERROR("[CommAHCAlignInfo][InitSliceInfo] cardOffset size [%u] cardSize size [%u] check error",
926 : logicCardSliceSize_.size(),logicCardSliceOffset_.size() ), HCCL_E_INTERNAL);
927 :
928 0 : return HCCL_SUCCESS;
929 0 : }
930 :
931 0 : HcclResult CommAHCAlignInfo::InitLogicCardInfo()
932 : {
933 : //计算 logicCardCommGroups_;
934 0 : for (std::vector<u32>::iterator it = (logicCardSliceOffset_.begin() + 1); it != logicCardSliceOffset_.end(); ++it) {
935 0 : std::vector<u32> logicGroup;
936 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
937 : u32 logicRank;
938 0 : if ((*it) % (totalSliceSegment_ / subGroups_[i].size()) != 0) {
939 0 : logicRank = (*it) / (totalSliceSegment_ / subGroups_[i].size()) + 1;
940 : } else {
941 0 : logicRank = (*it) / (totalSliceSegment_ / subGroups_[i].size());
942 : }
943 0 : logicGroup.push_back(subGroups_[i][logicRank - 1]);
944 : }
945 0 : logicCardCommGroups_.push_back(logicGroup);
946 0 : }
947 :
948 : //计算 logicCardGroup_
949 0 : u32 curRank = subGroups_[minSubGroupIdx_][0];
950 0 : u32 curOffset = 0;
951 0 : std::vector<u32>::iterator it = logicCardSliceOffset_.begin();
952 0 : u32 curIdx = 0;
953 0 : u32 curLogicIdx = 0;
954 0 : logicCardGroup_.resize(static_cast<u32>(subGroups_[minSubGroupIdx_].size()));
955 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
956 0 : if (logicCardCommGroups_[i][minSubGroupIdx_] != curRank) {
957 0 : logicCardGroup_[curLogicIdx].resize(i - curIdx);
958 0 : for (u32 j = 0; j < i - curIdx; j++) {
959 0 : logicCardGroup_[curLogicIdx][j] = curIdx + j;
960 : }
961 0 : curIdx = i;
962 0 : curLogicIdx++;
963 0 : curRank = logicCardCommGroups_[i][minSubGroupIdx_];
964 0 : curOffset = *it;
965 : }
966 0 : logicCardExecuteOffset_.push_back(*it - curOffset);
967 0 : it++;
968 : }
969 0 : if (curIdx != logicCardCommGroups_.size() - 1) {
970 0 : logicCardGroup_[curLogicIdx].resize(logicCardCommGroups_.size() - curIdx);
971 0 : for (u32 i = 0; i < logicCardCommGroups_.size() - curIdx; i++) {
972 0 : logicCardGroup_[curLogicIdx][i] = curIdx + i;
973 : }
974 : }
975 0 : return HCCL_SUCCESS;
976 : }
977 :
978 0 : bool CommAHCAlignInfo::CompareLogicCardExcuteOrder(u32 i, u32 j)
979 : {
980 0 : return logicCardExecuteOffset_[i] < logicCardExecuteOffset_[j];
981 : }
982 :
983 0 : HcclResult CommAHCAlignInfo::InitMapInfo()
984 : {
985 : //rank 到 logicCardOrder 初始化
986 0 : std::map<u32, u32> interRankOrder;
987 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
988 0 : interRankOrder.insert(std::make_pair(i, i));
989 : }
990 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
991 0 : interRankList_.push_back(interRankOrder);
992 0 : for (u32 j = 0; j < logicCardCommGroups_[i].size(); ++j) {
993 0 : rankLogicCardOrderMap_[logicCardCommGroups_[i][j]].push_back(i);
994 0 : rankLogicCardMap_[logicCardCommGroups_[i][j]].push_back(i);
995 : }
996 : }
997 :
998 : //定义 lambda 将对象指针传递到成员函数
999 0 : auto sortLambda = [this](u32 i, u32 j) {
1000 0 : return this->CompareLogicCardExcuteOrder(i,j);
1001 0 : };
1002 :
1003 : //rankLogicCardOrderMap_ 内的逻辑同号卡list按照 logicCardExecuteOffset_ 并发流开始时间排序
1004 0 : for (auto iter = rankLogicCardOrderMap_.begin(); iter != rankLogicCardOrderMap_.end(); iter++) {
1005 0 : std::vector<u32> &rankLogicCardList = iter->second;
1006 0 : std::sort(rankLogicCardList.begin(), rankLogicCardList.end(), sortLambda);
1007 : }
1008 :
1009 0 : return HCCL_SUCCESS;
1010 0 : }
1011 :
1012 : // 配置当前需要的 globalTotalSliceSegment_,用于 Multi-AllReduce 中
1013 0 : HcclResult CommAHCAlignInfo::SetGlobalTotalSliceSegment(u64 globalTotalSliceSegment)
1014 : {
1015 0 : globalTotalSliceSegment_ = globalTotalSliceSegment;
1016 0 : HCCL_DEBUG("[CommAHCAlignInfo][setGlobalTotalSliceSegment] globalTotalSliceSegment set to [%llu]", globalTotalSliceSegment_);
1017 0 : return HCCL_SUCCESS;
1018 : }
1019 :
1020 : //获取当前rank对应的多个逻辑同号卡,并且按照并发流的开始执行时间排序
1021 0 : HcclResult CommAHCAlignInfo::GetLogicCardExecuteOrder(u32 rank, std::vector<u32> &executeOrder)
1022 : {
1023 0 : executeOrder = rankLogicCardOrderMap_[rank];
1024 0 : return HCCL_SUCCESS;
1025 : }
1026 :
1027 0 : HcclResult CommAHCAlignInfo::SliceSizeAlignBound(Slice &slice, u64 offsetCount, u64 sliceSizeCalculated, const u64 boundSize, u32 boundOffsetCount, u32 &curOffset) const
1028 : {
1029 0 : u64 sliceSize = offsetCount * sliceSizeCalculated;
1030 0 : if (!isAlignBound_) {
1031 : // 对于 All-Reduce 中的 Reduce-Scatter 以及 All-Gather,不需要严格对齐bound
1032 0 : slice.size = slice.size + sliceSize;
1033 0 : curOffset = curOffset + offsetCount;
1034 0 : return HCCL_SUCCESS;
1035 : }
1036 0 : if (offsetCount < boundOffsetCount) {
1037 0 : if (sliceSize <= ((curOffset / boundOffsetCount + 1) * boundSize - (slice.size + slice.offset))) {
1038 0 : slice.size = slice.size + sliceSize;
1039 : } else {
1040 0 : slice.size = slice.size + ((curOffset / boundOffsetCount + 1) * boundSize - (slice.size + slice.offset));
1041 : }
1042 : } else {
1043 0 : slice.size = slice.size + (offsetCount / boundOffsetCount) * boundSize;
1044 : }
1045 0 : curOffset = curOffset + offsetCount;
1046 0 : return HCCL_SUCCESS;
1047 : }
1048 :
1049 : // Reduce-Scatter 及 All-Gather 组内切片逻辑
1050 0 : HcclResult CommAHCAlignInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
1051 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &intraLinksVector,
1052 : std::vector<std::vector<Slice>> &intraSlicesVector)
1053 : {
1054 : // Boundary 指在 RS 及 AG 中单卡应有的数据量的 offset, 如八卡跑8K,boundary 为 1024
1055 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
1056 :
1057 0 : u32 singleRankOffset = globalTotalSliceSegment_ / rankSize_; // 每个 rank 最后结果应有的小块数据份数
1058 0 : u64 sliceSizeCalculated = (totalSize_ / dataUnitSize + globalTotalSliceSegment_ - 1)
1059 0 : / globalTotalSliceSegment_ * dataUnitSize * (globalTotalSliceSegment_ / rankSize_ / totalSliceSegment_);
1060 0 : u64 totalSize = totalSize_;
1061 0 : u64 residueSize = totalSize;
1062 0 : HCCL_DEBUG("[CommAHCAlignInfo][AHCDEBUG] count[%u] ranksize[%u] sliceSizeCalculated[%u] totalSize[%u] globalTotalSliceSegment[%u] totalSliceSegment[%u] dataUnitSize[%u]",
1063 : count, rankSize_, sliceSizeCalculated, totalSize, globalTotalSliceSegment_, totalSliceSegment_, dataUnitSize);
1064 :
1065 0 : u32 curOffset = 0;
1066 0 : for (u32 k = 0; k < subGroups_.size(); ++k) {
1067 0 : std::vector<Slice> intraSlices;
1068 0 : std::vector<LINK> intraLinks;
1069 0 : std::vector<u32> curLogicCardGroup = rankLogicCardMap_[subGroups_[rankGroupMap_[rank]][0]]; // 获取当前rank对应的逻辑同号组
1070 0 : u32 singleSliceOffset = logicCardSliceSize_[curLogicCardGroup[0]];
1071 0 : for (u32 j = 1; j < curLogicCardGroup.size(); ++j) {
1072 0 : singleSliceOffset = singleSliceOffset + logicCardSliceSize_[curLogicCardGroup[j]];
1073 : }
1074 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
1075 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
1076 0 : intraLinks.push_back(links[curRank]);
1077 0 : Slice slice;
1078 0 : slice.size = 0;
1079 0 : slice.offset = totalSize - residueSize;
1080 0 : u64 targeOffset = singleSliceOffset * subGroups_[k].size();
1081 0 : u64 offsetCountBeforeBoundary = ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) < targeOffset ?
1082 : ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) : targeOffset;
1083 0 : SliceSizeAlignBound(slice, offsetCountBeforeBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1084 0 : u64 offsetCountCrossBoundary = (targeOffset - offsetCountBeforeBoundary) / singleRankOffset * singleRankOffset;
1085 0 : SliceSizeAlignBound(slice, offsetCountCrossBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1086 0 : u64 offsetCountBehindBoundary = (targeOffset - offsetCountBeforeBoundary - offsetCountCrossBoundary) % singleRankOffset;
1087 0 : SliceSizeAlignBound(slice, offsetCountBehindBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1088 0 : slice.size = (residueSize > slice.size) ? slice.size : residueSize;
1089 0 : residueSize -= slice.size;
1090 0 : intraSlices.push_back(slice);
1091 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] rank[%u], singleSliceOffset[%u], subGroups_[%u].size()[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
1092 : rank, singleSliceOffset, k, subGroups_[k].size(), i, slice.offset, i, slice.size);
1093 : }
1094 0 : intraLinksVector.push_back(intraLinks);
1095 0 : intraSlicesVector.push_back(intraSlices);
1096 0 : }
1097 :
1098 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
1099 0 : return HCCL_SUCCESS;
1100 : }
1101 :
1102 : // All-Reduce 组内切片逻辑
1103 0 : HcclResult CommAHCAlignInfo::CalcIntraSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
1104 : const std::vector<LINK> &links, std::vector<LINK> &intraLinks, std::vector<Slice> &intraSlices)
1105 : {
1106 : // 计算组内每个rank结果上的offset和size
1107 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
1108 :
1109 0 : u64 sliceSizeCalculated = (count + (totalSliceSegment_ * static_cast<u32>(subGroups_.size()) - 1))
1110 0 : / (totalSliceSegment_ * subGroups_.size()) * dataUnitSize;
1111 0 : u64 totalSize = count * dataUnitSize;
1112 0 : u64 residueSize = totalSize;
1113 0 : u64 sliceSizeAligned = sliceSizeCalculated;
1114 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
1115 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
1116 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1117 : } else {
1118 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
1119 : }
1120 0 : sliceSizeAligned = sliceSizeAligned * static_cast<u32>(subGroups_.size()) *
1121 0 : (totalSliceSegment_ / static_cast<u32>(subGroups_[rankGroupMap_[rank]].size()));
1122 :
1123 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
1124 0 : intraLinks.push_back(links[subGroups_[rankGroupMap_[rank]][i]]);
1125 0 : Slice slice;
1126 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
1127 0 : slice.offset = totalSize - residueSize;
1128 0 : residueSize -= slice.size;
1129 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
1130 : rank, i, slice.offset, i, slice.size);
1131 0 : intraSlices.push_back(slice);
1132 : }
1133 :
1134 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
1135 0 : return HCCL_SUCCESS;
1136 : }
1137 :
1138 : // 组间切片逻辑统一对外接口
1139 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinks(const u32 rank, const u32 dataUnitSize, const u64 count,
1140 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
1141 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
1142 : {
1143 0 : HcclResult ret = HCCL_SUCCESS;
1144 0 : switch (opType_) {
1145 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE:
1146 0 : ret = CalcInterSlicesAndLinksForAR(rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
1147 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
1148 : HCCL_ERROR("[CommAHCAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks step",
1149 : rank, count), ret);
1150 0 : break;
1151 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER:
1152 : case AHCOpType::AHC_OP_TYPE_ALLGATHER:
1153 0 : ret = CalcInterSlicesAndLinksForRS(rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
1154 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
1155 : HCCL_ERROR("[CommAHCAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks step",
1156 : rank, count), ret);
1157 0 : break;
1158 0 : default:
1159 0 : ret = HCCL_SUCCESS;
1160 : }
1161 0 : return ret;
1162 : }
1163 :
1164 0 : HcclResult CommAHCAlignInfo::PrepareIntraSlices(const u32 rank, const u32 dataUnitSize, const u64 count,
1165 : std::vector<std::vector<Slice>> &intraSlicesVector)
1166 : {
1167 : // 计算组内每个rank结果上的offset和size
1168 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareIntraSlices] begin calc intra slices and links rank[%u]", rank);
1169 :
1170 0 : u32 singleRankOffset = globalTotalSliceSegment_ / rankSize_;
1171 0 : u64 sliceSizeCalculated = (totalSize_ / dataUnitSize + globalTotalSliceSegment_ - 1)
1172 0 : / globalTotalSliceSegment_ * dataUnitSize * (globalTotalSliceSegment_ / rankSize_ / totalSliceSegment_);
1173 0 : u64 totalSize = totalSize_;
1174 0 : u64 residueSize = totalSize;
1175 0 : HCCL_DEBUG("[CommAHCAlignInfo][AHCDEBUG] count[%u] ranksize[%u] sliceSizeCalculated[%u] totalSize[%u]",
1176 : count, rankSize_, sliceSizeCalculated, totalSize);
1177 :
1178 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
1179 0 : std::vector<Slice> intraSlices;
1180 0 : intraSlicesVector.push_back(intraSlices);
1181 0 : }
1182 :
1183 0 : u32 curOffset = 0;
1184 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
1185 0 : for (u32 j = 0; j < logicCardCommGroups_.size(); ++j) {
1186 0 : Slice slice;
1187 0 : slice.size = 0;
1188 0 : slice.offset = totalSize - residueSize;
1189 0 : u64 targeOffset = logicCardSliceSize_[j] * subGroups_[i].size();
1190 0 : u64 offsetCountBeforeBoundary = ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) < targeOffset ?
1191 : ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) : targeOffset;
1192 0 : SliceSizeAlignBound(slice, offsetCountBeforeBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1193 0 : u64 offsetCountCrossBoundary = (targeOffset - offsetCountBeforeBoundary) / singleRankOffset * singleRankOffset;
1194 0 : SliceSizeAlignBound(slice, offsetCountCrossBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1195 0 : u64 offsetCountBehindBoundary = (targeOffset - offsetCountBeforeBoundary - offsetCountCrossBoundary) % singleRankOffset;
1196 0 : SliceSizeAlignBound(slice, offsetCountBehindBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset, curOffset);
1197 0 : slice.size = (residueSize > slice.size) ? slice.size : residueSize;
1198 0 : residueSize -= slice.size;
1199 0 : intraSlicesVector[j].push_back(slice);
1200 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareIntraSlices] rank[%u], round[%u], logicGroup[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
1201 : rank, i, j, j, slice.offset, j, slice.size);
1202 : }
1203 : }
1204 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareIntraSlices] end calc intra slices and links rank[%u]", rank);
1205 0 : return HCCL_SUCCESS;
1206 : }
1207 :
1208 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinksForRS(const u32 rank, const u32 dataUnitSize, const u64 count,
1209 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
1210 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
1211 : {
1212 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] begin calc inter slices and links rank[%u]", rank);
1213 0 : std::vector<std::vector<Slice>> intraSlicesVecotr;
1214 :
1215 0 : CHK_RET(PrepareIntraSlices(rank, dataUnitSize, count, intraSlicesVecotr));
1216 0 : GetLogicCardExecuteOrder(rank, logicCardList);
1217 :
1218 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1219 0 : u32 logicGroupIdx = logicCardList[i]; // 获取当前处理的目标逻辑同号卡组的下标
1220 0 : std::vector<Slice> curIntraSliceVector = intraSlicesVecotr[logicGroupIdx];
1221 0 : std::vector<Slice> interSlices;
1222 0 : std::vector<LINK> interLinks;
1223 0 : for (u32 j = 0; j < subGroups_.size(); j++) {
1224 0 : Slice curSlice = curIntraSliceVector[j];
1225 0 : interLinks.push_back(links[logicCardCommGroups_[logicGroupIdx][j]]); // 当前处理的逻辑同号卡
1226 0 : interSlices.push_back(curSlice);
1227 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], logicGroup[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
1228 : rank, logicCardCommGroups_[logicGroupIdx][j], logicGroupIdx, i, curSlice.offset, i, curSlice.size);
1229 : }
1230 0 : interSlicesVector.push_back(interSlices);
1231 0 : interLinksVector.push_back(interLinks);
1232 0 : }
1233 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] calc inter slices and links rank[%u] end", rank);
1234 0 : return HCCL_SUCCESS;
1235 0 : }
1236 :
1237 0 : HcclResult CommAHCAlignInfo::PrepareWholeLogicSlices(const Slice &intraSlice, const u64 sliceSizeAligned, const u32 originOffset,
1238 : std::vector<Slice> &logicGroupSlice, std::vector<u32> &logicCardList)
1239 : {
1240 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1241 0 : Slice logicSlice;
1242 0 : u32 logicRank = logicCardList[i];
1243 : // 计算当前逻辑同号组的offset大小
1244 0 : u32 offsetDiff = logicCardSliceOffset_[logicRank + 1] - logicCardSliceOffset_[logicRank];
1245 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice begin, logicRank : [%u]," \
1246 : "offsetDiff : [%u], offset_next : [%u], offset_cur[%u]", logicRank, offsetDiff,
1247 : logicCardSliceOffset_[logicRank + 1], logicCardSliceOffset_[logicRank]);
1248 :
1249 0 : logicSlice.size = sliceSizeAligned * offsetDiff;
1250 0 : logicSlice.offset = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1251 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareFullLogicSlices] logicGroupSlice end, logicRank : [%u] ," \
1252 : "size : [%u], offset : [%u] ", logicRank, logicSlice.size, logicSlice.offset);
1253 0 : logicGroupSlice.push_back(logicSlice);
1254 : }
1255 0 : return HCCL_SUCCESS;
1256 : }
1257 :
1258 0 : HcclResult CommAHCAlignInfo::PreparePartialLogicSlices(const Slice &intraSlice, const u64 sliceSizeAligned, const u32 originOffset,
1259 : std::vector<Slice> &logicGroupSlice, std::vector<u32> &logicCardList)
1260 : {
1261 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1262 0 : Slice logicSlice;
1263 0 : u32 logicRank = logicCardList[i];
1264 : // 计算当前逻辑同号组的offset大小
1265 0 : u32 offsetDiff = logicCardSliceOffset_[logicRank + 1] - logicCardSliceOffset_[logicRank];
1266 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice begin, logicRank : [%u]," \
1267 : "offsetDiff : [%u], offset_next : [%u], offset_cur[%u]", logicRank, offsetDiff,
1268 : logicCardSliceOffset_[logicRank + 1], logicCardSliceOffset_[logicRank]);
1269 :
1270 : // 当前rank在组内对应的offset能获取到完全的数据,即前几个逻辑同号卡
1271 0 : if ((logicCardSliceOffset_[logicRank + 1] - originOffset) <= intraSlice.size / sliceSizeAligned) {
1272 0 : logicSlice.size = sliceSizeAligned * offsetDiff;
1273 0 : logicSlice.offset = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1274 : // 当前rank在组内对应的offset能获取到部分的数据,即边界上的逻辑同号卡
1275 0 : } else if ((logicCardSliceOffset_[logicRank] - originOffset) <= intraSlice.size / sliceSizeAligned){
1276 0 : logicSlice.size = intraSlice.size - (logicCardSliceOffset_[logicRank] - originOffset) * sliceSizeAligned;
1277 0 : logicSlice.offset = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1278 : // 当前rank在组内对应的offset不能获取到数据,即最后的逻辑同号卡
1279 : } else {
1280 0 : logicSlice.size = 0;
1281 0 : logicSlice.offset = 0;
1282 : }
1283 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice end, logicRank : [%u] ," \
1284 : "size : [%u], offset : [%u] ",logicRank, logicSlice.size, logicSlice.offset);
1285 :
1286 0 : logicGroupSlice.push_back(logicSlice);
1287 : }
1288 0 : return HCCL_SUCCESS;
1289 : }
1290 :
1291 0 : HcclResult CommAHCAlignInfo::PrepareEmptyLogicSlices(std::vector<Slice> &logicGroupSlice,
1292 : const std::vector<u32> &logicCardList) const
1293 : {
1294 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1295 0 : Slice logicSlice;
1296 0 : logicSlice.size = 0;
1297 0 : logicSlice.offset = 0;
1298 0 : logicGroupSlice.push_back(logicSlice);
1299 : }
1300 0 : return HCCL_SUCCESS;
1301 : }
1302 :
1303 0 : HcclResult CommAHCAlignInfo::CalcLogicSlicesAndLinks(std::vector<Slice> &logicGroupSlice, std::vector<u32> &logicCardList,
1304 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
1305 : std::vector<std::vector<Slice>> &interSlicesVector)
1306 : {
1307 0 : for (u32 i = 0; i < logicGroupSlice.size(); i++) {
1308 0 : Slice logicSlice = logicGroupSlice[i];
1309 0 : std::vector<Slice> interSlices;
1310 0 : std::vector<LINK> interLinks;
1311 0 : u32 logicRank = logicCardList[i];
1312 0 : u64 totalSize = logicSlice.size;
1313 0 : u64 residueSize = totalSize;
1314 : u64 logicSliceSizeAligned;
1315 0 : if (logicSlice.size % subGroups_.size() == 0 && logicSlice.size % HCCL_MIN_SLICE_ALIGN == 0) {
1316 0 : logicSliceSizeAligned = logicSlice.size / subGroups_.size();
1317 : } else {
1318 0 : u64 sliceSizeCalculated = (logicSlice.size + static_cast<u32>(subGroups_.size()) - 1) / static_cast<u32>(subGroups_.size());
1319 0 : logicSliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1320 : }
1321 0 : for (u32 j = 0; j < subGroups_.size(); j++) {
1322 0 : u32 curRank = logicCardCommGroups_[logicRank][j];
1323 0 : Slice slice;
1324 0 : interLinks.push_back(links[curRank]);
1325 0 : slice.size = (residueSize > logicSliceSizeAligned) ? logicSliceSizeAligned : residueSize;
1326 0 : slice.offset = logicSlice.offset + totalSize - residueSize;
1327 0 : residueSize -= slice.size;
1328 0 : interSlices.push_back(slice);
1329 : }
1330 0 : interLinksVector.push_back(interLinks);
1331 0 : interSlicesVector.push_back(interSlices);
1332 0 : }
1333 0 : return HCCL_SUCCESS;
1334 : }
1335 :
1336 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinksForAR(const u32 rank, const u32 dataUnitSize, const u64 count,
1337 : const std::vector<LINK> &links, std::vector<std::vector<LINK>> &interLinksVector,
1338 : std::vector<std::vector<Slice>> &interSlicesVector, std::vector<u32> &logicCardList)
1339 : {
1340 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] begin calc inter slices and links rank[%u]", rank);
1341 :
1342 0 : u32 intraRank = GetIntraRank(rank);
1343 :
1344 0 : std::vector<Slice> intraSlices;
1345 0 : std::vector<LINK> intraLinks;
1346 :
1347 0 : CHK_RET(CalcIntraSlicesAndLinks(rank, dataUnitSize, count, links, intraLinks, intraSlices));
1348 0 : GetLogicCardExecuteOrder(rank, logicCardList);
1349 :
1350 : // 计算当前rank逻辑同号卡之间最小slice的大小
1351 0 : u64 sliceSizeCalculated = (count + (totalSliceSegment_ * static_cast<u32>(subGroups_.size()) - 1))
1352 0 : / (totalSliceSegment_ * subGroups_.size()) * dataUnitSize;
1353 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
1354 0 : u64 sliceSizeAligned = sliceSizeCalculated;
1355 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
1356 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1357 : } else {
1358 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
1359 : }
1360 :
1361 0 : sliceSizeAligned = sliceSizeAligned * static_cast<u32>(subGroups_.size());
1362 0 : u32 originOffset = intraRank * totalSliceSegment_ / subGroups_[rankGroupMap_[rank]].size(); // 当前rank起始offset
1363 :
1364 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] rank : [%u], intraslice.size : [%u], intraslice.offset : [%u]," \
1365 : "sliceSizeAligned : [%u], originOffset : [%u]", rank, intraSlices[intraRank].size, intraSlices[intraRank].offset,
1366 : sliceSizeAligned, originOffset);
1367 :
1368 : // 进行逻辑同号组对应的slice切分
1369 0 : std::vector<Slice> logicGroupSlice; // 逻辑同号组对应的slice
1370 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] check rank : [%u], intraSlices[%u].size : [%u], sliceSizeAligned : [%u]," \
1371 : "totalSliceSegment_ : [%u], subGroups_[rankGroupMap_[rank]].size : [%u]", rank, intraRank, intraSlices[intraRank].size,
1372 : sliceSizeAligned, totalSliceSegment_, subGroups_[rankGroupMap_[rank]].size());
1373 :
1374 : // 当前rank有完整的对齐后的数据量
1375 0 : if (intraSlices[intraRank].size / sliceSizeAligned == totalSliceSegment_ / subGroups_[rankGroupMap_[rank]].size()) {
1376 0 : CHK_RET(PrepareWholeLogicSlices(intraSlices[intraRank], sliceSizeAligned, originOffset, logicGroupSlice, logicCardList));
1377 : // 当前rank有不完整的数据量
1378 0 : } else if (intraSlices[intraRank].size != 0) {
1379 0 : CHK_RET(PreparePartialLogicSlices(intraSlices[intraRank], sliceSizeAligned, originOffset, logicGroupSlice, logicCardList));
1380 : } else {
1381 0 : CHK_RET(PrepareEmptyLogicSlices(logicGroupSlice, logicCardList));
1382 : }
1383 :
1384 : // 计算当前rank逻辑同号组之间的slice大小
1385 0 : CHK_RET(CalcLogicSlicesAndLinks(logicGroupSlice, logicCardList, links, interLinksVector, interSlicesVector));
1386 :
1387 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinks] end calc inter slices and links rank[%u]", rank);
1388 0 : return HCCL_SUCCESS;
1389 0 : }
1390 :
1391 : } // ~~ namespace hccl
|