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_all_reduce_ring_zerocopy_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 : constexpr u32 TWO_RING = 2;
17 :
18 0 : CollAllReduceRingZerocopyExecutor::CollAllReduceRingZerocopyExecutor(
19 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
20 0 : : CollAllReduceExecutor(dispatcher, topoMatcher)
21 : {
22 0 : DMAReduceFlag_ = true; // 设为true,以禁用RunLoop中的本地拷贝
23 0 : desc_.isZeroCopy = true;
24 0 : desc_.deterministic = 1;
25 : desc_.level1SupportedAlgos
26 0 : = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
27 0 : AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
28 : desc_.level2SupportedAlgos
29 0 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING,
30 0 : AlgTypeLevel2::ALG_LEVEL2_HD};
31 0 : }
32 :
33 0 : HcclResult CollAllReduceRingZerocopyExecutor::CalcStreamNum(u32& streamNum)
34 : {
35 0 : u32 totalStreamNum
36 0 : = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
37 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
38 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
39 0 : totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
40 : }
41 0 : streamNum = totalStreamNum - 1;
42 0 : HCCL_INFO("[CollAllReduceRingZerocopyExecutor][CalcStreamNum] tag[%s] streamNum_[%u]", tag_.c_str(), streamNum);
43 0 : return HCCL_SUCCESS;
44 : }
45 :
46 0 : HcclResult CollAllReduceRingZerocopyExecutor::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 : HcclResult
58 0 : CollAllReduceRingZerocopyExecutor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
59 : {
60 0 : inputType = TransportMemType::CCL_INPUT;
61 0 : outputType = TransportMemType::CCL_OUTPUT;
62 0 : HCCL_INFO(
63 : "[CollAllReduceRingZerocopyExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]", tag_.c_str(),
64 : inputType, outputType);
65 0 : return HCCL_SUCCESS;
66 : }
67 :
68 0 : HcclResult CollAllReduceRingZerocopyExecutor::CalcLevel0CommInfo(
69 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
70 : {
71 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
72 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
73 0 : LevelNSubCommTransport& commTransportLevel0 = opTransport[COMM_LEVEL0];
74 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
75 0 : commTransportLevel0[subCommIndex].isZeroCopy = true;
76 : }
77 0 : return HCCL_SUCCESS;
78 0 : }
79 :
80 0 : HcclResult CollAllReduceRingZerocopyExecutor::DoubleRingReduceScatter(
81 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
82 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream,
83 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
84 : const std::vector<std::vector<Slice>> multRingsUserMemSlice)
85 : {
86 : (void)tag;
87 0 : HCCL_INFO("[CollAllReduceRingZerocopyExecutor][DoubleRingReduceScatter] DoubleRingReduceScatter starts");
88 :
89 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
90 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
91 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
92 :
93 : // 适配AlignedDoubleRing的入参
94 0 : u32 ringSize = multRingsSliceZero[0].size();
95 0 : std::vector<std::vector<u32>> rankOrders(TWO_RING, std::vector<u32>(ringSize));
96 0 : for (u32 i = 0; i < ringSize; i++) {
97 0 : rankOrders[0][i] = i;
98 0 : rankOrders[1][i] = (i == 0) ? 0 : (ringSize - i);
99 : }
100 :
101 : // 执行算法编排
102 : std::unique_ptr<AlgTemplateBase> tempAlg
103 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_DB_RING, dispatcher_);
104 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_DB_RING in COMM_LEVEL0", __func__);
105 0 : CHK_SMART_PTR_NULL(tempAlg);
106 0 : CHK_RET(tempAlg->Prepare(
107 : inputMem, inputMem, outputMem, count, dataType, stream, multRingsSliceZero, reductionOp, LEVEL0_BRIDGE_RANK_ID,
108 : baseOffset, false, reduceAttr, opInfo, topoAttr_.userRank, algResResp_->slaveStreams, algResResp_->notifiesMain,
109 : algResResp_->notifiesAux, rankOrders, multRingsUserMemSlice));
110 0 : u32 ringIndexOp = COMM_INDEX_0;
111 0 : u32 rankSize = level0RingCommInfo.localRankSize;
112 0 : CHK_RET(tempAlg->RegisterProfiler(
113 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
114 : + level0RingCommInfo.localRank,
115 : profStage, HCCL_EXEC_STEP_NOT_SET, stream));
116 0 : CHK_RET(RunTemplate(tempAlg, level0RingCommInfo));
117 :
118 0 : return HCCL_SUCCESS;
119 0 : }
120 :
121 0 : HcclResult CollAllReduceRingZerocopyExecutor::KernelRunIntraServerPre(const OpParam& param, ExecMem& execMem)
122 : {
123 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollAllReduceRingZerocopyExecutor][Run]The CollAllReduceRingZerocopyExecutor starts");
124 0 : bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
125 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
126 0 : CHK_RET(GetCommRankInfoNormal(
127 : level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
128 :
129 0 : u32 perDataSize = 0;
130 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
131 :
132 0 : std::vector<Slice> level0Datalices;
133 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(param.DataDes.count, perDataSize, level0RankSize_, 0, level0Datalices));
134 :
135 0 : HcomCollOpInfo reduceScatterOpInfo
136 0 : = {"", execMem.inputMem.ptr(), nullptr, execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
137 :
138 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
139 : // 构造slice数据
140 0 : level0MultiRingDataSlices_ = {level0Datalices};
141 : // 执行算法编排
142 0 : CHK_RET(MultiRingReduceScatter(
143 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, param.reduceType,
144 : level0MultiRingDataSlices_, param.stream, PROF_STAGE_0, 0, &reduceScatterOpInfo,
145 : level0MultiRingDataSlices_));
146 : } else {
147 0 : CHK_PRT_RET(
148 : topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
149 : HCCL_ERROR("[%s] unknown topoType: %u", __func__, topoType_), HCCL_E_NOT_SUPPORT);
150 : // 构造slice数据(适配AlignedDoubleRing算法)
151 0 : level0MultiRingDataSlices_.resize(TWO_RING);
152 0 : level0MultiRingDataSlices_[0].resize(level0RankSize_);
153 0 : level0MultiRingDataSlices_[1].resize(level0RankSize_);
154 0 : for (u32 i = 0; i < level0Datalices.size(); i++) {
155 0 : level0MultiRingDataSlices_[0][i].offset = level0Datalices[i].offset;
156 0 : level0MultiRingDataSlices_[0][i].size = level0Datalices[i].size / perDataSize / TWO_RING * perDataSize;
157 0 : u32 j = (i == 0) ? 0 : (level0RankSize_ - i);
158 0 : level0MultiRingDataSlices_[1][j].offset
159 0 : = level0MultiRingDataSlices_[0][i].offset + level0MultiRingDataSlices_[0][i].size;
160 0 : level0MultiRingDataSlices_[1][j].size = level0Datalices[i].size - level0MultiRingDataSlices_[0][i].size;
161 : }
162 : // 执行算法编排
163 0 : CHK_RET(DoubleRingReduceScatter(
164 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, param.reduceType,
165 : level0MultiRingDataSlices_, param.stream, PROF_STAGE_0, 0, &reduceScatterOpInfo,
166 : level0MultiRingDataSlices_));
167 : }
168 :
169 0 : HCCL_INFO("AllReduce double ring stage0 run success");
170 0 : return HCCL_SUCCESS;
171 0 : }
172 :
173 0 : HcclResult CollAllReduceRingZerocopyExecutor::KernelRunInterServer(const OpParam& param, ExecMem& execMem)
174 : {
175 0 : u32 perDataSize = 0;
176 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
177 0 : Stream stream = param.stream;
178 :
179 : // copy data from user_in -> ccl_in
180 0 : std::vector<Slice> level0Datalices;
181 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(param.DataDes.count, perDataSize, level0RankSize_, 0, level0Datalices));
182 0 : u64 level1DataSize = execMem.count * perDataSize;
183 0 : DeviceMem dstMem = execMem.inputMem.range(0, level1DataSize);
184 : DeviceMem srcMem
185 0 : = DeviceMem::create(static_cast<u8*>(execMem.inputPtr) + level0Datalices[level0Rank_].offset, level1DataSize);
186 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
187 0 : bool isSelectAHC
188 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
189 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
190 0 : if (topoAttr_.superPodNum <= 1 || isSelectAHC) {
191 0 : CHK_RET(KernelRunInterServerAllReduceSingleSuperpod(param, execMem, level1DataSize));
192 0 : } else {
193 0 : CHK_RET(KernelRunInterServerAllReduceMultiSuperpod(param, execMem, level1DataSize));
194 : }
195 :
196 : // copy results from ccl_out -> user_out
197 0 : srcMem = execMem.outputMem.range(0, level1DataSize);
198 : dstMem
199 0 : = DeviceMem::create(static_cast<u8*>(execMem.outputPtr) + level0Datalices[level0Rank_].offset, level1DataSize);
200 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
201 :
202 0 : HCCL_INFO("AllReduce double ring stage2 run success");
203 0 : return HCCL_SUCCESS;
204 0 : }
205 :
206 : // 单超节点场景,Level1直接执行AllReduce编排
207 0 : HcclResult CollAllReduceRingZerocopyExecutor::KernelRunInterServerAllReduceSingleSuperpod(
208 : const OpParam& param, const ExecMem& execMem, const u64 level1DataSize)
209 : {
210 : // 获取Level1通信域
211 0 : bool isSelectAHC
212 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
213 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
214 0 : CommPlane commPlaneLevel1 = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
215 0 : CHK_RET(CheckCommSize(commPlaneLevel1, level0Rank_ + 1));
216 0 : SubCommInfo level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0Rank_);
217 :
218 : // 获取Template
219 0 : u32 perDataSize = 0;
220 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
221 0 : u64 level1DataCount = level1DataSize / perDataSize;
222 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
223 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
224 : level1TempAlg
225 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
226 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL1", __func__);
227 0 : CHK_SMART_PTR_NULL(level1TempAlg);
228 0 : } else if (
229 0 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
230 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
231 : // 获取通信域分组信息
232 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
233 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
234 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
235 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
236 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
237 : level1TempAlg
238 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_AHC, dispatcher_);
239 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_AHC in COMM_LEVEL1", __func__);
240 : } else {
241 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
242 0 : TemplateType::TEMPLATE_ALL_REDUCE_AHC_BROKE, dispatcher_);
243 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_AHC_BROKE in COMM_LEVEL1", __func__);
244 : }
245 0 : CHK_SMART_PTR_NULL(level1TempAlg);
246 0 : CHK_RET(level1TempAlg->Prepare(level1DataCount, globalSubGroups, ahcAlgOption));
247 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
248 : level1TempAlg
249 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NB, dispatcher_);
250 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NB in COMM_LEVEL1", __func__);
251 0 : CHK_SMART_PTR_NULL(level1TempAlg);
252 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
253 0 : HCCL_DEBUG(
254 : "AllReduce ring: level1DataSize[%llu] deviceNumPerAggregation[%u] commLevel0Size[%u]", level1DataSize,
255 : topoAttr_.deviceNumPerAggregation, level0RankSize_);
256 0 : if (level1DataSize / topoAttr_.deviceNumPerAggregation <= NHR_ALLREDUCE_SMALL_SIZE) {
257 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
258 0 : TemplateType::TEMPLATE_ALL_REDUCE_NHR_ONESHOT, dispatcher_);
259 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR_ONESHOT in COMM_LEVEL1", __func__);
260 : } else {
261 : level1TempAlg
262 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NHR, dispatcher_);
263 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR in COMM_LEVEL1", __func__);
264 : }
265 0 : CHK_SMART_PTR_NULL(level1TempAlg);
266 : } else {
267 0 : HCCL_ERROR("AllReduce ring: unsupported level1 algtype [%s]", AlgTypeToStr(algType_).c_str());
268 0 : return HCCL_E_NOT_SUPPORT;
269 : }
270 :
271 : // 执行算法编排
272 0 : DeviceMem allreduceInput = execMem.inputMem.range(0, level1DataSize);
273 0 : DeviceMem allreduceOutput = execMem.outputMem.range(0, level1DataSize);
274 0 : u64 reduceAttr = GetReduceAttr(allreduceInput, allreduceOutput, param.DataDes.dataType, param.reduceType);
275 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
276 0 : CHK_RET(level1TempAlg->Prepare(
277 : allreduceInput, allreduceOutput, allreduceOutput, level1DataCount, param.DataDes.dataType, param.stream,
278 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0), 0));
279 0 : CHK_RET(level1TempAlg->RegisterProfiler(
280 : (level1CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank, PROF_STAGE_1,
281 : HCCL_EXEC_STEP_NOT_SET, param.stream));
282 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
283 :
284 0 : return HCCL_SUCCESS;
285 0 : }
286 :
287 : // 单超节点场景,Level1先ReduceScatter,Level2再AllReduce,最后Level1再做AllGather
288 0 : HcclResult CollAllReduceRingZerocopyExecutor::KernelRunInterServerAllReduceMultiSuperpod(
289 : const OpParam& param, const ExecMem& execMem, const u64 level1DataSize)
290 : {
291 0 : u32 perDataSize = 0;
292 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
293 :
294 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, level0Rank_ + 1));
295 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, level0Rank_);
296 :
297 : // 根据数据量计算level1的数据切分
298 0 : u64 level1DataCount = level1DataSize / perDataSize;
299 0 : std::vector<Slice> level1DataSlices;
300 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(level1DataCount, perDataSize, level1RankSize_, 0, level1DataSlices));
301 0 : u64 level2DataCount = level1DataSlices[level1Rank_].size / perDataSize;
302 0 : DeviceMem level1InputMem = execMem.inputMem.range(0, level1DataSize);
303 0 : DeviceMem level1OutputMem = execMem.outputMem.range(0, level1DataSize);
304 : DeviceMem level2InputMem
305 0 : = level1InputMem.range(level1DataSlices[level1Rank_].offset, level1DataSlices[level1Rank_].size);
306 : DeviceMem level2OutputMem
307 0 : = level1OutputMem.range(level1DataSlices[level1Rank_].offset, level1DataSlices[level1Rank_].size);
308 :
309 : // Step1:超节点内、节点间做ReduceScatter
310 0 : if (level1RankSize_ > 1) {
311 : // 获取Template
312 0 : u64 reduceAttr = GetReduceAttr(level1InputMem, level1OutputMem, param.DataDes.dataType, param.reduceType);
313 0 : std::unique_ptr<AlgTemplateBase> level1RSTempAlg;
314 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
315 0 : level1RSTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
316 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
317 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
318 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr));
319 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
320 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
321 : level1RSTempAlg
322 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
323 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
324 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr));
325 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
326 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
327 : level1RSTempAlg
328 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
329 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
330 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr, false));
331 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
332 : } else {
333 0 : HCCL_ERROR("AllReduce ring: unsupported level1 algtype [%s]", AlgTypeToStr(algType_).c_str());
334 0 : return HCCL_E_NOT_SUPPORT;
335 : }
336 :
337 : // 执行算法编排
338 0 : CHK_RET(level1RSTempAlg->Prepare(
339 : level1InputMem, level1InputMem, level1OutputMem, level1DataCount, param.DataDes.dataType, param.stream,
340 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSlices, 0));
341 0 : CHK_RET(level1RSTempAlg->RegisterProfiler(
342 : (level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_, PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET,
343 : param.stream));
344 0 : CHK_RET(RunTemplate(level1RSTempAlg, level1CommInfo));
345 0 : HCCL_INFO("AllReduce double ring [superpod] level1 ReduceScatter run success");
346 0 : }
347 :
348 : // Step2:超节点间做allreduce
349 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
350 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
351 : // ==> 获取Template
352 0 : std::unique_ptr<AlgTemplateBase> level2ARTempAlg;
353 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
354 : level2ARTempAlg
355 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NB, dispatcher_);
356 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NB in COMM_LEVEL2", __func__);
357 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
358 : level2ARTempAlg
359 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NHR, dispatcher_);
360 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR in COMM_LEVEL2", __func__);
361 0 : if (algoAttr_.isSupportAtomicWrite) {
362 0 : CHK_SMART_PTR_NULL(level2ARTempAlg);
363 0 : level2ARTempAlg->CloseBarrier();
364 : }
365 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
366 : level2ARTempAlg
367 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
368 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL2", __func__);
369 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_HD) {
370 0 : level2ARTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
371 0 : TemplateType::TEMPLATE_ALL_REDUCE_RECURSIVE_HALVING_DOUBLING, dispatcher_);
372 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RECURSIVE_HALVING_DOUBLING in COMM_LEVEL2", __func__);
373 : } else {
374 0 : HCCL_ERROR("AllReduce ring: unsupported level2 algtype [%s]", AlgTypeToStr(algType_).c_str());
375 0 : return HCCL_E_NOT_SUPPORT;
376 : }
377 0 : CHK_SMART_PTR_NULL(level2ARTempAlg);
378 : // ==> 执行算法编排
379 0 : u64 reduceAttr = GetReduceAttr(level2InputMem, level2OutputMem, param.DataDes.dataType, param.reduceType);
380 0 : CHK_RET(level2ARTempAlg->Prepare(reduceAttr));
381 0 : CHK_RET(level2ARTempAlg->Prepare(
382 : level2InputMem, level2OutputMem, level2OutputMem, level2DataCount, param.DataDes.dataType, param.stream,
383 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0), level1DataSlices[level1Rank_].offset));
384 0 : CHK_RET(level2ARTempAlg->RegisterProfiler(
385 : (level2RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank_, PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET,
386 : param.stream));
387 0 : CHK_RET(RunTemplate(level2ARTempAlg, level2CommInfo));
388 0 : HCCL_INFO("AllReduce double ring [superpod] level2 AllReduce run success");
389 :
390 : // Step3:超节点内、节点间做allgather
391 0 : if (level1RankSize_ > 1) {
392 : // 获取Template
393 0 : std::unique_ptr<AlgTemplateBase> level1AGTempAlg;
394 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
395 : level1AGTempAlg
396 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
397 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
398 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
399 : level1AGTempAlg
400 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
401 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
402 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
403 : level1AGTempAlg
404 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
405 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
406 : } else {
407 0 : HCCL_ERROR("AllReduce ring: algType_[%u] is not supported", algType_.algoLevel1);
408 0 : return HCCL_E_NOT_SUPPORT;
409 : }
410 0 : CHK_SMART_PTR_NULL(level1AGTempAlg);
411 : // 执行算法编排
412 0 : CHK_RET(level1AGTempAlg->Prepare(
413 : level1OutputMem, level1OutputMem, level1OutputMem, level1DataCount, param.DataDes.dataType, param.stream,
414 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, level1DataSlices, 0));
415 0 : CHK_RET(level1AGTempAlg->RegisterProfiler(
416 : (level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_, PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET,
417 : param.stream));
418 0 : CHK_RET(RunTemplate(level1AGTempAlg, level1CommInfo));
419 0 : HCCL_INFO("AllReduce double ring [superpod] level1 AllGather run success");
420 0 : }
421 :
422 0 : return HCCL_SUCCESS;
423 0 : }
424 :
425 0 : HcclResult CollAllReduceRingZerocopyExecutor::DoubleRingAllGather(
426 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
427 : const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream, s32 profStage, const u64 baseOffset,
428 : HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>> multRingsUserMemSlice)
429 : {
430 : (void)tag;
431 0 : HCCL_INFO("[CollAllReduceRingZerocopyExecutor][DoubleRingAllGather] DoubleRingAllGather starts");
432 :
433 : // 拿到ring环映射关系
434 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
435 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
436 :
437 : // 适配AlignedDoubleRing的入参
438 0 : u32 ringSize = multRingsSliceZero[0].size();
439 0 : std::vector<std::vector<u32>> rankOrders(TWO_RING, std::vector<u32>(ringSize));
440 0 : for (u32 i = 0; i < ringSize; i++) {
441 0 : rankOrders[1][i] = (i == 0) ? 0 : (ringSize - i);
442 0 : rankOrders[0][i] = i;
443 : }
444 :
445 : // 执行算法编排
446 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALIGNED_ALL_GATHER_DOUBLE_RING in COMM_LEVEL0", __func__);
447 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
448 0 : TemplateType::TEMPLATE_ALIGNED_ALL_GATHER_DOUBLE_RING, dispatcher_);
449 0 : CHK_SMART_PTR_NULL(tempAlg);
450 0 : CHK_RET(tempAlg->Prepare(
451 : opInfo, topoAttr_.userRank, algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux,
452 : rankOrders, multRingsUserMemSlice));
453 0 : CHK_RET(tempAlg->Prepare(
454 : outputMem, outputMem, inputMem, count, dataType, stream, multRingsSliceZero, HCCL_REDUCE_RESERVED,
455 : LEVEL0_BRIDGE_RANK_ID, baseOffset));
456 0 : CHK_RET(tempAlg->RegisterProfiler(
457 : ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (level0RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID)
458 : + level0Rank_,
459 : profStage, HCCL_EXEC_STEP_NOT_SET, stream));
460 0 : CHK_RET(RunTemplate(tempAlg, level0RingCommInfo));
461 :
462 0 : return HCCL_SUCCESS;
463 0 : }
464 :
465 0 : HcclResult CollAllReduceRingZerocopyExecutor::KernelRunIntraServerPost(const OpParam& param, ExecMem& execMem)
466 : {
467 0 : HcomCollOpInfo allgatherOpInfo = {
468 0 : "", nullptr, execMem.outputMem.ptr(), execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
469 :
470 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
471 0 : CHK_RET(MultiRingAllGather(
472 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType,
473 : level0MultiRingDataSlices_, param.stream, PROF_STAGE_2, 0, &allgatherOpInfo, level0MultiRingDataSlices_));
474 : } else {
475 0 : CHK_PRT_RET(
476 : topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
477 : HCCL_ERROR("[%s] unknown topoType: %u", __func__, topoType_), HCCL_E_NOT_SUPPORT);
478 0 : CHK_RET(DoubleRingAllGather(
479 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType,
480 : level0MultiRingDataSlices_, param.stream, PROF_STAGE_2, 0, &allgatherOpInfo, level0MultiRingDataSlices_));
481 : }
482 0 : return HCCL_SUCCESS;
483 : }
484 :
485 : REGISTER_EXEC("AllReduceRingZerocopyExecutor", AllReduceRingZerocopy, CollAllReduceRingZerocopyExecutor);
486 :
487 : } // namespace hccl
|