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