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_broadcast_ring_for_910_93_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 0 : CollBroadCastRingFor91093::CollBroadCastRingFor91093(
17 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
18 0 : : CollBroadcastExecutor(dispatcher, topoMatcher)
19 : {
20 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
21 0 : DMAReduceFlag_ = true;
22 : } else {
23 0 : DMAReduceFlag_ = false;
24 : }
25 : desc_.level2SupportedAlgos
26 0 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_HD};
27 0 : desc_.level1SupportedAlgos = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB};
28 0 : }
29 :
30 0 : HcclResult CollBroadCastRingFor91093::CalcStreamNum(u32& streamNum)
31 : {
32 0 : u32 totalStreamNum = 0U;
33 0 : u32 ringFactor = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ? (LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE) :
34 : (LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
35 0 : HCCL_DEBUG("[CollBroadCastRingFor91093][CalcStreamNum]ringFactor is [%u]", ringFactor);
36 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
37 0 : totalStreamNum = ringFactor * STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
38 : } else {
39 0 : totalStreamNum = ringFactor;
40 : }
41 0 : streamNum = totalStreamNum - 1;
42 0 : HCCL_INFO("[CollBroadCastRingFor91093][CalcStreamNum] tag[%s] streamNum_[%u]", tag_.c_str(), streamNum);
43 0 : return HCCL_SUCCESS;
44 : }
45 :
46 0 : HcclResult CollBroadCastRingFor91093::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
47 : {
48 0 : TransportMemType inputType = TransportMemType::RESERVED;
49 0 : TransportMemType outputType = TransportMemType::RESERVED;
50 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
51 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
52 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
53 0 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
54 0 : return HCCL_SUCCESS;
55 : }
56 :
57 0 : HcclResult CollBroadCastRingFor91093::CalcLevel0CommInfo(
58 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
59 : {
60 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
61 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
62 0 : return HCCL_SUCCESS;
63 0 : }
64 :
65 0 : HcclResult CollBroadCastRingFor91093::CalcLevel2CommInfo(
66 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
67 : {
68 0 : HCCL_DEBUG("[CollBroadCastRingFor91093][CalcLevel2CommInfo]cal for level2CommInfo");
69 0 : CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX, root_);
70 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
71 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
72 0 : HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
73 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
74 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
75 0 : HCCL_INFO("[%s]Calc NBCommInfo", __func__);
76 : } else {
77 0 : commParaLevel2.commType = CommType::COMM_TAG_HALVING_DOUBLING;
78 0 : HCCL_INFO("[%s]Calc HDCommInfo", __func__);
79 : }
80 :
81 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
82 0 : return HCCL_SUCCESS;
83 0 : }
84 :
85 0 : HcclResult CollBroadCastRingFor91093::KernelRun(const OpParam& param, ExecMem& execMem)
86 : {
87 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[BroadCastOperator][CollBroadCastRingFor91093] The CollBroadCastRingFor91093 starts.");
88 0 : u32 perDataSize = 0;
89 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
90 0 : CHK_PRT_RET(
91 : perDataSize == 0,
92 : HCCL_ERROR(
93 : "[CollBroadCastRingFor91093][KernelRun]errNo[0x%01611x] datatype[%d] is invalid",
94 : HCCL_ERROR_CODE(HCCL_E_PARA), param.DataDes.dataType),
95 : HCCL_E_PARA);
96 0 : std::vector<Slice> dataSegsSlice; // 数据分成ranksize份,每份的起始偏移和大小
97 0 : std::vector<std::vector<Slice>> mulRingSlice; // 数据基于该rank上环0的偏移
98 :
99 0 : u32 ringNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ? (LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE) :
100 : (LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
101 :
102 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
103 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
104 :
105 : // 按ranksize得到内存切分slice数
106 0 : u32 sliceNum = level0CommInfo.localRankSize;
107 : // 将根节点数据切分成sliceNum份
108 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(execMem.count, perDataSize, sliceNum, 0, dataSegsSlice));
109 0 : HCCL_DEBUG("[CollBroadCastRingFor91093][KernelRun]ringNum[%u] sliceNum[%u]", ringNum, sliceNum);
110 :
111 : /* 节点内 scatter */
112 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
113 : // 将每slice再切分成2份,按各ring的dev顺序排列
114 0 : mulRingSlice = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
115 0 : CHK_PRT_RET(
116 : mulRingSlice.size() != ringNum,
117 : HCCL_ERROR(
118 : "[CollBroadCastRingFor91093][KernelRun] ringNum[%u] !=mulRingSlice size[%zu]", ringNum,
119 : mulRingSlice.size()),
120 : HCCL_E_INTERNAL);
121 : } else {
122 0 : mulRingSlice.push_back(dataSegsSlice); // 应该offset全为0,而大小和dataSegsSlice中一样,里面的offset不使用
123 : }
124 :
125 0 : HcomCollOpInfo* scatterOpInfoPtr = nullptr;
126 0 : HcomCollOpInfo scatterOpInfo
127 0 : = {"", execMem.inputPtr, nullptr, param.DataDes.count, param.DataDes.dataType, param.root, param.reduceType, 0};
128 :
129 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
130 0 : scatterOpInfoPtr = &scatterOpInfo;
131 0 : CHK_RET(ActiveSlaveStreams(param.stream));
132 0 : CHK_RET(DoubleRingScatter(
133 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, mulRingSlice,
134 : param.root, param.stream, scatterOpInfoPtr));
135 : } else {
136 0 : if (DMAReduceFlag_) {
137 0 : scatterOpInfoPtr = &scatterOpInfo;
138 : }
139 0 : CHK_RET(MultiRingScatter(
140 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, mulRingSlice,
141 : param.root, param.stream, scatterOpInfoPtr));
142 : }
143 0 : HCCL_INFO("[CollBroadCastRingFor91093][KernelRun] level0-scatter run success");
144 :
145 0 : u64 level1DataSize = 0;
146 0 : u32 commIndex = 0;
147 0 : u32 segmentIdx = 0;
148 0 : CHK_RET(PrepareLevel1CommInfo(segmentIdx, commIndex, level1DataSize, level0CommInfo, mulRingSlice, param.tag));
149 0 : u64 level1DataCount = level1DataSize / perDataSize;
150 0 : HCCL_DEBUG(
151 : "[CollBroadCastRingFor91093][KernelRun]usrRank[%u] level1 use level1DataCount[%llu]", topoAttr_.userRank,
152 : level1DataCount);
153 :
154 : // level 1 通信域获取
155 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
156 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
157 0 : HCCL_DEBUG(
158 : "[CollBroadCastRingFor91093][KernelRun]commIdx:%u TagCommInfo[%s].commLevel1.size():%u", commIndex,
159 : param.tag.c_str(), level1CommInfo.localRankSize);
160 0 : if (topoAttr_.superPodNum <= 1) {
161 0 : HCCL_INFO("Broadcast double ring No level2.");
162 : /* step2: server间 broadcast */
163 0 : bool isUsedRegister = false;
164 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
165 0 : u64 curSize = execMem.count * SIZE_TABLE[param.DataDes.dataType];
166 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
167 0 : HCCL_DEBUG(
168 : "broadcast ring: curSize[%llu] deviceNumPerAggregation[%u] commLevel0Size[%u]", curSize,
169 : topoAttr_.deviceNumPerAggregation, level0CommInfo.localRankSize);
170 0 : if (curSize / topoAttr_.deviceNumPerAggregation <= NHR_BCAST_SMALL_SIZE) {
171 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
172 0 : TemplateType::TEMPLATE_BROADCAST_NHR_ONESHOT, dispatcher_);
173 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR_ONESHOT in COMM_LEVEL1", __func__);
174 : } else {
175 : level1TempAlg
176 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
177 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR in COMM_LEVEL1", __func__);
178 : }
179 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
180 0 : isUsedRegister = true;
181 : level1TempAlg
182 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR_V1, dispatcher_);
183 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR_V1 in COMM_LEVEL1", __func__);
184 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
185 0 : const u32 level1RankSize = level1CommInfo.localRankSize;
186 0 : if (ShouldUseBinaryBroadcastOfNB(
187 0 : curSize / topoAttr_.deviceNumPerAggregation, level1RankSize, topoAttr_.userRankSize,
188 0 : topoAttr_.deviceNumPerAggregation)) {
189 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
190 0 : TemplateType::TEMPLATE_BROADCAST_NB_BINARY, dispatcher_);
191 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NB_BINARY in COMM_LEVEL1", __func__);
192 : } else {
193 : level1TempAlg
194 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
195 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NB in COMM_LEVEL1", __func__);
196 : }
197 : } else {
198 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
199 0 : TemplateType::TEMPLATE_BROADCAST_RECURSIVE_HD, dispatcher_);
200 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_RECURSIVE_HD in COMM_LEVEL1", __func__);
201 : }
202 0 : CHK_SMART_PTR_NULL(level1TempAlg);
203 :
204 0 : u32 subUserrankRoot = topoMatcher_->GetSubRootUserRank(topoAttr_.userRank, param.root);
205 0 : CHK_PRT_RET(
206 : subUserrankRoot == INVALID_VALUE_RANKID,
207 : HCCL_ERROR(
208 : "[CollBroadCastRingFor91093][KernelRun]subUserrankRoot[%u] is invalid,userRank[%u],root[%u]",
209 : subUserrankRoot, topoAttr_.userRank, param.root),
210 : HCCL_E_INTERNAL);
211 0 : u32 planeRoot = 0;
212 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL1, commIndex, subUserrankRoot, planeRoot));
213 0 : u32 ranksize = level1CommInfo.localRankSize;
214 : // 节点间的hd 使用环0来记录
215 0 : if (isUsedRegister) {
216 0 : PrepareData prepareData;
217 0 : prepareData.inputMem = execMem.inputMem;
218 0 : prepareData.outputMem = execMem.inputMem;
219 0 : prepareData.scratchMem = execMem.outputMem;
220 0 : prepareData.count = level1DataCount;
221 0 : prepareData.dataType = param.DataDes.dataType;
222 0 : prepareData.stream = param.stream;
223 0 : prepareData.reductionOp = HCCL_REDUCE_RESERVED;
224 0 : prepareData.root = planeRoot;
225 0 : prepareData.baseOffset = dataSegsSlice[segmentIdx].offset;
226 0 : CHK_RET(level1TempAlg->Prepare(prepareData));
227 0 : } else {
228 0 : CHK_RET(level1TempAlg->Prepare(
229 : execMem.inputMem, execMem.inputMem, execMem.outputMem, level1DataCount, param.DataDes.dataType,
230 : param.stream, HCCL_REDUCE_RESERVED, planeRoot, std::vector<Slice>(0),
231 : dataSegsSlice[segmentIdx].offset));
232 : }
233 :
234 0 : CHK_RET(level1TempAlg->RegisterProfiler(
235 : (ranksize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank, PROF_STAGE_1,
236 : HCCL_EXEC_STEP_NOT_SET, param.stream));
237 :
238 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
239 :
240 0 : HCCL_INFO("Broadcast double ring stage1 run success");
241 0 : } else {
242 0 : HCCL_INFO("Broadcast double ring with Level2.");
243 : /* step2: 节点间 scatter */
244 : // 按level1RankSize得到内存切分slice数
245 0 : u32 level1RankSize = level1CommInfo.localRankSize;
246 0 : u64 level1Offset = dataSegsSlice[segmentIdx].offset;
247 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(level1DataCount, perDataSize, level1RankSize, 0, dataSegsSlice));
248 :
249 0 : DeviceMem level1InputMem = execMem.inputMem.range(level1Offset, level1DataSize);
250 0 : CHK_SMART_PTR_NULL(level1InputMem);
251 0 : DeviceMem level1OutputMem = execMem.outputMem.range(level1Offset, level1DataSize);
252 0 : CHK_SMART_PTR_NULL(level1OutputMem);
253 :
254 0 : if (level1RankSize > 1) {
255 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
256 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
257 : level1TempAlg
258 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NHR, dispatcher_);
259 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NHR in COMM_LEVEL1", __func__);
260 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
261 : level1TempAlg
262 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_NB, dispatcher_);
263 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_NB in COMM_LEVEL1", __func__);
264 : } else {
265 0 : HCCL_ERROR("broadcast level1 only supports NB/NHR algo. not support algType_[%u]", algType_.algoLevel1);
266 0 : return HCCL_E_NOT_SUPPORT;
267 : }
268 0 : CHK_SMART_PTR_NULL(level1TempAlg);
269 :
270 : /* 获取每个超节点内的subroot */
271 0 : u32 subPodRoot = topoMatcher_->GetSubRootWithSuperPod(topoAttr_.userRank, param.root);
272 : /* 获取超节点内的节点在每个level1通信域的对应卡subServerRootUsrRank */
273 0 : u32 subServerRootUsrRank = topoMatcher_->GetSubRootUserRank(topoAttr_.userRank, subPodRoot);
274 0 : u32 level1RootRank = INVALID_VALUE_RANKID;
275 : /* 用此卡subServerRootUsrRank 获取每个level1通信域的相对root idx */
276 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL1, commIndex, subServerRootUsrRank, level1RootRank));
277 0 : CHK_PRT_RET(
278 : level1RootRank == INVALID_VALUE_RANKID,
279 : HCCL_ERROR("[CollBroadCastRingFor91093][KernelRun] get rootRank IDX in level1 failed."), HCCL_E_PARA);
280 :
281 0 : CHK_RET(level1TempAlg->Prepare(
282 : level1InputMem, level1InputMem, level1InputMem, level1DataCount, param.DataDes.dataType, param.stream,
283 : HCCL_REDUCE_RESERVED, level1RootRank, dataSegsSlice, level1Offset));
284 0 : CHK_RET(level1TempAlg->RegisterProfiler(
285 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank, PROF_STAGE_1,
286 : HCCL_EXEC_STEP_NOT_SET, param.stream));
287 :
288 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
289 0 : HCCL_INFO("Broadcast double ring [superpod] level1 run success");
290 0 : }
291 :
292 : /* step3: 超节点间 broadcast */
293 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
294 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
295 0 : u32 level2RankSize = level2CommInfo.localRankSize;
296 0 : u32 localRank = level1CommInfo.localRank;
297 0 : u32 subUserrankRootSupperPod = topoMatcher_->GetSubRootUserRankWithSuperPod(topoAttr_.userRank, param.root);
298 0 : CHK_PRT_RET(
299 : subUserrankRootSupperPod == INVALID_VALUE_RANKID,
300 : HCCL_ERROR(
301 : "[CollBroadCastRingFor91093][KernelRun]subUserrankRootSupperPod[%u] is invalid,userRank[%u],"
302 : "root[%u]",
303 : subUserrankRootSupperPod, topoAttr_.userRank, param.root),
304 : HCCL_E_INTERNAL);
305 0 : u32 planeRootSupperPod = 0;
306 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL2, COMM_INDEX_0, subUserrankRootSupperPod, planeRootSupperPod));
307 0 : HCCL_DEBUG(
308 : "level2 get root info as: subUserrankRootSupperPod[%u], planeRootSupperPod[%u]", subUserrankRootSupperPod,
309 : planeRootSupperPod);
310 :
311 0 : std::unique_ptr<AlgTemplateBase> level2TempAlg;
312 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
313 : level2TempAlg
314 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
315 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NB in COMM_LEVEL2", __func__);
316 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
317 : level2TempAlg
318 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
319 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_NHR in COMM_LEVEL2", __func__);
320 : } else {
321 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
322 0 : TemplateType::TEMPLATE_BROADCAST_RECURSIVE_HD, dispatcher_);
323 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_BROADCAST_RECURSIVE_HD in COMM_LEVEL2", __func__);
324 : }
325 0 : CHK_SMART_PTR_NULL(level2TempAlg);
326 0 : u64 bcastCount = dataSegsSlice[localRank].size / perDataSize;
327 :
328 0 : CHK_RET(level2TempAlg->Prepare(
329 : execMem.inputMem, execMem.inputMem, execMem.outputMem, bcastCount, param.DataDes.dataType, param.stream,
330 : HcclReduceOp::HCCL_REDUCE_RESERVED, planeRootSupperPod, std::vector<Slice>(0),
331 : dataSegsSlice[localRank].offset + level1Offset));
332 0 : HCCL_DEBUG(
333 : "[superpod]Broadcast level2-broadcast : dataSegsSlice[localRank].offset[%llu]"
334 : "dataSegsSlice[localRank].size[%llu] level1Offset[%llu]",
335 : dataSegsSlice[localRank].offset, dataSegsSlice[localRank].size, level1Offset);
336 :
337 0 : CHK_RET(level2TempAlg->RegisterProfiler(
338 : (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_1,
339 : HCCL_EXEC_STEP_NOT_SET, param.stream));
340 0 : CHK_RET(RunTemplate(level2TempAlg, level2CommInfo));
341 0 : HCCL_INFO("[CollBroadCastRingFor91093][superpod]Broadcast level2-broadcast run success");
342 :
343 : /* step4: 节点间 allgather */
344 0 : if (level1RankSize > 1) {
345 0 : std::unique_ptr<AlgTemplateBase> level1AGTempAlg;
346 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
347 : level1AGTempAlg
348 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
349 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
350 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
351 0 : level1AGTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
352 0 : TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
353 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
354 : } else {
355 0 : HCCL_ERROR("AllGather ring: algType_[%u] is not supported.", algType_.algoLevel1);
356 0 : return HCCL_E_NOT_SUPPORT;
357 : }
358 :
359 0 : CHK_SMART_PTR_NULL(level1AGTempAlg);
360 0 : CHK_RET(level1AGTempAlg->Prepare(
361 : level1InputMem, level1OutputMem, level1OutputMem, bcastCount, param.DataDes.dataType, param.stream,
362 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, dataSegsSlice, level1Offset));
363 0 : CHK_RET(level1AGTempAlg->RegisterProfiler(
364 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank, PROF_STAGE_1,
365 : HCCL_EXEC_STEP_NOT_SET, param.stream));
366 0 : CHK_RET(RunTemplate(level1AGTempAlg, level1CommInfo));
367 0 : HCCL_INFO("[CollBroadCastRingFor91093]broadcast [superpod] level1 allgather run success");
368 0 : }
369 0 : }
370 :
371 : /* step 3 or 5: 节点内 allgatherring */
372 0 : HcomCollOpInfo allgatherOpInfo = {
373 0 : "", nullptr, execMem.outputPtr, param.DataDes.count, param.DataDes.dataType, param.root, param.reduceType, 0};
374 0 : HcomCollOpInfo* allgatherOpInfoPtr = (DMAReduceFlag_) ? (&allgatherOpInfo) : (nullptr);
375 :
376 0 : CHK_RET(MultiRingAllGather(
377 : param.tag, execMem.inputMem, execMem.outputMem, level1DataCount, param.DataDes.dataType, mulRingSlice,
378 : param.stream, PROF_STAGE_2, 0, allgatherOpInfoPtr));
379 0 : HCCL_INFO("Broadcast double ring stage2 run success");
380 :
381 0 : return HCCL_SUCCESS;
382 0 : }
383 0 : HcclResult CollBroadCastRingFor91093::DoubleRingScatter(
384 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
385 : const std::vector<std::vector<Slice>> multRingsSliceZero, u32 root, Stream stream, const HcomCollOpInfo* opInfo,
386 : const u64 baseOffset)
387 : {
388 0 : HCCL_INFO("[BroadCastOperator][CollBroadCastRingFor91093] DoubleRingScatter starts.");
389 0 : HcclResult ret = HCCL_SUCCESS;
390 0 : u32 ringNum = multRingsSliceZero.size();
391 :
392 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
393 :
394 0 : std::vector<std::vector<u32>> ringNics;
395 0 : CHK_RET(GetRingNics(tag, ringNics));
396 :
397 : // 拿到ring环映射关系
398 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
399 0 : SubCommInfo level0CommInfo1 = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_1);
400 :
401 0 : auto nicList = topoAttr_.nicList;
402 : std::vector<std::vector<u32>> multiRingsOrder
403 0 : = GetRingsOrderByTopoType(level0CommInfo.localRankSize, topoType_, nicList);
404 :
405 0 : std::vector<std::vector<u32>> doubleRingsOrders;
406 0 : std::vector<std::vector<Slice>> doubleRingUserMemInputSlices;
407 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
408 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
409 0 : CHK_PRT_RET(
410 : singleRingSliceZero.empty(),
411 : HCCL_ERROR("[CollBroadCastRingFor91093][DoubleRingScatter]singleRingSliceZero is empty"), HCCL_E_INTERNAL);
412 :
413 : // 生成userMemIn_上对应的slices
414 0 : std::vector<Slice> userMemInputSlices;
415 0 : CHK_RET(
416 : CalUserMemSlices(dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemInputSlices));
417 0 : doubleRingUserMemInputSlices.push_back(userMemInputSlices);
418 0 : std::vector<u32> rankOrder;
419 0 : CHK_RET(GetRankOrder(multiRingsOrder, ringIndex, rankOrder));
420 0 : doubleRingsOrders.push_back(rankOrder);
421 0 : }
422 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
423 0 : TemplateType::TEMPLATE_SCATTER_DOUBLE_RING_DIRECT, dispatcher_);
424 0 : CHK_SMART_PTR_NULL(tempAlg);
425 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_SCATTER_DOUBLE_RING_DIRECT in COMM_LEVEL0", __func__);
426 :
427 0 : CHK_RET(tempAlg->Prepare(
428 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, level0CommInfo1.localRank, algResResp_->slaveStreams,
429 : algResResp_->notifiesMain, algResResp_->notifiesAux, doubleRingsOrders, multRingsSliceZero,
430 : doubleRingUserMemInputSlices));
431 :
432 0 : u32 rootRank = 0;
433 0 : ret = GetRankByUserRank(COMM_LEVEL0, COMM_INDEX_0, root, rootRank);
434 0 : CHK_PRT_RET(
435 : ret != HCCL_SUCCESS,
436 : HCCL_ERROR("[CollBroadCastRingFor91093][DoubleRingScatter]invalid root [%u] to get userrank", root), ret);
437 0 : ret = tempAlg->Prepare(
438 0 : inputMem, inputMem, outputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, rootRank, std::vector<Slice>(0),
439 : baseOffset);
440 0 : CHK_PRT_RET(
441 : ret != HCCL_SUCCESS,
442 : HCCL_ERROR("[CollBroadCastRingFor91093][DoubleRingScatter]scatter(ring) prepare failed, return[%d]", ret), ret);
443 :
444 0 : u32 rankSize = level0CommInfo.localRankSize;
445 0 : ret = tempAlg->RegisterProfiler(
446 0 : (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank, PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET,
447 : stream);
448 0 : CHK_PRT_RET(
449 : ret != HCCL_SUCCESS,
450 : HCCL_ERROR(
451 : "[CollBroadCastRingFor91093][DoubleRingScatter]scatter(ring) register profiler failed,return[%d]", ret),
452 : ret);
453 :
454 0 : ret = RunTemplate(tempAlg, level0CommInfo);
455 0 : CHK_PRT_RET(
456 : ret != HCCL_SUCCESS,
457 : HCCL_ERROR("[CollBroadCastRingFor91093][DoubleRingScatter]scatter(ring) run failed, return[%d]", ret), ret);
458 :
459 0 : HCCL_INFO("[CollBroadCastRingFor91093] double ring scatter run success");
460 0 : return HCCL_SUCCESS;
461 0 : }
462 :
463 0 : HcclResult CollBroadCastRingFor91093::Getlevel1CommRank(SubCommInfo& level1CommInfo)
464 : {
465 0 : if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
466 0 : return HCCL_E_UNAVAIL;
467 : }
468 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
469 :
470 0 : return HCCL_SUCCESS;
471 : }
472 :
473 0 : HcclResult CollBroadCastRingFor91093::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
474 : {
475 0 : if (level1RankSize > 1) {
476 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
477 : level1TempAlg
478 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
479 0 : HCCL_INFO("[superpod]Broadcast level2-broadcast: using nonuniform-bruck algo inter-superPod.");
480 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
481 : level1TempAlg
482 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
483 0 : HCCL_INFO("[superpod]Broadcast level2-broadcast: using nonuniform-hierarchical-ring algo inter-superPod.");
484 : } else {
485 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
486 0 : TemplateType::TEMPLATE_BROADCAST_RECURSIVE_HD, dispatcher_);
487 0 : HCCL_INFO("[superpod]Broadcast level2-broadcast: using Recursive halving-doubling algo inter-superPod.");
488 : }
489 0 : CHK_SMART_PTR_NULL(level1TempAlg);
490 0 : return HCCL_SUCCESS;
491 : }
492 0 : return HCCL_E_UNAVAIL;
493 : }
494 :
495 : REGISTER_EXEC("BroadCastRingFor91093Executor", BroadCastRingFor91093, CollBroadCastRingFor91093);
496 :
497 : } // namespace hccl
|