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_for_910_93_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 5 : CollAllReduceRingFor91093Executor::CollAllReduceRingFor91093Executor(
17 5 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
18 5 : : CollAllReduceExecutor(dispatcher, topoMatcher)
19 : {
20 5 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
21 4 : DMAReduceFlag_ = true;
22 : } else {
23 1 : DMAReduceFlag_ = false;
24 : }
25 5 : desc_.deterministic = 1;
26 : desc_.level1SupportedAlgos
27 5 : = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
28 5 : AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
29 : desc_.level2SupportedAlgos
30 5 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING,
31 5 : AlgTypeLevel2::ALG_LEVEL2_HD};
32 5 : }
33 :
34 5 : HcclResult CollAllReduceRingFor91093Executor::CalcStreamNum(u32& streamNum)
35 : {
36 5 : u32 totalStreamNum
37 5 : = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
38 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
39 5 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
40 4 : totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
41 : }
42 5 : streamNum = totalStreamNum - 1;
43 5 : HCCL_INFO("[CollAllReduceRingFor91093Executor][CalcStreamNum] tag[%s] streamNum_[%u].", tag_.c_str(), streamNum);
44 5 : return HCCL_SUCCESS;
45 : }
46 :
47 5 : HcclResult CollAllReduceRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
48 : {
49 5 : TransportMemType inputType = TransportMemType::RESERVED;
50 5 : TransportMemType outputType = TransportMemType::RESERVED;
51 5 : CHK_RET(CalcTransportMemType(inputType, outputType));
52 5 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
53 5 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
54 5 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
55 5 : return HCCL_SUCCESS;
56 : }
57 :
58 : HcclResult
59 5 : CollAllReduceRingFor91093Executor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
60 : {
61 5 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
62 4 : inputType = TransportMemType::CCL_INPUT;
63 4 : outputType = TransportMemType::CCL_OUTPUT;
64 : } else {
65 1 : inputType = TransportMemType::PARAM_INPUT;
66 1 : outputType = TransportMemType::PARAM_OUTPUT;
67 : }
68 5 : HCCL_INFO(
69 : "[CollAllReduceRingFor91093Executor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d].",
70 : tag_.c_str(), inputType, outputType);
71 5 : return HCCL_SUCCESS;
72 : }
73 :
74 5 : HcclResult CollAllReduceRingFor91093Executor::CalcLevel0CommInfo(
75 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
76 : {
77 5 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
78 5 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
79 5 : return HCCL_SUCCESS;
80 5 : }
81 :
82 16 : bool CollAllReduceRingFor91093Executor::IsSmallData(const u64 totalSize, const u64 curSize)
83 : {
84 : (void)totalSize;
85 16 : bool smallData = IsAllReduceSmallData(curSize);
86 16 : return smallData;
87 : }
88 :
89 16 : bool CollAllReduceRingFor91093Executor::IsHugeData(const u64 curSize)
90 : {
91 32 : bool hugeData = curSize / topoAttr_.deviceNumPerAggregation / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE
92 16 : || curSize > SDMA_SEND_MAX_SIZE;
93 16 : return hugeData;
94 : }
95 :
96 5 : HcclResult CollAllReduceRingFor91093Executor::CalcLevel2CommInfo(
97 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
98 : {
99 5 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
100 5 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
101 0 : HCCL_INFO("[CollAllReduceRingFor91093Executor][CalcLevel2CommInfo] select AHC bypass level2 comm calculate");
102 0 : return HCCL_SUCCESS;
103 : }
104 :
105 5 : CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX);
106 5 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
107 1 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
108 1 : HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
109 4 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
110 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
111 0 : HCCL_INFO("[%s]Calc NBCommInfo", __func__);
112 4 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
113 4 : commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
114 4 : HCCL_INFO("[%s]Calc RingCommInfo", __func__);
115 : } else {
116 0 : commParaLevel2.commType = CommType::COMM_TAG_HALVING_DOUBLING;
117 0 : HCCL_INFO("[%s]Calc HDCommInfo", __func__);
118 : }
119 5 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
120 5 : return HCCL_SUCCESS;
121 5 : }
122 :
123 0 : HcclResult CollAllReduceRingFor91093Executor::RunIntraSeverReduceScatter(
124 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
125 : const HcclReduceOp& reductionOp, const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream,
126 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
127 : const std::vector<std::vector<Slice>>& multRingsUserMemSlice, [[maybe_unused]] const bool disableDMAReduce)
128 : {
129 0 : CHK_RET(MultiRingReduceScatter(
130 : tag, inputMem, outputMem, count, dataType, reductionOp, multRingsSliceZero, stream, profStage, baseOffset,
131 : opInfo, multRingsUserMemSlice, logicalLevel0plane_));
132 0 : return HCCL_SUCCESS;
133 : }
134 :
135 0 : HcclResult CollAllReduceRingFor91093Executor::RunIntraSeverAllGather(
136 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
137 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream, s32 profStage,
138 : const u64 baseOffset, const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>>& multRingsUserMemSlice)
139 : {
140 0 : CHK_RET(MultiRingAllGather(
141 : tag, inputMem, outputMem, count, dataType, multRingsSliceZero, stream, profStage, baseOffset, opInfo,
142 : multRingsUserMemSlice, logicalLevel0plane_));
143 0 : return HCCL_SUCCESS;
144 : }
145 :
146 16 : HcclResult CollAllReduceRingFor91093Executor::GetLevelCommInfo()
147 : {
148 16 : logicalLevel0plane_ = COMM_LEVEL0;
149 16 : CHK_RET(CheckCommSize(logicalLevel0plane_, COMM_INDEX_0 + 1));
150 16 : logicalLevel0CommInfo_ = GetSubCommInfo(logicalLevel0plane_, COMM_INDEX_0);
151 16 : u32 commIndex = logicalLevel0CommInfo_.localRank;
152 16 : bool isSelectAHC
153 16 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
154 16 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
155 16 : logicalLevel1plane_ = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
156 16 : CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
157 16 : logicalLevel1CommInfo_ = GetSubCommInfo(logicalLevel1plane_, commIndex);
158 16 : return HCCL_SUCCESS;
159 : }
160 :
161 0 : HcclResult CollAllReduceRingFor91093Executor::PrepareARSLevel1CommInfo(
162 : u32& segmentIdx, u32& commIndex, u64& hdSize, const SubCommInfo& commInfo,
163 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const std::string& tag, const std::vector<u32>& nicList)
164 : {
165 0 : segmentIdx = logicalLevel0CommInfo_.localRank;
166 0 : commIndex = logicalLevel0CommInfo_.localRank;
167 0 : CHK_PRT_RET(multRingsSliceZero.empty(), HCCL_ERROR("[Prepare][Level1CommInfo]slice map is empty"), HCCL_E_PARA);
168 :
169 0 : if (multRingsSliceZero.size() > 1) {
170 : std::vector<u32>::const_iterator iterNic
171 0 : = std::find(nicList.begin(), nicList.end(), logicalLevel0CommInfo_.localRank);
172 0 : if (iterNic != nicList.end()) { // 如果当前rank为通信网口
173 0 : u32 nicIdx = std::distance(nicList.begin(), iterNic);
174 0 : std::unique_lock<std::mutex> lock(nicSendSizeListLock_);
175 0 : auto iter = nicSendSizeList_.find(tag);
176 0 : CHK_PRT_RET(
177 : iter == nicSendSizeList_.end(),
178 : HCCL_ERROR(
179 : "[Prepare][Level1CommInfo]find tag[%s] in "
180 : "nicSendSizeList_ failed",
181 : tag.c_str()),
182 : HCCL_E_INTERNAL);
183 0 : CHK_PRT_RET(
184 : nicIdx >= iter->second.size(),
185 : HCCL_ERROR(
186 : "[Prepare][Level1CommInfo]tag[%s] nicIdx[%u] "
187 : "invalid, expect less than %zu",
188 : tag.c_str(), nicIdx, iter->second.size()),
189 : HCCL_E_INTERNAL);
190 0 : hdSize = iter->second[nicIdx]; // 通过nicSendSizeList_得到该网口传输数据量
191 0 : u32 ringRanks = multRingsSliceZero[0].size(); // 获取单个 ring 上设备的数量
192 0 : segmentIdx = ringRanks / nicList.size() * nicIdx; // 通过网口位置得到该网口传输数据的起始位置
193 0 : commIndex = segmentIdx;
194 0 : } else { // 如果当前rank不是通信网口,则不发送数据
195 0 : hdSize = 0;
196 : }
197 0 : } else if (multRingsSliceZero.size() == 1) {
198 0 : segmentIdx = commInfo.localRank;
199 0 : CHK_PRT_RET(
200 : segmentIdx >= multRingsSliceZero[0].size(),
201 : HCCL_ERROR(
202 : "[Prepare][Level1CommInfo]index is out of "
203 : "range. Idx[%u] Slice size[%zu]",
204 : segmentIdx, multRingsSliceZero[0].size()),
205 : HCCL_E_PARA);
206 0 : hdSize = multRingsSliceZero[0][segmentIdx].size;
207 0 : commIndex = segmentIdx;
208 : } else {
209 0 : return HCCL_E_PARA;
210 : }
211 0 : HCCL_INFO(
212 : "[CollAllReduceRingFor91093Executor][PrepareARSLevel1CommInfo]userRank[%u] segmentIdx[%u] commIndex[%u] "
213 : "hdSize[%llu]",
214 : topoAttr_.userRank, segmentIdx, commIndex, hdSize);
215 0 : return HCCL_SUCCESS;
216 : }
217 :
218 16 : HcclResult CollAllReduceRingFor91093Executor::KernelRun(const OpParam& param, ExecMem& execMem)
219 : {
220 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] The CollAllReduceRingFor91093Executor starts", __func__);
221 16 : CHK_RET(ActiveSlaveStreams(param.stream));
222 16 : CHK_RET(GetLevelCommInfo()); // 获取通信域
223 16 : u32 perDataSize = 0;
224 16 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
225 16 : std::vector<Slice> dataSegsSlice; // 数据分成ranksize份,每份的起始偏移和大小
226 16 : std::vector<std::vector<Slice>> multRingsSliceZero; // 数据基于该rank上环0的偏移
227 16 : u32 sliceNum = logicalLevel0CommInfo_.localRankSize;
228 : // 根据数据量计算每个环上数据的偏移和大小
229 16 : CHK_RET(AlgTemplateBase::PrepareSliceData(execMem.count, perDataSize, sliceNum, 0, dataSegsSlice));
230 :
231 : /* 三步算法step1:外层 - 节点内 reduce-scatter */
232 : // 构造ring algorithm对应的reduce-scatter实例
233 16 : std::vector<u32> mockNicList = topoAttr_.nicList;
234 16 : CHK_RET(GetNicList(mockNicList));
235 16 : u32 level0RankSize = logicalLevel0CommInfo_.localRankSize;
236 16 : bool ARSFlag = topoMatcher_->GetARSFlag();
237 16 : bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
238 :
239 : // 多环数据切分
240 16 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING || ARSDoubleRing) {
241 16 : multRingsSliceZero = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, mockNicList, logicalLevel0plane_);
242 : } else {
243 0 : multRingsSliceZero.push_back(dataSegsSlice);
244 : }
245 :
246 : // 第一步的reducescatter输出放在CCL buffer上,通过设置nullptr指示不做最后一步的DMA削减动作
247 16 : HcomCollOpInfo reduceScatterOpInfo
248 16 : = {"", execMem.inputPtr, nullptr, execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
249 16 : HcomCollOpInfo reduceScatterGraphModeOpInfo
250 16 : = {"", execMem.inputMem.ptr(), nullptr, execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
251 16 : HcomCollOpInfo* reduceScatterOpInfoPtr = nullptr;
252 16 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
253 16 : reduceScatterOpInfoPtr = &reduceScatterGraphModeOpInfo;
254 : }
255 16 : if (DMAReduceFlag_) {
256 16 : reduceScatterOpInfoPtr = &reduceScatterOpInfo;
257 : }
258 16 : bool disableDMAReduce = algOpContext_.opRetryHandler.retryEnable
259 16 : && (algOpContext_.opRetryHandler.inPlaceSupportRetryStatus
260 : == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE1
261 0 : || algOpContext_.opRetryHandler.inPlaceSupportRetryStatus
262 : == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE2);
263 16 : const std::vector<std::vector<Slice>> multRingsUserMemSliceDefault = std::vector<std::vector<Slice>>(0);
264 16 : CHK_RET(RunIntraSeverReduceScatter(
265 : param.tag, execMem.inputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, param.reduceType,
266 : multRingsSliceZero, param.stream, PROF_STAGE_0, 0, reduceScatterOpInfoPtr, multRingsUserMemSliceDefault,
267 : disableDMAReduce));
268 16 : HCCL_INFO("AllReduce double ring stage0 run success.");
269 :
270 16 : bool isSelectAHC
271 16 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
272 16 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
273 :
274 16 : if (ARSFlag
275 0 : && (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
276 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE)) {
277 0 : algType_.algoLevel1 = AlgTypeLevel1::ALG_LEVEL1_NHR;
278 : }
279 :
280 : /* 三步算法step2: 内层 - 节点间 allreduce */
281 : u64 hdSize;
282 : u32 segmentIdx;
283 : u32 commIndex;
284 16 : if (ARSFlag) {
285 0 : CHK_RET(PrepareARSLevel1CommInfo(
286 : segmentIdx, commIndex, hdSize, logicalLevel0CommInfo_, multRingsSliceZero, param.tag, mockNicList));
287 : } else {
288 16 : CHK_RET(PrepareLevel1CommInfo(
289 : segmentIdx, commIndex, hdSize, logicalLevel0CommInfo_, multRingsSliceZero, param.tag));
290 : }
291 16 : if (ARSDoubleRing && reduceScatterOpInfoPtr == nullptr) {
292 0 : DeviceMem srcMem = execMem.inputMem.range(dataSegsSlice[segmentIdx].offset, hdSize);
293 0 : DeviceMem dstMem = execMem.outputMem.range(dataSegsSlice[segmentIdx].offset, hdSize);
294 0 : HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, const_cast<Stream&>(param.stream));
295 0 : }
296 :
297 16 : u64 hdCount = hdSize / perDataSize;
298 16 : if (topoAttr_.superPodNum <= 1 || isSelectAHC) {
299 16 : DeviceMem allreduceInput = execMem.inputMem.range(dataSegsSlice[segmentIdx].offset, hdSize);
300 16 : CHK_SMART_PTR_NULL(allreduceInput);
301 16 : DeviceMem allreduceOutput = execMem.outputMem.range(dataSegsSlice[segmentIdx].offset, hdSize);
302 16 : CHK_SMART_PTR_NULL(allreduceOutput);
303 :
304 16 : u64 reduceAttr = GetReduceAttr(allreduceInput, allreduceOutput, param.DataDes.dataType, param.reduceType);
305 16 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
306 16 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
307 : level1TempAlg
308 16 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
309 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL1", __func__);
310 16 : CHK_SMART_PTR_NULL(level1TempAlg);
311 16 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
312 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
313 : level1TempAlg
314 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NHR_V1, dispatcher_);
315 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR_V1 in COMM_LEVEL1", __func__);
316 0 : CHK_SMART_PTR_NULL(level1TempAlg);
317 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
318 0 : } else if (isSelectAHC) {
319 : // 获取通信域分组信息
320 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
321 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
322 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(logicalLevel1plane_, globalSubGroups));
323 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
324 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
325 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
326 0 : TemplateType::TEMPLATE_ALL_REDUCE_AHC, dispatcher_);
327 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_AHC in COMM_LEVEL1", __func__);
328 : } else {
329 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
330 0 : TemplateType::TEMPLATE_ALL_REDUCE_AHC_BROKE, dispatcher_);
331 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_AHC_BROKE in COMM_LEVEL1", __func__);
332 : }
333 0 : CHK_SMART_PTR_NULL(level1TempAlg);
334 0 : CHK_RET(level1TempAlg->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
335 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
336 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
337 : level1TempAlg
338 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NB, dispatcher_);
339 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NB in COMM_LEVEL1", __func__);
340 0 : CHK_SMART_PTR_NULL(level1TempAlg);
341 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
342 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
343 0 : u64 curSize = execMem.count * SIZE_TABLE[param.DataDes.dataType]; // 单位 byte
344 0 : HCCL_DEBUG(
345 : "allreduce ring: curSize[%llu] deviceNumPerAggregation[%u] commLevel0Size[%u]", curSize,
346 : logicalLevel0CommInfo_.localRankSize, logicalLevel0CommInfo_.localRankSize);
347 0 : if (curSize / logicalLevel0CommInfo_.localRankSize <= NHR_ALLREDUCE_SMALL_SIZE) {
348 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
349 0 : TemplateType::TEMPLATE_ALL_REDUCE_NHR_ONESHOT, dispatcher_);
350 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR_ONESHOT in COMM_LEVEL1", __func__);
351 : } else {
352 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
353 0 : TemplateType::TEMPLATE_ALL_REDUCE_NHR, dispatcher_);
354 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR in COMM_LEVEL1", __func__);
355 : }
356 0 : CHK_SMART_PTR_NULL(level1TempAlg);
357 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
358 : } else {
359 0 : HCCL_ERROR("AllReduce ring: algType_[%u] is not supported.", algType_.algoLevel1);
360 0 : return HCCL_E_NOT_SUPPORT;
361 : }
362 16 : CHK_SMART_PTR_NULL(level1TempAlg);
363 16 : u32 rankSize = logicalLevel1CommInfo_.localRankSize;
364 : // 节点间的hd 使用环0来记录
365 80 : CHK_RET(level1TempAlg->Prepare(
366 : allreduceInput, allreduceOutput, allreduceOutput, hdCount, param.DataDes.dataType, param.stream,
367 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0), dataSegsSlice[segmentIdx].offset));
368 16 : CHK_RET(level1TempAlg->RegisterProfiler(
369 : (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + logicalLevel1CommInfo_.localRank, PROF_STAGE_1,
370 : HCCL_EXEC_STEP_NOT_SET, param.stream));
371 16 : CHK_RET(RunTemplate(level1TempAlg, logicalLevel1CommInfo_));
372 :
373 16 : HCCL_INFO("AllReduce double ring stage1 run success");
374 32 : } else {
375 : // 超节点内做reducescatter
376 0 : CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
377 0 : u32 level1RankSize = logicalLevel1CommInfo_.localRankSize;
378 0 : u64 level1Offset = dataSegsSlice[segmentIdx].offset;
379 :
380 : // 根据数据量计算每个环上数据的偏移和大小
381 0 : CHK_RET(AlgTemplateBase::PrepareSliceData(hdCount, perDataSize, level1RankSize, 0, dataSegsSlice));
382 0 : DeviceMem reducescatterInput = execMem.inputMem.range(level1Offset, hdSize);
383 0 : CHK_SMART_PTR_NULL(reducescatterInput);
384 0 : DeviceMem reducescatterOutput = execMem.outputMem.range(level1Offset, hdSize);
385 0 : CHK_SMART_PTR_NULL(reducescatterOutput);
386 0 : if (level1RankSize > 1) {
387 : u64 reduceAttr
388 0 : = GetReduceAttr(reducescatterInput, reducescatterOutput, param.DataDes.dataType, param.reduceType);
389 0 : std::unique_ptr<AlgTemplateBase> level1RSTempAlg;
390 :
391 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
392 0 : level1RSTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
393 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
394 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
395 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
396 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr));
397 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
398 0 : level1RSTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
399 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
400 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
401 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
402 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr));
403 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
404 0 : level1RSTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
405 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
406 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
407 0 : CHK_SMART_PTR_NULL(level1RSTempAlg);
408 0 : CHK_RET(level1RSTempAlg->Prepare(reduceAttr, false));
409 : } else {
410 0 : HCCL_ERROR("ReduceScatter ring: algType_[%u] is not supported.", algType_.algoLevel1);
411 0 : return HCCL_E_NOT_SUPPORT;
412 : }
413 0 : CHK_RET(level1RSTempAlg->Prepare(
414 : reducescatterInput, reducescatterInput, reducescatterOutput, hdCount, param.DataDes.dataType,
415 : param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, dataSegsSlice, level1Offset));
416 :
417 0 : CHK_RET(level1RSTempAlg->RegisterProfiler(
418 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + logicalLevel1CommInfo_.localRank, PROF_STAGE_1,
419 : HCCL_EXEC_STEP_NOT_SET, param.stream));
420 0 : CHK_RET(RunTemplate(level1RSTempAlg, logicalLevel1CommInfo_));
421 0 : HCCL_INFO("AllReduce double ring [superpod] level1 ReduceScatter run success");
422 0 : }
423 :
424 : // 超节点间做allreduce
425 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
426 0 : u32 rankSize = level2CommInfo.localRankSize;
427 0 : u32 localRank = logicalLevel1CommInfo_.localRank;
428 :
429 : DeviceMem allreduceInput
430 0 : = reducescatterInput.range(dataSegsSlice[localRank].offset, dataSegsSlice[localRank].size);
431 0 : CHK_SMART_PTR_NULL(allreduceInput);
432 : DeviceMem allreduceOutput
433 0 : = reducescatterOutput.range(dataSegsSlice[localRank].offset, dataSegsSlice[localRank].size);
434 0 : CHK_SMART_PTR_NULL(allreduceOutput);
435 :
436 0 : u64 reduceAttr = GetReduceAttr(allreduceInput, allreduceOutput, param.DataDes.dataType, param.reduceType);
437 :
438 0 : std::unique_ptr<AlgTemplateBase> level2ARTempAlg;
439 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
440 : level2ARTempAlg
441 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NB, dispatcher_);
442 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NB in COMM_LEVEL2", __func__);
443 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
444 : level2ARTempAlg
445 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NHR, dispatcher_);
446 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_NHR in COMM_LEVEL2", __func__);
447 0 : if (algoAttr_.isSupportAtomicWrite) {
448 0 : CHK_SMART_PTR_NULL(level2ARTempAlg);
449 0 : level2ARTempAlg->CloseBarrier();
450 : }
451 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
452 : level2ARTempAlg
453 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
454 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL2", __func__);
455 : } else {
456 0 : level2ARTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
457 0 : TemplateType::TEMPLATE_ALL_REDUCE_RECURSIVE_HALVING_DOUBLING, dispatcher_);
458 0 : HCCL_CONFIG_INFO(
459 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RECURSIVE_HALVING_DOUBLING in COMM_LEVEL2", __func__);
460 : }
461 0 : CHK_SMART_PTR_NULL(level2ARTempAlg);
462 0 : CHK_RET(level2ARTempAlg->Prepare(reduceAttr));
463 :
464 0 : u64 arCount = dataSegsSlice[localRank].size / perDataSize;
465 0 : CHK_RET(level2ARTempAlg->Prepare(
466 : allreduceInput, allreduceOutput, allreduceOutput, arCount, param.DataDes.dataType, param.stream,
467 : param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0),
468 : dataSegsSlice[localRank].offset + level1Offset));
469 0 : CHK_RET(level2ARTempAlg->RegisterProfiler(
470 : (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_1,
471 : HCCL_EXEC_STEP_NOT_SET, param.stream));
472 0 : CHK_RET(RunTemplate(level2ARTempAlg, level2CommInfo));
473 0 : HCCL_INFO("AllReduce double ring [superpod] level2 AllReduce run success");
474 :
475 : // 超节点内做allgather
476 0 : if (level1RankSize > 1) {
477 0 : std::unique_ptr<AlgTemplateBase> level1AGTempAlg;
478 0 : DeviceMem allgatherInput = execMem.outputMem.range(level1Offset, hdSize);
479 0 : DeviceMem allgatherOutput = execMem.outputMem.range(level1Offset, hdSize);
480 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
481 0 : level1AGTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
482 0 : TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
483 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
484 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
485 : level1AGTempAlg
486 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
487 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
488 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
489 0 : level1AGTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
490 0 : TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
491 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
492 : } else {
493 0 : HCCL_ERROR("AllGather ring: algType_[%u] is not supported.", algType_.algoLevel1);
494 0 : return HCCL_E_NOT_SUPPORT;
495 : }
496 0 : CHK_SMART_PTR_NULL(level1AGTempAlg);
497 0 : CHK_RET(level1AGTempAlg->Prepare(
498 : allgatherInput, allgatherOutput, allgatherOutput, arCount, param.DataDes.dataType, param.stream,
499 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, dataSegsSlice, level1Offset));
500 0 : CHK_RET(level1AGTempAlg->RegisterProfiler(
501 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + logicalLevel1CommInfo_.localRank, PROF_STAGE_1,
502 : HCCL_EXEC_STEP_NOT_SET, param.stream));
503 0 : CHK_RET(RunTemplate(level1AGTempAlg, logicalLevel1CommInfo_));
504 0 : HCCL_INFO("AllReduce double ring [superpod] level1 AllGather run success");
505 0 : }
506 0 : }
507 : /* 三步算法step3:外层 - 节点内 allgather */
508 : // 第三步的allgather输入放在CCL buffer上,通过设置nullptr指示要从CCL buffer获取输入
509 16 : HcomCollOpInfo allgatherOpInfo
510 16 : = {"", nullptr, execMem.outputPtr, execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
511 16 : HcomCollOpInfo allgatherOpInfoGraphModeOpInfo = {
512 16 : "", nullptr, execMem.outputMem.ptr(), execMem.count, param.DataDes.dataType, param.root, param.reduceType, 0};
513 16 : HcomCollOpInfo* allgatherOpInfoPtr = nullptr;
514 16 : if (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) {
515 16 : allgatherOpInfoPtr = &allgatherOpInfoGraphModeOpInfo;
516 : }
517 16 : if (DMAReduceFlag_) {
518 16 : allgatherOpInfoPtr = &allgatherOpInfo;
519 : }
520 48 : CHK_RET(RunIntraSeverAllGather(
521 : param.tag, execMem.inputMem, execMem.outputMem, hdCount, param.DataDes.dataType, multRingsSliceZero,
522 : param.stream, PROF_STAGE_2, 0, allgatherOpInfoPtr));
523 16 : HCCL_INFO("AllReduce double ring stage2 run success");
524 16 : return HCCL_SUCCESS;
525 16 : }
526 0 : HcclResult CollAllReduceRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
527 : {
528 0 : bool isSelectAHC
529 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
530 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
531 0 : if (isSelectAHC) {
532 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
533 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
534 :
535 0 : u32 commIndex = level0CommInfo.localRank;
536 :
537 0 : CommPlane commPlaneLevel1 = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
538 0 : CHK_RET(CheckCommSize(commPlaneLevel1, commIndex + 1));
539 0 : level1CommInfo = GetSubCommInfo(commPlaneLevel1, commIndex);
540 0 : return HCCL_SUCCESS;
541 0 : }
542 0 : if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
543 0 : return HCCL_E_UNAVAIL;
544 : }
545 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
546 :
547 0 : return HCCL_SUCCESS;
548 : }
549 :
550 : HcclResult
551 0 : CollAllReduceRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
552 : {
553 0 : bool isSelectAHC
554 0 : = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
555 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
556 0 : if (isSelectAHC) {
557 0 : CommPlane commPlaneLevel1 = COMM_LEVEL1_AHC;
558 : // 获取通信域分组信息
559 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
560 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
561 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
562 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
563 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
564 : level1TempAlg
565 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_AHC, dispatcher_);
566 0 : HCCL_INFO("allreduce ring: using ahc algo inter-server.");
567 : } else {
568 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
569 0 : TemplateType::TEMPLATE_ALL_REDUCE_AHC_BROKE, dispatcher_);
570 0 : HCCL_INFO("allreduce ring: using ahc-broke algo inter-server.");
571 : }
572 0 : CHK_SMART_PTR_NULL(level1TempAlg);
573 0 : CHK_RET(level1TempAlg->Prepare(NSLBDP_MIN_COUNT, globalSubGroups, ahcAlgOption));
574 0 : return HCCL_SUCCESS;
575 0 : }
576 0 : if (level1RankSize > 1) {
577 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
578 : level1TempAlg
579 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NB, dispatcher_);
580 0 : HCCL_INFO("AllReduce ring: using nonuniform-bruck algo inter-superPod.");
581 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
582 : level1TempAlg
583 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_NHR, dispatcher_);
584 0 : HCCL_INFO("AllReduce ring: using nonuniform-hierarchical-ring algo inter-superPod.");
585 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
586 : level1TempAlg
587 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
588 0 : HCCL_INFO("AllReduce ring: using ring algo inter-superPod.");
589 : } else {
590 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
591 0 : TemplateType::TEMPLATE_ALL_REDUCE_RECURSIVE_HALVING_DOUBLING, dispatcher_);
592 0 : HCCL_INFO("AllReduce ring: using halving-doubling algo inter-superPod.");
593 : }
594 0 : CHK_SMART_PTR_NULL(level1TempAlg);
595 0 : return HCCL_SUCCESS;
596 : }
597 0 : return HCCL_E_UNAVAIL;
598 : }
599 :
600 16 : HcclResult CollAllReduceRingFor91093Executor::GetNicList(std::vector<u32>& mockNicList)
601 : {
602 16 : mockNicList.clear();
603 16 : if (logicalLevel0plane_ == COMM_LEVEL0_LOGICAL) {
604 0 : mockNicList.reserve(logicalLevel0CommInfo_.localRankSize);
605 0 : for (u32 rankIndex = 0; rankIndex < logicalLevel0CommInfo_.localRankSize; rankIndex++) {
606 0 : mockNicList.push_back(rankIndex);
607 : }
608 : } else {
609 16 : mockNicList = topoAttr_.nicList;
610 : }
611 16 : return HCCL_SUCCESS;
612 : }
613 :
614 : REGISTER_EXEC("AllReduceRingFor91093Executor", AllReduceRingFor91093, CollAllReduceRingFor91093Executor);
615 :
616 : } // namespace hccl
|