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