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_gather_ring_zerocopy_executor.h"
12 :
13 : namespace hccl {
14 0 : CollAllGatherRingZerocopyExecutor::CollAllGatherRingZerocopyExecutor(
15 0 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
16 0 : : CollAllGatherExecutor(dispatcher, topoMatcher)
17 : {
18 0 : DMAReduceFlag_ = true; // 设为true,以禁用RunLoop中的本地拷贝
19 0 : desc_.isZeroCopy = true;
20 : desc_.level1SupportedAlgos
21 0 : = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
22 0 : AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
23 : desc_.level2SupportedAlgos
24 0 : = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING};
25 0 : }
26 :
27 0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcStreamNum(u32& streamNum)
28 : {
29 0 : u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ? (LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE + 1) :
30 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
31 0 : streamNum = totalStreamNum - 1;
32 0 : HCCL_INFO("[%s] tag[%s] streamNum_[%u]", __func__, tag_.c_str(), streamNum);
33 0 : return HCCL_SUCCESS;
34 : }
35 :
36 0 : void CollAllGatherRingZerocopyExecutor::ParseParam(const OpParam& param) { tag_ = param.tag; }
37 :
38 0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
39 : {
40 0 : TransportMemType inputType = TransportMemType::RESERVED;
41 0 : TransportMemType outputType = TransportMemType::RESERVED;
42 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
43 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
44 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
45 0 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
46 0 : return HCCL_SUCCESS;
47 : }
48 :
49 : HcclResult
50 0 : CollAllGatherRingZerocopyExecutor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
51 : {
52 0 : inputType = TransportMemType::CCL_INPUT;
53 0 : outputType = TransportMemType::CCL_OUTPUT;
54 0 : return HCCL_SUCCESS;
55 : }
56 :
57 0 : HcclResult CollAllGatherRingZerocopyExecutor::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 : LevelNSubCommTransport& commTransportLevel0 = opTransport[COMM_LEVEL0];
63 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
64 0 : commTransportLevel0[subCommIndex].isZeroCopy = true;
65 : }
66 0 : return HCCL_SUCCESS;
67 0 : }
68 :
69 0 : u64 CollAllGatherRingZerocopyExecutor::CalcLoopMaxCount(const u64 cclBuffSize, const u32 unitSize)
70 : {
71 0 : u64 maxCountPerLoop = cclBuffSize / topoAttr_.serverNum / HCCL_MIN_SLICE_ALIGN * HCCL_MIN_SLICE_ALIGN / unitSize;
72 0 : return maxCountPerLoop;
73 : }
74 :
75 0 : HcclResult CollAllGatherRingZerocopyExecutor::SemiRingAllGather(
76 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
77 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream, s32 profStage,
78 : const u64 baseOffset, const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>>& multRingsUserMemSlice)
79 : {
80 : (void)multRingsSliceZero;
81 : (void)tag;
82 : (void)baseOffset;
83 : (void)opInfo;
84 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
85 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
86 :
87 : // 执行
88 : std::unique_ptr<AlgTemplateBase> executor
89 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_UNIFIED_MARCH, dispatcher_);
90 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_UNIFIED_MARCH in COMM_LEVEL0", __func__);
91 0 : CHK_SMART_PTR_NULL(executor);
92 :
93 0 : CHK_RET(executor->Prepare(
94 : stream, level0CommInfo, algResResp_->paramInputMem, algResResp_->paramOutputMem, inputMem, outputMem,
95 : count * SIZE_TABLE[dataType], algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux,
96 : multRingsUserMemSlice));
97 0 : HcclResult ret = executor->RegisterProfiler(
98 : ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
99 0 : + (level0CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank,
100 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
101 0 : CHK_PRT_RET(
102 : ret != HCCL_SUCCESS,
103 : HCCL_ERROR(
104 : "[CollAllGatherRingZerocopyExecutor][SemiRingAllGather]Double ring "
105 : "AllGather failed, return[%d]",
106 : ret),
107 : ret);
108 0 : CHK_RET(executor->RunAsync());
109 0 : return ret;
110 0 : }
111 :
112 0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunIntraServerPost(const OpParam& param, ExecMem& execMem)
113 : {
114 0 : bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
115 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
116 0 : CHK_RET(GetCommRankInfoNormal(
117 : level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
118 :
119 : // 计算slice信息
120 0 : std::vector<Slice> dataSegsSlice;
121 0 : CHK_RET(CalcLevel0DataSlices(param, execMem, dataSegsSlice));
122 : // 执行AllGather
123 0 : u64 level0Count = (dataSegsSlice.size() > level0RankSize_) ? // 如果是非连续数据通信
124 : (execMem.count) :
125 0 : (execMem.count * level1RankSize_ * level2RankSize_);
126 0 : std::vector<std::vector<Slice>> multRingsUserMemSlice = {dataSegsSlice};
127 0 : if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
128 0 : HCCL_INFO("[%s] single ring AllGather", __func__);
129 0 : CHK_RET(MultiRingAllGather(
130 : param.tag, execMem.inputMem, execMem.outputMem, level0Count, param.DataDes.dataType, multRingsUserMemSlice,
131 : param.stream, PROF_STAGE_0, 0, nullptr, multRingsUserMemSlice));
132 : } else {
133 0 : CHK_PRT_RET(
134 : topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
135 : HCCL_ERROR("[%s] unknown topoType: %u", __func__, topoType_), HCCL_E_NOT_SUPPORT);
136 0 : HCCL_INFO("[%s] semi ring AllGather", __func__);
137 0 : CHK_RET(SemiRingAllGather(
138 : param.tag, execMem.inputMem, execMem.outputMem, level0Count, param.DataDes.dataType, multRingsUserMemSlice,
139 : param.stream, PROF_STAGE_0, 0, nullptr, multRingsUserMemSlice));
140 : }
141 :
142 0 : return HCCL_SUCCESS;
143 0 : }
144 :
145 : HcclResult
146 0 : CollAllGatherRingZerocopyExecutor::KernelRunInterServerPreProcess(const OpParam& param, const ExecMem& execMem)
147 : {
148 : // 将数据从User Input拷到CCL Output
149 0 : u32 dataIndex = level1Rank_ * level2RankSize_ + level2Rank_;
150 0 : u64 curSize = execMem.inputMem.size();
151 0 : DeviceMem dstMem = execMem.outputMem.range(curSize * dataIndex, curSize);
152 0 : DeviceMem srcMem = DeviceMem::create(static_cast<u8*>(execMem.inputPtr), curSize);
153 0 : Stream stream = param.stream;
154 0 : return HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream);
155 0 : }
156 :
157 0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunInterServer(const OpParam& param, ExecMem& execMem)
158 : {
159 0 : HCCL_CONFIG_INFO(
160 : HCCL_ALG, "[CollAllGatherRingZerocopyExecutor][KernelRunInterServer] The AllGatherDoubleRingExecutor starts");
161 0 : bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
162 0 : || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
163 0 : CHK_RET(GetCommRankInfoNormal(
164 : level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
165 :
166 : // 前处理
167 0 : CHK_RET(KernelRunInterServerPreProcess(param, execMem));
168 :
169 : // 计算slice
170 0 : std::vector<Slice> level1DataSegsSlice;
171 0 : CalcLevel1DataSlices(execMem.inputMem.size(), level1RankSize_, level2RankSize_, level1DataSegsSlice);
172 :
173 : // 超节点间通信
174 0 : if (level2RankSize_ > 1 && !isAHCAlgo) {
175 : // 获取对应算法的Template
176 0 : std::unique_ptr<AlgTemplateBase> level2AGTemplage;
177 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
178 : level2AGTemplage
179 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
180 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL2", __func__);
181 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
182 : level2AGTemplage
183 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
184 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL2", __func__);
185 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING) {
186 : level2AGTemplage
187 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
188 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL2", __func__);
189 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_HD) {
190 0 : level2AGTemplage = AlgTemplateRegistry::Instance().GetAlgTemplate(
191 0 : TemplateType::TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING, dispatcher_);
192 0 : HCCL_CONFIG_INFO(
193 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING in COMM_LEVEL2", __func__);
194 : } else {
195 0 : HCCL_ERROR("AllGather ring: unsupported level2 algtype [%s]", AlgTypeToStr(algType_).c_str());
196 0 : return HCCL_E_NOT_SUPPORT;
197 : }
198 0 : CHK_SMART_PTR_NULL(level2AGTemplage);
199 : // 执行算法编排
200 : DeviceMem level2OutputMem
201 0 : = execMem.outputMem.range(level1DataSegsSlice[level1Rank_].offset, level1DataSegsSlice[level1Rank_].size);
202 0 : CHK_RET(level2AGTemplage->Prepare(
203 : level2OutputMem, level2OutputMem, execMem.inputMem, execMem.count, param.DataDes.dataType, param.stream,
204 : HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, std::vector<Slice>(0),
205 : level1DataSegsSlice[level1Rank_].offset));
206 0 : CHK_RET(level2AGTemplage->RegisterProfiler(
207 : (level2RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank_, PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET,
208 : param.stream));
209 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
210 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
211 0 : CHK_RET(RunTemplate(level2AGTemplage, level2CommInfo));
212 0 : HCCL_INFO("AllGather double ring [superpod] level2 AllGather run success");
213 0 : }
214 :
215 : // 超节点内、节点间通信
216 0 : if (level1RankSize_ > 1) {
217 : // 获取对应算法的Template
218 0 : std::unique_ptr<AlgTemplateBase> level1AGTemplate;
219 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
220 : level1AGTemplate
221 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
222 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
223 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
224 : level1AGTemplate
225 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
226 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
227 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
228 : level1AGTemplate
229 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
230 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
231 0 : } else if (isAHCAlgo) {
232 : // 获取通信域分组信息
233 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
234 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
235 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(COMM_LEVEL1_AHC, globalSubGroups));
236 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
237 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
238 0 : level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(
239 0 : TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
240 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC in COMM_LEVEL1", __func__);
241 : } else {
242 0 : level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(
243 0 : TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
244 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC_BROKE in COMM_LEVEL1", __func__);
245 : }
246 0 : CHK_SMART_PTR_NULL(level1AGTemplate);
247 0 : CHK_RET(level1AGTemplate->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
248 0 : } else {
249 0 : HCCL_ERROR("AllGather ring: unsupported level1 algtype [%s]", AlgTypeToStr(algType_).c_str());
250 0 : return HCCL_E_NOT_SUPPORT;
251 : }
252 0 : CHK_SMART_PTR_NULL(level1AGTemplate);
253 : // 执行算法编排
254 0 : CHK_RET(level1AGTemplate->Prepare(
255 : execMem.outputMem, execMem.outputMem, execMem.inputMem, INVALID_U64, param.DataDes.dataType, param.stream,
256 : HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level1DataSegsSlice));
257 0 : CHK_RET(level1AGTemplate->RegisterProfiler(
258 : (level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_, PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET,
259 : param.stream));
260 0 : CommPlane commPlaneLevel1 = isAHCAlgo ? COMM_LEVEL1_AHC : COMM_LEVEL1;
261 0 : CHK_RET(CheckCommSize(commPlaneLevel1, level0Rank_ + 1));
262 0 : SubCommInfo level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0Rank_);
263 0 : CHK_RET(RunTemplate(level1AGTemplate, level1CommInfo));
264 0 : HCCL_INFO("AllGather double ring [superpod] level1 AllGather run success");
265 0 : }
266 :
267 : // 后处理
268 0 : CHK_RET(KernelRunInterServerPostProcess(param, execMem));
269 :
270 0 : return HCCL_SUCCESS;
271 0 : }
272 :
273 : HcclResult
274 0 : CollAllGatherRingZerocopyExecutor::KernelRunInterServerPostProcess(const OpParam& param, const ExecMem& execMem)
275 : {
276 0 : u32 unitSize = 0;
277 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, unitSize));
278 :
279 0 : DeviceMem dstMem;
280 0 : DeviceMem srcMem;
281 0 : u64 curSize = execMem.inputMem.size();
282 0 : Stream stream = param.stream;
283 0 : for (u32 i = 0; i < level1RankSize_; i++) {
284 0 : for (u32 j = 0; j < level2RankSize_; j++) {
285 : // 拷贝input上每个slice的数据到中转内存,源端每个slice的size固定为output的size
286 0 : u32 dstIndex = i * level2RankSize_ + j;
287 0 : u32 srcIndex = j * level1RankSize_ + i;
288 0 : srcMem = execMem.outputMem.range(dstIndex * curSize, curSize);
289 0 : dstMem = DeviceMem::create(
290 0 : static_cast<u8*>(execMem.outputPtr) + param.DataDes.count * unitSize * level0RankSize_ * srcIndex
291 0 : + param.DataDes.count * unitSize * level0Rank_,
292 0 : curSize);
293 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
294 : }
295 : }
296 0 : return HCCL_SUCCESS;
297 0 : }
298 :
299 0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcLevel0DataSlices(
300 : const OpParam& param, const ExecMem& execMem, std::vector<Slice>& dataSegsSlice)
301 : {
302 0 : return CalcIntraServerDataSlicesDiscontinuous(
303 0 : param, execMem, level0RankSize_, level1RankSize_, level2RankSize_, dataSegsSlice);
304 : }
305 :
306 : REGISTER_EXEC("AllGatherRingZerocopyExecutor", AllGatherRingZerocopy, CollAllGatherRingZerocopyExecutor);
307 :
308 : } // namespace hccl
|