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_reduce_scatter_ring_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 3 : CollReduceScatterRingExecutor::CollReduceScatterRingExecutor(const HcclDispatcher dispatcher,
17 3 : std::unique_ptr<TopoMatcher> &topoMatcher)
18 3 : : CollReduceScatterExecutor(dispatcher, topoMatcher)
19 : {
20 3 : DMAReduceFlag_ = (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE &&
21 0 : topoAttr_.deviceType == DevType::DEV_TYPE_910_93);
22 3 : }
23 :
24 6 : void CollReduceScatterRingExecutor::ParseParam(const OpParam& param)
25 : {
26 6 : tag_ = param.tag;
27 :
28 : // 是否需要scratch memory
29 12 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE &&
30 0 : (topoAttr_.deviceType == DevType::DEV_TYPE_910B || topoAttr_.deviceType == DevType::DEV_TYPE_910_93) &&
31 6 : isSupportSDMAReduce_ && IsSupportRDMAReduce(param.DataDes.dataType, param.reduceType)) {
32 0 : scratchMemFlag_ = false;
33 : } else {
34 6 : scratchMemFlag_ = true;
35 : }
36 :
37 : // 记录图模式总数据量
38 6 : HCCL_DEBUG("[CollReduceScatterRingExecutor][ParseParam]scratchMemFlag is %d", scratchMemFlag_);
39 6 : totalSize_ = topoAttr_.userRankSize * param.DataDes.count * SIZE_TABLE[param.DataDes.dataType];
40 6 : aicpuUnfoldMode_ = param.aicpuUnfoldMode;
41 6 : }
42 :
43 3 : HcclResult CollReduceScatterRingExecutor::CalcScratchMemSize(u64& scratchMemSize)
44 : {
45 3 : if (scratchMemFlag_) {
46 3 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
47 0 : scratchMemSize = inCCLbufferSize_;
48 : } else {
49 3 : scratchMemSize = totalSize_;
50 : }
51 : } else {
52 0 : scratchMemSize = 0U;
53 : }
54 3 : HCCL_INFO("[CollReduceScatterRingExecutor][CalcScratchMemSize] tag[%s] scratchMemSize[%llu]",
55 : tag_.c_str(), scratchMemSize);
56 3 : return HCCL_SUCCESS;
57 : }
58 :
59 3 : HcclResult CollReduceScatterRingExecutor::CalcStreamNum(u32& streamNum)
60 : {
61 3 : u32 totalStreamNum = 1U;
62 3 : if (algType_.algoLevel0 == AlgTypeLevel0::ALG_LEVEL0_8P_RING) {
63 0 : totalStreamNum = LEVEL0_PLANE_NUM_IN_8PRING;
64 : }
65 3 : streamNum = totalStreamNum - 1;
66 3 : HCCL_INFO("[CollReduceScatterRingExecutor][CalcStreamNum] tag[%s] streamNum[%u]", tag_.c_str(), streamNum);
67 3 : return HCCL_SUCCESS;
68 : }
69 :
70 3 : HcclResult CollReduceScatterRingExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
71 : {
72 3 : TransportMemType inputType = TransportMemType::RESERVED;
73 3 : TransportMemType outputType = TransportMemType::RESERVED;
74 3 : CHK_RET(CalcTransportMemType(inputType, outputType));
75 3 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
76 3 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
77 3 : return HCCL_SUCCESS;
78 : }
79 :
80 3 : HcclResult CollReduceScatterRingExecutor::CalcTransportMemType(TransportMemType &inputType,
81 : TransportMemType &outputType)
82 : {
83 3 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
84 0 : inputType = TransportMemType::CCL_INPUT;
85 0 : if (scratchMemFlag_) {
86 0 : outputType = TransportMemType::SCRATCH;
87 : } else {
88 0 : outputType = TransportMemType::CCL_OUTPUT;
89 : }
90 : } else {
91 3 : inputType = TransportMemType::PARAM_INPUT;
92 3 : if (scratchMemFlag_) {
93 3 : outputType = TransportMemType::SCRATCH;
94 : } else {
95 0 : outputType = TransportMemType::PARAM_OUTPUT;
96 : }
97 : }
98 3 : HCCL_INFO("[CollReduceScatterRingExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
99 : tag_.c_str(), inputType, outputType);
100 3 : return HCCL_SUCCESS;
101 : }
102 :
103 3 : HcclResult CollReduceScatterRingExecutor::CalcLevel0CommInfo(TransportMemType inputType,
104 : TransportMemType outputType,
105 : std::vector<LevelNSubCommTransport>& opTransport)
106 : {
107 3 : HCCL_INFO("[CollReduceScatterRingExecutor][CalcLevel0CommInfo]tag[%s] start", tag_.c_str());
108 3 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
109 3 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
110 3 : HCCL_INFO("[CollReduceScatterRingExecutor][CalcLevel0CommInfo]tag[%s] Calc RingComm finish", tag_.c_str());
111 3 : return HCCL_SUCCESS;
112 3 : }
113 :
114 0 : u64 CollReduceScatterRingExecutor::CalcLoopMaxCount(const u32 unitSize)
115 : {
116 : // 中转内存单次最多能够接受的output count,放开ranksize限制
117 0 : u64 maxCountPerLoop = inCCLbufferSize_ / (topoAttr_.userRankSize * unitSize);
118 0 : return maxCountPerLoop;
119 : }
120 :
121 0 : bool CollReduceScatterRingExecutor::IsHugeData(const u64 curSize, OpParam *param)
122 : {
123 : bool hugeData;
124 0 : if (DMAReduceFlag_) {
125 0 : hugeData = curSize > SDMA_SEND_MAX_SIZE;
126 : } else {
127 0 : hugeData = (curSize * topoAttr_.userRankSize / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE) ||
128 : (curSize > SDMA_SEND_MAX_SIZE);
129 : }
130 :
131 0 : return hugeData;
132 : }
133 :
134 3 : HcclResult CollReduceScatterRingExecutor::KernelRun(const OpParam ¶m, ExecMem &execMem)
135 : {
136 3 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollReduceScatterRingExecutor][KernelRun] userRank[%u] starts.", topoAttr_.userRank);
137 3 : u32 perDataSize = 0;
138 3 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
139 :
140 3 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
141 3 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
142 :
143 3 : u32 ringNum = (topoType_ == TopoType::TOPO_TYPE_8P_RING) ? LEVEL0_PLANE_NUM_IN_8PRING :
144 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
145 :
146 3 : u32 commIndex = (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) ? topoAttr_.devicePhyId : level0CommInfo.localRank;
147 :
148 3 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
149 3 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
150 :
151 : /* ******************网口裁剪步骤: 节点内allreduce *******************************/
152 3 : std::vector<Slice> dataSegsSlice; // 数据分成ranksize份,每份的起始偏移和大小
153 3 : std::vector<std::vector<Slice> > multiStreamSlice; // 每个stream使用的数据基于用户buffer的偏移
154 3 : u32 sliceNum = level0CommInfo.localRankSize;
155 : // Slice sliceTemp;
156 3 : bool isMultiNic = topoType_ == TopoType::TOPO_TYPE_8P_RING && topoAttr_.nicList.size() != DEVICE_EIGHT;
157 3 : if (isMultiNic) {
158 0 : u64 inputDataCount = execMem.inputMem.size() / perDataSize;
159 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(inputDataCount, perDataSize, sliceNum, 0, dataSegsSlice));
160 0 : multiStreamSlice = PrepareMultiRingSlice(dataSegsSlice, param.tag);
161 0 : CHK_PRT_RET(multiStreamSlice.size() != ringNum,
162 : HCCL_ERROR("[CollReduceScatterRingExecutor][KernelRun]ringNum[%u] != multiStreamSlice size[%zu]",
163 : ringNum, multiStreamSlice.size()), HCCL_E_INTERNAL);
164 :
165 0 : CHK_RET(MultiRingAllReduce(param.tag, execMem.inputMem, execMem.scratchMem, inputDataCount,
166 : param.DataDes.dataType, param.reduceType, multiStreamSlice, param.stream, PROF_STAGE_0));
167 :
168 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.inputMem, execMem.scratchMem,
169 : const_cast<Stream&>(param.stream)));
170 : }
171 :
172 3 : std::vector<u32> &nicList = const_cast<std::vector<u32> &>(topoAttr_.nicList);
173 3 : std::vector<u32>::iterator iterNic = std::find(nicList.begin(), nicList.end(), topoAttr_.devicePhyId);
174 3 : bool innRunRet = isMultiNic && (iterNic == nicList.end());
175 3 : if (!innRunRet) { // 1. 8P ring的拓扑。2. 网口不满配。3. 当前device不出网口。 的情况下不进行节点间的reduce scatter
176 : /* ******************第一步: 节点间reducescatter *******************************/
177 3 : u32 level1RankSize = level1CommInfo.localRankSize;
178 3 : if (level1RankSize > 1) {
179 6 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, param.DataDes.dataType,
180 3 : param.reduceType);
181 3 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
182 :
183 3 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
184 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
185 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
186 0 : HCCL_INFO("ReduceScatter ring: using ring algo inter-server.");
187 0 : CHK_SMART_PTR_NULL(level1TempAlg);
188 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
189 :
190 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
191 0 : u64 ringCount = ringSize / perDataSize;
192 :
193 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
194 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID,
195 : std::vector<Slice>(0)));
196 3 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
197 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
198 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
199 0 : HCCL_INFO("ReduceScatter ring: using nhr algo inter-server.");
200 0 : CHK_SMART_PTR_NULL(level1TempAlg);
201 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
202 :
203 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
204 0 : u64 ringCount = ringSize / perDataSize;
205 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
206 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID,
207 : std::vector<Slice>(0)));
208 0 : level1TempAlg->CloseBarrier();
209 3 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
210 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
211 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
212 0 : HCCL_INFO("ReduceScatter ring: using nhr_v1 algo inter-server.");
213 0 : CHK_SMART_PTR_NULL(level1TempAlg);
214 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
215 :
216 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
217 0 : u64 ringCount = ringSize / perDataSize;
218 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
219 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID,
220 : std::vector<Slice>(0)));
221 3 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
222 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
223 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
224 0 : HCCL_INFO("ReduceScatter ring: using nonuniform-bruck algo inter-server.");
225 0 : CHK_SMART_PTR_NULL(level1TempAlg);
226 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
227 :
228 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
229 0 : u64 ringCount = ringSize / perDataSize;
230 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
231 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID,
232 : std::vector<Slice>(0)));
233 : } else {
234 6 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
235 3 : TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
236 3 : HCCL_INFO("ReduceScatter ring: using halving-doubling algo inter-server.");
237 :
238 3 : CHK_SMART_PTR_NULL(level1TempAlg);
239 3 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
240 3 : u64 inputDataCount = execMem.inputMem.size() / perDataSize; // count是output的数据个数
241 15 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, inputDataCount,
242 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID,
243 : std::vector<Slice>(0)));
244 : }
245 3 : CHK_RET(level1TempAlg->RegisterProfiler(
246 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank,
247 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, param.stream));
248 3 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
249 3 : }
250 : }
251 :
252 : /* ***********第二步: 节点内reducescatter(正常场景), 节点内多根结点scatter(网口裁剪)*****************************/
253 3 : CHK_RET(ActiveSlaveStreams(param.stream));
254 :
255 3 : bool useInlineRduce = false;
256 3 : bool isInlineReduce = IsSupportSDMAReduce(execMem.inputMem.ptr(), execMem.scratchMem.ptr(), param.DataDes.dataType,
257 3 : param.reduceType);
258 3 : useInlineRduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
259 3 : multiStreamSlice = ReduceScatterRingSlicePrepare(ringNum, sliceNum, useInlineRduce, execMem.outputMem,
260 3 : dataSegsSlice, param.tag);
261 3 : bool bRet = (multiStreamSlice.size() != ringNum);
262 3 : CHK_PRT_RET(bRet,
263 : HCCL_ERROR("[CollReduceScatterRingExecutor][KernelRun]sliceNum-1[%u] != multiStreamSlice size[%zu]", \
264 : sliceNum - 1, multiStreamSlice.size()), HCCL_E_INTERNAL);
265 :
266 3 : if (isMultiNic) { // 网口裁剪情况下需要改变slice最终在rank上位置
267 0 : PrepareMultiRingSlice(dataSegsSlice, param.tag, false, nicList); // 刷新多环ringRankList信息
268 0 : std::vector<std::vector<u32>> ringNics;
269 0 : CHK_RET(GetRingNics(param.tag, ringNics));
270 :
271 0 : for (u32 ringIdx = 0; ringIdx < ringNum; ringIdx++) { // 按第一个网口位置改变slice最终在rank上的位置
272 0 : u32 firstNicIdx = ringNics[ringIdx][0];
273 0 : std::rotate(multiStreamSlice[ringIdx].begin(), multiStreamSlice[ringIdx].begin() + firstNicIdx,
274 0 : multiStreamSlice[ringIdx].end());
275 : }
276 0 : }
277 :
278 3 : DeviceMem srcMem;
279 3 : if (isMultiNic) {
280 0 : u32 level1RankSize = topoAttr_.userRankSize / DEVICE_EIGHT; // currComm->commLevel0[0]->UserRankSize();
281 : // 每个server分配的slice大小
282 0 : CHK_PRT_RET(level1RankSize == 0,
283 : HCCL_ERROR("[CollReduceScatterRingExecutor][KernelRun]level1RankSize is illegal"), HCCL_E_PARA);
284 0 : u64 serverSliceSize = execMem.inputMem.size() / level1RankSize;
285 : // 每个服务器对应的偏移
286 0 : u32 serverIndex = level1CommInfo.localRank;
287 0 : CHK_PRT_RET(serverIndex == INVALID_VALUE_RANKID,
288 : HCCL_ERROR("[CollReduceScatterRingExecutor][KernelRun]get rank of "
289 : "bridgeRank failed, commIdx[%u]", commIndex), HCCL_E_PARA);
290 0 : u64 serverSliceOffset = serverSliceSize * serverIndex;
291 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
292 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.scratchMem, execMem.inputMem,
293 : const_cast<Stream&>(param.stream)));
294 : }
295 0 : DeviceMem reduceScatterRingOutput = execMem.scratchMem.range(serverSliceOffset, serverSliceSize);
296 0 : CHK_SMART_PTR_NULL(reduceScatterRingOutput.ptr());
297 0 : u64 countLocal = serverSliceSize / perDataSize;
298 0 : CHK_RET(MultiRingMultiRootScatter(param.tag, reduceScatterRingOutput, reduceScatterRingOutput, countLocal,
299 : param.DataDes.dataType, multiStreamSlice, serverIndex * DEVICE_EIGHT, param.stream, serverSliceOffset));
300 :
301 0 : srcMem = reduceScatterRingOutput.range(dataSegsSlice[topoAttr_.devicePhyId].offset,
302 0 : execMem.count * perDataSize);
303 0 : CHK_SMART_PTR_NULL(srcMem.ptr());
304 0 : } else {
305 3 : u32 level1RankSize = level1CommInfo.localRankSize;
306 : // 每个server分配的slice大小
307 3 : u64 serverSliceSize = execMem.inputMem.size() / level1RankSize;
308 : // 每个服务器对应的偏移
309 3 : u32 serverIndex = level1CommInfo.localRank;
310 3 : u64 serverSliceOffset = serverSliceSize * serverIndex;
311 3 : HCCL_DEBUG("inputMem.size=%llu, level0CommInfo.localRankSize=%u, serverSliceSize=%llu, serverSliceOffset=%llu "\
312 : "commIndex=%u commLevel1[commIndex]->rank=%u", execMem.inputMem.size(), level0CommInfo.localRankSize,
313 : serverSliceSize, serverSliceOffset, commIndex, level1CommInfo.localRank);
314 3 : DeviceMem reduceScatterRingInput = execMem.inputMem.range(serverSliceOffset, serverSliceSize);
315 3 : CHK_SMART_PTR_NULL(reduceScatterRingInput.ptr());
316 3 : DeviceMem reduceScatterRingOutput = execMem.scratchMem.range(serverSliceOffset, serverSliceSize);
317 3 : CHK_SMART_PTR_NULL(reduceScatterRingOutput.ptr());
318 3 : u64 countLocal = serverSliceSize / perDataSize;
319 :
320 3 : HcomCollOpInfo opInfo = {"", execMem.inputPtr, execMem.outputPtr, param.DataDes.count, param.DataDes.dataType,
321 3 : param.root, param.reduceType};
322 3 : HcomCollOpInfo *opInfoPtr = nullptr;
323 3 : if (DMAReduceFlag_) {
324 0 : opInfoPtr = &opInfo;
325 : }
326 :
327 9 : CHK_RET(MultiRingReduceScatter(param.tag, reduceScatterRingInput, reduceScatterRingOutput, countLocal,
328 : param.DataDes.dataType, param.reduceType, multiStreamSlice, param.stream, PROF_STAGE_1, serverSliceOffset,
329 : opInfoPtr));
330 :
331 6 : srcMem = execMem.inputMem.range(serverSliceOffset + dataSegsSlice[commIndex].offset,
332 6 : execMem.count * perDataSize);
333 3 : CHK_SMART_PTR_NULL(srcMem.ptr());
334 3 : }
335 :
336 3 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
337 :
338 3 : return HCCL_SUCCESS;
339 3 : }
340 :
341 :
342 0 : HcclResult CollReduceScatterRingExecutor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
343 : {
344 0 : if (CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
345 0 : return HCCL_E_UNAVAIL;
346 : }
347 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
348 0 : u32 ringNum = (topoType_ == TopoType::TOPO_TYPE_8P_RING) ? LEVEL0_PLANE_NUM_IN_8PRING :
349 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
350 0 : u32 commIndex = (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) ? topoAttr_.devicePhyId : level0CommInfo.localRank;
351 :
352 0 : if (CheckCommSize(COMM_LEVEL1, commIndex + 1) != HCCL_SUCCESS) {
353 0 : return HCCL_E_UNAVAIL;
354 : }
355 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
356 :
357 0 : return HCCL_SUCCESS;
358 0 : }
359 :
360 0 : HcclResult CollReduceScatterRingExecutor::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
361 : {
362 0 : if (level1RankSize > 1) {
363 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
364 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
365 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
366 0 : CHK_SMART_PTR_NULL(level1TempAlg);
367 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
368 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
369 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
370 0 : CHK_SMART_PTR_NULL(level1TempAlg);
371 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
372 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
373 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
374 0 : CHK_SMART_PTR_NULL(level1TempAlg);
375 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
376 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
377 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
378 0 : CHK_SMART_PTR_NULL(level1TempAlg);
379 : } else {
380 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
381 0 : TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
382 0 : CHK_SMART_PTR_NULL(level1TempAlg);
383 : }
384 0 : return HCCL_SUCCESS;
385 : }
386 0 : return HCCL_E_UNAVAIL;
387 : }
388 :
389 :
390 : REGISTER_EXEC("ReduceScatterRingExecutor", ReduceScatterRing, CollReduceScatterRingExecutor);
391 :
392 : }
393 :
|