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