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