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 47 : AHCCommCalcFuncRegistry::AHCCommCalcFuncRegistry()
22 : {
23 47 : commCalcFuncCreators_.resize(static_cast<u32>(AHCTemplateType::AHC_TEMPLATE_RESERVED), nullptr);
24 47 : }
25 :
26 141 : AHCCommCalcFuncRegistry& AHCCommCalcFuncRegistry::Instance()
27 : {
28 141 : static AHCCommCalcFuncRegistry globalAlgTemplateRegistry;
29 141 : return globalAlgTemplateRegistry;
30 : }
31 :
32 141 : HcclResult AHCCommCalcFuncRegistry::Register(AHCTemplateType type, AHCCommCalcFuncPtr funPtr)
33 : {
34 141 : 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 141 : const std::lock_guard<std::mutex> lock(mu_);
40 141 : 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 141 : commCalcFuncCreators_[static_cast<u32>(type)] = funPtr;
45 141 : return HcclResult::HCCL_SUCCESS;
46 141 : }
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, [[maybe_unused]] 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, [[maybe_unused]] const u32 dataUnitSize, [[maybe_unused]] const u64 count,
666 : const std::vector<LINK>& links, std::vector<std::vector<LINK>>& intraLinksVector,
667 : std::vector<std::vector<Slice>>& intraSlicesVector)
668 : {
669 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
670 :
671 0 : u64 sliceSizeAligned = totalSize_ / rankSize_;
672 0 : u64 curoffset = 0;
673 :
674 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] calculate sliceSizeAligned[%llu]", sliceSizeAligned);
675 :
676 0 : for (u32 k = 0; k < subGroups_.size(); ++k) {
677 : // 满片分组处理过程
678 0 : for (u32 j = 0; j < subGroups_[k].size() / subGroups_[rankGroupMap_[rank]].size(); ++j) {
679 0 : std::vector<Slice> intraSlices;
680 0 : std::vector<LINK> intraLinks;
681 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
682 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
683 0 : intraLinks.push_back(links[curRank]);
684 0 : Slice slice;
685 0 : slice.size = sliceSizeAligned;
686 0 : slice.offset = curoffset;
687 0 : curoffset = curoffset + slice.size;
688 0 : intraSlices.push_back(slice);
689 0 : HCCL_DEBUG(
690 : "[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], link[%u] slices[%u].offset=%llu, "
691 : "slices[%u].size=%llu",
692 : rank, curRank, i, slice.offset, i, slice.size);
693 : }
694 0 : intraLinksVector.push_back(intraLinks);
695 0 : intraSlicesVector.push_back(intraSlices);
696 0 : }
697 0 : std::vector<Slice> intraSlices;
698 0 : std::vector<LINK> intraLinks;
699 : // 涉及空片分组非零切片处理过程
700 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
701 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
702 0 : intraLinks.push_back(links[curRank]);
703 0 : Slice slice;
704 0 : slice.size = i < subGroups_[k].size() % subGroups_[rankGroupMap_[rank]].size() ? sliceSizeAligned : 0;
705 0 : slice.offset = curoffset;
706 0 : curoffset = curoffset + slice.size;
707 0 : intraSlices.push_back(slice);
708 0 : HCCL_DEBUG(
709 : "[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], link[%u] slices[%u].offset=%llu, "
710 : "slices[%u].size=%llu",
711 : rank, curRank, i, slice.offset, i, slice.size);
712 : }
713 0 : intraLinksVector.push_back(intraLinks);
714 0 : intraSlicesVector.push_back(intraSlices);
715 0 : }
716 :
717 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
718 0 : return HCCL_SUCCESS;
719 : }
720 :
721 : // All-Reduce 组内切片逻辑
722 0 : HcclResult CommBrokeAlignInfo::CalcIntraSlicesAndLinks(
723 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
724 : std::vector<LINK>& intraLinks, std::vector<Slice>& intraSlices)
725 : {
726 : // 计算组内每个rank结果上的offset和size
727 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
728 :
729 0 : u64 sliceSizeCalculated = (count + (static_cast<u32>(subGroups_[minSubGroupIdx_].size()) - 1))
730 0 : / subGroups_[minSubGroupIdx_].size() * dataUnitSize;
731 0 : u64 totalSize = count * dataUnitSize;
732 0 : u64 residueSize = totalSize;
733 : u64 sliceSizeAligned;
734 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
735 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
736 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
737 : } else {
738 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
739 : }
740 :
741 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
742 0 : intraLinks.push_back(links[subGroups_[rankGroupMap_[rank]][i]]);
743 0 : Slice slice;
744 0 : if (i < subGroups_[minSubGroupIdx_].size()) {
745 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
746 0 : slice.offset = totalSize - residueSize;
747 0 : residueSize -= slice.size;
748 : } else {
749 0 : slice.size = 0;
750 0 : slice.offset = totalSize - residueSize;
751 : }
752 0 : HCCL_DEBUG(
753 : "[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
754 : rank, i, slice.offset, i, slice.size);
755 0 : intraSlices.push_back(slice);
756 : }
757 :
758 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
759 :
760 0 : return HCCL_SUCCESS;
761 : }
762 :
763 : // 组间切片逻辑统一对外接口
764 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinks(
765 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
766 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector,
767 : std::vector<u32>& logicCardList)
768 : {
769 0 : HcclResult ret = HCCL_SUCCESS;
770 0 : switch (opType_) {
771 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE:
772 0 : ret = CalcInterSlicesAndLinksForAR(rank, dataUnitSize, count, links, interLinksVector, interSlicesVector);
773 0 : CHK_PRT_RET(
774 : ret != HCCL_SUCCESS,
775 : HCCL_ERROR(
776 : "[CommBrokeAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in "
777 : "CalcInterSlicesAndLinks step",
778 : rank, count),
779 : ret);
780 0 : break;
781 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER:
782 : case AHCOpType::AHC_OP_TYPE_ALLGATHER:
783 0 : ret = CalcInterSlicesAndLinksForRS(
784 : rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
785 0 : CHK_PRT_RET(
786 : ret != HCCL_SUCCESS,
787 : HCCL_ERROR(
788 : "[CommBrokeAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in "
789 : "CalcInterSlicesAndLinks step",
790 : rank, count),
791 : ret);
792 0 : break;
793 0 : default:
794 0 : ret = HCCL_SUCCESS;
795 : }
796 0 : return ret;
797 : }
798 :
799 0 : HcclResult CommBrokeAlignInfo::PrepareIntraSlices(
800 : const u32 rank, const u32 dataUnitSize, const u64 count, std::vector<Slice>& intraSlices) const
801 : {
802 : (void)dataUnitSize;
803 : (void)count;
804 :
805 : // 计算组内每个rank结果上的offset和size
806 0 : HCCL_DEBUG(
807 : "[CommBrokeAlignInfo][PrepareIntraSlices] begin calc intra slices and links rank[%u] ranksize[%u]", rank,
808 : rankSize_);
809 :
810 0 : u64 sliceSizeAligned = totalSize_ / rankSize_;
811 0 : u64 curoffset = 0;
812 :
813 0 : for (u32 i = 0; i < rankSize_; ++i) {
814 0 : Slice slice;
815 0 : slice.size = sliceSizeAligned;
816 0 : slice.offset = curoffset;
817 0 : curoffset = curoffset + slice.size;
818 0 : intraSlices.push_back(slice);
819 0 : HCCL_DEBUG(
820 : "[CommBrokeAlignInfo][PrepareIntraSlices] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu", rank, i,
821 : slice.offset, i, slice.size);
822 : }
823 0 : HCCL_DEBUG("[CommBrokeAlignInfo][PrepareIntraSlices] end calc intra slices and links rank[%u]", rank);
824 0 : return HCCL_SUCCESS;
825 : }
826 :
827 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinksForRS(
828 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
829 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector,
830 : std::vector<u32>& logicCardList)
831 : {
832 0 : std::vector<Slice> intraSlices;
833 :
834 0 : CHK_RET(PrepareIntraSlices(rank, dataUnitSize, count, intraSlices));
835 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u] begin inter", rank);
836 0 : u32 intraRank = GetIntraRank(rank);
837 0 : u32 groupCountForRank = subGroups_[maxSubGroupIdx_].size() / subGroups_[rankGroupMap_[rank]].size();
838 0 : if (subGroups_[maxSubGroupIdx_].size() % subGroups_[rankGroupMap_[rank]].size() > intraRank) {
839 0 : groupCountForRank++;
840 : }
841 :
842 0 : for (u32 k = 0; k < groupCountForRank; ++k) {
843 0 : std::vector<Slice> interSlices;
844 0 : std::vector<LINK> interLinks;
845 0 : u32 curGroupIdx = intraRank + k * subGroups_[rankGroupMap_[rank]].size();
846 0 : if (curGroupIdx < subGroups_[minSubGroupIdx_].size()) { // 参与运算的所有 slice 都是有数据的
847 0 : logicCardList.push_back(rankGroupMap_[rank]);
848 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
849 0 : Slice curSlice = intraSlices[groupOriginOffset_[i] + curGroupIdx];
850 0 : interLinks.push_back(links[subGroups_[i][intraRank]]);
851 0 : interSlices.push_back(curSlice);
852 0 : HCCL_DEBUG(
853 : "[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], curIdx[%u], subGroup[%u], "
854 : "groupIdx[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
855 : rank, subGroups_[i][intraRank], groupOriginOffset_[i] + curGroupIdx, i, curGroupIdx,
856 : groupOriginOffset_[i], curSlice.offset, groupOriginOffset_[i], curSlice.size);
857 : }
858 : } else { // 部分空片参与运算
859 0 : Slice emptySlice;
860 0 : emptySlice.size = 0;
861 0 : emptySlice.offset = 0;
862 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
863 0 : u32 curSubgroupsIdx = i < completeGroupOrder_[curGroupIdx].size() ?
864 0 : completeGroupOrder_[curGroupIdx][i] :
865 0 : emptyGroupOrder_[curGroupIdx][i - completeGroupOrder_[curGroupIdx].size()];
866 0 : if (curSubgroupsIdx == rankGroupMap_[rank]) {
867 0 : logicCardList.push_back(i);
868 : }
869 0 : Slice curSlice = i < completeGroupOrder_[curGroupIdx].size() ?
870 0 : intraSlices[groupOriginOffset_[curSubgroupsIdx] + curGroupIdx] :
871 0 : emptySlice;
872 0 : interLinks.push_back(
873 0 : links[subGroups_[curSubgroupsIdx][curGroupIdx % subGroups_[curSubgroupsIdx].size()]]);
874 0 : interSlices.push_back(curSlice);
875 0 : HCCL_DEBUG(
876 : "[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], curIdx[%u], subGroup[%u], "
877 : "groupIdx[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
878 : rank, subGroups_[curSubgroupsIdx][curGroupIdx % subGroups_[curSubgroupsIdx].size()],
879 : groupOriginOffset_[curSubgroupsIdx] + curGroupIdx, curSubgroupsIdx, curGroupIdx,
880 : groupOriginOffset_[curSubgroupsIdx], curSlice.offset, groupOriginOffset_[curSubgroupsIdx],
881 : curSlice.size);
882 : }
883 : }
884 0 : interLinksVector.push_back(interLinks);
885 0 : interSlicesVector.push_back(interSlices);
886 0 : }
887 :
888 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u] end inter", rank);
889 0 : return HCCL_SUCCESS;
890 0 : }
891 :
892 0 : HcclResult CommBrokeAlignInfo::CalcInterSlicesAndLinksForAR(
893 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
894 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector)
895 : {
896 : // 查找自己位于组内的第几个rank
897 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] begin calc inter slices and links rank[%u]", rank);
898 :
899 0 : u32 intraRank = GetIntraRank(rank);
900 :
901 0 : std::vector<Slice> intraSlices;
902 0 : std::vector<LINK> intraLinks;
903 0 : CHK_RET(CalcIntraSlicesAndLinks(rank, dataUnitSize, count, links, intraLinks, intraSlices));
904 :
905 : // 计算组间每个rank结果上的offset和size
906 0 : u64 sliceSizeCalculated = (intraSlices[intraRank].size / dataUnitSize + (static_cast<u32>(subGroups_.size()) - 1))
907 0 : / subGroups_.size() * dataUnitSize;
908 0 : u64 totalSize = intraSlices[intraRank].size;
909 0 : u64 residueSize = totalSize;
910 : u64 sliceSizeAligned;
911 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
912 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
913 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
914 : } else {
915 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
916 : }
917 :
918 0 : std::vector<LINK> interLinks;
919 0 : std::vector<Slice> interSlices;
920 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
921 0 : interLinks.push_back(links[subGroups_[i][intraRank]]);
922 0 : Slice slice;
923 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
924 0 : slice.offset = intraSlices[intraRank].offset + totalSize - residueSize;
925 0 : residueSize -= slice.size;
926 0 : HCCL_DEBUG(
927 : "[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
928 : rank, i, slice.offset, i, slice.size);
929 0 : interSlices.push_back(slice);
930 : }
931 0 : interLinksVector.push_back(interLinks);
932 0 : interSlicesVector.push_back(interSlices);
933 :
934 0 : HCCL_DEBUG("[CommBrokeAlignInfo][CalcInterSlicesAndLinksForAR] end calc inter slices and links rank[%u]", rank);
935 0 : return HCCL_SUCCESS;
936 0 : }
937 :
938 0 : CommAHCAlignInfo::CommAHCAlignInfo(const std::vector<std::vector<u32>>& subGroups) : CommAHCBaseInfo(subGroups) {}
939 :
940 0 : CommAHCAlignInfo::~CommAHCAlignInfo() {}
941 :
942 0 : HcclResult CommAHCAlignInfo::Init(AHCOpType opType, std::map<AHCConcOpType, TemplateType>& ahcAlgOption)
943 : {
944 0 : ahcAlgOption_ = ahcAlgOption;
945 :
946 : // 参数检查
947 0 : opType_ = opType;
948 0 : CHK_RET(CheckSubGroups(subGroups_));
949 :
950 : // 初始化slice相关信息
951 0 : CHK_RET(InitSliceInfo());
952 :
953 : // 计算 logicCard 相关信息;
954 0 : InitLogicCardInfo();
955 :
956 : // 初始化相关Map信息
957 0 : CHK_RET(InitMapInfo());
958 :
959 0 : return HCCL_SUCCESS;
960 : }
961 :
962 0 : HcclResult CommAHCAlignInfo::InitSliceInfo()
963 : {
964 : // 计算 totalSliceSegment_ ,即所有分组大小的最小公倍数, 以及 interRankOrder
965 0 : totalSliceSegment_ = subGroups_[0].size();
966 : // u32 groupSizeGcd;
967 0 : for (u32 i = 1; i < subGroups_.size(); ++i) {
968 0 : u32 groupSize = static_cast<u32>(subGroups_[i].size());
969 0 : totalSliceSegment_ = totalSliceSegment_ * groupSize / std::__gcd(totalSliceSegment_, groupSize);
970 : }
971 0 : globalTotalSliceSegment_ = rankSize_ * totalSliceSegment_;
972 0 : HCCL_DEBUG("[CommAHCAlignInfo][InitSliceInfo] totalSliceSegment [%u]", totalSliceSegment_);
973 :
974 : // 计算 logicCardSliceSize_ ;
975 0 : std::set<u32> sliceOffset;
976 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
977 0 : for (u32 j = 0; j < subGroups_[i].size(); ++j) {
978 0 : u32 rankSliceSize = (totalSliceSegment_ / subGroups_[i].size()) * (j + 1);
979 0 : sliceOffset.insert(rankSliceSize);
980 0 : HCCL_DEBUG("[CommAHCAlignInfo][InitSliceInfo] sliceOffset [%u]", rankSliceSize);
981 : }
982 : }
983 0 : sliceOffset.insert(static_cast<u32>(0));
984 0 : logicCardSliceOffset_.resize(sliceOffset.size());
985 0 : std::copy(sliceOffset.begin(), sliceOffset.end(), logicCardSliceOffset_.begin());
986 :
987 0 : std::vector<u32>::iterator itPre = logicCardSliceOffset_.begin();
988 0 : std::vector<u32>::iterator itNext = logicCardSliceOffset_.begin();
989 0 : itNext++;
990 0 : while (itNext != logicCardSliceOffset_.end()) {
991 0 : auto boundDiff = (*itNext) - (*itPre);
992 0 : logicCardSliceSize_.push_back(boundDiff);
993 0 : itPre++;
994 0 : itNext++;
995 : }
996 :
997 0 : CHK_PRT_RET(
998 : logicCardSliceSize_.size() != (logicCardSliceOffset_.size() - 1),
999 : HCCL_ERROR(
1000 : "[CommAHCAlignInfo][InitSliceInfo] cardOffset size [%u] cardSize size [%u] check error",
1001 : logicCardSliceSize_.size(), logicCardSliceOffset_.size()),
1002 : HCCL_E_INTERNAL);
1003 :
1004 0 : return HCCL_SUCCESS;
1005 0 : }
1006 :
1007 0 : HcclResult CommAHCAlignInfo::InitLogicCardInfo()
1008 : {
1009 : // 计算 logicCardCommGroups_;
1010 0 : for (std::vector<u32>::iterator it = (logicCardSliceOffset_.begin() + 1); it != logicCardSliceOffset_.end(); ++it) {
1011 0 : std::vector<u32> logicGroup;
1012 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
1013 : u32 logicRank;
1014 0 : if ((*it) % (totalSliceSegment_ / subGroups_[i].size()) != 0) {
1015 0 : logicRank = (*it) / (totalSliceSegment_ / subGroups_[i].size()) + 1;
1016 : } else {
1017 0 : logicRank = (*it) / (totalSliceSegment_ / subGroups_[i].size());
1018 : }
1019 0 : logicGroup.push_back(subGroups_[i][logicRank - 1]);
1020 : }
1021 0 : logicCardCommGroups_.push_back(logicGroup);
1022 0 : }
1023 :
1024 : // 计算 logicCardGroup_
1025 0 : u32 curRank = subGroups_[minSubGroupIdx_][0];
1026 0 : u32 curOffset = 0;
1027 0 : std::vector<u32>::iterator it = logicCardSliceOffset_.begin();
1028 0 : u32 curIdx = 0;
1029 0 : u32 curLogicIdx = 0;
1030 0 : logicCardGroup_.resize(static_cast<u32>(subGroups_[minSubGroupIdx_].size()));
1031 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
1032 0 : if (logicCardCommGroups_[i][minSubGroupIdx_] != curRank) {
1033 0 : logicCardGroup_[curLogicIdx].resize(i - curIdx);
1034 0 : for (u32 j = 0; j < i - curIdx; j++) {
1035 0 : logicCardGroup_[curLogicIdx][j] = curIdx + j;
1036 : }
1037 0 : curIdx = i;
1038 0 : curLogicIdx++;
1039 0 : curRank = logicCardCommGroups_[i][minSubGroupIdx_];
1040 0 : curOffset = *it;
1041 : }
1042 0 : logicCardExecuteOffset_.push_back(*it - curOffset);
1043 0 : it++;
1044 : }
1045 0 : if (curIdx != logicCardCommGroups_.size() - 1) {
1046 0 : logicCardGroup_[curLogicIdx].resize(logicCardCommGroups_.size() - curIdx);
1047 0 : for (u32 i = 0; i < logicCardCommGroups_.size() - curIdx; i++) {
1048 0 : logicCardGroup_[curLogicIdx][i] = curIdx + i;
1049 : }
1050 : }
1051 0 : return HCCL_SUCCESS;
1052 : }
1053 :
1054 0 : bool CommAHCAlignInfo::CompareLogicCardExcuteOrder(u32 i, u32 j)
1055 : {
1056 0 : return logicCardExecuteOffset_[i] < logicCardExecuteOffset_[j];
1057 : }
1058 :
1059 0 : HcclResult CommAHCAlignInfo::InitMapInfo()
1060 : {
1061 : // rank 到 logicCardOrder 初始化
1062 0 : std::map<u32, u32> interRankOrder;
1063 0 : for (u32 i = 0; i < subGroups_.size(); ++i) {
1064 0 : interRankOrder.insert(std::make_pair(i, i));
1065 : }
1066 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
1067 0 : interRankList_.push_back(interRankOrder);
1068 0 : for (u32 j = 0; j < logicCardCommGroups_[i].size(); ++j) {
1069 0 : rankLogicCardOrderMap_[logicCardCommGroups_[i][j]].push_back(i);
1070 0 : rankLogicCardMap_[logicCardCommGroups_[i][j]].push_back(i);
1071 : }
1072 : }
1073 :
1074 : // 定义 lambda 将对象指针传递到成员函数
1075 0 : auto sortLambda = [this](u32 i, u32 j) {
1076 0 : return this->CompareLogicCardExcuteOrder(i, j);
1077 0 : };
1078 :
1079 : // rankLogicCardOrderMap_ 内的逻辑同号卡list按照 logicCardExecuteOffset_ 并发流开始时间排序
1080 0 : for (auto iter = rankLogicCardOrderMap_.begin(); iter != rankLogicCardOrderMap_.end(); iter++) {
1081 0 : std::vector<u32>& rankLogicCardList = iter->second;
1082 0 : std::sort(rankLogicCardList.begin(), rankLogicCardList.end(), sortLambda);
1083 : }
1084 :
1085 0 : return HCCL_SUCCESS;
1086 0 : }
1087 :
1088 : // 配置当前需要的 globalTotalSliceSegment_,用于 Multi-AllReduce 中
1089 0 : HcclResult CommAHCAlignInfo::SetGlobalTotalSliceSegment(u64 globalTotalSliceSegment)
1090 : {
1091 0 : globalTotalSliceSegment_ = globalTotalSliceSegment;
1092 0 : HCCL_DEBUG(
1093 : "[CommAHCAlignInfo][setGlobalTotalSliceSegment] globalTotalSliceSegment set to [%llu]",
1094 : globalTotalSliceSegment_);
1095 0 : return HCCL_SUCCESS;
1096 : }
1097 :
1098 : // 获取当前rank对应的多个逻辑同号卡,并且按照并发流的开始执行时间排序
1099 0 : HcclResult CommAHCAlignInfo::GetLogicCardExecuteOrder(u32 rank, std::vector<u32>& executeOrder)
1100 : {
1101 0 : executeOrder = rankLogicCardOrderMap_[rank];
1102 0 : return HCCL_SUCCESS;
1103 : }
1104 :
1105 0 : HcclResult CommAHCAlignInfo::SliceSizeAlignBound(
1106 : Slice& slice, u64 offsetCount, u64 sliceSizeCalculated, const u64 boundSize, u32 boundOffsetCount,
1107 : u32& curOffset) const
1108 : {
1109 0 : u64 sliceSize = offsetCount * sliceSizeCalculated;
1110 0 : if (!isAlignBound_) {
1111 : // 对于 All-Reduce 中的 Reduce-Scatter 以及 All-Gather,不需要严格对齐bound
1112 0 : slice.size = slice.size + sliceSize;
1113 0 : curOffset = curOffset + offsetCount;
1114 0 : return HCCL_SUCCESS;
1115 : }
1116 0 : if (offsetCount < boundOffsetCount) {
1117 0 : if (sliceSize <= ((curOffset / boundOffsetCount + 1) * boundSize - (slice.size + slice.offset))) {
1118 0 : slice.size = slice.size + sliceSize;
1119 : } else {
1120 0 : slice.size = slice.size + ((curOffset / boundOffsetCount + 1) * boundSize - (slice.size + slice.offset));
1121 : }
1122 : } else {
1123 0 : slice.size = slice.size + (offsetCount / boundOffsetCount) * boundSize;
1124 : }
1125 0 : curOffset = curOffset + offsetCount;
1126 0 : return HCCL_SUCCESS;
1127 : }
1128 :
1129 : // Reduce-Scatter 及 All-Gather 组内切片逻辑
1130 0 : HcclResult CommAHCAlignInfo::CalcIntraSlicesAndLinks(
1131 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
1132 : std::vector<std::vector<LINK>>& intraLinksVector, std::vector<std::vector<Slice>>& intraSlicesVector)
1133 : {
1134 : // Boundary 指在 RS 及 AG 中单卡应有的数据量的 offset, 如八卡跑8K,boundary 为 1024
1135 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
1136 :
1137 0 : u32 singleRankOffset = globalTotalSliceSegment_ / rankSize_; // 每个 rank 最后结果应有的小块数据份数
1138 0 : u64 sliceSizeCalculated = (totalSize_ / dataUnitSize + globalTotalSliceSegment_ - 1) / globalTotalSliceSegment_
1139 0 : * dataUnitSize * (globalTotalSliceSegment_ / rankSize_ / totalSliceSegment_);
1140 0 : u64 totalSize = totalSize_;
1141 0 : u64 residueSize = totalSize;
1142 0 : HCCL_DEBUG(
1143 : "[CommAHCAlignInfo][AHCDEBUG] count[%u] ranksize[%u] sliceSizeCalculated[%u] totalSize[%u] "
1144 : "globalTotalSliceSegment[%u] totalSliceSegment[%u] dataUnitSize[%u]",
1145 : count, rankSize_, sliceSizeCalculated, totalSize, globalTotalSliceSegment_, totalSliceSegment_, dataUnitSize);
1146 :
1147 0 : u32 curOffset = 0;
1148 0 : for (u32 k = 0; k < subGroups_.size(); ++k) {
1149 0 : std::vector<Slice> intraSlices;
1150 0 : std::vector<LINK> intraLinks;
1151 : std::vector<u32> curLogicCardGroup
1152 0 : = rankLogicCardMap_[subGroups_[rankGroupMap_[rank]][0]]; // 获取当前rank对应的逻辑同号组
1153 0 : u32 singleSliceOffset = logicCardSliceSize_[curLogicCardGroup[0]];
1154 0 : for (u32 j = 1; j < curLogicCardGroup.size(); ++j) {
1155 0 : singleSliceOffset = singleSliceOffset + logicCardSliceSize_[curLogicCardGroup[j]];
1156 : }
1157 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
1158 0 : u32 curRank = subGroups_[rankGroupMap_[rank]][i];
1159 0 : intraLinks.push_back(links[curRank]);
1160 0 : Slice slice;
1161 0 : slice.size = 0;
1162 0 : slice.offset = totalSize - residueSize;
1163 0 : u64 targeOffset = singleSliceOffset * subGroups_[k].size();
1164 0 : u64 offsetCountBeforeBoundary
1165 0 : = ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) < targeOffset ?
1166 : ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) :
1167 : targeOffset;
1168 0 : SliceSizeAlignBound(
1169 0 : slice, offsetCountBeforeBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1170 : curOffset);
1171 0 : u64 offsetCountCrossBoundary
1172 0 : = (targeOffset - offsetCountBeforeBoundary) / singleRankOffset * singleRankOffset;
1173 0 : SliceSizeAlignBound(
1174 0 : slice, offsetCountCrossBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1175 : curOffset);
1176 0 : u64 offsetCountBehindBoundary
1177 0 : = (targeOffset - offsetCountBeforeBoundary - offsetCountCrossBoundary) % singleRankOffset;
1178 0 : SliceSizeAlignBound(
1179 0 : slice, offsetCountBehindBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1180 : curOffset);
1181 0 : slice.size = (residueSize > slice.size) ? slice.size : residueSize;
1182 0 : residueSize -= slice.size;
1183 0 : intraSlices.push_back(slice);
1184 0 : HCCL_DEBUG(
1185 : "[CommAHCAlignInfo][CalcIntraSlicesAndLinks] rank[%u], singleSliceOffset[%u], "
1186 : "subGroups_[%u].size()[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
1187 : rank, singleSliceOffset, k, subGroups_[k].size(), i, slice.offset, i, slice.size);
1188 : }
1189 0 : intraLinksVector.push_back(intraLinks);
1190 0 : intraSlicesVector.push_back(intraSlices);
1191 0 : }
1192 :
1193 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
1194 0 : return HCCL_SUCCESS;
1195 : }
1196 :
1197 : // All-Reduce 组内切片逻辑
1198 0 : HcclResult CommAHCAlignInfo::CalcIntraSlicesAndLinks(
1199 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
1200 : std::vector<LINK>& intraLinks, std::vector<Slice>& intraSlices)
1201 : {
1202 : // 计算组内每个rank结果上的offset和size
1203 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] begin calc intra slices and links rank[%u]", rank);
1204 :
1205 0 : u64 sliceSizeCalculated = (count + (totalSliceSegment_ * static_cast<u32>(subGroups_.size()) - 1))
1206 0 : / (totalSliceSegment_ * subGroups_.size()) * dataUnitSize;
1207 0 : u64 totalSize = count * dataUnitSize;
1208 0 : u64 residueSize = totalSize;
1209 0 : u64 sliceSizeAligned = sliceSizeCalculated;
1210 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
1211 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
1212 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1213 : } else {
1214 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
1215 : }
1216 0 : sliceSizeAligned = sliceSizeAligned * static_cast<u32>(subGroups_.size())
1217 0 : * (totalSliceSegment_ / static_cast<u32>(subGroups_[rankGroupMap_[rank]].size()));
1218 :
1219 0 : for (u32 i = 0; i < subGroups_[rankGroupMap_[rank]].size(); ++i) {
1220 0 : intraLinks.push_back(links[subGroups_[rankGroupMap_[rank]][i]]);
1221 0 : Slice slice;
1222 0 : slice.size = (residueSize > sliceSizeAligned) ? sliceSizeAligned : residueSize;
1223 0 : slice.offset = totalSize - residueSize;
1224 0 : residueSize -= slice.size;
1225 0 : HCCL_DEBUG(
1226 : "[CommAHCAlignInfo][CalcIntraSlicesAndLinks] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu", rank,
1227 : i, slice.offset, i, slice.size);
1228 0 : intraSlices.push_back(slice);
1229 : }
1230 :
1231 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcIntraSlicesAndLinks] end calc intra slices and links rank[%u]", rank);
1232 0 : return HCCL_SUCCESS;
1233 : }
1234 :
1235 : // 组间切片逻辑统一对外接口
1236 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinks(
1237 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
1238 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector,
1239 : std::vector<u32>& logicCardList)
1240 : {
1241 0 : HcclResult ret = HCCL_SUCCESS;
1242 0 : switch (opType_) {
1243 0 : case AHCOpType::AHC_OP_TYPE_ALLREDUCE:
1244 0 : ret = CalcInterSlicesAndLinksForAR(
1245 : rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
1246 0 : CHK_PRT_RET(
1247 : ret != HCCL_SUCCESS,
1248 : HCCL_ERROR(
1249 : "[CommAHCAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks "
1250 : "step",
1251 : rank, count),
1252 : ret);
1253 0 : break;
1254 0 : case AHCOpType::AHC_OP_TYPE_REDUCE_SCATTER:
1255 : case AHCOpType::AHC_OP_TYPE_ALLGATHER:
1256 0 : ret = CalcInterSlicesAndLinksForRS(
1257 : rank, dataUnitSize, count, links, interLinksVector, interSlicesVector, logicCardList);
1258 0 : CHK_PRT_RET(
1259 : ret != HCCL_SUCCESS,
1260 : HCCL_ERROR(
1261 : "[CommAHCAlignInfo][CalcInterSlicesAndLinks]rank[%u] count[%llu] failed in CalcInterSlicesAndLinks "
1262 : "step",
1263 : rank, count),
1264 : ret);
1265 0 : break;
1266 0 : default:
1267 0 : ret = HCCL_SUCCESS;
1268 : }
1269 0 : return ret;
1270 : }
1271 :
1272 0 : HcclResult CommAHCAlignInfo::PrepareIntraSlices(
1273 : const u32 rank, const u32 dataUnitSize, const u64 count, std::vector<std::vector<Slice>>& intraSlicesVector)
1274 : {
1275 : // 计算组内每个rank结果上的offset和size
1276 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareIntraSlices] begin calc intra slices and links rank[%u]", rank);
1277 :
1278 0 : u32 singleRankOffset = globalTotalSliceSegment_ / rankSize_;
1279 0 : u64 sliceSizeCalculated = (totalSize_ / dataUnitSize + globalTotalSliceSegment_ - 1) / globalTotalSliceSegment_
1280 0 : * dataUnitSize * (globalTotalSliceSegment_ / rankSize_ / totalSliceSegment_);
1281 0 : u64 totalSize = totalSize_;
1282 0 : u64 residueSize = totalSize;
1283 0 : HCCL_DEBUG(
1284 : "[CommAHCAlignInfo][AHCDEBUG] count[%u] ranksize[%u] sliceSizeCalculated[%u] totalSize[%u]", count, rankSize_,
1285 : sliceSizeCalculated, totalSize);
1286 :
1287 0 : for (u32 i = 0; i < logicCardCommGroups_.size(); ++i) {
1288 0 : std::vector<Slice> intraSlices;
1289 0 : intraSlicesVector.push_back(intraSlices);
1290 0 : }
1291 :
1292 0 : u32 curOffset = 0;
1293 0 : for (u32 i = 0; i < subGroups_.size(); i++) {
1294 0 : for (u32 j = 0; j < logicCardCommGroups_.size(); ++j) {
1295 0 : Slice slice;
1296 0 : slice.size = 0;
1297 0 : slice.offset = totalSize - residueSize;
1298 0 : u64 targeOffset = logicCardSliceSize_[j] * subGroups_[i].size();
1299 0 : u64 offsetCountBeforeBoundary
1300 0 : = ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) < targeOffset ?
1301 : ((curOffset + singleRankOffset - 1) / singleRankOffset * singleRankOffset - curOffset) :
1302 : targeOffset;
1303 0 : SliceSizeAlignBound(
1304 0 : slice, offsetCountBeforeBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1305 : curOffset);
1306 0 : u64 offsetCountCrossBoundary
1307 0 : = (targeOffset - offsetCountBeforeBoundary) / singleRankOffset * singleRankOffset;
1308 0 : SliceSizeAlignBound(
1309 0 : slice, offsetCountCrossBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1310 : curOffset);
1311 0 : u64 offsetCountBehindBoundary
1312 0 : = (targeOffset - offsetCountBeforeBoundary - offsetCountCrossBoundary) % singleRankOffset;
1313 0 : SliceSizeAlignBound(
1314 0 : slice, offsetCountBehindBoundary, sliceSizeCalculated, totalSize_ / rankSize_, singleRankOffset,
1315 : curOffset);
1316 0 : slice.size = (residueSize > slice.size) ? slice.size : residueSize;
1317 0 : residueSize -= slice.size;
1318 0 : intraSlicesVector[j].push_back(slice);
1319 0 : HCCL_DEBUG(
1320 : "[CommAHCAlignInfo][PrepareIntraSlices] rank[%u], round[%u], logicGroup[%u], slices[%u].offset=%llu, "
1321 : "slices[%u].size=%llu",
1322 : rank, i, j, j, slice.offset, j, slice.size);
1323 : }
1324 : }
1325 0 : HCCL_DEBUG("[CommAHCAlignInfo][PrepareIntraSlices] end calc intra slices and links rank[%u]", rank);
1326 0 : return HCCL_SUCCESS;
1327 : }
1328 :
1329 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinksForRS(
1330 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
1331 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector,
1332 : std::vector<u32>& logicCardList)
1333 : {
1334 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] begin calc inter slices and links rank[%u]", rank);
1335 0 : std::vector<std::vector<Slice>> intraSlicesVecotr;
1336 :
1337 0 : CHK_RET(PrepareIntraSlices(rank, dataUnitSize, count, intraSlicesVecotr));
1338 0 : GetLogicCardExecuteOrder(rank, logicCardList);
1339 :
1340 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1341 0 : u32 logicGroupIdx = logicCardList[i]; // 获取当前处理的目标逻辑同号卡组的下标
1342 0 : std::vector<Slice> curIntraSliceVector = intraSlicesVecotr[logicGroupIdx];
1343 0 : std::vector<Slice> interSlices;
1344 0 : std::vector<LINK> interLinks;
1345 0 : for (u32 j = 0; j < subGroups_.size(); j++) {
1346 0 : Slice curSlice = curIntraSliceVector[j];
1347 0 : interLinks.push_back(links[logicCardCommGroups_[logicGroupIdx][j]]); // 当前处理的逻辑同号卡
1348 0 : interSlices.push_back(curSlice);
1349 0 : HCCL_DEBUG(
1350 : "[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] rank[%u], link[%u], logicGroup[%u], "
1351 : "slices[%u].offset=%llu, slices[%u].size=%llu",
1352 : rank, logicCardCommGroups_[logicGroupIdx][j], logicGroupIdx, i, curSlice.offset, i, curSlice.size);
1353 : }
1354 0 : interSlicesVector.push_back(interSlices);
1355 0 : interLinksVector.push_back(interLinks);
1356 0 : }
1357 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForRS] calc inter slices and links rank[%u] end", rank);
1358 0 : return HCCL_SUCCESS;
1359 0 : }
1360 :
1361 0 : HcclResult CommAHCAlignInfo::PrepareWholeLogicSlices(
1362 : const Slice& intraSlice, const u64 sliceSizeAligned, const u32 originOffset, std::vector<Slice>& logicGroupSlice,
1363 : std::vector<u32>& logicCardList)
1364 : {
1365 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1366 0 : Slice logicSlice;
1367 0 : u32 logicRank = logicCardList[i];
1368 : // 计算当前逻辑同号组的offset大小
1369 0 : u32 offsetDiff = logicCardSliceOffset_[logicRank + 1] - logicCardSliceOffset_[logicRank];
1370 0 : HCCL_DEBUG(
1371 : "[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice begin, logicRank : [%u],"
1372 : "offsetDiff : [%u], offset_next : [%u], offset_cur[%u]",
1373 : logicRank, offsetDiff, logicCardSliceOffset_[logicRank + 1], logicCardSliceOffset_[logicRank]);
1374 :
1375 0 : logicSlice.size = sliceSizeAligned * offsetDiff;
1376 0 : logicSlice.offset = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1377 0 : HCCL_DEBUG(
1378 : "[CommAHCAlignInfo][PrepareFullLogicSlices] logicGroupSlice end, logicRank : [%u] ,"
1379 : "size : [%u], offset : [%u] ",
1380 : logicRank, logicSlice.size, logicSlice.offset);
1381 0 : logicGroupSlice.push_back(logicSlice);
1382 : }
1383 0 : return HCCL_SUCCESS;
1384 : }
1385 :
1386 0 : HcclResult CommAHCAlignInfo::PreparePartialLogicSlices(
1387 : const Slice& intraSlice, const u64 sliceSizeAligned, const u32 originOffset, std::vector<Slice>& logicGroupSlice,
1388 : std::vector<u32>& logicCardList)
1389 : {
1390 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1391 0 : Slice logicSlice;
1392 0 : u32 logicRank = logicCardList[i];
1393 : // 计算当前逻辑同号组的offset大小
1394 0 : u32 offsetDiff = logicCardSliceOffset_[logicRank + 1] - logicCardSliceOffset_[logicRank];
1395 0 : HCCL_DEBUG(
1396 : "[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice begin, logicRank : [%u],"
1397 : "offsetDiff : [%u], offset_next : [%u], offset_cur[%u]",
1398 : logicRank, offsetDiff, logicCardSliceOffset_[logicRank + 1], logicCardSliceOffset_[logicRank]);
1399 :
1400 : // 当前rank在组内对应的offset能获取到完全的数据,即前几个逻辑同号卡
1401 0 : if ((logicCardSliceOffset_[logicRank + 1] - originOffset) <= intraSlice.size / sliceSizeAligned) {
1402 0 : logicSlice.size = sliceSizeAligned * offsetDiff;
1403 : logicSlice.offset
1404 0 : = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1405 : // 当前rank在组内对应的offset能获取到部分的数据,即边界上的逻辑同号卡
1406 0 : } else if ((logicCardSliceOffset_[logicRank] - originOffset) <= intraSlice.size / sliceSizeAligned) {
1407 0 : logicSlice.size = intraSlice.size - (logicCardSliceOffset_[logicRank] - originOffset) * sliceSizeAligned;
1408 : logicSlice.offset
1409 0 : = intraSlice.offset + sliceSizeAligned * (logicCardSliceOffset_[logicRank] - originOffset);
1410 : // 当前rank在组内对应的offset不能获取到数据,即最后的逻辑同号卡
1411 : } else {
1412 0 : logicSlice.size = 0;
1413 0 : logicSlice.offset = 0;
1414 : }
1415 0 : HCCL_DEBUG(
1416 : "[CommAHCAlignInfo][CalcInterSlicesAndLinks] logicGroupSlice end, logicRank : [%u] ,"
1417 : "size : [%u], offset : [%u] ",
1418 : logicRank, logicSlice.size, logicSlice.offset);
1419 :
1420 0 : logicGroupSlice.push_back(logicSlice);
1421 : }
1422 0 : return HCCL_SUCCESS;
1423 : }
1424 :
1425 0 : HcclResult CommAHCAlignInfo::PrepareEmptyLogicSlices(
1426 : std::vector<Slice>& logicGroupSlice, const std::vector<u32>& logicCardList) const
1427 : {
1428 0 : for (u32 i = 0; i < logicCardList.size(); i++) {
1429 0 : Slice logicSlice;
1430 0 : logicSlice.size = 0;
1431 0 : logicSlice.offset = 0;
1432 0 : logicGroupSlice.push_back(logicSlice);
1433 : }
1434 0 : return HCCL_SUCCESS;
1435 : }
1436 :
1437 0 : HcclResult CommAHCAlignInfo::CalcLogicSlicesAndLinks(
1438 : std::vector<Slice>& logicGroupSlice, std::vector<u32>& logicCardList, const std::vector<LINK>& links,
1439 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector)
1440 : {
1441 0 : for (u32 i = 0; i < logicGroupSlice.size(); i++) {
1442 0 : Slice logicSlice = logicGroupSlice[i];
1443 0 : std::vector<Slice> interSlices;
1444 0 : std::vector<LINK> interLinks;
1445 0 : u32 logicRank = logicCardList[i];
1446 0 : u64 totalSize = logicSlice.size;
1447 0 : u64 residueSize = totalSize;
1448 : u64 logicSliceSizeAligned;
1449 0 : if (logicSlice.size % subGroups_.size() == 0 && logicSlice.size % HCCL_MIN_SLICE_ALIGN == 0) {
1450 0 : logicSliceSizeAligned = logicSlice.size / subGroups_.size();
1451 : } else {
1452 : u64 sliceSizeCalculated
1453 0 : = (logicSlice.size + static_cast<u32>(subGroups_.size()) - 1) / static_cast<u32>(subGroups_.size());
1454 0 : logicSliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1455 : }
1456 0 : for (u32 j = 0; j < subGroups_.size(); j++) {
1457 0 : u32 curRank = logicCardCommGroups_[logicRank][j];
1458 0 : Slice slice;
1459 0 : interLinks.push_back(links[curRank]);
1460 0 : slice.size = (residueSize > logicSliceSizeAligned) ? logicSliceSizeAligned : residueSize;
1461 0 : slice.offset = logicSlice.offset + totalSize - residueSize;
1462 0 : residueSize -= slice.size;
1463 0 : interSlices.push_back(slice);
1464 : }
1465 0 : interLinksVector.push_back(interLinks);
1466 0 : interSlicesVector.push_back(interSlices);
1467 0 : }
1468 0 : return HCCL_SUCCESS;
1469 : }
1470 :
1471 0 : HcclResult CommAHCAlignInfo::CalcInterSlicesAndLinksForAR(
1472 : const u32 rank, const u32 dataUnitSize, const u64 count, const std::vector<LINK>& links,
1473 : std::vector<std::vector<LINK>>& interLinksVector, std::vector<std::vector<Slice>>& interSlicesVector,
1474 : std::vector<u32>& logicCardList)
1475 : {
1476 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] begin calc inter slices and links rank[%u]", rank);
1477 :
1478 0 : u32 intraRank = GetIntraRank(rank);
1479 :
1480 0 : std::vector<Slice> intraSlices;
1481 0 : std::vector<LINK> intraLinks;
1482 :
1483 0 : CHK_RET(CalcIntraSlicesAndLinks(rank, dataUnitSize, count, links, intraLinks, intraSlices));
1484 0 : GetLogicCardExecuteOrder(rank, logicCardList);
1485 :
1486 : // 计算当前rank逻辑同号卡之间最小slice的大小
1487 0 : u64 sliceSizeCalculated = (count + (totalSliceSegment_ * static_cast<u32>(subGroups_.size()) - 1))
1488 0 : / (totalSliceSegment_ * subGroups_.size()) * dataUnitSize;
1489 0 : const u64 sizeAlignedMinSize = 128 * 1024; // 优化小包性能,小于128k不切片
1490 0 : u64 sliceSizeAligned = sliceSizeCalculated;
1491 0 : if (sliceSizeCalculated > sizeAlignedMinSize) {
1492 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, HCCL_MIN_SLICE_ALIGN);
1493 : } else {
1494 0 : sliceSizeAligned = AlgTemplateBase::RoundUpWithDivisor(sliceSizeCalculated, sizeAlignedMinSize);
1495 : }
1496 :
1497 0 : sliceSizeAligned = sliceSizeAligned * static_cast<u32>(subGroups_.size());
1498 0 : u32 originOffset = intraRank * totalSliceSegment_ / subGroups_[rankGroupMap_[rank]].size(); // 当前rank起始offset
1499 :
1500 0 : HCCL_DEBUG(
1501 : "[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] rank : [%u], intraslice.size : [%u], intraslice.offset : "
1502 : "[%u],"
1503 : "sliceSizeAligned : [%u], originOffset : [%u]",
1504 : rank, intraSlices[intraRank].size, intraSlices[intraRank].offset, sliceSizeAligned, originOffset);
1505 :
1506 : // 进行逻辑同号组对应的slice切分
1507 0 : std::vector<Slice> logicGroupSlice; // 逻辑同号组对应的slice
1508 0 : HCCL_DEBUG(
1509 : "[CommAHCAlignInfo][CalcInterSlicesAndLinksForAR] check rank : [%u], intraSlices[%u].size : [%u], "
1510 : "sliceSizeAligned : [%u],"
1511 : "totalSliceSegment_ : [%u], subGroups_[rankGroupMap_[rank]].size : [%u]",
1512 : rank, intraRank, intraSlices[intraRank].size, sliceSizeAligned, totalSliceSegment_,
1513 : subGroups_[rankGroupMap_[rank]].size());
1514 :
1515 : // 当前rank有完整的对齐后的数据量
1516 0 : if (intraSlices[intraRank].size / sliceSizeAligned == totalSliceSegment_ / subGroups_[rankGroupMap_[rank]].size()) {
1517 0 : CHK_RET(PrepareWholeLogicSlices(
1518 : intraSlices[intraRank], sliceSizeAligned, originOffset, logicGroupSlice, logicCardList));
1519 : // 当前rank有不完整的数据量
1520 0 : } else if (intraSlices[intraRank].size != 0) {
1521 0 : CHK_RET(PreparePartialLogicSlices(
1522 : intraSlices[intraRank], sliceSizeAligned, originOffset, logicGroupSlice, logicCardList));
1523 : } else {
1524 0 : CHK_RET(PrepareEmptyLogicSlices(logicGroupSlice, logicCardList));
1525 : }
1526 :
1527 : // 计算当前rank逻辑同号组之间的slice大小
1528 0 : CHK_RET(CalcLogicSlicesAndLinks(logicGroupSlice, logicCardList, links, interLinksVector, interSlicesVector));
1529 :
1530 0 : HCCL_DEBUG("[CommAHCAlignInfo][CalcInterSlicesAndLinks] end calc inter slices and links rank[%u]", rank);
1531 0 : return HCCL_SUCCESS;
1532 0 : }
1533 :
1534 : } // namespace hccl
|