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 "coll_all_gather_ring_for_910_93_executor.h"
12 :
13 : namespace hccl {
14 2 : CollAllGatherRingFor91093Executor::CollAllGatherRingFor91093Executor(
15 2 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
16 2 : : CollAllGatherExecutor(dispatcher, topoMatcher)
17 : {
18 2 : DMAReduceFlag_ = workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
19 : desc_.level1SupportedAlgos
20 2 : = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
21 2 : AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
22 : desc_.level2SupportedAlgos
23 2 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING};
24 2 : }
25 :
26 2 : HcclResult CollAllGatherRingFor91093Executor::CalcStreamNum(u32& streamNum)
27 : {
28 2 : u32 totalStreamNum
29 2 : = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
30 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
31 2 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
32 2 : totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
33 : }
34 :
35 2 : streamNum = totalStreamNum - 1;
36 2 : HCCL_INFO("[CollAllGatherRingFor91093Executor][CalcStreamNum] tag[%s] streamNum_[%u]", tag_.c_str(), streamNum);
37 2 : return HCCL_SUCCESS;
38 : }
39 :
40 2 : HcclResult CollAllGatherRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
41 : {
42 2 : TransportMemType inputType = TransportMemType::RESERVED;
43 2 : TransportMemType outputType = TransportMemType::RESERVED;
44 2 : CHK_RET(CalcTransportMemType(inputType, outputType));
45 2 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
46 2 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
47 2 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
48 2 : return HCCL_SUCCESS;
49 : }
50 :
51 : HcclResult
52 2 : CollAllGatherRingFor91093Executor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
53 : {
54 2 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
55 2 : inputType = TransportMemType::CCL_INPUT;
56 2 : outputType = TransportMemType::CCL_OUTPUT;
57 : } else {
58 0 : inputType = TransportMemType::PARAM_INPUT;
59 0 : outputType = TransportMemType::PARAM_OUTPUT;
60 : }
61 2 : HCCL_INFO(
62 : "[CollAllGatherRingFor91093Executor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]", tag_.c_str(),
63 : inputType, outputType);
64 2 : return HCCL_SUCCESS;
65 : }
66 :
67 2 : HcclResult CollAllGatherRingFor91093Executor::CalcLevel0CommInfo(
68 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
69 : {
70 2 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
71 2 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
72 2 : return HCCL_SUCCESS;
73 2 : }
74 :
75 2 : HcclResult CollAllGatherRingFor91093Executor::CalcLevel2CommInfo(
76 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
77 : {
78 2 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
79 2 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
80 0 : HCCL_INFO("[CollAllGatherRingFor91093Executor][CalcLevel2CommInfo] select AHC bypass level2 comm calculate");
81 0 : return HCCL_SUCCESS;
82 : }
83 :
84 2 : CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX);
85 2 : HCCL_DEBUG("[CollAllGatherRingFor91093Executor][CalcLevel2CommInfo]Level2CommInfo start set");
86 2 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
87 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
88 0 : HCCL_INFO("[%s]Calc NHRCommInfo.", __func__);
89 2 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
90 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
91 0 : HCCL_INFO("[%s]Calc NBCommInfo.", __func__);
92 : } else {
93 2 : commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
94 2 : HCCL_INFO("[%s]Calc RingCommInfo.", __func__);
95 : }
96 2 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
97 2 : return HCCL_SUCCESS;
98 2 : }
99 :
100 2 : u64 CollAllGatherRingFor91093Executor::CalcLoopMaxCount(const u64 cclBuffSize, const u32 unitSize)
101 : {
102 2 : u64 maxCountPerLoop = cclBuffSize / topoAttr_.userRankSize / HCCL_MIN_SLICE_ALIGN * HCCL_MIN_SLICE_ALIGN / unitSize;
103 2 : return maxCountPerLoop;
104 : }
105 :
106 0 : HcclResult CollAllGatherRingFor91093Executor::RunIntraSeverAllGather(
107 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
108 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream, s32 profStage,
109 : const u64 baseOffset, const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>>& multRingsUserMemSlice)
110 : {
111 0 : CHK_RET(MultiRingAllGather(
112 : tag, inputMem, outputMem, count, dataType, multRingsSliceZero, stream, profStage, baseOffset, opInfo,
113 : multRingsUserMemSlice, logicalLevel0plane_));
114 0 : return HCCL_SUCCESS;
115 : }
116 :
117 18 : u64 CollAllGatherRingFor91093Executor::CalcDstMemOffset(
118 : [[maybe_unused]] const OpParam& param, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
119 : {
120 18 : return topoAttr_.userRank * inputMemSize;
121 : }
122 :
123 18 : HcomCollOpInfo CollAllGatherRingFor91093Executor::GetHcomCollOpInfo(const OpParam& param, const ExecMem& execMem) const
124 : {
125 18 : HcomCollOpInfo opInfo
126 18 : = {"", execMem.inputPtr, execMem.outputPtr, param.DataDes.count, param.DataDes.dataType,
127 18 : 0, HCCL_REDUCE_RESERVED, param.DataDes.strideCount};
128 18 : if (!DMAReduceFlag_ && (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING)) {
129 0 : opInfo.inputAddr = execMem.inputMem.ptr();
130 0 : opInfo.outputAddr = execMem.outputMem.ptr();
131 : }
132 18 : return opInfo;
133 : }
134 :
135 18 : HcclResult CollAllGatherRingFor91093Executor::PrepareSlicesL0(
136 : std::vector<std::vector<Slice>>& multRingsSlice, const OpParam& param, const SubCommInfo& level2CommInfo,
137 : const SubCommInfo& level1CommInfo, const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize,
138 : u64 inputMemSize)
139 : {
140 18 : const u32 level0RankSize = level0CommInfo.localRankSize;
141 18 : const u32 level1RankSize = level1CommInfo.localRankSize;
142 18 : const u32 level2RankSize = level2CommInfo.localRankSize;
143 :
144 18 : std::vector<Slice> dataSegsSlice;
145 18 : CHK_RET(PrepareAllgatherSlice(level0RankSize, inputMemSize, dataSegsSlice));
146 :
147 : // 多环数据切分
148 18 : std::vector<std::vector<Slice>> multRingsSliceZero; // 数据基于该rank上环0的偏移
149 18 : bool ARSFlag = topoMatcher_->GetARSFlag();
150 18 : bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
151 :
152 18 : if (ARSDoubleRing) {
153 0 : std::vector<u32> mockNicList;
154 0 : mockNicList.reserve(level0RankSize);
155 0 : for (u32 rankIndex = 0; rankIndex < level0RankSize; rankIndex++) {
156 0 : mockNicList.push_back(rankIndex);
157 : }
158 0 : multRingsSliceZero = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, mockNicList);
159 0 : } else if (
160 18 : topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING
161 18 : && !IsSupportUnifiedMarch(param, topoType_, topoAttr_.serverNum, topoAttr_.superPodNum)) {
162 18 : multRingsSliceZero = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
163 : } else {
164 0 : multRingsSliceZero.push_back(dataSegsSlice);
165 : }
166 54 : for (u32 ringIndex = 0; ringIndex < multRingsSliceZero.size(); ringIndex++) {
167 36 : std::vector<Slice> level2DataSlice;
168 36 : CHK_RET(CalculateLevel2AllgatherSlice(
169 : inputMemSize, level0RankSize, level1RankSize, level2RankSize, multRingsSliceZero, level2DataSlice,
170 : ringIndex));
171 36 : multRingsSlice.push_back(level2DataSlice);
172 36 : }
173 :
174 18 : return HCCL_SUCCESS;
175 18 : }
176 :
177 18 : std::vector<Slice> CollAllGatherRingFor91093Executor::PrepareSlicesL1(
178 : [[maybe_unused]] const OpParam& param, const SubCommInfo& level2CommInfo, const SubCommInfo& level1CommInfo,
179 : const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
180 : {
181 18 : const u32 level0RankSize = level0CommInfo.localRankSize;
182 18 : const u32 level0ServerIndex = level0CommInfo.localRank;
183 18 : const u32 level1RankSize = level1CommInfo.localRankSize;
184 18 : const u32 level2RankSize = level2CommInfo.localRankSize;
185 18 : std::vector<Slice> level1DataSegsSlice;
186 54 : for (u32 j = 0; j < level1RankSize; j++) {
187 72 : for (u32 i = 0; i < level2RankSize; i++) {
188 36 : Slice level1Slice;
189 36 : level1Slice.size = inputMemSize;
190 : level1Slice.offset
191 36 : = inputMemSize * (i * level1RankSize * level0RankSize + j * level0RankSize + level0ServerIndex);
192 :
193 36 : HCCL_DEBUG(
194 : "[CollAllGatherRingFor91093Executor][PrepareSlicesL1] rank[%u], level1index[%u], level2index[%u], "
195 : "slices.offset=%llu, slices.size=%llu",
196 : level0CommInfo.localRank, j, i, level1Slice.offset, level1Slice.size);
197 :
198 36 : level1DataSegsSlice.push_back(level1Slice);
199 : }
200 : }
201 18 : return level1DataSegsSlice;
202 0 : }
203 :
204 0 : std::vector<Slice> CollAllGatherRingFor91093Executor::PrepareSlicesL2(
205 : [[maybe_unused]] const OpParam& param, const SubCommInfo& level2CommInfo, const SubCommInfo& level1CommInfo,
206 : const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
207 : {
208 0 : const u32 level0RankSize = level0CommInfo.localRankSize;
209 0 : const u32 level0ServerIndex = level0CommInfo.localRank;
210 0 : const u32 level1RankSize = level1CommInfo.localRankSize;
211 0 : const u32 level1ServerIndex = level1CommInfo.localRank;
212 0 : const u32 level2RankSize = level2CommInfo.localRankSize;
213 0 : std::vector<Slice> level2DataSegsSlice;
214 0 : for (u32 i = 0; i < level2RankSize; i++) {
215 0 : Slice sliceTemp;
216 0 : sliceTemp.size = inputMemSize;
217 : sliceTemp.offset
218 0 : = inputMemSize
219 0 : * (i * level1RankSize * level0RankSize + level1ServerIndex * level0RankSize + level0ServerIndex);
220 0 : level2DataSegsSlice.push_back(sliceTemp);
221 : }
222 0 : return level2DataSegsSlice;
223 0 : }
224 :
225 18 : HcclResult CollAllGatherRingFor91093Executor::PrepareUserMemSlices(
226 : std::vector<std::vector<Slice>>& userMemSlices, const std::vector<std::vector<Slice>>& multRingsSlice,
227 : const OpParam& param, [[maybe_unused]] const SubCommInfo& level2CommInfo,
228 : [[maybe_unused]] const SubCommInfo& level1CommInfo, [[maybe_unused]] const SubCommInfo& level0CommInfo,
229 : u32 perDataSize, u64 inputMemSize)
230 : {
231 18 : CHK_PRT_RET(
232 : 0 < param.DataDes.strideCount && param.DataDes.strideCount < param.DataDes.count,
233 : HCCL_ERROR(
234 : "[CollAllGatherRingFor91093Executor][KernelRun]strideCount[%llu] is smaller than opCount[%llu]",
235 : param.DataDes.strideCount, param.DataDes.count),
236 : HCCL_E_PARA);
237 18 : HCCL_DEBUG(
238 : "[CollAllGatherRingFor91093Executor][KernelRun]strideCount[%llu], opCount[%llu]", param.DataDes.strideCount,
239 : param.DataDes.count);
240 :
241 18 : if (!DMAReduceFlag_) {
242 0 : userMemSlices = multRingsSlice;
243 : // 图模式,根据strideCount更新slice的offset
244 0 : if (param.DataDes.strideCount != 0) {
245 0 : CHK_RET(UpdateOffsetBasedOnStrideCount(param, userMemSlices));
246 : }
247 : } else {
248 54 : for (u32 ringIndex = 0; ringIndex < multRingsSlice.size(); ringIndex++) {
249 36 : std::vector<Slice> userMemSlice;
250 180 : for (const auto& cclSlice : multRingsSlice[ringIndex]) {
251 144 : Slice tmpSlice;
252 144 : u64 count = (param.DataDes.strideCount == 0) ? param.DataDes.count : param.DataDes.strideCount;
253 144 : tmpSlice.size = cclSlice.size;
254 : tmpSlice.offset
255 144 : = (cclSlice.offset / inputMemSize) * count * perDataSize + multRingsSlice[ringIndex][0].offset;
256 144 : userMemSlice.push_back(tmpSlice);
257 144 : HCCL_DEBUG(
258 : "rank[%u], ringIndex[%u], tmpSlice.offset=[%llu], size=[%llu]", topoAttr_.userRank, ringIndex,
259 : tmpSlice.offset, tmpSlice.size);
260 : }
261 36 : userMemSlices.push_back(userMemSlice);
262 36 : }
263 : }
264 18 : return HCCL_SUCCESS;
265 : }
266 :
267 18 : HcclResult CollAllGatherRingFor91093Executor::GetLevelCommInfo()
268 : {
269 18 : logicalLevel0plane_ = COMM_LEVEL0;
270 18 : CHK_RET(CheckCommSize(logicalLevel0plane_, COMM_INDEX_0 + 1));
271 18 : logicalLevel0CommInfo_ = GetSubCommInfo(logicalLevel0plane_, COMM_INDEX_0);
272 18 : u32 commIndex = logicalLevel0CommInfo_.localRank;
273 18 : bool isSelectAHC
274 18 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
275 18 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
276 18 : logicalLevel1plane_ = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
277 18 : CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
278 18 : logicalLevel1CommInfo_ = GetSubCommInfo(logicalLevel1plane_, commIndex);
279 18 : return HCCL_SUCCESS;
280 : }
281 :
282 18 : HcclResult CollAllGatherRingFor91093Executor::KernelRun(const OpParam& param, ExecMem& execMem)
283 : {
284 18 : HCCL_CONFIG_INFO(
285 : HCCL_ALG, "[%s] The AllGatherRingExecutor starts, topoType_[%u], agv[%u]", __func__, topoType_, isAllGatherV_);
286 18 : CHK_RET(GetLevelCommInfo()); // 设置逻辑通信域
287 18 : CHK_RET(ActiveSlaveStreams(param.stream));
288 18 : const HcclDataType dataType = param.GetDataType();
289 18 : u32 perDataSize = 0;
290 18 : CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
291 18 : CHK_PRT_RET(
292 : perDataSize == 0,
293 : HCCL_ERROR(
294 : "[CollAllGatherRingFor91093Executor][KernelRun]errNo[0x%016llx] datatype[%s] is invalid",
295 : HCCL_ERROR_CODE(HCCL_E_PARA), GetDataTypeEnumStr(dataType).c_str()),
296 : HCCL_E_PARA);
297 :
298 18 : bool isSelectAHC
299 18 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
300 18 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
301 :
302 18 : u32 level1RankSize = logicalLevel1CommInfo_.localRankSize;
303 :
304 18 : SubCommInfo level2CommInfo;
305 18 : if (isSelectAHC) {
306 0 : level2CommInfo = logicalLevel1CommInfo_;
307 0 : level2CommInfo.localRankSize = 1; // AHC bypass level2
308 : } else {
309 18 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
310 18 : level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
311 : }
312 18 : const u32 level2RankSize = level2CommInfo.localRankSize;
313 :
314 : // 第一步,将数据从input内存拷贝到output内存的对应位置
315 18 : u64 inputMemSize = execMem.inputMem.size();
316 18 : u64 dstMemOffset = CalcDstMemOffset(param, perDataSize, inputMemSize);
317 18 : DeviceMem dstMem = execMem.outputMem.range(dstMemOffset, inputMemSize);
318 18 : CHK_SMART_PTR_NULL(dstMem);
319 :
320 18 : HcomCollOpInfo opInfo = GetHcomCollOpInfo(param, execMem);
321 18 : HcomCollOpInfo* opInfoPtr
322 18 : = (DMAReduceFlag_ || (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING)) ? &opInfo : nullptr;
323 :
324 : // 图模式opinfo不为空,但需要将数据从ccl input拷贝到ccl output上
325 18 : HcclResult ret = HCCL_SUCCESS;
326 18 : if (!DMAReduceFlag_) {
327 0 : ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, execMem.inputMem, const_cast<Stream&>(param.stream));
328 0 : CHK_PRT_RET(
329 : ret != HCCL_SUCCESS,
330 : HCCL_ERROR(
331 : "[CollAllGatherRingFor91093Executor][KernelRun]AllGather double "
332 : "ring memcpy Failed, Offset[%llu], Size[%llu]",
333 : dstMemOffset, inputMemSize),
334 : ret);
335 : } else {
336 : // 先做server间算法,带有消减拷贝场景数据需要从user input取,拷贝到ccl output上
337 18 : if (level1RankSize > 1 || level2RankSize > 1) {
338 18 : DeviceMem srcMem = DeviceMem::create(static_cast<u8*>(execMem.inputPtr), inputMemSize);
339 18 : ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, const_cast<Stream&>(param.stream));
340 18 : CHK_PRT_RET(
341 : ret != HCCL_SUCCESS,
342 : HCCL_ERROR(
343 : "[CollAllGatherRingFor91093Executor][KernelRun]AllGather double "
344 : "ring user memcpy Failed, Offset[%llu], Size[%llu]",
345 : dstMemOffset, inputMemSize),
346 : ret);
347 18 : }
348 : }
349 18 : if (level2RankSize > 1) {
350 0 : std::unique_ptr<AlgTemplateBase> level2AGExecutor;
351 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
352 : level2AGExecutor
353 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
354 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL2", __func__);
355 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
356 : level2AGExecutor
357 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
358 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL2", __func__);
359 : } else {
360 : level2AGExecutor
361 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
362 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL2", __func__);
363 : }
364 0 : CHK_SMART_PTR_NULL(level2AGExecutor);
365 :
366 0 : std::vector<Slice> level2DataSegsSlice = PrepareSlicesL2(
367 0 : param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize, inputMemSize);
368 0 : CHK_RET(level2AGExecutor->Prepare(
369 : execMem.outputMem, execMem.outputMem, execMem.inputMem, execMem.count, dataType, param.stream,
370 : HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level2DataSegsSlice, 0));
371 :
372 0 : CHK_RET(level2AGExecutor->RegisterProfiler(
373 : (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_0,
374 : HCCL_EXEC_STEP_NOT_SET, param.stream));
375 :
376 0 : CHK_RET(RunTemplate(level2AGExecutor, level2CommInfo));
377 0 : HCCL_INFO(
378 : "AllGather ring [superpod] level2 AllGather run successtopoType_[%u], agv[%u]", topoType_, isAllGatherV_);
379 0 : }
380 18 : if (level1RankSize > 1) {
381 : // 计算slice, 不同超节点相同slice
382 18 : std::vector<Slice> level1DataSegsSlice = PrepareSlicesL1(
383 18 : param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize, inputMemSize);
384 :
385 18 : std::unique_ptr<AlgTemplateBase> level1AGExecutor;
386 18 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
387 : level1AGExecutor
388 18 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
389 18 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
390 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
391 : level1AGExecutor
392 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
393 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
394 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
395 : level1AGExecutor
396 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
397 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
398 0 : } else if (isSelectAHC) {
399 : // 获取通信域分组信息
400 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
401 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
402 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(logicalLevel1plane_, globalSubGroups));
403 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
404 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
405 0 : level1AGExecutor = AlgTemplateRegistry::Instance().GetAlgTemplate(
406 0 : TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
407 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC in COMM_LEVEL1", __func__);
408 : } else {
409 0 : level1AGExecutor = AlgTemplateRegistry::Instance().GetAlgTemplate(
410 0 : TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
411 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC_BROKE in COMM_LEVEL1", __func__);
412 : }
413 0 : CHK_SMART_PTR_NULL(level1AGExecutor);
414 0 : CHK_RET(level1AGExecutor->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
415 0 : } else {
416 0 : HCCL_ERROR("AllGather ring: unsupported algtype [%s].", AlgTypeToStr(algType_).c_str());
417 0 : return HCCL_E_NOT_SUPPORT;
418 : }
419 18 : CHK_SMART_PTR_NULL(level1AGExecutor);
420 54 : CHK_RET(level1AGExecutor->Prepare(
421 : execMem.outputMem, execMem.outputMem, execMem.inputMem, execMem.count, dataType, param.stream,
422 : HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level1DataSegsSlice, 0));
423 :
424 18 : CHK_RET(level1AGExecutor->RegisterProfiler(
425 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_1,
426 : HCCL_EXEC_STEP_NOT_SET, param.stream));
427 :
428 18 : CHK_RET(RunTemplate(level1AGExecutor, logicalLevel1CommInfo_));
429 18 : HCCL_INFO(
430 : "AllGather ring [superpod] level1 AllGather run successtopoType_[%u], agv[%u]", topoType_, isAllGatherV_);
431 18 : }
432 : // 节点内做AllGather ring
433 18 : std::vector<std::vector<Slice>> multRingsSlice;
434 18 : CHK_RET(PrepareSlicesL0(
435 : multRingsSlice, param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize,
436 : inputMemSize));
437 :
438 18 : std::vector<std::vector<Slice>> multRingsUserMemSlice;
439 18 : CHK_RET(PrepareUserMemSlices(
440 : multRingsUserMemSlice, multRingsSlice, param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_,
441 : perDataSize, inputMemSize));
442 :
443 18 : if (DMAReduceFlag_ && (level1RankSize > 1 || level2RankSize > 1)) {
444 : // allgather输入放在CCL buffer上,通过设置nullptr指示要从CCL buffer获取输入
445 18 : opInfo.inputAddr = nullptr;
446 : }
447 18 : CHK_RET(RunIntraSeverAllGather(
448 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, dataType, multRingsSlice, param.stream,
449 : PROF_STAGE_2, 0, opInfoPtr, multRingsUserMemSlice));
450 18 : HCCL_INFO("AllGather ring run success. topoType_[%u], agv[%u]", topoType_, isAllGatherV_);
451 18 : return HCCL_SUCCESS;
452 18 : }
453 :
454 0 : HcclResult CollAllGatherRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
455 : {
456 0 : HCCL_INFO("[CollAllGatherRingFor91093Executor][Getlevel1CommRank] Entry Getlevel1CommRank.");
457 0 : bool isSelectAHC
458 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
459 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
460 0 : if (isSelectAHC) {
461 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
462 0 : u32 level0ServerIndex = level0CommInfo.localRank;
463 :
464 0 : CommPlane commPlaneLevel1 = COMM_LEVEL1;
465 0 : CHK_RET(CheckCommSize(commPlaneLevel1, level0ServerIndex + 1));
466 0 : level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0ServerIndex);
467 0 : u32 level1RankSize = level1CommInfo.localRankSize;
468 0 : HCCL_INFO("Getlevel1CommRank. level1RankSize[%u]", level1RankSize);
469 0 : return HCCL_SUCCESS;
470 0 : }
471 0 : if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
472 0 : HCCL_INFO("[nslbdp] Getlevel1CommRank size not match.");
473 0 : return HCCL_E_UNAVAIL;
474 : }
475 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
476 :
477 0 : return HCCL_SUCCESS;
478 : }
479 :
480 : HcclResult
481 0 : CollAllGatherRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
482 : {
483 0 : HCCL_INFO("[nslbdp] Entry SelectTempAlg, level1RankSize = [%u].", level1RankSize);
484 0 : bool isSelectAHC
485 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
486 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
487 0 : if (isSelectAHC) {
488 0 : CommPlane commPlaneLevel1 = COMM_LEVEL1;
489 : // 获取通信域分组信息
490 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
491 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
492 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
493 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
494 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
495 : level1TempAlg
496 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
497 0 : HCCL_INFO("allgather comm: using ahc algo inter-server.");
498 : } else {
499 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
500 0 : TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
501 0 : HCCL_INFO("allgather comm: using ahc-broke algo inter-server.");
502 : }
503 0 : CHK_SMART_PTR_NULL(level1TempAlg);
504 0 : CHK_RET(level1TempAlg->Prepare(NSLBDP_MIN_COUNT, globalSubGroups, ahcAlgOption));
505 0 : return HCCL_SUCCESS;
506 0 : }
507 0 : if (level1RankSize > 1) {
508 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
509 : level1TempAlg
510 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
511 0 : HCCL_INFO("AllGather ring: using nonuniform-bruck algo inter-superPod.");
512 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
513 : level1TempAlg
514 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
515 0 : HCCL_INFO("AllGather ring: using nonuniform-hierarchical-ring algo inter-superPod.");
516 : } else {
517 : level1TempAlg
518 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
519 0 : HCCL_INFO("AllGather ring: using ring algo inter-superPod.");
520 : }
521 0 : CHK_SMART_PTR_NULL(level1TempAlg);
522 0 : return HCCL_SUCCESS;
523 : }
524 0 : return HCCL_E_UNAVAIL;
525 : }
526 :
527 : REGISTER_EXEC("AllGatherRingFor91093Executor", AllGatherRingFor91093, CollAllGatherRingFor91093Executor);
528 :
529 : } // namespace hccl
|