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