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_reduce_scatter_mesh_executor.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 7 : CollReduceScatterMeshExecutor::CollReduceScatterMeshExecutor(const HcclDispatcher dispatcher,
17 7 : std::unique_ptr<TopoMatcher> &topoMatcher)
18 7 : : CollReduceScatterExecutor(dispatcher, topoMatcher)
19 : {
20 10 : DMAReduceFlag_ = false;
21 10 : }
22 :
23 12 : void CollReduceScatterMeshExecutor::ParseParam(const OpParam& param)
24 : {
25 12 : tag_ = param.tag;
26 :
27 : // 910B 图模式非确定计算,inlineReduce使能,MESH拓扑场景下,创建一个mesh平面
28 24 : bool isInlineReduce = IsSupportSDMAReduce(param.inputPtr, param.outputPtr, param.DataDes.dataType,
29 12 : param.reduceType);
30 32 : meshSinglePlane_ = (topoAttr_.deviceType == DevType::DEV_TYPE_910B) &&
31 8 : topoMatcher_->GetExternalInputHcclDeterministic() == DETERMINISTIC_DISABLE &&
32 20 : isInlineReduce && (workflowMode_ != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE);
33 12 : HCCL_DEBUG("[CollReduceScatterMeshExecutor][ParseParam]meshSinglePlane is %d", meshSinglePlane_);
34 :
35 : // 是否需要scratch memory
36 32 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE &&
37 8 : (topoAttr_.deviceType == DevType::DEV_TYPE_910B || topoAttr_.deviceType == DevType::DEV_TYPE_910_93) &&
38 20 : isSupportSDMAReduce_ && IsSupportRDMAReduce(param.DataDes.dataType, param.reduceType)) {
39 0 : scratchMemFlag_ = false;
40 : } else {
41 12 : scratchMemFlag_ = true;
42 : }
43 :
44 : // 记录图模式总数据量
45 12 : totalSize_ = topoAttr_.userRankSize * param.DataDes.count * SIZE_TABLE[param.DataDes.dataType];
46 12 : aicpuUnfoldMode_ = param.aicpuUnfoldMode;
47 12 : }
48 :
49 10 : HcclResult CollReduceScatterMeshExecutor::CalcScratchMemSize(u64& scratchMemSize)
50 : {
51 10 : if (scratchMemFlag_) {
52 10 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
53 8 : scratchMemSize = inCCLbufferSize_;
54 : } else {
55 2 : scratchMemSize = totalSize_;
56 : }
57 : } else {
58 0 : scratchMemSize = 0U;
59 : }
60 10 : HCCL_INFO("[CollReduceScatterMeshExecutor][CalcScratchMemSize] tag[%s] scratchMemSize[%llu]",
61 : tag_.c_str(), scratchMemSize);
62 10 : return HCCL_SUCCESS;
63 : }
64 :
65 10 : HcclResult CollReduceScatterMeshExecutor::CalcStreamNum(u32& streamNum)
66 : {
67 10 : u32 totalStreamNum = topoAttr_.deviceNumPerAggregation > 1U ? topoAttr_.deviceNumPerAggregation - 1U : 1U;
68 10 : streamNum = totalStreamNum - 1U;
69 10 : HCCL_INFO("[CollReduceScatterMeshExecutor][CalcStreamNum] tag[%s] streamNum[%u]", tag_.c_str(), streamNum);
70 10 : return HCCL_SUCCESS;
71 : }
72 :
73 10 : HcclResult CollReduceScatterMeshExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
74 : {
75 10 : TransportMemType inputType = TransportMemType::RESERVED;
76 10 : TransportMemType outputType = TransportMemType::RESERVED;
77 10 : CHK_RET(CalcTransportMemType(inputType, outputType));
78 10 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
79 10 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
80 10 : return HCCL_SUCCESS;
81 : }
82 :
83 10 : HcclResult CollReduceScatterMeshExecutor::CalcTransportMemType(TransportMemType &inputType,
84 : TransportMemType &outputType)
85 : {
86 10 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
87 8 : inputType = TransportMemType::CCL_INPUT;
88 8 : if (scratchMemFlag_) {
89 8 : outputType = TransportMemType::SCRATCH;
90 : } else {
91 0 : outputType = TransportMemType::CCL_OUTPUT;
92 : }
93 : } else {
94 2 : inputType = TransportMemType::PARAM_INPUT;
95 2 : if (scratchMemFlag_) {
96 2 : outputType = TransportMemType::SCRATCH;
97 : } else {
98 0 : outputType = TransportMemType::PARAM_OUTPUT;
99 : }
100 : }
101 10 : HCCL_INFO("[CollReduceScatterMeshExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
102 : tag_.c_str(), inputType, outputType);
103 10 : return HCCL_SUCCESS;
104 : }
105 :
106 10 : HcclResult CollReduceScatterMeshExecutor::CalcLevel0CommInfo(TransportMemType inputType,
107 : TransportMemType outputType,
108 : std::vector<LevelNSubCommTransport>& opTransport)
109 : {
110 10 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_MESH);
111 10 : commParaLevel0.meshSinglePlane = meshSinglePlane_;
112 10 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
113 10 : return HCCL_SUCCESS;
114 10 : }
115 :
116 0 : u64 CollReduceScatterMeshExecutor::CalcLoopMaxCount(const u32 unitSize)
117 : {
118 : // 中转内存单次最多能够接受的output count
119 0 : u64 maxCountPerLoop = inCCLbufferSize_ / (topoAttr_.userRankSize * unitSize);
120 0 : return maxCountPerLoop;
121 : }
122 :
123 0 : bool CollReduceScatterMeshExecutor::IsHugeData(const u64 curSize, OpParam *param)
124 : {
125 0 : bool hugeData = (curSize * topoAttr_.userRankSize / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE) ||
126 : (curSize > SDMA_SEND_MAX_SIZE);
127 0 : return hugeData;
128 : }
129 :
130 2 : HcclResult CollReduceScatterMeshExecutor::KernelRun(const OpParam ¶m, ExecMem &execMem)
131 : {
132 2 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollReduceScatterMeshExecutor][KernelRun] userRank[%u] starts.", topoAttr_.userRank);
133 2 : u32 perDataSize = SIZE_TABLE[param.DataDes.dataType];
134 :
135 2 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
136 2 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
137 :
138 : /* ******************第一步: 节点间reducescatter *******************************/
139 2 : u32 commIndex = level0CommInfo.localRank; // 找到rank所在的节点间平面
140 :
141 2 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
142 2 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
143 :
144 2 : u32 level1RankSize = level1CommInfo.localRankSize;
145 2 : if (level1RankSize > 1) {
146 2 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.outputMem, param.DataDes.dataType, param.reduceType);
147 2 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
148 2 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
149 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
150 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
151 0 : CHK_SMART_PTR_NULL(level1TempAlg);
152 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
153 0 : HCCL_INFO("ReduceScatter mesh: using ring algo inter-server.");
154 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
155 0 : u64 ringCount = ringSize / perDataSize;
156 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
157 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
158 2 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
159 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
160 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
161 0 : HCCL_INFO("ReduceScatter mesh: using nhr algo inter-server.");
162 0 : CHK_SMART_PTR_NULL(level1TempAlg);
163 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
164 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
165 0 : u64 ringCount = ringSize / perDataSize;
166 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
167 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
168 0 : level1TempAlg->CloseBarrier();
169 2 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
170 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
171 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
172 0 : HCCL_INFO("ReduceScatter mesh: using nhr_v1 algo inter-server.");
173 0 : CHK_SMART_PTR_NULL(level1TempAlg);
174 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
175 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
176 0 : u64 ringCount = ringSize / perDataSize;
177 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
178 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
179 2 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
180 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
181 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
182 0 : HCCL_INFO("ReduceScatter mesh: using nonuniform-bruck algo inter-server.");
183 0 : CHK_SMART_PTR_NULL(level1TempAlg);
184 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
185 0 : u64 ringSize = execMem.inputMem.size() / level1RankSize;
186 0 : u64 ringCount = ringSize / perDataSize;
187 0 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount,
188 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
189 : } else {
190 4 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
191 2 : TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
192 2 : CHK_SMART_PTR_NULL(level1TempAlg);
193 2 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
194 2 : HCCL_INFO("ReduceScatter mesh: using halving-doubling algo inter-server.");
195 2 : u64 inputDataCount = execMem.inputMem.size() / perDataSize; // count是output的数据个数
196 10 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, inputDataCount,
197 : param.DataDes.dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
198 : }
199 2 : CHK_RET(level1TempAlg->RegisterProfiler(
200 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank,
201 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, param.stream));
202 :
203 2 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
204 2 : }
205 :
206 : /* *******************第二步: 节点内reducescatter ******************************************/
207 2 : CHK_RET(ActiveSlaveStreams(param.stream));
208 :
209 2 : u32 sliceNum = level0CommInfo.localRankSize;
210 : // 根据数据量算每个环上数据的偏移和大小,把做完hd的slice均分成RankSize份
211 2 : std::vector<Slice> dataSegsSlice;
212 2 : CHK_RET(PrepareReduceScatterSliceData(execMem.count, perDataSize, sliceNum, dataSegsSlice));
213 :
214 : // 每个server分配的slice大小
215 2 : u64 serverSliceSize = execMem.inputMem.size() / level1RankSize;
216 : // 每个服务器对应的偏移
217 2 : u64 serverSliceOffset = serverSliceSize * level1CommInfo.localRank;
218 :
219 2 : HCCL_DEBUG("inputMem.size=%llu, level0CommInfo.localRankSize=%u, serverSliceSize=%llu, serverSliceOffset=%llu "\
220 : "commIndex=%u level1CommInfo.localRank=%u", execMem.inputMem.size(), level0CommInfo.localRankSize,
221 : serverSliceSize, serverSliceOffset, commIndex, level1CommInfo.localRank);
222 :
223 2 : DeviceMem reduceScatterMeshInput = execMem.inputMem.range(serverSliceOffset, serverSliceSize);
224 2 : CHK_SMART_PTR_NULL(reduceScatterMeshInput);
225 2 : DeviceMem reduceScatterMeshOutput = execMem.scratchMem.range(serverSliceOffset, serverSliceSize);
226 2 : CHK_SMART_PTR_NULL(reduceScatterMeshOutput);
227 :
228 2 : HcomCollOpInfo *opInfoPtr = nullptr;
229 :
230 2 : if (topoMatcher_->GetExternalInputHcclDeterministic() == DETERMINISTIC_DISABLE &&
231 4 : (param.DataDes.dataType != HCCL_DATA_TYPE_INT64) &&
232 2 : (topoAttr_.deviceType == DevType::DEV_TYPE_910B && param.reduceType != HCCL_REDUCE_PROD)) {
233 0 : CHK_RET(MultiStreamReduceScatterMeshAtomic(param.tag, reduceScatterMeshInput, reduceScatterMeshOutput, // 非确定性
234 : execMem.count, param.DataDes.dataType, param.reduceType, dataSegsSlice, const_cast<Stream&>(param.stream),
235 : COMM_LEVEL0, serverSliceOffset, opInfoPtr));
236 : } else {
237 2 : std::vector<std::vector<Slice> > multiStreamSlice; // 每个stream使用的数据基于用户buffer的偏移
238 : // mesh算法stream数量为rank数减1
239 2 : CHK_RET(AlgTemplateBase::PrepareSliceMeshStreams(dataSegsSlice, sliceNum - 1, multiStreamSlice));
240 2 : CHK_RET(MultiStreamReduceScatterMesh(param.tag, reduceScatterMeshInput, reduceScatterMeshOutput, // 确定性
241 : execMem.count, param.DataDes.dataType, param.reduceType, multiStreamSlice,
242 : const_cast<Stream&>(param.stream), COMM_LEVEL0, serverSliceOffset));
243 2 : }
244 :
245 4 : DeviceMem srcMem = execMem.inputMem.range(serverSliceOffset + dataSegsSlice[commIndex].offset,
246 2 : execMem.count * perDataSize);
247 2 : CHK_SMART_PTR_NULL(srcMem);
248 2 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
249 :
250 2 : return HCCL_SUCCESS;
251 2 : }
252 :
253 :
254 0 : HcclResult CollReduceScatterMeshExecutor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
255 : {
256 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
257 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
258 :
259 0 : u32 commIndex = level0CommInfo.localRank; // 找到rank所在的节点间平面
260 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
261 :
262 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
263 0 : return HCCL_SUCCESS;
264 0 : }
265 :
266 0 : HcclResult CollReduceScatterMeshExecutor::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
267 : {
268 0 : if (level1RankSize > 1) {
269 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
270 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
271 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
272 0 : CHK_SMART_PTR_NULL(level1TempAlg);
273 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
274 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
275 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
276 0 : HCCL_INFO("ReduceScatter mesh: using nhr algo inter-server.");
277 0 : CHK_SMART_PTR_NULL(level1TempAlg);
278 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
279 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
280 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
281 0 : HCCL_INFO("ReduceScatter mesh: using nhr_v1 algo inter-server.");
282 0 : CHK_SMART_PTR_NULL(level1TempAlg);
283 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
284 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
285 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
286 0 : HCCL_INFO("ReduceScatter mesh: using nonuniform-bruck algo inter-server.");
287 0 : CHK_SMART_PTR_NULL(level1TempAlg);
288 : } else {
289 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
290 0 : TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
291 0 : CHK_SMART_PTR_NULL(level1TempAlg);
292 : }
293 : }
294 0 : return HCCL_SUCCESS;
295 : }
296 :
297 : REGISTER_EXEC("ReduceScatterMeshExecutor", ReduceScatterMesh, CollReduceScatterMeshExecutor);
298 : }
|