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_scatter_ring_for_910_93_executor.h"
12 :
13 : namespace hccl {
14 :
15 0 : CollScatterRingFor91093Executor::CollScatterRingFor91093Executor(const HcclDispatcher dispatcher,
16 0 : std::unique_ptr<TopoMatcher> &topoMatcher)
17 0 : : CollScatterExecutor(dispatcher, topoMatcher)
18 : {
19 0 : DMAReduceFlag_ = workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
20 0 : }
21 :
22 0 : HcclResult CollScatterRingFor91093Executor::CalcStreamNum(u32& streamNum)
23 : {
24 0 : u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
25 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
26 : // scatter在910_93场景仅支持单算子模式,已有mainstream需要-1
27 0 : streamNum = totalStreamNum - 1;
28 0 : HCCL_INFO("[CollScatterRingFor91093Executor][CalcStreamNum] tag[%s] streamNum[%u]",
29 : tag_.c_str(), streamNum);
30 0 : return HCCL_SUCCESS;
31 : }
32 :
33 0 : HcclResult CollScatterRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
34 : {
35 0 : TransportMemType inputType = TransportMemType::RESERVED;
36 0 : TransportMemType outputType = TransportMemType::RESERVED;
37 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
38 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
39 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
40 0 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
41 0 : return HCCL_SUCCESS;
42 : }
43 :
44 0 : HcclResult CollScatterRingFor91093Executor::CalcLevel0CommInfo(TransportMemType inputType,
45 : TransportMemType outputType,
46 : std::vector<LevelNSubCommTransport>& opTransport)
47 : {
48 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
49 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
50 0 : return HCCL_SUCCESS;
51 0 : }
52 :
53 0 : HcclResult CollScatterRingFor91093Executor::CalcLevel1CommInfo(TransportMemType inputType,
54 : TransportMemType outputType,
55 : std::vector<LevelNSubCommTransport>& opTransport)
56 : {
57 0 : CommParaInfo commParaLevel1(COMM_LEVEL1, CommType::COMM_TAG_MAX);
58 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
59 0 : commParaLevel1.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
60 0 : HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
61 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
62 0 : commParaLevel1.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
63 0 : HCCL_INFO("[%s]Calc NBCommInfo", __func__);
64 : } else {
65 0 : commParaLevel1.commType = CommType::COMM_TAG_RING_INNER;
66 0 : HCCL_INFO("[%s]Calc RingCommInfo", __func__);
67 : }
68 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel1, opTransport[COMM_LEVEL1], inputType, outputType));
69 :
70 0 : return HCCL_SUCCESS;
71 0 : }
72 :
73 0 : HcclResult CollScatterRingFor91093Executor::CalcLevel2CommInfo(TransportMemType inputType,
74 : TransportMemType outputType,
75 : std::vector<LevelNSubCommTransport>& opTransport)
76 : {
77 : // 910_93 level2当前仅支持nhr、nb、ring算法
78 0 : CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX, root_);
79 :
80 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
81 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
82 0 : HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
83 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
84 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
85 0 : HCCL_INFO("[%s]Calc NBCommInfo", __func__);
86 : } else {
87 0 : commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
88 0 : HCCL_INFO("[%s]Calc RingCommInfo", __func__);
89 : }
90 :
91 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
92 0 : return HCCL_SUCCESS;
93 0 : }
94 :
95 0 : HcclResult CollScatterRingFor91093Executor::KernelRun(const OpParam ¶m, ExecMem &execMem)
96 : {
97 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] starts.", __func__);
98 0 : Stream& stream = const_cast<Stream&>(param.stream);
99 :
100 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize_));
101 :
102 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
103 0 : level0CommInfo_ = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
104 :
105 0 : commIndex_ = level0CommInfo_.localRank;
106 :
107 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex_ + 1));
108 0 : level1CommInfo_ = GetSubCommInfo(COMM_LEVEL1, commIndex_);
109 :
110 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
111 0 : level2CommInfo_ = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
112 :
113 0 : CHK_RET(KernelRunLevel2(param, execMem, stream));
114 0 : CHK_RET(KernelRunLevel1(param, execMem, stream));
115 0 : CHK_RET(KernelRunLevel0(param, execMem, stream));
116 :
117 0 : if (!DMAReduceFlag_) {
118 0 : DeviceMem srcMem = execMem.inputMem.range(serverSliceOffset_ + execMem.outputMem.size() * commIndex_,
119 0 : execMem.count * perDataSize_);
120 0 : CHK_SMART_PTR_NULL(srcMem.ptr());
121 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, stream));
122 0 : }
123 0 : HCCL_INFO("scatter ring run success");
124 0 : return HCCL_SUCCESS;
125 : }
126 :
127 : /* ***********超节点间scatter*********** */
128 0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel2(const OpParam ¶m, ExecMem &execMem, Stream& stream)
129 : {
130 0 : u32 level2RankSize = level2CommInfo_.localRankSize;
131 0 : u32 level2Rank = level2CommInfo_.localRank;
132 0 : subUserRankRootSupperPod_ = topoMatcher_->GetSubRootWithSuperPod(topoAttr_.userRank, param.root);
133 :
134 0 : if (level2RankSize > 1 && subUserRankRootSupperPod_ == topoAttr_.userRank) {
135 0 : u32 planeRootSupperPod = 0;
136 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL2, COMM_INDEX_0, param.root, planeRootSupperPod));
137 0 : std::unique_ptr<AlgTemplateBase> level2TempAlg;
138 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
139 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
140 0 : TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
141 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NB in COMM_LEVEL2", __func__);
142 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
143 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
144 0 : TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
145 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NHR in COMM_LEVEL2", __func__);
146 : } else {
147 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
148 0 : TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
149 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_RING in COMM_LEVEL2", __func__);
150 : }
151 :
152 0 : CHK_SMART_PTR_NULL(level2TempAlg);
153 :
154 0 : u64 level2Count = execMem.inputMem.size() / perDataSize_;
155 0 : CHK_RET(level2TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, level2Count,
156 : param.DataDes.dataType, stream, HCCL_REDUCE_RESERVED, planeRootSupperPod, std::vector<Slice>(0)));
157 0 : CHK_RET(level2TempAlg->RegisterProfiler((level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank,
158 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, stream));
159 0 : CHK_RET(RunTemplate(level2TempAlg, level2CommInfo_));
160 0 : }
161 0 : return HCCL_SUCCESS;
162 : }
163 :
164 : /* ***********节点间scatter*********** */
165 0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel1(const OpParam ¶m, ExecMem &execMem, Stream& stream)
166 : {
167 0 : u32 level2RankSize = level2CommInfo_.localRankSize;
168 0 : u32 level2Rank = level2CommInfo_.localRank;
169 0 : u32 level1RankSize = level1CommInfo_.localRankSize;
170 0 : u32 level1Rank = level1CommInfo_.localRank;
171 0 : HCCL_DEBUG("level1RankSize:%u level1Rank:%u", level1RankSize, level1Rank);
172 :
173 0 : u64 level1SliceSize = execMem.inputMem.size() / level2RankSize;
174 0 : u64 level1SliceCount = level1SliceSize / perDataSize_;
175 0 : level1SliceOffset_ = level1SliceSize * level2Rank;
176 :
177 0 : CHK_RET(topoMatcher_->GetSubRootForScatter(subUserRankRootSupperPod_, subRoot_));
178 0 : CHK_PRT_RET(subRoot_ == INVALID_VALUE_RANKID, \
179 : HCCL_ERROR("[CollScatterRingFor91093Executor][KernelRun]GetSubRootForScatter failed, " \
180 : "userRank[%u], root[%u], subRoot[%u]", topoAttr_.userRank, param.root, subRoot_), HCCL_E_INTERNAL);
181 0 : HCCL_DEBUG("[CollScatterRingFor91093Executor][KernelRun]GetSubRootForScatter, userRank[%u], root[%u], subRoot[%u]",
182 : topoAttr_.userRank, param.root, subRoot_);
183 :
184 0 : if (level1RankSize > 1 && subRoot_ == topoAttr_.userRank) {
185 0 : u32 rootRankLevel1 = 0;
186 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL1, commIndex_, subUserRankRootSupperPod_, rootRankLevel1));
187 :
188 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
189 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
190 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
191 0 : TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
192 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NB in COMM_LEVEL1", __func__);
193 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
194 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
195 0 : TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
196 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NHR in COMM_LEVEL1", __func__);
197 : } else {
198 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
199 0 : TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
200 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_RING in COMM_LEVEL1", __func__);
201 : }
202 0 : CHK_SMART_PTR_NULL(level1TempAlg);
203 :
204 0 : DeviceMem level1InputMem = execMem.inputMem.range(level1SliceOffset_, level1SliceSize);
205 0 : CHK_SMART_PTR_NULL(level1InputMem.ptr());
206 :
207 0 : CHK_RET(level1TempAlg->Prepare(level1InputMem, level1InputMem, level1InputMem, level1SliceCount,
208 : param.DataDes.dataType, stream, HCCL_REDUCE_RESERVED, rootRankLevel1, std::vector<Slice>(0),
209 : level1SliceOffset_));
210 0 : CHK_RET(level1TempAlg->RegisterProfiler((level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank,
211 : PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET, stream));
212 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo_));
213 0 : }
214 0 : return HCCL_SUCCESS;
215 : }
216 :
217 : /* ***********节点内scatter*********** */
218 0 : HcclResult CollScatterRingFor91093Executor::KernelRunLevel0(const OpParam ¶m, ExecMem &execMem, Stream& stream)
219 : {
220 : // 每个server分配的slice大小
221 0 : u32 level0RankSize = level0CommInfo_.localRankSize;
222 0 : u32 level2RankSize = level2CommInfo_.localRankSize;
223 0 : u32 level1RankSize = level1CommInfo_.localRankSize;
224 0 : u32 level1Rank = level1CommInfo_.localRank;
225 :
226 0 : u64 serverSliceSize = execMem.inputMem.size() / (level1RankSize * level2RankSize);
227 0 : serverSliceOffset_ = serverSliceSize * level1Rank + level1SliceOffset_;
228 0 : HCCL_DEBUG("inputMem.size()=%llu, commLevel0->RankSize()=%u, serverSliceSize=%llu, serverSliceOffset=%llu "\
229 : "commIndex=%u commLevel1[commIndex]->rank=%u", execMem.inputMem.size(), level0RankSize, serverSliceSize,
230 : serverSliceOffset_, commIndex_, level1Rank);
231 :
232 0 : DeviceMem scatterRingInput = execMem.inputMem.range(serverSliceOffset_, serverSliceSize);
233 0 : CHK_SMART_PTR_NULL(scatterRingInput);
234 :
235 : // 将根节点数据切分成level0RankSize份
236 0 : std::vector<Slice> dataSegsSlice; // 数据分成ranksize份,每份的起始偏移和大小
237 0 : std::vector<std::vector<Slice> > mulRingSlice; // 每个stream使用的数据基于用户buffer的偏移
238 : // 根据数据量算每个环上数据的偏移和大小
239 0 : CHK_RET(PrepareDataSlice(execMem.count, perDataSize_, level0RankSize, dataSegsSlice));
240 :
241 : u32 ringNum;
242 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
243 0 : ringNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
244 0 : mulRingSlice = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
245 : } else {
246 0 : ringNum = LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
247 0 : mulRingSlice.push_back(dataSegsSlice);
248 : }
249 0 : CHK_PRT_RET(mulRingSlice.size() != ringNum,
250 : HCCL_ERROR("[CollScatterRingFor91093Executor][KernelRunLevel0]ringNum[%u] != mulRingSlice size[%zu]",
251 : ringNum, mulRingSlice.size()),
252 : HCCL_E_INTERNAL);
253 0 : HCCL_INFO("scatter ring/scatter ring direct: using multiring algo inner-server.");
254 0 : HcomCollOpInfo *scatterOpInfoPtr = nullptr;
255 0 : HcomCollOpInfo scatterOpInfo = {"", nullptr, execMem.outputPtr, param.DataDes.count, param.DataDes.dataType,
256 0 : subRoot_};
257 0 : if (DMAReduceFlag_) {
258 0 : scatterOpInfoPtr = &scatterOpInfo;
259 : }
260 0 : CHK_RET(MultiRingScatter(param.tag, scatterRingInput, scatterRingInput, execMem.count, param.DataDes.dataType,
261 : mulRingSlice, subRoot_, stream, scatterOpInfoPtr, serverSliceOffset_));
262 0 : return HCCL_SUCCESS;
263 0 : }
264 0 : HcclResult CollScatterRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
265 : {
266 0 : if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
267 0 : return HCCL_E_UNAVAIL;
268 : }
269 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
270 :
271 0 : return HCCL_SUCCESS;
272 : }
273 :
274 0 : HcclResult CollScatterRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
275 : {
276 0 : if (level1RankSize > 1) {
277 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
278 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
279 0 : TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
280 0 : CHK_SMART_PTR_NULL(level1TempAlg);
281 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
282 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
283 0 : TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
284 0 : CHK_SMART_PTR_NULL(level1TempAlg);
285 : } else {
286 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
287 0 : TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
288 0 : CHK_SMART_PTR_NULL(level1TempAlg);
289 : }
290 0 : return HCCL_SUCCESS;
291 : }
292 0 : return HCCL_E_UNAVAIL;
293 : }
294 : REGISTER_EXEC("ScatterRingFor91093Executor", ScatterRingFor91093, CollScatterRingFor91093Executor);
295 : }
|