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_comm_executor.h"
12 : #include "stream_active_manager.h"
13 : #include "device_capacity.h"
14 :
15 : namespace hccl {
16 179 : CollCommExecutor::CollCommExecutor(const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
17 179 : : CollNativeExecutorBase(dispatcher, topoMatcher)
18 179 : {}
19 :
20 0 : HcclResult CollCommExecutor::GetSubStreamInfoOnOneRing(
21 : const u32 ringIndex, std::vector<Stream>& subStreamsInOneRing,
22 : std::vector<std::shared_ptr<LocalNotify>>& mainSignalsInOneRing,
23 : std::vector<std::shared_ptr<LocalNotify>>& subSignalsInOneRing)
24 : {
25 0 : u32 ringNum = GetLevel0RingNum();
26 0 : if (ringNum == LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE * STREAM_NUM_FOR_DMAREDUCE_ONE_RING) {
27 : // double ring
28 0 : subStreamsInOneRing.push_back(algResResp_->slaveStreams[ringIndex + 1]);
29 0 : mainSignalsInOneRing.push_back(algResResp_->notifiesMain[ringIndex + 1]);
30 0 : subSignalsInOneRing.push_back(algResResp_->notifiesAux[ringIndex + 1]);
31 0 : } else if (ringNum == LEVEL0_PLANE_NUM_IN_NPRING_SINGLE * STREAM_NUM_FOR_DMAREDUCE_ONE_RING) {
32 : // single ring
33 0 : subStreamsInOneRing.push_back(algResResp_->slaveStreams[ringIndex]);
34 0 : mainSignalsInOneRing.push_back(algResResp_->notifiesMain[ringIndex]);
35 0 : subSignalsInOneRing.push_back(algResResp_->notifiesAux[ringIndex]);
36 : }
37 0 : return HCCL_SUCCESS;
38 : }
39 :
40 0 : u32 CollCommExecutor::GetLevel0RingNum() const { return algResResp_->slaveStreams.size() + 1; }
41 :
42 0 : HcclResult CollCommExecutor::MultiRingAllReduce(
43 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType dataType,
44 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>>& multRingsSliceZero, Stream stream,
45 : s32 profStage, const u64 baseOffset)
46 : {
47 0 : HcclResult ret = HCCL_SUCCESS;
48 0 : u32 ringNum = multRingsSliceZero.size();
49 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
50 :
51 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
52 :
53 0 : std::vector<std::vector<u32>> ringNics;
54 0 : CHK_RET(GetRingNics(tag, ringNics));
55 :
56 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
57 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
58 0 : CHK_PRT_RET(
59 : singleRingSliceZero.empty(),
60 : HCCL_ERROR("[CollCommExecutor][MultiRingAllReduce]singleRingSliceZero is empty"), HCCL_E_INTERNAL);
61 :
62 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, ringIndex);
63 :
64 0 : u32 rankSize = level0RingCommInfo.localRankSize;
65 0 : u32 ringIndexOp = ringIndex;
66 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
67 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
68 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL0", __func__);
69 0 : CHK_SMART_PTR_NULL(tempAlg);
70 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
71 :
72 0 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
73 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
74 0 : CHK_RET(StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
75 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr()));
76 : }
77 :
78 0 : ret = LocalNotify::Wait(
79 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
80 0 : CHK_PRT_RET(
81 : ret != HCCL_SUCCESS,
82 : HCCL_ERROR("[CollCommExecutor][MultiRingAllReduce]stream[%u] wait failed", ringIndex), ret);
83 0 : ret = tempAlg->Prepare(
84 0 : inputMem, outputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex], reductionOp,
85 0 : LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset, ringNics[ringIndex]);
86 0 : CHK_PRT_RET(
87 : ret != HCCL_SUCCESS,
88 : HCCL_ERROR(
89 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) prepare failed,"
90 : "return[%d]",
91 : ringIndex, ret),
92 : ret);
93 :
94 0 : ret = tempAlg->RegisterProfiler(
95 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
96 0 : + level0RingCommInfo.localRank,
97 0 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
98 0 : CHK_PRT_RET(
99 : ret != HCCL_SUCCESS,
100 : HCCL_ERROR(
101 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) register Profiler "
102 : "failed,return[%d]",
103 : ringIndex, ret),
104 : ret);
105 :
106 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
107 0 : CHK_PRT_RET(
108 : ret != HCCL_SUCCESS,
109 : HCCL_ERROR(
110 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) run failed,"
111 : "return[%d]",
112 : ringIndex, ret),
113 : ret);
114 :
115 0 : ret = LocalNotify::Post(
116 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
117 0 : CHK_PRT_RET(
118 : ret != HCCL_SUCCESS,
119 : HCCL_ERROR("[CollCommExecutor][MultiRingAllReduce]stream[%u] record failed", ringIndex), ret);
120 :
121 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
122 0 : CHK_PRT_RET(
123 : ret != HCCL_SUCCESS,
124 : HCCL_ERROR("[CollCommExecutor][MultiRingAllReduce]stream[%u] record failed", ringIndex), ret);
125 : } else { // 主环
126 : tempAlg
127 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_REDUCE_RING, dispatcher_);
128 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_REDUCE_RING in COMM_LEVEL0", __func__);
129 0 : CHK_SMART_PTR_NULL(tempAlg);
130 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
131 :
132 0 : ret = tempAlg->Prepare(
133 : inputMem, outputMem, outputMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID,
134 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex]);
135 0 : CHK_PRT_RET(
136 : ret != HCCL_SUCCESS,
137 : HCCL_ERROR(
138 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) prepare failed, "
139 : "return[%d]",
140 : ringIndex, ret),
141 : ret);
142 :
143 0 : ret = tempAlg->RegisterProfiler(
144 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
145 0 : + level0RingCommInfo.localRank,
146 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
147 0 : CHK_PRT_RET(
148 : ret != HCCL_SUCCESS,
149 : HCCL_ERROR(
150 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) register Profiler "
151 : "failed,return[%d]",
152 : ringIndex, ret),
153 : ret);
154 :
155 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
156 0 : CHK_PRT_RET(
157 : ret != HCCL_SUCCESS,
158 : HCCL_ERROR(
159 : "[CollCommExecutor][MultiRingAllReduce]stream[%u], AllReduce(ring) run failed, "
160 : "return[%d]",
161 : ringIndex, ret),
162 : ret);
163 :
164 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
165 : /* 等待executor执行完毕 */
166 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
167 0 : CHK_PRT_RET(
168 : ret != HCCL_SUCCESS,
169 : HCCL_ERROR("[CollCommExecutor][MultiRingAllReduce]stream[%u] wait failed", ring), ret);
170 : }
171 : }
172 0 : }
173 : // 添加空task,保证执行时不乱序
174 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
175 0 : return HCCL_SUCCESS;
176 0 : }
177 :
178 0 : HcclResult CollCommExecutor::UpdateOffsetBasedOnStrideCount(
179 : const OpParam& param, std::vector<std::vector<Slice>>& multRingsUserMemSlice) const
180 : {
181 0 : u32 perDataSize = 0;
182 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
183 0 : for (u32 ringIndex = 0; ringIndex < multRingsUserMemSlice.size(); ringIndex++) {
184 0 : for (u32 sliceIndex = 0; sliceIndex < multRingsUserMemSlice[ringIndex].size(); sliceIndex++) {
185 0 : u64 selfRank = multRingsUserMemSlice[ringIndex][sliceIndex].offset / (param.DataDes.count * perDataSize);
186 0 : HCCL_DEBUG(
187 : "rank[%u], ringIndex[%u], sliceIndex[%u], slice.offset=[%llu], slice.size=[%llu], selfRank[%llu].",
188 : topoAttr_.userRank, ringIndex, sliceIndex, multRingsUserMemSlice[ringIndex][sliceIndex].offset,
189 : multRingsUserMemSlice[ringIndex][sliceIndex].size, selfRank);
190 0 : multRingsUserMemSlice[ringIndex][sliceIndex].offset
191 0 : = multRingsUserMemSlice[ringIndex][sliceIndex].offset
192 0 : + selfRank * ((param.DataDes.strideCount - param.DataDes.count) * perDataSize);
193 0 : HCCL_DEBUG(
194 : "rank[%u], ringIndex[%u], sliceIndex[%u], slice.offset=[%llu], slice.size=[%llu], selfRank[%llu] "
195 : "updated.",
196 : topoAttr_.userRank, ringIndex, sliceIndex, multRingsUserMemSlice[ringIndex][sliceIndex].offset,
197 : multRingsUserMemSlice[ringIndex][sliceIndex].size, selfRank);
198 : }
199 : }
200 0 : return HCCL_SUCCESS;
201 : }
202 :
203 0 : HcclResult CollCommExecutor::MultiRingAllGather(
204 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
205 : const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream, s32 profStage, const u64 baseOffset,
206 : const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>> multRingsUserMemSlice,
207 : const CommPlane leveIndex)
208 : {
209 0 : HcclResult ret = HCCL_SUCCESS;
210 0 : u32 ringNum = multRingsSliceZero.size();
211 0 : CHK_RET(CheckCommSize(leveIndex, ringNum));
212 :
213 0 : std::vector<std::vector<u32>> ringNics;
214 0 : CHK_RET(GetRingNics(tag, ringNics));
215 : // 拿到ring环映射关系
216 0 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(leveIndex, COMM_INDEX_0);
217 0 : auto nicList = topoAttr_.nicList;
218 0 : TopoType topoType = topoType_;
219 :
220 0 : if (leveIndex == COMM_LEVEL0_LOGICAL) {
221 0 : std::vector<u32> mockNicList;
222 0 : mockNicList.reserve(level0ZeroCommInfo.localRankSize);
223 0 : for (u32 rankIndex = 0; rankIndex < level0ZeroCommInfo.localRankSize; rankIndex++) {
224 0 : mockNicList.push_back(rankIndex);
225 : }
226 0 : nicList = mockNicList;
227 0 : u32 ARSRankSize = topoMatcher_->GetCommPlaneRanks(COMM_LEVEL0_LOGICAL)[0].size();
228 0 : bool ARSDoubleRing
229 0 : = ((ARSRankSize > FACTOR_TWO) && (ARSRankSize % FACTOR_TWO == 0) && topoAttr_.isARSDoubleRing);
230 :
231 0 : if (ARSDoubleRing) {
232 0 : topoType = TopoType::TOPO_TYPE_NP_DOUBLE_RING;
233 : } else {
234 0 : topoType = TopoType::TOPO_TYPE_NP_SINGLE_RING;
235 : }
236 0 : }
237 : std::vector<std::vector<u32>> multiRingsOrder
238 0 : = GetRingsOrderByTopoType(level0ZeroCommInfo.localRankSize, topoType, nicList);
239 :
240 : // 空拷贝用于后续操作附着
241 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
242 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
243 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
244 0 : CHK_PRT_RET(
245 : singleRingSliceZero.empty(),
246 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGather]"
247 : "singleRingSliceZero is empty"),
248 : HCCL_E_INTERNAL);
249 :
250 : // 910_93场景 生成userMemOut_上对应的slices
251 0 : std::vector<Slice> userMemOutputSlices;
252 0 : if (multRingsUserMemSlice.size() == 0) {
253 0 : CHK_RET(CalUserMemSlices(
254 : dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemOutputSlices));
255 : } else {
256 0 : userMemOutputSlices = multRingsUserMemSlice[ringIndex];
257 : }
258 0 : std::vector<u32> rankOrder;
259 0 : CHK_RET(GetRankOrder(multiRingsOrder, ringIndex, rankOrder));
260 :
261 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(leveIndex, ringIndex);
262 :
263 0 : u32 rankSize = level0RingCommInfo.localRankSize;
264 0 : u32 ringIndexOp = ringIndex;
265 :
266 : // 910_93场景 准备环中的从流
267 0 : std::vector<Stream> subStreamsInOneRing;
268 0 : std::vector<std::shared_ptr<LocalNotify>> mainSignalsInOneRing;
269 0 : std::vector<std::shared_ptr<LocalNotify>> subSignalsInOneRing;
270 0 : if (opInfo != nullptr) {
271 0 : CHK_RET(
272 : GetSubStreamInfoOnOneRing(ringIndex, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing));
273 : }
274 0 : if (ringIndex != (ringNum - 1)) { // 最后一个环是主stream,所以这里减1,符合条件的走从stream
275 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
276 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
277 0 : if (opInfo != nullptr) {
278 0 : algResResp_->threadManage[ringIndex]->Prepare(
279 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
280 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
281 0 : ringNics[ringIndex], tag, profStage, level0RingCommInfo, algResResp_->notifiesAux[ringIndex],
282 0 : algResResp_->notifiesMain[ringIndex], ringIndex, ExecutorType::ALLGATHER_RING_DIRECT, 0, opInfo,
283 : subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing, rankOrder, userMemOutputSlices);
284 : } else {
285 0 : algResResp_->threadManage[ringIndex]->Prepare(
286 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
287 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
288 0 : ringNics[ringIndex], tag, profStage, level0RingCommInfo, algResResp_->notifiesAux[ringIndex],
289 0 : algResResp_->notifiesMain[ringIndex], ringIndex, ExecutorType::ALLGATHER_RING);
290 : }
291 0 : algResResp_->threadManage[ringIndex]->NotifyStart(); // 给线程发信号启动处理
292 : } else {
293 0 : ret = LocalNotify::Wait(
294 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
295 0 : CHK_PRT_RET(
296 : ret != HCCL_SUCCESS,
297 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGather]stream[%u] wait failed", ringIndex), ret);
298 : // 如何判断是否环内是否有数据, 以ring的第一个rank的 size为判断依据
299 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
300 0 : if (opInfo != nullptr) {
301 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
302 0 : TemplateType::TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT, dispatcher_);
303 0 : HCCL_CONFIG_INFO(
304 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT in COMM_LEVEL0", __func__);
305 0 : CHK_SMART_PTR_NULL(tempAlg);
306 0 : CHK_RET(tempAlg->Prepare(
307 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing,
308 : mainSignalsInOneRing, subSignalsInOneRing, rankOrder, userMemOutputSlices));
309 : } else {
310 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
311 0 : TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
312 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL0", __func__);
313 0 : CHK_SMART_PTR_NULL(tempAlg);
314 : }
315 :
316 0 : ret = tempAlg->Prepare(
317 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
318 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
319 0 : ringNics[ringIndex]);
320 0 : CHK_PRT_RET(
321 : ret != HCCL_SUCCESS,
322 : HCCL_ERROR(
323 : "[CollCommExecutor][MultiRingAllGather]stream[%u],AllGather(ring) prepare "
324 : "failed,return[%d]",
325 : ringIndex, ret),
326 : ret);
327 0 : ret = tempAlg->RegisterProfiler(
328 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
329 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
330 0 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
331 0 : CHK_PRT_RET(
332 : ret != HCCL_SUCCESS,
333 : HCCL_ERROR(
334 : "[CollCommExecutor][MultiRingAllGather]stream[%u],AllGather(ring) register "
335 : "Profiler failed,return[%d]",
336 : ringIndex, ret),
337 : ret);
338 :
339 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
340 0 : CHK_PRT_RET(
341 : ret != HCCL_SUCCESS,
342 : HCCL_ERROR(
343 : "[CollCommExecutor][MultiRingAllGather]stream[%u],AllGather(ring) run failed, "
344 : "return[%d]",
345 : ringIndex, ret),
346 : ret);
347 :
348 0 : ret = LocalNotify::Post(
349 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
350 0 : CHK_PRT_RET(
351 : ret != HCCL_SUCCESS,
352 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGather]stream[%u] record failed", ringIndex), ret);
353 0 : }
354 :
355 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
356 0 : CHK_PRT_RET(
357 : ret != HCCL_SUCCESS,
358 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGather]stream[%u] record failed", ringIndex), ret);
359 : } else { // 主环
360 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
361 0 : if (opInfo != nullptr) {
362 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
363 0 : TemplateType::TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT, dispatcher_);
364 0 : HCCL_CONFIG_INFO(
365 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT in COMM_LEVEL0", __func__);
366 0 : CHK_SMART_PTR_NULL(tempAlg);
367 0 : CHK_RET(tempAlg->Prepare(
368 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
369 : subSignalsInOneRing, rankOrder, userMemOutputSlices));
370 : } else {
371 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
372 0 : TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
373 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL0", __func__);
374 0 : CHK_SMART_PTR_NULL(tempAlg);
375 : }
376 :
377 0 : ret = tempAlg->Prepare(
378 : outputMem, outputMem, inputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID,
379 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex]);
380 0 : CHK_PRT_RET(
381 : ret != HCCL_SUCCESS,
382 : HCCL_ERROR(
383 : "[CollCommExecutor][MultiRingAllGather]stream[%u],AllGather(ring) prepare failed,"
384 : "return[%d]",
385 : ringIndex, ret),
386 : ret);
387 :
388 0 : ret = tempAlg->RegisterProfiler(
389 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
390 0 : + level0RingCommInfo.localRank,
391 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
392 0 : CHK_PRT_RET(
393 : ret != HCCL_SUCCESS,
394 : HCCL_ERROR(
395 : "[CollCommExecutor][MultiRingAllGather]stream[%u], AllGather(ring) register Profiler "
396 : "failed,return[%d]",
397 : ringIndex, ret),
398 : ret);
399 :
400 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
401 0 : CHK_PRT_RET(
402 : ret != HCCL_SUCCESS,
403 : HCCL_ERROR(
404 : "[CollCommExecutor][MultiRingAllGather]stream[%u], AllGather(ring) run failed,"
405 : "return[%d]",
406 : ringIndex, ret),
407 : ret);
408 :
409 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
410 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
411 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
412 0 : algResResp_->threadManage[ring]->WaitDone(); // 单算子模式,等待线程处理完成信号
413 : }
414 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
415 0 : CHK_PRT_RET(
416 : ret != HCCL_SUCCESS,
417 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGather]stream[%u] wait failed", ring), ret);
418 : }
419 0 : }
420 0 : }
421 : // 添加空task,保证执行时不乱序
422 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
423 0 : return HCCL_SUCCESS;
424 0 : }
425 :
426 0 : HcclResult CollCommExecutor::MultiRingAllGatherConcurrent(
427 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
428 : const std::vector<std::pair<bool, std::vector<Slice>>> multRingsSliceZero, Stream stream, s32 profStage,
429 : const u64 baseOffset, const HcomCollOpInfo* opInfo,
430 : const std::vector<std::pair<bool, std::vector<Slice>>> multRingsUserMemSlice)
431 : {
432 0 : HcclResult ret = HCCL_SUCCESS;
433 0 : u32 ringNum = multRingsSliceZero.size(); // 环数, 当前为4环
434 :
435 0 : std::vector<std::vector<u32>> ringNics;
436 0 : CHK_RET(GetRingNics(tag, ringNics));
437 0 : auto halfRingSize = ringNum;
438 0 : if (ringNum > RDMA_PLANE_NUM_IN_NPRING_DOUBLE) {
439 0 : halfRingSize = ringNum / 2; // 2环
440 : }
441 : // 拿到ring环映射关系
442 0 : CHK_RET(CheckCommSize(COMM_LEVEL0_ANYPATH_SDMA, COMM_INDEX_1));
443 0 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(COMM_LEVEL0_ANYPATH_SDMA, COMM_INDEX_0);
444 0 : auto nicList = topoAttr_.nicList;
445 : std::vector<std::vector<u32>> multiRingsOrder
446 0 : = GetRingsOrderForAnyPath(level0ZeroCommInfo.localRankSize, topoType_, nicList);
447 :
448 : // 空拷贝用于后续操作附着
449 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
450 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
451 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex].second; // 取出sdma/rdma的数据块
452 0 : CHK_PRT_RET(
453 : singleRingSliceZero.empty(),
454 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGatherConcurrent]"
455 : "singleRingSliceZero is empty"),
456 : HCCL_E_INTERNAL);
457 :
458 : // 910_93场景 生成userMemOut_上对应的slices
459 0 : std::vector<Slice> userMemOutputSlices;
460 0 : if (multRingsUserMemSlice.size() == 0) {
461 0 : CHK_RET(CalUserMemSlices(
462 : dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemOutputSlices));
463 : } else {
464 0 : userMemOutputSlices = multRingsUserMemSlice[ringIndex].second;
465 : }
466 0 : std::vector<u32> rankOrder;
467 0 : u32 commIndex = ringIndex % halfRingSize;
468 0 : CHK_RET(GetRankOrder(multiRingsOrder, commIndex, rankOrder));
469 :
470 0 : SubCommInfo level0RingCommInfo = multRingsSliceZero[ringIndex].first ?
471 : GetSubCommInfo(COMM_LEVEL0_ANYPATH_SDMA, commIndex) :
472 0 : GetSubCommInfo(COMM_LEVEL0_ANYPATH_RDMA, commIndex);
473 :
474 0 : u32 rankSize = level0RingCommInfo.localRankSize;
475 0 : u32 ringIndexOp = ringIndex;
476 :
477 : // 910_93场景 准备环中的从流
478 0 : std::vector<Stream> subStreamsInOneRing;
479 0 : std::vector<std::shared_ptr<LocalNotify>> mainSignalsInOneRing;
480 0 : std::vector<std::shared_ptr<LocalNotify>> subSignalsInOneRing;
481 0 : if (opInfo != nullptr) {
482 0 : CHK_RET(
483 : GetSubStreamInfoOnOneRing(ringIndex, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing));
484 : }
485 0 : bool isSdma = multRingsSliceZero[ringIndex].first;
486 0 : if (ringIndex != (ringNum - 1)) { // 最后一个环是主stream,所以这里减1,符合条件的走从stream
487 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
488 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
489 0 : if (opInfo != nullptr) {
490 0 : ExecutorType type
491 0 : = isSdma ? ExecutorType::ALLGATHER_RING_DIRECT : ExecutorType::ALLGATHER_RING_DIRECT_RDMA;
492 0 : algResResp_->threadManage[ringIndex]->Prepare(
493 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
494 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
495 0 : ringNics[ringIndex % halfRingSize], tag, profStage, level0RingCommInfo,
496 0 : algResResp_->notifiesAux[ringIndex], algResResp_->notifiesMain[ringIndex], ringIndex, type, 0,
497 : opInfo, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing, rankOrder,
498 : userMemOutputSlices);
499 : } else {
500 0 : algResResp_->threadManage[ringIndex]->Prepare(
501 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
502 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
503 0 : ringNics[ringIndex % halfRingSize], tag, profStage, level0RingCommInfo,
504 0 : algResResp_->notifiesAux[ringIndex], algResResp_->notifiesMain[ringIndex], ringIndex,
505 : ExecutorType::ALLGATHER_RING);
506 : }
507 0 : algResResp_->threadManage[ringIndex]->NotifyStart(); // 给线程发信号启动处理
508 : } else {
509 0 : ret = LocalNotify::Wait(
510 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
511 0 : CHK_PRT_RET(
512 : ret != HCCL_SUCCESS,
513 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u] wait failed", ringIndex),
514 : ret);
515 : // 如何判断是否环内是否有数据, 以ring的第一个rank的 size为判断依据
516 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
517 0 : if (opInfo != nullptr) {
518 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
519 0 : TemplateType::TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT, dispatcher_);
520 0 : HCCL_CONFIG_INFO(
521 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT in COMM_LEVEL0", __func__);
522 0 : CHK_SMART_PTR_NULL(tempAlg);
523 0 : CHK_RET(tempAlg->Prepare(
524 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing,
525 : mainSignalsInOneRing, subSignalsInOneRing, rankOrder, userMemOutputSlices, isSdma));
526 : } else {
527 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
528 0 : TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
529 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL0", __func__);
530 0 : CHK_SMART_PTR_NULL(tempAlg);
531 : }
532 0 : ret = tempAlg->Prepare(
533 0 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
534 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
535 0 : ringNics[ringIndex % halfRingSize]);
536 0 : CHK_PRT_RET(
537 : ret != HCCL_SUCCESS,
538 : HCCL_ERROR(
539 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring) prepare "
540 : "failed,return[%d]",
541 : ringIndex, ret),
542 : ret);
543 0 : ret = tempAlg->RegisterProfiler(
544 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
545 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
546 0 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
547 0 : CHK_PRT_RET(
548 : ret != HCCL_SUCCESS,
549 : HCCL_ERROR(
550 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring) register "
551 : "Profiler failed,return[%d]",
552 : ringIndex, ret),
553 : ret);
554 :
555 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
556 0 : CHK_PRT_RET(
557 : ret != HCCL_SUCCESS,
558 : HCCL_ERROR(
559 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring)"
560 : " run failed,return[%d]",
561 : ringIndex, ret),
562 : ret);
563 :
564 0 : ret = LocalNotify::Post(
565 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
566 0 : CHK_PRT_RET(
567 : ret != HCCL_SUCCESS,
568 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u] record failed", ringIndex),
569 : ret);
570 0 : }
571 :
572 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
573 0 : CHK_PRT_RET(
574 : ret != HCCL_SUCCESS,
575 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u] record failed", ringIndex), ret);
576 : } else { // 主环
577 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
578 0 : if (opInfo != nullptr) {
579 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
580 0 : TemplateType::TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT, dispatcher_);
581 0 : HCCL_CONFIG_INFO(
582 : HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING_CONCURRENT_DIRECT in COMM_LEVEL0", __func__);
583 0 : CHK_SMART_PTR_NULL(tempAlg);
584 0 : CHK_RET(tempAlg->Prepare(
585 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
586 : subSignalsInOneRing, rankOrder, userMemOutputSlices, isSdma));
587 : } else {
588 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
589 0 : TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
590 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL0", __func__);
591 0 : CHK_SMART_PTR_NULL(tempAlg);
592 : }
593 0 : ret = tempAlg->Prepare(
594 : outputMem, outputMem, inputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID,
595 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex % halfRingSize]);
596 0 : CHK_PRT_RET(
597 : ret != HCCL_SUCCESS,
598 : HCCL_ERROR(
599 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring) prepare"
600 : " failed,return[%d]",
601 : ringIndex, ret),
602 : ret);
603 :
604 0 : ret = tempAlg->RegisterProfiler(
605 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
606 0 : + level0RingCommInfo.localRank,
607 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
608 0 : CHK_PRT_RET(
609 : ret != HCCL_SUCCESS,
610 : HCCL_ERROR(
611 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring) register "
612 : "Profiler failed, return[%d]",
613 : ringIndex, ret),
614 : ret);
615 :
616 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
617 0 : CHK_PRT_RET(
618 : ret != HCCL_SUCCESS,
619 : HCCL_ERROR(
620 : "[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u],AllGather(ring) run failed,"
621 : "return[%d]",
622 : ringIndex, ret),
623 : ret);
624 :
625 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
626 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
627 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
628 0 : algResResp_->threadManage[ring]->WaitDone(); // 单算子模式,等待线程处理完成信号
629 : }
630 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
631 0 : CHK_PRT_RET(
632 : ret != HCCL_SUCCESS,
633 : HCCL_ERROR("[CollCommExecutor][MultiRingAllGatherConcurrent]stream[%u] wait failed", ring), ret);
634 : }
635 0 : }
636 0 : }
637 : // 添加空task,保证执行时不乱序
638 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
639 0 : return HCCL_SUCCESS;
640 0 : }
641 :
642 0 : HcclResult CollCommExecutor::Level1AllGatherConcurrent(
643 : DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType, Stream stream, s32 profStage,
644 : std::vector<Slice>& level1DataSegsSlice, u32 syncTrans)
645 : {
646 0 : std::vector<std::pair<bool, std::vector<Slice>>> level1MultSlice;
647 0 : std::vector<Slice> level1DataSegsSliceSdma;
648 0 : std::vector<Slice> level1DataSegsSliceRdma;
649 0 : bool isAnyPathCommLevel0 = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING
650 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) ?
651 : true :
652 : false;
653 0 : CommPlane commplane = (isAnyPathCommLevel0) ? COMM_LEVEL0_ANYPATH_SDMA : COMM_LEVEL0;
654 0 : CHK_RET(CheckCommSize(commplane, COMM_INDEX_1));
655 0 : SubCommInfo level0CommInfo = GetSubCommInfo(commplane, COMM_INDEX_0);
656 0 : u32 level0ServerIndex = level0CommInfo.localRank;
657 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1_ANYPATH_SDMA, level0ServerIndex);
658 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
659 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
660 0 : HcclResult ret = HCCL_SUCCESS;
661 0 : level1MultSlice.resize(RDMA_PLANE_NUM_IN_NPRING_DOUBLE);
662 :
663 0 : for (u32 i = 0; i < level1CommInfo.localRankSize; i++) {
664 0 : Slice sdmaSlice;
665 0 : Slice rdmaSlice;
666 : u64 sdmaSliceSize
667 0 : = ((level1DataSegsSlice[i].size <= HCCL_MIN_SLICE_ALIGN_910_93) || (syncTrans == MAX_SPLIT_VALUE)) ?
668 0 : level1DataSegsSlice[i].size :
669 0 : ((syncTrans * level1DataSegsSlice[i].size / MAX_SPLIT_VALUE) / HCCL_MIN_SLICE_ALIGN_910_93)
670 0 : * HCCL_MIN_SLICE_ALIGN_910_93;
671 0 : sdmaSlice.size = sdmaSliceSize;
672 0 : sdmaSlice.offset = level1DataSegsSlice[i].offset;
673 0 : rdmaSlice.size = level1DataSegsSlice[i].size - sdmaSliceSize;
674 0 : rdmaSlice.offset = level1DataSegsSlice[i].offset + sdmaSliceSize;
675 0 : level1DataSegsSliceSdma.push_back(sdmaSlice);
676 0 : level1DataSegsSliceRdma.push_back(rdmaSlice);
677 0 : HCCL_DEBUG(
678 : "Level1 index:[%u], Original [offset %llu, size %llu], sdma [offset %llu, size %llu], "
679 : "rdma [offset %llu, size %llu]",
680 : i, level1DataSegsSlice[i].offset, level1DataSegsSlice[i].size, sdmaSlice.offset, sdmaSlice.size,
681 : rdmaSlice.offset, rdmaSlice.size);
682 : }
683 0 : level1MultSlice[0] = std::make_pair(true, level1DataSegsSliceSdma);
684 0 : level1MultSlice[1] = std::make_pair(false, level1DataSegsSliceRdma);
685 :
686 0 : u32 commPlaneNum = level1MultSlice.size();
687 0 : for (u32 planeIndex = 0; planeIndex < commPlaneNum; planeIndex++) {
688 0 : std::vector<Slice>& singleSlice = level1MultSlice[planeIndex].second;
689 0 : SubCommInfo level1RdmaCommInfo = GetSubCommInfo(COMM_LEVEL1_ANYPATH_RDMA, level0ServerIndex);
690 0 : SubCommInfo level1TempCommInfo = level1MultSlice[planeIndex].first ? level1CommInfo : level1RdmaCommInfo;
691 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
692 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
693 : level1TempAlg
694 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
695 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
696 : } else {
697 : level1TempAlg
698 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
699 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
700 : }
701 0 : CHK_SMART_PTR_NULL(level1TempAlg);
702 :
703 0 : if (planeIndex != (commPlaneNum - 1)) {
704 0 : ret = LocalNotify::Wait(
705 0 : algResResp_->slaveStreams[planeIndex], dispatcher_, algResResp_->notifiesAux[planeIndex], profStage);
706 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("stream[%u] wait failed", planeIndex), ret);
707 :
708 0 : CHK_RET(level1TempAlg->Prepare(
709 : outputMem, outputMem, inputMem, count, dataType, algResResp_->slaveStreams[planeIndex],
710 : HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, singleSlice, 0));
711 :
712 0 : CHK_RET(level1TempAlg->RegisterProfiler(
713 : (level1TempCommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank,
714 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[planeIndex]));
715 :
716 0 : CHK_RET(RunTemplate(level1TempAlg, level1TempCommInfo));
717 0 : ret = LocalNotify::Post(
718 0 : algResResp_->slaveStreams[planeIndex], dispatcher_, algResResp_->notifiesMain[planeIndex], profStage);
719 0 : CHK_PRT_RET(
720 : ret != HCCL_SUCCESS, HCCL_ERROR("[collAllGather]level1 stream[%u] record failed", planeIndex), ret);
721 : // 主环record启动从环
722 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[planeIndex], profStage);
723 0 : CHK_PRT_RET(
724 : ret != HCCL_SUCCESS, HCCL_ERROR("[collAllGather]level1 stream[%u] record failed", planeIndex), ret);
725 : } else {
726 0 : CHK_RET(level1TempAlg->Prepare(
727 : outputMem, outputMem, inputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID,
728 : singleSlice, 0));
729 0 : CHK_RET(level1TempAlg->RegisterProfiler(
730 : (level1TempCommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank,
731 : profStage, HCCL_EXEC_STEP_NOT_SET, stream));
732 :
733 0 : CHK_RET(RunTemplate(level1TempAlg, level1TempCommInfo));
734 0 : for (u32 ring = 0; ring < (commPlaneNum - 1); ring++) {
735 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
736 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("stream[%u] wait failed", ring), ret);
737 : }
738 : }
739 0 : }
740 0 : HCCL_INFO("Level1AllGatherConcurrent run success");
741 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
742 0 : return HCCL_SUCCESS;
743 0 : }
744 :
745 0 : u32 CollCommExecutor::CalcOptimalIntraRingsize(u64 count, HcclDataType dataType, HcclCMDType opType)
746 : {
747 0 : if (!topoMatcher_->GetARSFlag())
748 0 : return 0;
749 :
750 0 : u32 level0RankSize = topoMatcher_->GetCommPlaneRanks(COMM_LEVEL0)[0].size();
751 0 : u32 rankSizeInSuperPod = topoMatcher_->GetCommPlaneRanks(COMM_ARS)[0].size();
752 0 : u32 perDataSize = 0;
753 0 : CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
754 : // 不支持 ARS 或环内卡数不是 2 的倍数
755 0 : u32 level0RingSize = 1;
756 0 : if (!topoAttr_.isARSDoubleRing || (level0RankSize % FACTOR_TWO != 0)) {
757 0 : HCCL_INFO(
758 : "not Support ARS doubleRing, level0RingSize:[%u], level0RankSize[%u].", level0RingSize, level0RankSize);
759 0 : return level0RingSize;
760 : }
761 : // --- 1. 带宽 & 基本参数 ---
762 : float bwHCCS, bwHBM, bwSIO;
763 0 : constexpr u32 level0 = 0;
764 0 : constexpr u32 level2 = 2;
765 0 : constexpr u32 level3 = 3;
766 0 : CHK_RET(GetBandWidthPerNPU(level0, topoAttr_.userRankSize, topoAttr_.deviceNumPerAggregation, bwHCCS));
767 0 : CHK_RET(GetBandWidthPerNPU(level2, topoAttr_.userRankSize, topoAttr_.deviceNumPerAggregation, bwHBM));
768 0 : CHK_RET(GetBandWidthPerNPU(level3, topoAttr_.userRankSize, topoAttr_.deviceNumPerAggregation, bwSIO));
769 0 : float latency = BASE_COMM_LATENCY / MULTIPLIER_MS2US; // ms
770 : // --- 2. 数据总量 (GB) ---
771 0 : float baseSizeGB = static_cast<double>(count) * perDataSize / (1024 * 1024 * 1024);
772 0 : float totalSize = baseSizeGB;
773 0 : HCCL_INFO(
774 : "CalcOptimalIntraRingsize: count[%u], totalSize:[%lf]GB, perDataSize[%u].", count, totalSize, perDataSize);
775 0 : if (opType == HcclCMDType::HCCL_CMD_ALLGATHER || opType == HcclCMDType::HCCL_CMD_REDUCE_SCATTER) {
776 0 : totalSize *= rankSizeInSuperPod;
777 : }
778 : // --- 3. 枚举可能的环大小 ---
779 0 : std::vector<u32> factors;
780 0 : for (u32 i = 1; i <= rankSizeInSuperPod / i; ++i) {
781 0 : if (rankSizeInSuperPod % i == 0) {
782 0 : factors.push_back(i);
783 0 : if (i != rankSizeInSuperPod / i) {
784 0 : factors.push_back(rankSizeInSuperPod / i);
785 : }
786 : }
787 : }
788 0 : std::sort(factors.begin(), factors.end());
789 : // --- 4. 计算最优带宽 ---
790 0 : double maxBwARS = 0.0;
791 0 : for (u32 N1 : factors) {
792 0 : u32 N2 = rankSizeInSuperPod / N1;
793 : // 静态时延 (ms)
794 0 : double interStep = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) ? (N2 - 1) : log2(N2);
795 0 : double latencyStep = (interStep + (N1 - 1)) * latency;
796 : // 传输时延 (ms)
797 : double latencyIntra;
798 0 : if ((N1 % FACTOR_TWO == 0) && (N1 > FACTOR_TWO)) {
799 0 : latencyIntra = (N1 - 1) * totalSize * MULTIPLIER_S2MS / N1 / bwHCCS / FACTOR_TWO;
800 0 : } else if (N1 == FACTOR_TWO) {
801 0 : latencyIntra = totalSize * MULTIPLIER_S2MS / FACTOR_TWO / bwSIO;
802 : } else {
803 0 : latencyIntra = (N1 - 1) * totalSize * MULTIPLIER_S2MS / N1 / bwHCCS;
804 : }
805 0 : double latencyInter = (N2 - 1) * totalSize * MULTIPLIER_S2MS / N1 / N2 / bwHCCS;
806 : // HBM 拷贝时延 (ms)
807 0 : double latencyCopy = totalSize * MULTIPLIER_S2MS / bwHBM;
808 0 : u8 mul = (opType == HcclCMDType::HCCL_CMD_ALLREDUCE) ? FACTOR_TWO : 1;
809 0 : double timeCost = mul * (latencyStep + latencyIntra + latencyInter) + latencyCopy;
810 0 : double bwARS = totalSize / timeCost; // GB/ms
811 0 : if (bwARS > maxBwARS) {
812 0 : maxBwARS = bwARS;
813 0 : level0RingSize = N1;
814 : }
815 : }
816 0 : HCCL_INFO("level0RingSize:[%u], totalSize:[%lf]GB, level0RankSize[%u].", level0RingSize, totalSize, level0RankSize);
817 0 : return level0RingSize;
818 0 : }
819 :
820 67 : HcclResult CollCommExecutor::CollectMultiRingsUserMemSlices(
821 : u32 ringNum, const HcclDataType dataType, const HcomCollOpInfo* opInfo,
822 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const std::vector<std::vector<u32>>& multiRingsOrder,
823 : const std::vector<std::vector<Slice>>& multRingsUserMemSlice,
824 : std::vector<std::vector<Slice>>& userMemSlicesOfMultiRings)
825 : {
826 67 : CHK_PTR_NULL(opInfo);
827 66 : CHK_PRT_RET(
828 : 0 < opInfo->strideCount && opInfo->strideCount < opInfo->count,
829 : HCCL_ERROR(
830 : "[CollCommExecutor][CollectMultiRingsUserMemSlices]strideCount[%llu] is smaller than opCount[%llu]",
831 : opInfo->strideCount, opInfo->count),
832 : HCCL_E_PARA);
833 198 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
834 132 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
835 132 : CHK_PRT_RET(
836 : singleRingSliceZero.empty(),
837 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]singleRingSliceZero is empty"), HCCL_E_INTERNAL);
838 132 : std::vector<Slice> userMemSlices;
839 132 : HCCL_DEBUG(
840 : "[CollCommExecutor][CollectMultiRingsUserMemSlices]multRingsUserMemSlice.size()[%zu], strideCount[%llu], "
841 : "opCount[%llu]",
842 : multRingsUserMemSlice.size(), opInfo->strideCount, opInfo->count);
843 132 : if (multRingsUserMemSlice.size() == 0) {
844 64 : CHK_RET(CalUserMemSlices(dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemSlices));
845 : } else {
846 68 : userMemSlices = multRingsUserMemSlice[ringIndex];
847 : }
848 132 : userMemSlicesOfMultiRings.push_back(userMemSlices);
849 132 : }
850 66 : return HCCL_SUCCESS;
851 : }
852 :
853 66 : HcclResult CollCommExecutor::CollectMultiRingsRankOrder(
854 : u32 ringNum, const std::vector<std::vector<u32>>& multiRingsOrder, std::vector<std::vector<u32>>& rankOrders)
855 : {
856 198 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
857 132 : std::vector<u32> rankOrder;
858 132 : CHK_RET(GetRankOrder(multiRingsOrder, ringIndex, rankOrder));
859 132 : rankOrders.push_back(rankOrder);
860 132 : }
861 66 : return HCCL_SUCCESS;
862 : }
863 :
864 3 : HcclResult CollCommExecutor::MultiRingReduceScatter(
865 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
866 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>> multRingsSliceZero, Stream stream,
867 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
868 : const std::vector<std::vector<Slice>> multRingsUserMemSlice, const CommPlane levelIndex)
869 : {
870 3 : HCCL_INFO("[MultiRingReduceScatter] MultiRingReduceScatter starts");
871 3 : HcclResult ret = HCCL_SUCCESS;
872 3 : u32 ringNum = multRingsSliceZero.size();
873 3 : CHK_RET(CheckCommSize(levelIndex, ringNum));
874 :
875 3 : std::vector<std::vector<u32>> ringNics;
876 3 : CHK_RET(GetRingNics(tag, ringNics));
877 : // 拿到ring环映射关系
878 3 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(levelIndex, COMM_INDEX_0);
879 3 : auto nicList = topoAttr_.nicList;
880 :
881 3 : TopoType topoType = topoType_;
882 :
883 3 : if (levelIndex == COMM_LEVEL0_LOGICAL) {
884 0 : std::vector<u32> mockNicList;
885 0 : mockNicList.reserve(level0ZeroCommInfo.localRankSize);
886 0 : for (u32 rankIndex = 0; rankIndex < level0ZeroCommInfo.localRankSize; rankIndex++) {
887 0 : mockNicList.push_back(rankIndex);
888 : }
889 0 : nicList = mockNicList;
890 0 : u32 ARSRankSize = topoMatcher_->GetCommPlaneRanks(COMM_LEVEL0_LOGICAL)[0].size();
891 0 : bool ARSDoubleRing
892 0 : = ((ARSRankSize > FACTOR_TWO) && (ARSRankSize % FACTOR_TWO == 0) && topoAttr_.isARSDoubleRing);
893 0 : if (ARSDoubleRing) {
894 0 : topoType = TopoType::TOPO_TYPE_NP_DOUBLE_RING;
895 : } else {
896 0 : topoType = TopoType::TOPO_TYPE_NP_SINGLE_RING;
897 : }
898 0 : }
899 : std::vector<std::vector<u32>> multiRingsOrder
900 3 : = GetRingsOrderByTopoType(level0ZeroCommInfo.localRankSize, topoType, nicList);
901 :
902 3 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
903 :
904 : // 空拷贝用于后续操作附着
905 3 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
906 6 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
907 3 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
908 3 : CHK_PRT_RET(
909 : singleRingSliceZero.empty(),
910 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]singleRingSliceZero is empty"), HCCL_E_INTERNAL);
911 :
912 : // 生成userMemIn_上对应的slices
913 3 : std::vector<Slice> userMemInputSlices;
914 3 : if (multRingsUserMemSlice.size() == 0) {
915 3 : CHK_RET(CalUserMemSlices(
916 : dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemInputSlices));
917 : } else {
918 0 : userMemInputSlices = multRingsUserMemSlice[ringIndex];
919 : }
920 :
921 3 : std::vector<u32> rankOrder;
922 3 : CHK_RET(GetRankOrder(multiRingsOrder, ringIndex, rankOrder));
923 :
924 3 : SubCommInfo level0RingCommInfo = GetSubCommInfo(levelIndex, ringIndex);
925 3 : u32 rankSize = level0RingCommInfo.localRankSize;
926 3 : u32 ringIndexOp = ringIndex;
927 :
928 3 : std::vector<Stream> subStreamsInOneRing;
929 3 : std::vector<std::shared_ptr<LocalNotify>> mainSignalsInOneRing;
930 3 : std::vector<std::shared_ptr<LocalNotify>> subSignalsInOneRing;
931 3 : if (opInfo != nullptr) {
932 0 : CHK_RET(
933 : GetSubStreamInfoOnOneRing(ringIndex, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing));
934 : }
935 3 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
936 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
937 0 : ret = StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
938 0 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr());
939 0 : CHK_PRT_RET(
940 : ret != HCCL_SUCCESS,
941 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]active stream[%u], failed", ringIndex), ret);
942 : }
943 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
944 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
945 : /* 更新线程参数 */
946 0 : if (opInfo != nullptr) {
947 0 : algResResp_->threadManage[ringIndex]->Prepare(
948 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
949 0 : reductionOp, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset, ringNics[ringIndex], tag,
950 0 : profStage, level0RingCommInfo, algResResp_->notifiesAux[ringIndex],
951 0 : algResResp_->notifiesMain[ringIndex], ringIndex, ExecutorType::REDUCE_SCATTER_RING_DIRECT,
952 : reduceAttr, opInfo, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing, rankOrder,
953 : userMemInputSlices);
954 : } else {
955 0 : algResResp_->threadManage[ringIndex]->Prepare(
956 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
957 0 : reductionOp, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset, ringNics[ringIndex], tag,
958 0 : profStage, level0RingCommInfo, algResResp_->notifiesAux[ringIndex],
959 0 : algResResp_->notifiesMain[ringIndex], ringIndex, ExecutorType::REDUCE_SCATTER_RING, reduceAttr);
960 : }
961 :
962 0 : algResResp_->threadManage[ringIndex]->NotifyStart(); // 给线程发通知启动线程执行
963 : } else {
964 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
965 0 : if (opInfo != nullptr) {
966 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
967 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, dispatcher_);
968 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
969 0 : CHK_SMART_PTR_NULL(tempAlg);
970 0 : CHK_RET(tempAlg->Prepare(
971 : reduceAttr, opInfo, topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
972 : subSignalsInOneRing, rankOrder, userMemInputSlices));
973 : } else {
974 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
975 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
976 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL0", __func__);
977 0 : CHK_SMART_PTR_NULL(tempAlg);
978 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
979 : }
980 0 : CHK_SMART_PTR_NULL(tempAlg);
981 :
982 0 : ret = LocalNotify::Wait(
983 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
984 0 : CHK_PRT_RET(
985 : ret != HCCL_SUCCESS,
986 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]stream[%u] wait failed", ringIndex), ret);
987 0 : ret = tempAlg->Prepare(
988 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex], reductionOp,
989 0 : LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset, ringNics[ringIndex]);
990 0 : CHK_PRT_RET(
991 : ret != HCCL_SUCCESS,
992 : HCCL_ERROR(
993 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) "
994 : "prepare failed,return[%d]",
995 : ringIndex, ret),
996 : ret);
997 0 : ret = tempAlg->RegisterProfiler(
998 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
999 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
1000 0 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
1001 0 : CHK_PRT_RET(
1002 : ret != HCCL_SUCCESS,
1003 : HCCL_ERROR(
1004 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) "
1005 : "register Profiler failed,return[%d]",
1006 : ringIndex, ret),
1007 : ret);
1008 :
1009 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1010 0 : CHK_PRT_RET(
1011 : ret != HCCL_SUCCESS,
1012 : HCCL_ERROR(
1013 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) run "
1014 : "failed,return[%d]",
1015 : ringIndex, ret),
1016 : ret);
1017 :
1018 0 : ret = LocalNotify::Post(
1019 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
1020 0 : CHK_PRT_RET(
1021 : ret != HCCL_SUCCESS,
1022 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]stream[%u] record failed", ringIndex), ret);
1023 0 : }
1024 : /* 主环record启动从环 */
1025 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
1026 0 : CHK_PRT_RET(
1027 : ret != HCCL_SUCCESS,
1028 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]stream[%u] record failed", ringIndex), ret);
1029 : } else { // 主环 最后一个环
1030 3 : std::unique_ptr<AlgTemplateBase> tempAlg;
1031 3 : if (opInfo != nullptr) {
1032 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1033 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, dispatcher_);
1034 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
1035 0 : CHK_SMART_PTR_NULL(tempAlg);
1036 0 : CHK_RET(tempAlg->Prepare(
1037 : reduceAttr, opInfo, topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
1038 : subSignalsInOneRing, rankOrder, userMemInputSlices));
1039 : } else {
1040 6 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1041 3 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
1042 3 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL0", __func__);
1043 3 : CHK_SMART_PTR_NULL(tempAlg);
1044 3 : CHK_RET(tempAlg->Prepare(reduceAttr));
1045 : }
1046 3 : CHK_SMART_PTR_NULL(tempAlg);
1047 6 : ret = tempAlg->Prepare(
1048 : inputMem, inputMem, outputMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID,
1049 3 : singleRingSliceZero, baseOffset, ringNics[ringIndex]);
1050 3 : CHK_PRT_RET(
1051 : ret != HCCL_SUCCESS,
1052 : HCCL_ERROR(
1053 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) prepare "
1054 : "failed,return[%d]",
1055 : ringIndex, ret),
1056 : ret);
1057 :
1058 3 : ret = tempAlg->RegisterProfiler(
1059 3 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1060 3 : + level0RingCommInfo.localRank,
1061 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
1062 3 : CHK_PRT_RET(
1063 : ret != HCCL_SUCCESS,
1064 : HCCL_ERROR(
1065 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) register "
1066 : "Profiler failed,return[%d]",
1067 : ringIndex, ret),
1068 : ret);
1069 :
1070 3 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1071 3 : CHK_PRT_RET(
1072 : ret != HCCL_SUCCESS,
1073 : HCCL_ERROR(
1074 : "[CollCommExecutor][MultiRingReduceScatter]stream[%u],ReduceScatter(ring) run "
1075 : "failed,return[%d]",
1076 : ringIndex, ret),
1077 : ret);
1078 3 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
1079 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
1080 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
1081 0 : algResResp_->threadManage[ring]->WaitDone();
1082 : }
1083 : /* 等待executor执行完毕 */
1084 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
1085 :
1086 0 : CHK_PRT_RET(
1087 : ret != HCCL_SUCCESS,
1088 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatter]stream[%u] wait failed", ring), ret);
1089 : }
1090 3 : }
1091 3 : }
1092 :
1093 3 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1094 3 : return HCCL_SUCCESS;
1095 3 : }
1096 :
1097 0 : HcclResult CollCommExecutor::MultiRingGather(
1098 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
1099 : const std::vector<std::vector<Slice>> multRingsSliceZero, HcclReduceOp op, u32 root, Stream stream, s32 profStage)
1100 : {
1101 0 : u32 ringNum = multRingsSliceZero.size();
1102 0 : std::vector<std::vector<u32>> ringNics;
1103 0 : CHK_RET(GetRingNics(tag, ringNics));
1104 :
1105 : HcclResult ret;
1106 :
1107 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
1108 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
1109 0 : CHK_PRT_RET(
1110 : singleRingSliceZero.empty(), HCCL_ERROR("[CommonOperator][MultiRingGather]singleRingSliceZero is empty"),
1111 : HCCL_E_INTERNAL);
1112 :
1113 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, ringIndex);
1114 0 : u32 rankSize = level0RingCommInfo.localRankSize;
1115 0 : u32 rootRank = 0;
1116 0 : ret = GetRankByUserRank(COMM_LEVEL0, ringIndex, root, rootRank);
1117 0 : CHK_PRT_RET(
1118 : ret == HCCL_E_PARA,
1119 : HCCL_ERROR("[CommonOperator][MultiRingGather]invalid root rank[%u] to get user rank", root), ret);
1120 :
1121 0 : std::unique_ptr<AlgTemplateBase> tempAlg = nullptr;
1122 0 : EXCEPTION_CATCH(
1123 : (tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_GATHER_RING, dispatcher_)),
1124 : return HCCL_E_PTR);
1125 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_GATHER_RING in COMM_LEVEL0", __func__);
1126 0 : CHK_SMART_PTR_NULL(tempAlg);
1127 :
1128 0 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
1129 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
1130 0 : CHK_RET(StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
1131 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr()));
1132 : }
1133 0 : ret = LocalNotify::Wait(
1134 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
1135 0 : CHK_PRT_RET(
1136 : ret != HCCL_SUCCESS,
1137 : HCCL_ERROR("[CommonOperator][MultiRingGather]in stream[%u] wait failed", ringIndex), ret);
1138 0 : if (singleRingSliceZero[0].size != 0) {
1139 0 : ret = tempAlg->Prepare(
1140 0 : inputMem, outputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex], op, rootRank,
1141 0 : singleRingSliceZero, 0, ringNics[ringIndex]);
1142 0 : CHK_PRT_RET(
1143 : ret != HCCL_SUCCESS,
1144 : HCCL_ERROR(
1145 : "[CommonOperator][MultiRingGather]stream[%u],gather(ring) prepare failed, "
1146 : "return[%d]",
1147 : ringIndex, ret),
1148 : ret);
1149 :
1150 0 : ret = tempAlg->RegisterProfiler(
1151 0 : level0RingCommInfo.localRank, profStage, HCCL_EXEC_STEP_NOT_SET,
1152 0 : algResResp_->slaveStreams[ringIndex]);
1153 0 : CHK_PRT_RET(
1154 : ret != HCCL_SUCCESS,
1155 : HCCL_ERROR(
1156 : "[CommonOperator][MultiRingGather]stream[%u], gather(ring) register profiler "
1157 : "failed,return[%d]",
1158 : ringIndex, ret),
1159 : ret);
1160 :
1161 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1162 0 : CHK_PRT_RET(
1163 : ret != HCCL_SUCCESS,
1164 : HCCL_ERROR(
1165 : "[CommonOperator][MultiRingGather]stream[%u],gather(ring) run failed,return[%d]", ringIndex,
1166 : ret),
1167 : ret);
1168 : }
1169 0 : ret = LocalNotify::Post(
1170 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
1171 :
1172 0 : CHK_PRT_RET(
1173 : ret != HCCL_SUCCESS, HCCL_ERROR("[CommonOperator][MultiRingGather]stream[%u] record failed", ringIndex),
1174 : ret);
1175 :
1176 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
1177 0 : CHK_PRT_RET(
1178 : ret != HCCL_SUCCESS, HCCL_ERROR("[CommonOperator][MultiRingGather]stream[%u] record failed", ringIndex),
1179 : ret);
1180 : } else { // 主环
1181 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_GATHER_RING, dispatcher_);
1182 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_GATHER_RING in COMM_LEVEL0", __func__);
1183 0 : CHK_SMART_PTR_NULL(tempAlg);
1184 :
1185 0 : ret = tempAlg->Prepare(
1186 : inputMem, outputMem, outputMem, count, dataType, stream, op, rootRank, singleRingSliceZero, 0,
1187 0 : ringNics[ringIndex]);
1188 0 : CHK_PRT_RET(
1189 : ret != HCCL_SUCCESS,
1190 : HCCL_ERROR(
1191 : "[CommonOperator][MultiRingGather]stream[%u],gather(ring) prepare failed, "
1192 : "return[%d]",
1193 : ringIndex, ret),
1194 : ret);
1195 :
1196 0 : ret = tempAlg->RegisterProfiler(
1197 0 : ((ringIndex + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1198 0 : + level0RingCommInfo.localRank,
1199 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
1200 0 : CHK_PRT_RET(
1201 : ret != HCCL_SUCCESS,
1202 : HCCL_ERROR(
1203 : "[CommonOperator][MultiRingGather]stream[%u], gather(ring) register "
1204 : "profiler failed,return[%d]",
1205 : ringIndex, ret),
1206 : ret);
1207 :
1208 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1209 0 : CHK_PRT_RET(
1210 : ret != HCCL_SUCCESS,
1211 : HCCL_ERROR(
1212 : "[CommonOperator][MultiRingGather]stream[%u],gather(ring) run failed, "
1213 : "return[%d]",
1214 : ringIndex, ret),
1215 : ret);
1216 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
1217 : /* 等待executor执行完毕 , 当前环没有分配数据,跳过此环处理,继续下一个环 */
1218 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
1219 0 : CHK_PRT_RET(
1220 : ret != HCCL_SUCCESS, HCCL_ERROR("[CommonOperator][MultiRingGather]stream[%u] wait failed", ring),
1221 : ret);
1222 : }
1223 : }
1224 0 : }
1225 :
1226 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1227 0 : return HCCL_SUCCESS;
1228 0 : }
1229 :
1230 0 : HcclResult CollCommExecutor::MultiRingReduceScatterConcurrent(
1231 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
1232 : const HcclReduceOp reductionOp, const std::vector<std::pair<bool, std::vector<Slice>>> multRingsSliceZero,
1233 : Stream stream, s32 profStage, const u64 baseOffset, const HcomCollOpInfo* opInfo,
1234 : const std::vector<std::pair<bool, std::vector<Slice>>> multRingsUserMemSlice)
1235 : {
1236 0 : HcclResult ret = HCCL_SUCCESS;
1237 0 : u32 ringNum = multRingsSliceZero.size();
1238 :
1239 0 : std::vector<std::vector<u32>> ringNics;
1240 0 : CHK_RET(GetRingNics(tag, ringNics));
1241 0 : u32 halfRingSize = ringNum;
1242 0 : u32 DoubleRing = 2;
1243 0 : if (ringNum > RDMA_PLANE_NUM_IN_NPRING_DOUBLE) {
1244 0 : halfRingSize = ringNum / DoubleRing;
1245 : }
1246 :
1247 : // 拿到ring环映射关系
1248 0 : SubCommInfo level0ZeroCommInfo = GetSubCommInfo(COMM_LEVEL0_ANYPATH_SDMA, COMM_INDEX_0);
1249 0 : auto nicList = topoAttr_.nicList;
1250 : std::vector<std::vector<u32>> multiRingsOrder
1251 0 : = GetRingsOrderForAnyPath(level0ZeroCommInfo.localRankSize, topoType_, nicList);
1252 :
1253 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
1254 :
1255 : // 空拷贝用于后续操作附着
1256 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1257 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
1258 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex].second;
1259 0 : CHK_PRT_RET(
1260 : singleRingSliceZero.empty(),
1261 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatterConcurrent]singleRingSliceZero is empty"),
1262 : HCCL_E_INTERNAL);
1263 :
1264 : // 生成userMemIn_上对应的slices
1265 0 : std::vector<Slice> userMemInputSlices;
1266 0 : u32 commIndex = ringIndex % halfRingSize;
1267 0 : if (multRingsUserMemSlice.size() == 0) {
1268 0 : CHK_RET(CalUserMemSlices(
1269 : dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemInputSlices));
1270 : } else {
1271 0 : userMemInputSlices = multRingsUserMemSlice[ringIndex].second;
1272 : }
1273 0 : std::vector<u32> rankOrder;
1274 0 : CHK_RET(GetRankOrder(multiRingsOrder, commIndex, rankOrder));
1275 :
1276 0 : SubCommInfo level0RingCommInfo = multRingsSliceZero[ringIndex].first ?
1277 : GetSubCommInfo(COMM_LEVEL0_ANYPATH_SDMA, commIndex) :
1278 0 : GetSubCommInfo(COMM_LEVEL0_ANYPATH_RDMA, commIndex);
1279 0 : u32 rankSize = level0RingCommInfo.localRankSize;
1280 0 : u32 ringIndexOp = ringIndex;
1281 :
1282 0 : std::vector<Stream> subStreamsInOneRing;
1283 0 : std::vector<std::shared_ptr<LocalNotify>> mainSignalsInOneRing;
1284 0 : std::vector<std::shared_ptr<LocalNotify>> subSignalsInOneRing;
1285 0 : if (opInfo != nullptr) {
1286 0 : CHK_RET(
1287 : GetSubStreamInfoOnOneRing(ringIndex, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing));
1288 : }
1289 0 : bool isSdma = multRingsSliceZero[ringIndex].first;
1290 0 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
1291 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
1292 0 : ret = StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
1293 0 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr());
1294 0 : CHK_PRT_RET(
1295 : ret != HCCL_SUCCESS,
1296 : HCCL_ERROR(
1297 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]active stream[%u], failed", ringIndex),
1298 : ret);
1299 : }
1300 :
1301 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
1302 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
1303 : /* 更新线程参数 */
1304 0 : if (opInfo != nullptr) {
1305 0 : ExecutorType type = isSdma ? ExecutorType::REDUCE_SCATTER_RING_DIRECT :
1306 : ExecutorType::REDUCE_SCATTER_RING_DIRECT_RDMA;
1307 0 : algResResp_->threadManage[ringIndex]->Prepare(
1308 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
1309 : reductionOp, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
1310 0 : ringNics[ringIndex % halfRingSize], tag, profStage, level0RingCommInfo,
1311 0 : algResResp_->notifiesAux[ringIndex], algResResp_->notifiesMain[ringIndex], ringIndex, type,
1312 : reduceAttr, opInfo, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing, rankOrder,
1313 : userMemInputSlices);
1314 : } else {
1315 0 : algResResp_->threadManage[ringIndex]->Prepare(
1316 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
1317 : reductionOp, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
1318 0 : ringNics[ringIndex % halfRingSize], tag, profStage, level0RingCommInfo,
1319 0 : algResResp_->notifiesAux[ringIndex], algResResp_->notifiesMain[ringIndex], ringIndex,
1320 : ExecutorType::REDUCE_SCATTER_RING, reduceAttr);
1321 : }
1322 :
1323 0 : algResResp_->threadManage[ringIndex]->NotifyStart(); // 给线程发通知启动线程执行
1324 : } else {
1325 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
1326 0 : if (opInfo != nullptr) {
1327 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1328 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, dispatcher_);
1329 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
1330 0 : CHK_SMART_PTR_NULL(tempAlg);
1331 0 : CHK_RET(tempAlg->Prepare(
1332 : reduceAttr, opInfo, topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
1333 : subSignalsInOneRing, rankOrder, userMemInputSlices, isSdma));
1334 0 : HCCL_DEBUG("[MultiRingReduceScatterConcurrent]run in COMM_LEVEL0 ends");
1335 : } else {
1336 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1337 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
1338 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL0", __func__);
1339 0 : CHK_SMART_PTR_NULL(tempAlg);
1340 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
1341 : }
1342 0 : HCCL_DEBUG("[MultiRingReduceScatterConcurrent]run in COMM_LEVEL0 ends");
1343 0 : CHK_SMART_PTR_NULL(tempAlg);
1344 :
1345 0 : ret = LocalNotify::Wait(
1346 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
1347 0 : CHK_PRT_RET(
1348 : ret != HCCL_SUCCESS,
1349 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u] wait failed", ringIndex),
1350 : ret);
1351 0 : ret = tempAlg->Prepare(
1352 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex], reductionOp,
1353 0 : LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset, ringNics[ringIndex % halfRingSize]);
1354 0 : CHK_PRT_RET(
1355 : ret != HCCL_SUCCESS,
1356 : HCCL_ERROR(
1357 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring) "
1358 : "prepare failed,return[%d]",
1359 : ringIndex, ret),
1360 : ret);
1361 0 : ret = tempAlg->RegisterProfiler(
1362 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
1363 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
1364 0 : profStage, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
1365 0 : CHK_PRT_RET(
1366 : ret != HCCL_SUCCESS,
1367 : HCCL_ERROR(
1368 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring) "
1369 : "register Profiler failed,return[%d]",
1370 : ringIndex, ret),
1371 : ret);
1372 :
1373 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1374 0 : CHK_PRT_RET(
1375 : ret != HCCL_SUCCESS,
1376 : HCCL_ERROR(
1377 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring)"
1378 : " run failed,return[%d]",
1379 : ringIndex, ret),
1380 : ret);
1381 :
1382 0 : ret = LocalNotify::Post(
1383 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], profStage);
1384 0 : CHK_PRT_RET(
1385 : ret != HCCL_SUCCESS,
1386 : HCCL_ERROR(
1387 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u] record failed", ringIndex),
1388 : ret);
1389 0 : }
1390 : /* 主环record启动从环 */
1391 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], profStage);
1392 0 : CHK_PRT_RET(
1393 : ret != HCCL_SUCCESS,
1394 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u] record failed", ringIndex),
1395 : ret);
1396 : } else { // 主环 最后一个环
1397 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
1398 0 : if (opInfo != nullptr) {
1399 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1400 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, dispatcher_);
1401 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
1402 0 : CHK_SMART_PTR_NULL(tempAlg);
1403 0 : CHK_RET(tempAlg->Prepare(
1404 : reduceAttr, opInfo, topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
1405 : subSignalsInOneRing, rankOrder, userMemInputSlices, isSdma));
1406 : } else {
1407 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1408 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
1409 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL0", __func__);
1410 0 : CHK_SMART_PTR_NULL(tempAlg);
1411 0 : CHK_RET(tempAlg->Prepare(reduceAttr));
1412 : }
1413 0 : CHK_SMART_PTR_NULL(tempAlg);
1414 0 : ret = tempAlg->Prepare(
1415 : inputMem, inputMem, outputMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID,
1416 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex % halfRingSize]);
1417 0 : CHK_PRT_RET(
1418 : ret != HCCL_SUCCESS,
1419 : HCCL_ERROR(
1420 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring) "
1421 : "prepare failed,return[%d]",
1422 : ringIndex, ret),
1423 : ret);
1424 :
1425 0 : ret = tempAlg->RegisterProfiler(
1426 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1427 0 : + level0RingCommInfo.localRank,
1428 : profStage, HCCL_EXEC_STEP_NOT_SET, stream);
1429 0 : CHK_PRT_RET(
1430 : ret != HCCL_SUCCESS,
1431 : HCCL_ERROR(
1432 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring) "
1433 : "register Profiler failed,return[%d]",
1434 : ringIndex, ret),
1435 : ret);
1436 :
1437 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1438 0 : CHK_PRT_RET(
1439 : ret != HCCL_SUCCESS,
1440 : HCCL_ERROR(
1441 : "[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u],ReduceScatter(ring) run "
1442 : "failed,return[%d]",
1443 : ringIndex, ret),
1444 : ret);
1445 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
1446 0 : if (!static_cast<bool>(topoMatcher_->GetExternalInputHcclEnableFfts())
1447 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
1448 0 : algResResp_->threadManage[ring]->WaitDone();
1449 : }
1450 : /* 等待executor执行完毕 */
1451 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], profStage);
1452 0 : CHK_PRT_RET(
1453 : ret != HCCL_SUCCESS,
1454 : HCCL_ERROR("[CollCommExecutor][MultiRingReduceScatterConcurrent]stream[%u] wait failed", ring),
1455 : ret);
1456 : }
1457 0 : }
1458 0 : }
1459 :
1460 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1461 0 : return HCCL_SUCCESS;
1462 0 : }
1463 :
1464 0 : HcclResult CollCommExecutor::Level1ReduceScatterConcurrent(
1465 : DeviceMem inputMem, DeviceMem scratchMem, const u64 count, const HcclDataType dataType,
1466 : const HcclReduceOp reductionOp, Stream stream, s32 profStage, std::vector<Slice>& level1DataSegsSlice,
1467 : u32 syncTrans, u64 reduceAttr)
1468 : {
1469 : (void)profStage;
1470 0 : std::vector<std::pair<bool, std::vector<Slice>>> level1MultSlice;
1471 0 : level1MultSlice.resize(RDMA_PLANE_NUM_IN_NPRING_DOUBLE);
1472 0 : std::vector<Slice> sdmaSlice;
1473 0 : std::vector<Slice> rdmaSlice;
1474 0 : for (u32 segsIndex = 0; segsIndex < level1DataSegsSlice.size(); segsIndex++) {
1475 0 : u64 totalSize = level1DataSegsSlice[segsIndex].size;
1476 0 : u64 sdmaSliceOffset = level1DataSegsSlice[segsIndex].offset;
1477 0 : u64 sdmaSliceSize = ((totalSize <= HCCL_MIN_SLICE_ALIGN_910_93) || (syncTrans == MAX_SPLIT_VALUE)) ?
1478 : totalSize :
1479 0 : ((syncTrans * totalSize / MAX_SPLIT_VALUE) / HCCL_MIN_SLICE_ALIGN_910_93)
1480 0 : * HCCL_MIN_SLICE_ALIGN_910_93;
1481 0 : Slice sdmaSliceTmp;
1482 0 : sdmaSliceTmp.offset = sdmaSliceOffset;
1483 0 : sdmaSliceTmp.size = sdmaSliceSize;
1484 0 : Slice rdmaSliceTmp;
1485 0 : rdmaSliceTmp.offset = sdmaSliceOffset + sdmaSliceSize;
1486 0 : rdmaSliceTmp.size = totalSize - sdmaSliceSize;
1487 0 : sdmaSlice.push_back(sdmaSliceTmp);
1488 0 : rdmaSlice.push_back(rdmaSliceTmp);
1489 0 : HCCL_DEBUG(
1490 : "Level1 data segId:%u, Original [offset %llu, size %llu], sdma [offset %llu, size %llu], "
1491 : "rdma [offset %llu, size %llu]",
1492 : segsIndex, sdmaSliceOffset, totalSize, sdmaSliceTmp.offset, sdmaSliceTmp.size, rdmaSliceTmp.offset,
1493 : rdmaSliceTmp.size);
1494 : }
1495 0 : level1MultSlice[0] = std::make_pair(true, sdmaSlice); // true表示使用sdma
1496 0 : level1MultSlice[1] = std::make_pair(false, rdmaSlice); // false表示rdma
1497 :
1498 0 : u32 commPlaneNum = level1MultSlice.size();
1499 0 : bool isAnyPathCommLevel0 = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING
1500 0 : && workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) ?
1501 : true :
1502 : false;
1503 0 : CommPlane commplane = (isAnyPathCommLevel0) ? COMM_LEVEL0_ANYPATH_SDMA : COMM_LEVEL0;
1504 0 : u32 commIndex = GetSubCommInfo(commplane, COMM_INDEX_0).localRank;
1505 0 : CHK_RET(CheckCommSize(COMM_LEVEL1_ANYPATH_SDMA, commIndex + 1));
1506 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1_ANYPATH_SDMA, commIndex);
1507 0 : CHK_RET(CheckCommSize(COMM_LEVEL1_ANYPATH_RDMA, commIndex + 1));
1508 0 : SubCommInfo level1RdmaCommInfo = GetSubCommInfo(COMM_LEVEL1_ANYPATH_RDMA, commIndex);
1509 0 : for (u32 planeIndex = 0; planeIndex < commPlaneNum; planeIndex++) {
1510 0 : std::vector<Slice>& singleSlice = level1MultSlice[planeIndex].second;
1511 0 : SubCommInfo level1TempCommInfo = level1MultSlice[planeIndex].first ? level1CommInfo : level1RdmaCommInfo;
1512 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
1513 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
1514 : level1TempAlg
1515 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
1516 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
1517 : } else {
1518 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1519 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
1520 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
1521 : }
1522 0 : CHK_SMART_PTR_NULL(level1TempAlg);
1523 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
1524 0 : HcclResult ret = HCCL_SUCCESS;
1525 :
1526 0 : if (planeIndex != (commPlaneNum - 1)) {
1527 0 : ret = LocalNotify::Wait(
1528 0 : algResResp_->slaveStreams[planeIndex], dispatcher_, algResResp_->notifiesAux[planeIndex], reductionOp);
1529 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("stream[%u] wait failed", planeIndex), ret);
1530 :
1531 0 : CHK_RET(level1TempAlg->Prepare(
1532 : inputMem, inputMem, scratchMem, count, dataType, algResResp_->slaveStreams[planeIndex], reductionOp,
1533 : LEVEL0_BRIDGE_RANK_ID, singleSlice));
1534 :
1535 0 : CHK_RET(level1TempAlg->RegisterProfiler(
1536 : (level1TempCommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1TempCommInfo.localRank,
1537 : reductionOp, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[planeIndex]));
1538 :
1539 0 : CHK_RET(RunTemplate(level1TempAlg, level1TempCommInfo));
1540 0 : ret = LocalNotify::Post(
1541 0 : algResResp_->slaveStreams[planeIndex], dispatcher_, algResResp_->notifiesMain[planeIndex], reductionOp);
1542 0 : CHK_PRT_RET(
1543 : ret != HCCL_SUCCESS, HCCL_ERROR("[collAllGather]level1 stream[%u] record failed", planeIndex), ret);
1544 : // 主环record启动从环
1545 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[planeIndex], reductionOp);
1546 0 : CHK_PRT_RET(
1547 : ret != HCCL_SUCCESS, HCCL_ERROR("[collAllGather]level1 stream[%u] record failed", planeIndex), ret);
1548 : } else {
1549 0 : CHK_RET(level1TempAlg->Prepare(
1550 : inputMem, inputMem, scratchMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID,
1551 : singleSlice));
1552 0 : CHK_RET(level1TempAlg->RegisterProfiler(
1553 : (level1TempCommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1TempCommInfo.localRank,
1554 : reductionOp, HCCL_EXEC_STEP_NOT_SET, stream));
1555 :
1556 0 : CHK_RET(RunTemplate(level1TempAlg, level1TempCommInfo));
1557 0 : for (u32 ring = 0; ring < (commPlaneNum - 1); ring++) {
1558 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], reductionOp);
1559 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("param.stream[%u] wait failed", ring), ret);
1560 : }
1561 : }
1562 0 : }
1563 0 : HCCL_INFO("Level1ReduceScatterConcurrent run success");
1564 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, scratchMem, stream, dispatcher_));
1565 0 : return HCCL_SUCCESS;
1566 0 : }
1567 :
1568 0 : HcclResult CollCommExecutor::MultiRingMultiRootScatter(
1569 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType dataType,
1570 : const std::vector<std::vector<Slice>>& multRingsSliceZero, u32 root, Stream stream, const u64 baseOffset)
1571 : {
1572 0 : HcclResult ret = HCCL_SUCCESS;
1573 0 : u32 ringNum = multRingsSliceZero.size();
1574 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
1575 :
1576 0 : std::vector<std::vector<u32>> ringNics;
1577 0 : CHK_RET(GetRingNics(tag, ringNics));
1578 :
1579 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
1580 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
1581 0 : CHK_PRT_RET(
1582 : singleRingSliceZero.empty(),
1583 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]singleRingSliceZero is empty"), HCCL_E_INTERNAL);
1584 :
1585 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, ringIndex);
1586 :
1587 0 : u32 rankSize = level0RingCommInfo.localRankSize;
1588 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1589 0 : TemplateType::TEMPLATE_MULTI_ROOT_SCATTER_RING, dispatcher_);
1590 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_MULTI_ROOT_SCATTER_RING in COMM_LEVEL0", __func__);
1591 0 : CHK_SMART_PTR_NULL(tempAlg);
1592 :
1593 0 : if (ringIndex != (ringNum - 1)) {
1594 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
1595 0 : CHK_RET(StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
1596 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr()));
1597 : }
1598 : }
1599 :
1600 0 : u32 rootRank = 0;
1601 0 : ret = GetRankByUserRank(COMM_LEVEL0, ringIndex, root, rootRank);
1602 0 : CHK_PRT_RET(
1603 : ret == HCCL_E_PARA,
1604 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]invalid root [%u] to get userrank", root), ret);
1605 :
1606 0 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
1607 0 : ret = LocalNotify::Wait(
1608 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex], PROF_STAGE_0);
1609 0 : CHK_PRT_RET(
1610 : ret != HCCL_SUCCESS,
1611 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]in stream[%u] wait failed", ringIndex), ret);
1612 :
1613 0 : ret = tempAlg->Prepare(
1614 0 : inputMem, outputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
1615 : HcclReduceOp::HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, singleRingSliceZero, baseOffset,
1616 0 : ringNics[ringIndex]);
1617 0 : CHK_PRT_RET(
1618 : ret != HCCL_SUCCESS,
1619 : HCCL_ERROR(
1620 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u],multirootscatter(ring) "
1621 : "prepare failed,return[%d]",
1622 : ringIndex, ret),
1623 : ret);
1624 :
1625 0 : ret = tempAlg->RegisterProfiler(
1626 0 : ((ringIndex + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1627 0 : + level0RingCommInfo.localRank,
1628 0 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
1629 0 : CHK_PRT_RET(
1630 : ret != HCCL_SUCCESS,
1631 : HCCL_ERROR(
1632 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u], multirootscatter(ring) "
1633 : "register profiler failed,return[%d]",
1634 : ringIndex, ret),
1635 : ret);
1636 :
1637 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1638 0 : CHK_PRT_RET(
1639 : ret != HCCL_SUCCESS,
1640 : HCCL_ERROR(
1641 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u],multirootscatter(ring) "
1642 : "failed,return[%d]",
1643 : ringIndex, ret),
1644 : ret);
1645 :
1646 0 : ret = LocalNotify::Post(
1647 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex], PROF_STAGE_0);
1648 :
1649 0 : CHK_PRT_RET(
1650 : ret != HCCL_SUCCESS,
1651 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]stream[%u] record failed", ringIndex), ret);
1652 :
1653 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], PROF_STAGE_0);
1654 0 : CHK_PRT_RET(
1655 : ret != HCCL_SUCCESS,
1656 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]stream[%u] record failed", ringIndex), ret);
1657 : } else { // 主环
1658 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1659 0 : TemplateType::TEMPLATE_MULTI_ROOT_SCATTER_RING, dispatcher_);
1660 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_MULTI_ROOT_SCATTER_RING in COMM_LEVEL0", __func__);
1661 0 : CHK_SMART_PTR_NULL(tempAlg);
1662 0 : ret = tempAlg->Prepare(
1663 : inputMem, outputMem, outputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID,
1664 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex]);
1665 0 : CHK_PRT_RET(
1666 : ret != HCCL_SUCCESS,
1667 : HCCL_ERROR(
1668 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u],multirootscatter(ring) "
1669 : "prepare failed,return[%d]",
1670 : ringIndex, ret),
1671 : ret);
1672 :
1673 0 : ret = tempAlg->RegisterProfiler(
1674 0 : ((ringIndex + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1675 0 : + level0RingCommInfo.localRank,
1676 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, stream);
1677 0 : CHK_PRT_RET(
1678 : ret != HCCL_SUCCESS,
1679 : HCCL_ERROR(
1680 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u], multirootscatter(ring) "
1681 : "register profiler failed,return[%d]",
1682 : ringIndex, ret),
1683 : ret);
1684 :
1685 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
1686 0 : CHK_PRT_RET(
1687 : ret != HCCL_SUCCESS,
1688 : HCCL_ERROR(
1689 : "[CollCommExecutor][MultiRingMultiRootScatter]stream[%u],multirootscatter(ring) run "
1690 : "failed,return[%d]",
1691 : ringIndex, ret),
1692 : ret);
1693 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
1694 : /* 等待executor执行完毕 , 当前环没有分配数据,跳过此环处理,继续下一个环 */
1695 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], PROF_STAGE_0);
1696 0 : CHK_PRT_RET(
1697 : ret != HCCL_SUCCESS,
1698 : HCCL_ERROR("[CollCommExecutor][MultiRingMultiRootScatter]stream[%u] wait failed", ring), ret);
1699 : }
1700 : }
1701 0 : }
1702 :
1703 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1704 0 : return HCCL_SUCCESS;
1705 0 : }
1706 :
1707 0 : HcclResult CollCommExecutor::MultiStreamReduceScatterMeshAtomic(
1708 : const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType dataType,
1709 : const HcclReduceOp reductionOp, const std::vector<Slice>& dataSliceVct, Stream& stream,
1710 : const CommPlane commLevelIndex, const u64 baseOffset, HcomCollOpInfo* opInfo)
1711 : {
1712 : (void)tag;
1713 0 : u32 unitSize = SIZE_TABLE[dataType];
1714 :
1715 0 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
1716 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
1717 0 : DeviceMem deviceOutputMem = inputMem;
1718 0 : if (topoAttr_.isSingleMeshAggregation && (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE)
1719 0 : && static_cast<bool>((reduceAttr & INLINE_REDUCE_BITMASK)) && (opInfo != nullptr)) {
1720 0 : if (((opInfo->count) * unitSize <= HCCL_SMALL_COUNT_32_KB)
1721 0 : && (topoAttr_.deviceNumPerAggregation == DEVICE_EIGHT)) {
1722 0 : deviceOutputMem = outputMem;
1723 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1724 0 : TemplateType::TEMPLATE_REDUCESCATTER_HDSTAGE, dispatcher_);
1725 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_HDSTAGE in COMM_LEVEL0", __func__);
1726 0 : } else {
1727 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1728 0 : TemplateType::TEMPLATE_REDUCESCATTER_MESH_DIRECT, dispatcher_);
1729 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_MESH_DIRECT in COMM_LEVEL0", __func__);
1730 : }
1731 0 : } else {
1732 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1733 0 : TemplateType::TEMPLATE_REDUCESCATTER_MESH_ATOMIC, dispatcher_);
1734 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_MESH_ATOMIC in COMM_LEVEL0", __func__);
1735 : }
1736 0 : CHK_SMART_PTR_NULL(tempAlg);
1737 :
1738 0 : CHK_RET(CheckCommSize(commLevelIndex, COMM_INDEX_0 + 1));
1739 0 : const SubCommInfo subCommInfo = GetSubCommInfo(commLevelIndex, COMM_INDEX_0);
1740 0 : CHK_RET(tempAlg->Prepare(
1741 : inputMem, deviceOutputMem, outputMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID, dataSliceVct,
1742 : baseOffset, reduceAttr, algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux,
1743 : topoAttr_.userRank, opInfo));
1744 :
1745 0 : CHK_RET(tempAlg->RegisterProfiler(
1746 : (subCommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + subCommInfo.localRank, PROF_STAGE_0,
1747 : HCCL_EXEC_STEP_NOT_SET, stream));
1748 :
1749 0 : CHK_RET(RunTemplate(tempAlg, subCommInfo));
1750 :
1751 0 : return HCCL_SUCCESS;
1752 0 : }
1753 :
1754 2 : HcclResult CollCommExecutor::MultiStreamReduceScatterMesh(
1755 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
1756 : const HcclReduceOp reductionOp, const std::vector<std::vector<Slice>>& multStreamsSlice, Stream stream,
1757 : const CommPlane commLevelIndex, const u64 baseOffset)
1758 : {
1759 : (void)tag;
1760 2 : HcclResult ret = HCCL_SUCCESS;
1761 2 : u64 streamNum = multStreamsSlice.size();
1762 2 : HCCL_INFO("MultiStreamReduceScatterMesh streamNum[%llu]", streamNum);
1763 2 : CHK_RET(CheckCommSize(commLevelIndex, streamNum));
1764 2 : const SubCommInfo zeroCommInfo = GetSubCommInfo(commLevelIndex, COMM_INDEX_0);
1765 :
1766 2 : u64 reduceAttr = GetReduceAttr(inputMem, outputMem, dataType, reductionOp);
1767 :
1768 2 : for (u32 streamIndex = 0; streamIndex < streamNum; streamIndex++) {
1769 0 : std::vector<Slice> singleStreamSlice = multStreamsSlice[streamIndex];
1770 0 : CHK_PRT_RET(
1771 : singleStreamSlice.size() <= 0,
1772 : HCCL_ERROR("[CollCommExecutor][MultiStreamReduceScatterMesh]singleStreamSlice is empty"), HCCL_E_INTERNAL);
1773 :
1774 0 : const SubCommInfo subCommInfo = GetSubCommInfo(commLevelIndex, streamIndex);
1775 0 : u32 commIndex = subCommInfo.localRank;
1776 0 : CHK_PRT_RET(
1777 : commIndex >= singleStreamSlice.size(),
1778 : HCCL_ERROR(
1779 : "[CollCommExecutor][MultiStreamReduceScatterMesh]commIndex[%u] => "
1780 : "singleStreamSlice size[%zu]",
1781 : commIndex, singleStreamSlice.size()),
1782 : HCCL_E_INTERNAL);
1783 :
1784 0 : u32 rankSize = subCommInfo.localRankSize;
1785 0 : u32 ringIndexOp = streamIndex;
1786 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
1787 :
1788 0 : if (topoAttr_.isDiffDeviceType) {
1789 0 : HCCL_DEBUG("[CollCommExecutor][MultiStreamReduceScatterMesh]isDiffDeviceType");
1790 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1791 0 : TemplateType::TEMPLATE_REDUCESCATTER_MESH_MIX_SS, dispatcher_);
1792 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_MESH_MIX_SS in COMM_LEVEL0", __func__);
1793 : } else {
1794 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
1795 0 : TemplateType::TEMPLATE_REDUCESCATTER_MESH, dispatcher_);
1796 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_MESH in COMM_LEVEL0", __func__);
1797 : }
1798 0 : CHK_SMART_PTR_NULL(tempAlg);
1799 0 : CHK_RET(tempAlg->Prepare(reduceAttr, streamIndex));
1800 :
1801 0 : if (streamIndex != (streamNum - 1)) { // 0~ringNum-2的环
1802 0 : HCCL_INFO("MultiStreamReduceScatterMesh step into subStream");
1803 0 : ret = LocalNotify::Wait(
1804 0 : algResResp_->slaveStreams[streamIndex], dispatcher_, algResResp_->notifiesAux[streamIndex],
1805 : PROF_STAGE_0);
1806 : // 等待executor执行完毕
1807 0 : CHK_PRT_RET(
1808 : ret != HCCL_SUCCESS,
1809 : HCCL_ERROR("[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u] wait failed", streamIndex), ret);
1810 :
1811 0 : ret = tempAlg->Prepare(
1812 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[streamIndex], reductionOp,
1813 : LEVEL0_BRIDGE_RANK_ID, singleStreamSlice, baseOffset);
1814 0 : CHK_PRT_RET(
1815 : ret != HCCL_SUCCESS,
1816 : HCCL_ERROR(
1817 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u],ReduceScatter(mesh) "
1818 : "prepare failed,return[%d]",
1819 : streamIndex, ret),
1820 : ret);
1821 :
1822 0 : ret = tempAlg->RegisterProfiler(
1823 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1824 0 : + zeroCommInfo.localRank,
1825 0 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[streamIndex]);
1826 0 : CHK_PRT_RET(
1827 : ret != HCCL_SUCCESS,
1828 : HCCL_ERROR(
1829 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u],ReduceScatter(mesh) "
1830 : "register Profiler failed,return[%d]",
1831 : streamIndex, ret),
1832 : ret);
1833 :
1834 0 : ret = RunTemplate(tempAlg, subCommInfo);
1835 0 : CHK_PRT_RET(
1836 : ret != HCCL_SUCCESS,
1837 : HCCL_ERROR(
1838 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u],ReduceScatter(mesh) run "
1839 : "failed,return[%d]",
1840 : streamIndex, ret),
1841 : ret);
1842 :
1843 0 : ret = LocalNotify::Post(
1844 0 : algResResp_->slaveStreams[streamIndex], dispatcher_, algResResp_->notifiesMain[streamIndex],
1845 : PROF_STAGE_0);
1846 0 : CHK_PRT_RET(
1847 : ret != HCCL_SUCCESS,
1848 : HCCL_ERROR("[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u] record failed", streamIndex),
1849 : ret);
1850 :
1851 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[streamIndex], PROF_STAGE_0);
1852 0 : CHK_PRT_RET(
1853 : ret != HCCL_SUCCESS,
1854 : HCCL_ERROR("[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u] record failed", streamIndex),
1855 : ret);
1856 : } else { // 主环
1857 0 : HCCL_INFO("MultiStreamReduceScatterMesh step into mainStream");
1858 :
1859 0 : ret = tempAlg->Prepare(
1860 : inputMem, inputMem, outputMem, count, dataType, stream, reductionOp, LEVEL0_BRIDGE_RANK_ID,
1861 : singleStreamSlice, baseOffset);
1862 0 : CHK_PRT_RET(
1863 : ret != HCCL_SUCCESS,
1864 : HCCL_ERROR(
1865 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u], "
1866 : "ReduceScatter(mesh) prepare failed, return[%d]",
1867 : streamIndex, ret),
1868 : ret);
1869 :
1870 0 : ret = tempAlg->RegisterProfiler(
1871 0 : ((ringIndexOp + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
1872 0 : + zeroCommInfo.localRank,
1873 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, stream);
1874 0 : CHK_PRT_RET(
1875 : ret != HCCL_SUCCESS,
1876 : HCCL_ERROR(
1877 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u], ReduceScatter(mesh) "
1878 : "register Profiler failed, return[%d]",
1879 : streamIndex, ret),
1880 : ret);
1881 :
1882 0 : ret = RunTemplate(tempAlg, subCommInfo);
1883 0 : CHK_PRT_RET(
1884 : ret != HCCL_SUCCESS,
1885 : HCCL_ERROR(
1886 : "[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u], "
1887 : "ReduceScatter(mesh) run failed, return[%d]",
1888 : streamIndex, ret),
1889 : ret);
1890 :
1891 0 : for (u32 streamIndex = 0; streamIndex < (streamNum - 1); streamIndex++) {
1892 : // 等待executor执行完毕
1893 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[streamIndex], PROF_STAGE_0);
1894 0 : CHK_PRT_RET(
1895 : ret != HCCL_SUCCESS,
1896 : HCCL_ERROR("[CollCommExecutor][MultiStreamReduceScatterMesh]stream[%u] wait failed", streamIndex),
1897 : ret);
1898 : }
1899 : }
1900 0 : }
1901 :
1902 2 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
1903 2 : return ret;
1904 2 : }
1905 :
1906 2 : HcclResult CollCommExecutor::PrepareReduceScatterSliceData(
1907 : u64 dataCount, u32 unitSize, u32 sliceNum, std::vector<Slice>& dataSlice)
1908 : {
1909 2 : CHK_PRT_RET(
1910 : (sliceNum == 0), HCCL_ERROR("[CollCommExecutor][PrepareReduceScatterSliceData]sliceNum is zero."), HCCL_E_PARA);
1911 :
1912 2 : dataSlice.resize(sliceNum);
1913 2 : u64 sliceSize = dataCount * unitSize;
1914 4 : for (u32 i = 0; i < sliceNum; i++) {
1915 2 : dataSlice[i].size = sliceSize;
1916 2 : dataSlice[i].offset = (i * sliceSize);
1917 : }
1918 2 : return HCCL_SUCCESS;
1919 : }
1920 :
1921 : std::vector<std::vector<u32>>
1922 120 : CollCommExecutor::GetRingsOrderByTopoType(u32 ranksSize, TopoType topoType, std::vector<u32>& nicList)
1923 : {
1924 120 : std::vector<std::vector<u32>> multiRingOrder;
1925 120 : if (topoType == TopoType::TOPO_TYPE_8P_RING) { // 4 ring 场景
1926 : // 每个环的排序是按照设备物理ID进行的
1927 0 : std::vector<u32> tmpLevel00 = {0, 1, 2, 6, 5, 4, 7, 3}; // 环0
1928 0 : std::vector<u32> tmpLevel01 = {0, 3, 7, 4, 5, 6, 2, 1}; // 环1
1929 0 : std::vector<u32> tmpLevel02 = {0, 2, 3, 1, 5, 7, 6, 4}; // 环2
1930 0 : std::vector<u32> tmpLevel03 = {0, 4, 6, 7, 5, 1, 3, 2}; // 环3
1931 :
1932 : // 填充8pring 多环的comm level0 四个环的顺序
1933 0 : multiRingOrder.push_back(tmpLevel00);
1934 0 : multiRingOrder.push_back(tmpLevel01);
1935 0 : multiRingOrder.push_back(tmpLevel02);
1936 0 : multiRingOrder.push_back(tmpLevel03);
1937 120 : } else if (topoType == TopoType::TOPO_TYPE_NP_DOUBLE_RING) { // 2 ring 场景
1938 116 : std::vector<u32> tmpLevel00; // 环0
1939 116 : std::vector<u32> tmpLevel01; // 环1
1940 116 : tmpLevel00 = nicList; // { 0, 1, 2, 3, 4, 5, 6, 7 };
1941 116 : tmpLevel01.reserve(ranksSize);
1942 116 : tmpLevel01.push_back(nicList[0]);
1943 116 : tmpLevel01.insert(tmpLevel01.end(), tmpLevel00.rbegin(), tmpLevel00.rend() - 1);
1944 : // 填充 double ring 两环的comm level0的顺序
1945 116 : multiRingOrder.push_back(tmpLevel00);
1946 116 : multiRingOrder.push_back(tmpLevel01);
1947 116 : } else { // 1 ring 场景
1948 4 : std::vector<u32> tmpLevel00 = nicList; // 环0
1949 :
1950 : // 填充 single ring 单环的comm level0的顺序
1951 4 : multiRingOrder.push_back(tmpLevel00);
1952 4 : }
1953 : // 打印多个环
1954 120 : if (UNLIKELY(HcclCheckLogLevel(DLOG_DEBUG))) {
1955 356 : for (size_t i = 0; i < multiRingOrder.size(); i++) {
1956 236 : auto ring = multiRingOrder[i];
1957 236 : std::ostringstream stringRepresentation;
1958 753 : for (std::vector<uint32_t>::iterator it = ring.begin(); it != ring.end(); it++) {
1959 517 : stringRepresentation << *it << " ";
1960 : }
1961 236 : std::string ringString = stringRepresentation.str();
1962 236 : const char* charRing = ringString.c_str();
1963 236 : HCCL_DEBUG("[GetRingsOrderByTopoType] The No.%zu ring: %s", i, charRing);
1964 236 : }
1965 : }
1966 120 : return multiRingOrder;
1967 0 : }
1968 :
1969 : std::vector<std::vector<u32>>
1970 0 : CollCommExecutor::GetRingsOrderForAnyPath(u32 ranksSize, TopoType topoType, std::vector<u32>& nicList)
1971 : {
1972 0 : std::vector<std::vector<u32>> multiRingOrder;
1973 0 : if (topoType == TopoType::TOPO_TYPE_NP_DOUBLE_RING) { // 2 ring 场景
1974 0 : std::vector<u32> tmpLevel00; // 环0
1975 0 : std::vector<u32> tmpLevel01; // 环1
1976 0 : std::vector<u32> rohLevel0;
1977 0 : if (topoMatcher_->CheckSdmaWithRohTopo(nicList, rohLevel0)) {
1978 0 : tmpLevel00 = rohLevel0; // 环0, 8卡 { 0, 1, 3, 2, 4, 5, 7, 6 };
1979 0 : tmpLevel01.reserve(ranksSize); // 环1, 8卡 { 0, 6, 7, 5, 4, 2, 3, 1 };
1980 0 : tmpLevel01.push_back(rohLevel0[0]);
1981 0 : tmpLevel01.insert(tmpLevel01.end(), rohLevel0.rbegin(), rohLevel0.rend() - 1);
1982 : } else {
1983 0 : tmpLevel00 = nicList; // { 0, 1, 2, 3, 4, 5, 6, 7 };
1984 0 : tmpLevel01.reserve(ranksSize);
1985 0 : tmpLevel01.push_back(nicList[0]);
1986 0 : tmpLevel01.insert(tmpLevel01.end(), tmpLevel00.rbegin(), tmpLevel00.rend() - 1);
1987 : }
1988 : // 填充 double ring 两环的comm level0的顺序
1989 0 : multiRingOrder.push_back(tmpLevel00);
1990 0 : multiRingOrder.push_back(tmpLevel01);
1991 0 : } else { // 1 ring 场景
1992 0 : std::vector<u32> tmpLevel00 = nicList; // 环0
1993 :
1994 : // 填充 single ring 单环的comm level0的顺序
1995 0 : multiRingOrder.push_back(tmpLevel00);
1996 0 : }
1997 : // 打印多个环
1998 0 : for (size_t i = 0; i < multiRingOrder.size(); i++) {
1999 0 : auto ring = multiRingOrder[i];
2000 0 : std::ostringstream stringRepresentation;
2001 0 : for (std::vector<uint32_t>::iterator it = ring.begin(); it != ring.end(); it++) {
2002 0 : stringRepresentation << *it << " ";
2003 : }
2004 0 : std::string ringString = stringRepresentation.str();
2005 0 : const char* charRing = ringString.c_str();
2006 0 : HCCL_INFO("[GetRingsOrderByRdmaSdmaConcurrent] The No.%zu ring: %s", i, charRing);
2007 0 : }
2008 0 : return multiRingOrder;
2009 0 : }
2010 :
2011 50 : HcclResult CollCommExecutor::MutliSegSlicePrepare(
2012 : const std::vector<Slice>& dataSegsSlice, std::vector<std::vector<Slice>>& mutliSegsSlices, u32 ringCount)
2013 : {
2014 50 : std::vector<Slice> singleSegSlices;
2015 50 : singleSegSlices.reserve(ringCount);
2016 158 : for (u32 rankId = 0; rankId < dataSegsSlice.size(); rankId++) {
2017 108 : Slice rankSliceTemp;
2018 108 : u64 rankDataSize = dataSegsSlice[rankId].size;
2019 108 : u32 ringIndex = 0;
2020 108 : u64 offsetStart = dataSegsSlice[rankId].offset;
2021 108 : if (rankDataSize > 0 && ringCount != 0) {
2022 108 : u64 sizeTemp = (rankDataSize + ringCount - 1) / ringCount; /* 1是为了向上取整 */
2023 108 : u64 sizePerRing = AlgTemplateBase::RoundUpWithDivisor(sizeTemp, HCCL_MIN_SLICE_ALIGN);
2024 108 : u64 residueSize = rankDataSize;
2025 :
2026 322 : while (residueSize > 0) {
2027 214 : u64 singleRingSize = sizePerRing < residueSize ? sizePerRing : residueSize;
2028 214 : rankSliceTemp.size = singleRingSize;
2029 214 : rankSliceTemp.offset = offsetStart + rankDataSize - residueSize;
2030 214 : ringIndex++;
2031 214 : if (singleRingSize == 0) {
2032 0 : HCCL_ERROR(
2033 : "[CollCommExecutor][MutliSegSlicePrepare]"
2034 : "Multrings slices prepare: singleRingSize[%llu]",
2035 : singleRingSize);
2036 0 : return HCCL_E_INTERNAL;
2037 : }
2038 214 : residueSize -= singleRingSize;
2039 214 : singleSegSlices.push_back(rankSliceTemp);
2040 : }
2041 : }
2042 110 : while (ringIndex < ringCount) {
2043 2 : rankSliceTemp.size = 0;
2044 2 : rankSliceTemp.offset = offsetStart;
2045 2 : ringIndex++;
2046 2 : singleSegSlices.push_back(rankSliceTemp);
2047 : }
2048 108 : mutliSegsSlices.push_back(singleSegSlices); // rings_slice 判断大小不为 8 则异常
2049 108 : singleSegSlices.clear();
2050 : }
2051 50 : return HCCL_SUCCESS;
2052 50 : }
2053 :
2054 0 : HcclResult CollCommExecutor::MutliSegSlicePrepareAvoidCceRewrite(
2055 : const std::vector<Slice>& dataSegsSlice, std::vector<std::vector<Slice>>& mutliSegsSlices, u32 ringCount) const
2056 : {
2057 0 : for (u32 rankId = 0; rankId < dataSegsSlice.size(); rankId++) {
2058 0 : Slice rankSliceTemp;
2059 0 : std::vector<Slice> singleSegSlices;
2060 0 : for (u32 ringIndex = 0; ringIndex < ringCount; ringIndex++) {
2061 0 : if (ringIndex < ringCount - 1) {
2062 0 : rankSliceTemp.size = 0;
2063 0 : rankSliceTemp.offset = dataSegsSlice[rankId].offset;
2064 : } else {
2065 0 : rankSliceTemp.size = dataSegsSlice[rankId].size;
2066 0 : rankSliceTemp.offset = dataSegsSlice[rankId].offset;
2067 : }
2068 0 : singleSegSlices.push_back(rankSliceTemp);
2069 : }
2070 0 : mutliSegsSlices.push_back(singleSegSlices); // rings_slice 判断大小不为 8 则异常
2071 0 : }
2072 0 : return HCCL_SUCCESS;
2073 : }
2074 :
2075 50 : void CollCommExecutor::NicSendSizeCal(
2076 : const std::vector<std::vector<Slice>>& mutliSegsSlices, u32 ringCount, u32 chunkSize,
2077 : const std::vector<u32>& nicList, const std::string& tag)
2078 : {
2079 : // 计算每个网口最终会发送的数据量大小
2080 50 : std::vector<u64> sizeList;
2081 50 : sizeList.reserve(nicList.size());
2082 158 : for (u32 nicIdx = 0; nicIdx < nicList.size(); nicIdx++) {
2083 108 : u64 tempSize = 0;
2084 216 : for (u32 chunkIdx = 0; chunkIdx < chunkSize; chunkIdx++) {
2085 324 : for (u32 ringIdx = 0; ringIdx < ringCount; ringIdx++) {
2086 216 : tempSize += mutliSegsSlices[nicIdx * chunkSize + chunkIdx][ringIdx].size;
2087 : }
2088 : }
2089 108 : sizeList.push_back(tempSize);
2090 : }
2091 50 : SetNicSendSize(tag, sizeList);
2092 50 : }
2093 :
2094 50 : std::vector<std::vector<Slice>> CollCommExecutor::PrepareMultiRingSlice(
2095 : const std::vector<Slice>& dataSegsSlice, const std::string& tag, bool avoidCceRewrite, std::vector<u32> nicList,
2096 : CommPlane commLevelIndex)
2097 : {
2098 : // get ranksSize
2099 50 : u32 ranksSize = GetSubCommInfo(commLevelIndex, COMM_INDEX_0).localRankSize;
2100 : // 获取每个ring上设备的排布顺序,顺序均为deviceID
2101 50 : sort(nicList.begin(), nicList.end());
2102 50 : std::vector<std::vector<u32>> multiRingsOrder;
2103 50 : if (topoMatcher_->GetARSFlag()) {
2104 0 : multiRingsOrder = GetRingsOrderByTopoType(nicList.size(), TopoType::TOPO_TYPE_NP_DOUBLE_RING, nicList);
2105 : } else {
2106 50 : multiRingsOrder = GetRingsOrderByTopoType(ranksSize, topoType_, nicList);
2107 : }
2108 50 : HCCL_INFO("[%s], multiRingsOrder.size() = %u", __func__, multiRingsOrder.size());
2109 50 : std::vector<std::vector<Slice>> mutliRingsSlices;
2110 50 : std::vector<std::vector<Slice>> mutliSegsSlices;
2111 50 : u32 ringCount = multiRingsOrder.size();
2112 : // 单环场景不应该走入此流程,需要在函数外校验
2113 50 : CHK_PRT_RET(
2114 : ringCount <= 1, HCCL_ERROR("[CollCommExecutor][PrepareMultiRingSlice] ringCount[%u] <= 1", ringCount),
2115 : mutliRingsSlices);
2116 :
2117 50 : u32 ringRanks = multiRingsOrder[0].size(); // 获取单个 ring 上设备的数量
2118 :
2119 : // 将数每块据切分为 ringCount 份
2120 50 : mutliSegsSlices.reserve(dataSegsSlice.size());
2121 : HcclResult ret;
2122 50 : if (avoidCceRewrite) {
2123 0 : ret = MutliSegSlicePrepareAvoidCceRewrite(dataSegsSlice, mutliSegsSlices, ringCount);
2124 : } else {
2125 50 : ret = MutliSegSlicePrepare(dataSegsSlice, mutliSegsSlices, ringCount);
2126 : }
2127 50 : if (ret != HCCL_SUCCESS) {
2128 0 : return mutliRingsSlices;
2129 : }
2130 50 : u32 chunkSize = ringRanks / nicList.size();
2131 50 : HCCL_DEBUG("[CollCommExecutor][PrepareMultiRingSlice]chunkSize is %u", chunkSize);
2132 50 : (void)NicSendSizeCal(mutliSegsSlices, ringCount, chunkSize, nicList, tag);
2133 50 : std::vector<u32> rankList;
2134 50 : std::vector<Slice> singleRingSlices;
2135 50 : std::vector<std::vector<u32>> ringRankList;
2136 :
2137 50 : ringRankList.reserve(ringCount);
2138 50 : singleRingSlices.reserve(ringRanks);
2139 50 : rankList.reserve(ringRanks);
2140 :
2141 150 : for (u32 ringIndex = 0; ringIndex < ringCount; ringIndex++) {
2142 316 : for (u32 segsIndex = 0; segsIndex < ringRanks; segsIndex++) {
2143 216 : u32 deviceIdx = multiRingsOrder[ringIndex][segsIndex];
2144 216 : std::vector<u32>::iterator iterRank = std::find(nicList.begin(), nicList.end(), deviceIdx);
2145 216 : if (iterRank != nicList.end()) {
2146 216 : rankList.push_back(segsIndex);
2147 216 : u32 nicPosition = distance(nicList.begin(), iterRank);
2148 432 : for (u32 chunkIdx = 0; chunkIdx < chunkSize; chunkIdx++) {
2149 216 : Slice tempSlice = mutliSegsSlices[nicPosition * chunkSize + chunkIdx][ringIndex];
2150 216 : singleRingSlices.push_back(tempSlice);
2151 : }
2152 : }
2153 : }
2154 100 : mutliRingsSlices.push_back(singleRingSlices);
2155 100 : ringRankList.push_back(rankList);
2156 100 : singleRingSlices.clear();
2157 100 : rankList.clear();
2158 : }
2159 :
2160 50 : ret = SetRingNics(tag, ringRankList);
2161 50 : if (ret != HCCL_SUCCESS) {
2162 0 : std::vector<std::vector<Slice>> emptySlice;
2163 0 : HCCL_ERROR("[Prepare][MultiRingSlice]set nics in ring failed, ret[%u]", ret);
2164 0 : return emptySlice;
2165 0 : }
2166 50 : return mutliRingsSlices;
2167 50 : }
2168 :
2169 0 : std::vector<std::vector<Slice>> CollCommExecutor::AnyPathPrepareMultiRingSlice(
2170 : const std::vector<Slice>& dataSegsSlice, const std::string& tag, bool avoidCceRewrite, std::vector<u32> nicList)
2171 : {
2172 0 : CheckCommSize(COMM_LEVEL0_ANYPATH_SDMA, COMM_INDEX_1);
2173 0 : u32 ranksSize = GetSubCommInfo(COMM_LEVEL0_ANYPATH_SDMA, COMM_INDEX_0).localRankSize;
2174 : // 获取每个ring上设备的排布顺序,顺序均为deviceID
2175 0 : sort(nicList.begin(), nicList.end());
2176 0 : std::vector<std::vector<u32>> multiRingsOrder = GetRingsOrderForAnyPath(ranksSize, topoType_, nicList);
2177 0 : std::vector<std::vector<Slice>> mutliRingsSlices;
2178 0 : std::vector<std::vector<Slice>> mutliSegsSlices;
2179 0 : u32 ringCount = multiRingsOrder.size();
2180 : // 单环场景不应该走入此流程,需要在函数外校验
2181 0 : CHK_PRT_RET(
2182 : ringCount <= 1, HCCL_ERROR("[CollCommExecutor][PrepareMultiRingSlice] ringCount[%u] <= 1", ringCount),
2183 : mutliRingsSlices);
2184 :
2185 0 : u32 ringRanks = multiRingsOrder[0].size(); // 获取单个 ring 上设备的数量
2186 :
2187 : // 将数每块据切分为 ringCount 份
2188 : HcclResult ret;
2189 0 : mutliSegsSlices.reserve(dataSegsSlice.size());
2190 0 : if (avoidCceRewrite) {
2191 0 : ret = MutliSegSlicePrepareAvoidCceRewrite(dataSegsSlice, mutliSegsSlices, ringCount);
2192 : } else {
2193 0 : ret = MutliSegSlicePrepare(dataSegsSlice, mutliSegsSlices, ringCount);
2194 : }
2195 0 : if (ret != HCCL_SUCCESS) {
2196 0 : return mutliRingsSlices;
2197 : }
2198 0 : u32 chunkSize = ringRanks / nicList.size();
2199 0 : (void)NicSendSizeCal(mutliSegsSlices, ringCount, chunkSize, nicList, tag);
2200 0 : std::vector<std::vector<u32>> ringRankList;
2201 0 : std::vector<Slice> singleRingSlices;
2202 0 : std::vector<u32> rankList;
2203 :
2204 0 : ringRankList.reserve(ringCount);
2205 0 : rankList.reserve(ringRanks);
2206 0 : singleRingSlices.reserve(ringRanks);
2207 :
2208 0 : for (u32 ringIndex = 0; ringIndex < ringCount; ringIndex++) {
2209 0 : for (u32 segsIndex = 0; segsIndex < ringRanks; ++segsIndex) {
2210 0 : u32 deviceIdx = multiRingsOrder[ringIndex][segsIndex];
2211 0 : std::vector<u32>::iterator iterRank = std::find(nicList.begin(), nicList.end(), deviceIdx);
2212 0 : if (iterRank != nicList.end()) {
2213 0 : u32 nicPosition = distance(nicList.begin(), iterRank);
2214 0 : for (u32 chunkIdx = 0; chunkIdx < chunkSize; chunkIdx++) {
2215 0 : Slice tempSlice = mutliSegsSlices[nicPosition * chunkSize + chunkIdx][ringIndex];
2216 0 : singleRingSlices.push_back(tempSlice);
2217 : }
2218 0 : rankList.push_back(segsIndex);
2219 : }
2220 : }
2221 0 : mutliRingsSlices.push_back(singleRingSlices);
2222 0 : singleRingSlices.clear();
2223 0 : ringRankList.push_back(rankList);
2224 0 : rankList.clear();
2225 : }
2226 :
2227 0 : ret = SetRingNics(tag, ringRankList);
2228 0 : if (ret != HCCL_SUCCESS) {
2229 0 : HCCL_ERROR("[Prepare][MultiRingSlice]set nics in ring failed, ret[%u]", ret);
2230 0 : std::vector<std::vector<Slice>> emptySlice;
2231 0 : return emptySlice;
2232 0 : }
2233 0 : return mutliRingsSlices;
2234 0 : }
2235 :
2236 109 : u64 CollCommExecutor::GetReduceAttr(DeviceMem& inputMem, DeviceMem& outputMem, HcclDataType dataType, HcclReduceOp op)
2237 : {
2238 109 : u64 reduceAttr = 0;
2239 109 : bool isInlineReduce = IsSupportSDMAReduce(inputMem.ptr(), outputMem.ptr(), dataType, op);
2240 109 : if (isInlineReduce && algoAttr_.inlineReduceSwitchOn) {
2241 97 : SalSetBitOne(reduceAttr, ATTR_POS_INLINE_REDUCE);
2242 : }
2243 :
2244 109 : bool isRdmaReduce = IsSupportRDMAReduce(dataType, op);
2245 109 : if (isRdmaReduce) {
2246 106 : SalSetBitOne(reduceAttr, ATTR_POS_SUPPORT_RDMA_REDUCE);
2247 : }
2248 :
2249 109 : return reduceAttr;
2250 : }
2251 :
2252 67 : HcclResult CollCommExecutor::CalUserMemSlices(
2253 : const HcclDataType dataType, const HcomCollOpInfo* opInfo, const std::vector<Slice>& singleRingSliceZero,
2254 : u32 ringIndex, const std::vector<std::vector<u32>>& multiRingsOrder, std::vector<Slice>& userMemSlices)
2255 : {
2256 67 : if (opInfo == nullptr || opInfo->inputAddr == nullptr || opInfo->outputAddr == nullptr) {
2257 : // 910_93场景下,allreduce算子的userMem上的slice信息
2258 67 : userMemSlices = singleRingSliceZero;
2259 67 : return HCCL_SUCCESS;
2260 : }
2261 : // 910_93场景下,reduce scatter和AllGather算子的userMem上的slice信息
2262 0 : std::vector<u32> ring0 = multiRingsOrder[0];
2263 0 : for (u32 sliceIdx = 0; sliceIdx < singleRingSliceZero.size(); sliceIdx++) {
2264 0 : Slice userMemSlice;
2265 : u32 deviceId;
2266 0 : if (ringIndex >= SLICES_FACTOR) {
2267 0 : deviceId = multiRingsOrder[ringIndex % SLICES_FACTOR][sliceIdx];
2268 : } else {
2269 0 : deviceId = multiRingsOrder[ringIndex][sliceIdx];
2270 : }
2271 :
2272 0 : u32 pos = distance(ring0.begin(), find(ring0.begin(), ring0.end(), deviceId));
2273 : // 专用于MC2调用的 strideCount 特性
2274 0 : u64 count = (opInfo->strideCount == 0) ? opInfo->count : opInfo->strideCount;
2275 0 : userMemSlice.offset = pos * count * SIZE_TABLE[dataType] + singleRingSliceZero[0].offset;
2276 0 : userMemSlice.size = singleRingSliceZero[sliceIdx].size;
2277 0 : userMemSlices.push_back(userMemSlice);
2278 0 : HCCL_DEBUG(
2279 : "[CollCommExecutor][CalUserMemSlices] Push back userMemSlice offset[%llu], size[%llu] at rank[%u]",
2280 : userMemSlice.offset, userMemSlice.size, topoAttr_.userRank);
2281 : }
2282 0 : return HCCL_SUCCESS;
2283 0 : }
2284 :
2285 135 : HcclResult CollCommExecutor::GetRankOrder(
2286 : const std::vector<std::vector<u32>>& multiRingsOrder, u32 ringIndex, std::vector<u32>& rankOrder)
2287 : {
2288 135 : std::vector<u32> ring0 = multiRingsOrder[0];
2289 135 : std::vector<u32> ringOrder = multiRingsOrder[ringIndex];
2290 434 : for (u32 i = 0; i < ringOrder.size(); i++) {
2291 299 : u32 deviceId = ringOrder[i];
2292 299 : u32 pos = distance(ring0.begin(), find(ring0.begin(), ring0.end(), deviceId));
2293 299 : rankOrder.push_back(pos);
2294 : }
2295 135 : return HCCL_SUCCESS;
2296 135 : }
2297 :
2298 0 : HcclResult CollCommExecutor::MultiRingScatter(
2299 : const std::string& tag, DeviceMem inputMem, DeviceMem outputMem, const u64 count, const HcclDataType dataType,
2300 : const std::vector<std::vector<Slice>> multRingsSliceZero, u32 root, Stream stream, const HcomCollOpInfo* opInfo,
2301 : const u64 baseOffset)
2302 : {
2303 0 : HcclResult ret = HCCL_SUCCESS;
2304 0 : u32 ringNum = multRingsSliceZero.size();
2305 :
2306 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, ringNum));
2307 :
2308 0 : std::vector<std::vector<u32>> ringNics;
2309 0 : CHK_RET(GetRingNics(tag, ringNics));
2310 :
2311 : // 拿到ring环映射关系
2312 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
2313 0 : auto nicList = topoAttr_.nicList;
2314 : std::vector<std::vector<u32>> multiRingsOrder
2315 0 : = GetRingsOrderByTopoType(level0CommInfo.localRankSize, topoType_, nicList);
2316 :
2317 : // 空拷贝用于后续操作附着
2318 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
2319 0 : for (u32 ringIndex = 0; ringIndex < ringNum; ringIndex++) {
2320 0 : std::vector<Slice> singleRingSliceZero = multRingsSliceZero[ringIndex];
2321 0 : CHK_PRT_RET(
2322 : singleRingSliceZero.empty(), HCCL_ERROR("[CollCommExecutor][MultiRingScatter]singleRingSliceZero is empty"),
2323 : HCCL_E_INTERNAL);
2324 :
2325 : // 生成userMemIn_上对应的slices
2326 0 : std::vector<Slice> userMemInputSlices;
2327 0 : CHK_RET(
2328 : CalUserMemSlices(dataType, opInfo, singleRingSliceZero, ringIndex, multiRingsOrder, userMemInputSlices));
2329 0 : std::vector<u32> rankOrder;
2330 0 : CHK_RET(GetRankOrder(multiRingsOrder, ringIndex, rankOrder));
2331 0 : SubCommInfo level0RingCommInfo = GetSubCommInfo(COMM_LEVEL0, ringIndex);
2332 0 : u32 rankSize = level0RingCommInfo.localRankSize;
2333 :
2334 0 : std::vector<Stream> subStreamsInOneRing;
2335 0 : std::vector<std::shared_ptr<LocalNotify>> mainSignalsInOneRing;
2336 0 : std::vector<std::shared_ptr<LocalNotify>> subSignalsInOneRing;
2337 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
2338 0 : if (opInfo == nullptr) {
2339 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
2340 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING in COMM_LEVEL0", __func__);
2341 0 : CHK_SMART_PTR_NULL(tempAlg);
2342 0 : } else if (opInfo->inputAddr != nullptr) {
2343 0 : CHK_RET(
2344 : GetSubStreamInfoOnOneRing(ringIndex, subStreamsInOneRing, mainSignalsInOneRing, subSignalsInOneRing));
2345 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
2346 0 : TemplateType::TEMPLATE_SCATTER_RING_CONCURRENT_DIRECT, dispatcher_);
2347 0 : HCCL_CONFIG_INFO(
2348 : HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING_CONCURRENT_DIRECT in COMM_LEVEL0", __func__);
2349 0 : CHK_SMART_PTR_NULL(tempAlg);
2350 0 : CHK_RET(tempAlg->Prepare(
2351 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing, mainSignalsInOneRing,
2352 : subSignalsInOneRing, rankOrder, userMemInputSlices));
2353 : } else {
2354 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
2355 0 : TemplateType::TEMPLATE_SCATTER_RING_DIRECT, dispatcher_);
2356 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
2357 0 : CHK_SMART_PTR_NULL(tempAlg);
2358 0 : CHK_RET(tempAlg->Prepare(
2359 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, rankOrder, userMemInputSlices));
2360 : }
2361 :
2362 0 : if (ringIndex != (ringNum - 1)) {
2363 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
2364 0 : ret = StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
2365 0 : .StreamActive(algResResp_->slaveStreams[ringIndex].ptr(), stream.ptr());
2366 0 : CHK_PRT_RET(
2367 : ret != HCCL_SUCCESS,
2368 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]stream[%u],active stream failed", ringIndex), ret);
2369 : }
2370 : }
2371 :
2372 0 : u32 rootRank = 0;
2373 0 : ret = GetRankByUserRank(COMM_LEVEL0, ringIndex, root, rootRank);
2374 0 : CHK_PRT_RET(
2375 : ret == HCCL_E_PARA,
2376 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]invalid root [%u] to get userrank", root), ret);
2377 :
2378 0 : if (ret == HCCL_SUCCESS) {
2379 0 : if (ringIndex != (ringNum - 1)) { // 0~ringNum-2的环
2380 0 : ret = LocalNotify::Wait(
2381 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesAux[ringIndex],
2382 : PROF_STAGE_0);
2383 0 : CHK_PRT_RET(
2384 : ret != HCCL_SUCCESS,
2385 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]in stream[%u] wait failed", ringIndex), ret);
2386 :
2387 0 : ret = tempAlg->Prepare(
2388 0 : inputMem, inputMem, outputMem, count, dataType, algResResp_->slaveStreams[ringIndex],
2389 0 : HCCL_REDUCE_RESERVED, rootRank, singleRingSliceZero, baseOffset, ringNics[ringIndex]);
2390 0 : CHK_PRT_RET(
2391 : ret != HCCL_SUCCESS,
2392 : HCCL_ERROR(
2393 : "[CollCommExecutor][MultiRingScatter]stream[%u],scatter(ring) prepare failed, "
2394 : "return[%d]",
2395 : ringIndex, ret),
2396 : ret);
2397 :
2398 0 : ret = tempAlg->RegisterProfiler(
2399 0 : ((ringIndex + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
2400 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
2401 0 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, algResResp_->slaveStreams[ringIndex]);
2402 0 : CHK_PRT_RET(
2403 : ret != HCCL_SUCCESS,
2404 : HCCL_ERROR(
2405 : "[CollCommExecutor][MultiRingScatter]stream[%u], scatter(ring) register profiler "
2406 : "failed,return[%d]",
2407 : ringIndex, ret),
2408 : ret);
2409 :
2410 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
2411 0 : CHK_PRT_RET(
2412 : ret != HCCL_SUCCESS,
2413 : HCCL_ERROR(
2414 : "[CollCommExecutor][MultiRingScatter]stream[%u],scatter(ring) run failed, "
2415 : "return[%d]",
2416 : ringIndex, ret),
2417 : ret);
2418 :
2419 0 : ret = LocalNotify::Post(
2420 0 : algResResp_->slaveStreams[ringIndex], dispatcher_, algResResp_->notifiesMain[ringIndex],
2421 : PROF_STAGE_0);
2422 0 : CHK_PRT_RET(
2423 : ret != HCCL_SUCCESS,
2424 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]stream[%u] record failed", ringIndex), ret);
2425 : /* 主环record启动从环 */
2426 0 : ret = LocalNotify::Post(stream, dispatcher_, algResResp_->notifiesAux[ringIndex], PROF_STAGE_0);
2427 0 : CHK_PRT_RET(
2428 : ret != HCCL_SUCCESS,
2429 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]stream[%u] record failed", ringIndex), ret);
2430 : } else { // 主环
2431 0 : std::unique_ptr<AlgTemplateBase> tempAlg;
2432 0 : if (opInfo == nullptr) {
2433 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
2434 0 : TemplateType::TEMPLATE_SCATTER_RING, dispatcher_);
2435 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING in COMM_LEVEL0", __func__);
2436 0 : CHK_SMART_PTR_NULL(tempAlg);
2437 0 : } else if (opInfo->inputAddr != nullptr) {
2438 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
2439 0 : TemplateType::TEMPLATE_SCATTER_RING_CONCURRENT_DIRECT, dispatcher_);
2440 0 : HCCL_CONFIG_INFO(
2441 : HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING_CONCURRENT_DIRECT in COMM_LEVEL0",
2442 : __func__);
2443 0 : CHK_SMART_PTR_NULL(tempAlg);
2444 0 : CHK_RET(tempAlg->Prepare(
2445 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, subStreamsInOneRing,
2446 : mainSignalsInOneRing, subSignalsInOneRing, rankOrder, userMemInputSlices));
2447 : } else {
2448 0 : tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
2449 0 : TemplateType::TEMPLATE_SCATTER_RING_DIRECT, dispatcher_);
2450 0 : HCCL_CONFIG_INFO(
2451 : HCCL_ALG, "[%s][KernelRun] Run TEMPLATE_SCATTER_RING_DIRECT in COMM_LEVEL0", __func__);
2452 0 : CHK_SMART_PTR_NULL(tempAlg);
2453 0 : CHK_RET(tempAlg->Prepare(
2454 : const_cast<HcomCollOpInfo*>(opInfo), topoAttr_.userRank, rankOrder, userMemInputSlices));
2455 : }
2456 :
2457 0 : ret = tempAlg->Prepare(
2458 : inputMem, inputMem, outputMem, count, dataType, stream, HCCL_REDUCE_RESERVED, rootRank,
2459 0 : singleRingSliceZero, baseOffset, ringNics[ringIndex]);
2460 0 : CHK_PRT_RET(
2461 : ret != HCCL_SUCCESS,
2462 : HCCL_ERROR(
2463 : "[CollCommExecutor][MultiRingScatter]stream[%u],scatter(ring) prepare failed, "
2464 : "return[%d]",
2465 : ringIndex, ret),
2466 : ret);
2467 0 : ret = tempAlg->RegisterProfiler(
2468 0 : ((ringIndex + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID)
2469 0 : + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0RingCommInfo.localRank,
2470 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, stream);
2471 0 : CHK_PRT_RET(
2472 : ret != HCCL_SUCCESS,
2473 : HCCL_ERROR(
2474 : "[CollCommExecutor][MultiRingScatter]stream[%u], scatter(ring) register profiler "
2475 : "failed,return[%d]",
2476 : ringIndex, ret),
2477 : ret);
2478 :
2479 0 : ret = RunTemplate(tempAlg, level0RingCommInfo);
2480 0 : CHK_PRT_RET(
2481 : ret != HCCL_SUCCESS,
2482 : HCCL_ERROR(
2483 : "[CollCommExecutor][MultiRingScatter]stream[%u],scatter(ring) run failed, "
2484 : "return[%d]",
2485 : ringIndex, ret),
2486 : ret);
2487 :
2488 0 : for (u32 ring = 0; ring < (ringNum - 1); ring++) {
2489 : /* 等待executor执行完毕 , 当前环没有分配数据,跳过此环处理,继续下一个环 */
2490 0 : ret = LocalNotify::Wait(stream, dispatcher_, algResResp_->notifiesMain[ring], PROF_STAGE_0);
2491 0 : CHK_PRT_RET(
2492 : ret != HCCL_SUCCESS,
2493 : HCCL_ERROR("[CollCommExecutor][MultiRingScatter]stream[%u] wait failed", ring), ret);
2494 : }
2495 0 : }
2496 : }
2497 0 : }
2498 :
2499 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem, outputMem, stream, dispatcher_));
2500 0 : return HCCL_SUCCESS;
2501 0 : }
2502 :
2503 50 : HcclResult CollCommExecutor::SetRingNics(const std::string& tag, const std::vector<std::vector<u32>>& ringNics)
2504 : {
2505 50 : std::unique_lock<std::mutex> lock(ringNicListLock_);
2506 50 : ringNicList_[tag] = ringNics;
2507 50 : return HCCL_SUCCESS;
2508 50 : }
2509 3 : HcclResult CollCommExecutor::GetRingNics(const std::string& tag, std::vector<std::vector<u32>>& ringNics)
2510 : {
2511 3 : std::unique_lock<std::mutex> lock(ringNicListLock_);
2512 3 : auto iterRingNic = ringNicList_.find(tag);
2513 3 : if (iterRingNic == ringNicList_.end()) {
2514 9 : ringNics = {{0, 1, 2, 3, 4, 5, 6, 7}};
2515 : } else {
2516 0 : ringNics = iterRingNic->second;
2517 : }
2518 3 : return HCCL_SUCCESS;
2519 9 : }
2520 50 : HcclResult CollCommExecutor::SetNicSendSize(const std::string& tag, std::vector<u64>& sizeList)
2521 : {
2522 50 : std::unique_lock<std::mutex> lock(nicSendSizeListLock_);
2523 50 : nicSendSizeList_[tag] = sizeList;
2524 50 : return HCCL_SUCCESS;
2525 50 : }
2526 16 : HcclResult CollCommExecutor::PrepareLevel1CommInfo(
2527 : u32& segmentIdx, u32& commIndex, u64& hdSize, const SubCommInfo& commInfo,
2528 : const std::vector<std::vector<Slice>>& multRingsSliceZero, const std::string& tag)
2529 : {
2530 16 : segmentIdx = topoAttr_.devicePhyId;
2531 16 : commIndex = topoAttr_.devicePhyId;
2532 16 : CHK_PRT_RET(multRingsSliceZero.empty(), HCCL_ERROR("[Prepare][Level1CommInfo]slice map is empty"), HCCL_E_PARA);
2533 16 : if (multRingsSliceZero.size() > 1) {
2534 : std::vector<u32>::const_iterator iterNic
2535 16 : = std::find(topoAttr_.nicList.begin(), topoAttr_.nicList.end(), topoAttr_.devicePhyId);
2536 16 : if (iterNic != topoAttr_.nicList.end()) { // 如果当前rank为通信网口
2537 16 : u32 nicIdx = std::distance(topoAttr_.nicList.begin(), iterNic);
2538 16 : std::unique_lock<std::mutex> lock(nicSendSizeListLock_);
2539 16 : auto iter = nicSendSizeList_.find(tag);
2540 16 : CHK_PRT_RET(
2541 : iter == nicSendSizeList_.end(),
2542 : HCCL_ERROR(
2543 : "[Prepare][Level1CommInfo]find tag[%s] in "
2544 : "nicSendSizeList_ failed",
2545 : tag.c_str()),
2546 : HCCL_E_INTERNAL);
2547 16 : CHK_PRT_RET(
2548 : nicIdx >= iter->second.size(),
2549 : HCCL_ERROR(
2550 : "[Prepare][Level1CommInfo]tag[%s] nicIdx[%u] "
2551 : "invalid, expect less than %zu",
2552 : tag.c_str(), nicIdx, iter->second.size()),
2553 : HCCL_E_INTERNAL);
2554 16 : hdSize = iter->second[nicIdx]; // 通过nicSendSizeList_得到该网口传输数据量
2555 16 : u32 ringRanks = multRingsSliceZero[0].size(); // 获取单个 ring 上设备的数量
2556 16 : segmentIdx = ringRanks / topoAttr_.nicList.size() * nicIdx; // 通过网口位置得到该网口传输数据的起始位置
2557 16 : if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
2558 16 : segmentIdx = commInfo.localRank;
2559 16 : hdSize = iter->second[segmentIdx];
2560 16 : commIndex = segmentIdx;
2561 : }
2562 16 : } else { // 如果当前rank不是通信网口,则不发送数据
2563 0 : hdSize = 0;
2564 : }
2565 0 : } else if (multRingsSliceZero.size() == 1) {
2566 0 : segmentIdx = commInfo.localRank; // 针对0、4device下
2567 0 : CHK_PRT_RET(
2568 : segmentIdx >= multRingsSliceZero[0].size(),
2569 : HCCL_ERROR(
2570 : "[Prepare][Level1CommInfo]index is out of "
2571 : "range. Idx[%u] Slice size[%zu]",
2572 : segmentIdx, multRingsSliceZero[0].size()),
2573 : HCCL_E_PARA);
2574 0 : hdSize = multRingsSliceZero[0][segmentIdx].size;
2575 0 : commIndex = segmentIdx;
2576 : } else {
2577 0 : return HCCL_E_PARA;
2578 : }
2579 16 : HCCL_INFO(
2580 : "[CollCommExecutor][PrepareLevel1CommInfo]userRank[%u] segmentIdx[%u] commIndex[%u] hdSize[%llu]",
2581 : topoAttr_.userRank, segmentIdx, commIndex, hdSize);
2582 16 : return HCCL_SUCCESS;
2583 : }
2584 :
2585 : /* ↓↓ ====================== 用于ZerocopyExecutor ====================== ↓↓ */
2586 0 : HcclResult CollCommExecutor::CalcIntraServerDataSlicesDiscontinuous(
2587 : const OpParam& param, const ExecMem& execMem, u32 level0RankSize, u32 level1RankSize, u32 level2RankSize,
2588 : std::vector<Slice>& dataSegsSlice)
2589 : {
2590 0 : u32 perDataSize = 0;
2591 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
2592 :
2593 0 : u64 level0Count = execMem.count * level0RankSize;
2594 0 : u64 level0StrideCount = param.DataDes.strideCount * level0RankSize;
2595 0 : u64 sliceSize = perDataSize * execMem.count;
2596 0 : u64 strideSize = perDataSize * ((level0StrideCount != 0) ? level0StrideCount : level0Count);
2597 0 : dataSegsSlice.resize(topoAttr_.userRankSize);
2598 0 : for (u32 i = 0; i < level0RankSize; i++) {
2599 0 : for (u32 j = 0; j < level1RankSize * level2RankSize; j++) {
2600 0 : u32 index = i * level1RankSize * level2RankSize + j;
2601 0 : dataSegsSlice[index].size = sliceSize;
2602 0 : dataSegsSlice[index].offset = j * strideSize + i * sliceSize;
2603 : }
2604 : }
2605 0 : return HCCL_SUCCESS;
2606 : }
2607 :
2608 0 : HcclResult CollCommExecutor::CalcIntraServerDataSlicesContinuous(
2609 : const OpParam& param, const ExecMem& execMem, u32 level0RankSize, u32 level1RankSize, u32 level2RankSize,
2610 : std::vector<Slice>& dataSegsSlice)
2611 : {
2612 0 : u32 perDataSize = 0;
2613 0 : CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
2614 :
2615 0 : u64 level0Count = execMem.count * level1RankSize * level2RankSize;
2616 0 : u64 level0StrideCount = param.DataDes.strideCount * level1RankSize * level1RankSize;
2617 0 : u64 sliceSize = perDataSize * level0Count;
2618 0 : u64 strideSize = perDataSize * ((level0StrideCount != 0) ? level0StrideCount : level0Count);
2619 0 : dataSegsSlice.resize(level0RankSize);
2620 0 : for (u32 i = 0; i < level0RankSize; i++) {
2621 0 : dataSegsSlice[i].size = sliceSize;
2622 0 : dataSegsSlice[i].offset = (i * strideSize);
2623 : }
2624 0 : return HCCL_SUCCESS;
2625 : }
2626 :
2627 0 : void CollCommExecutor::CalcLevel1DataSlices(
2628 : u64 sliceSize, u32 level1RankSize, u32 level2RankSize, std::vector<Slice>& level1DataSegsSlice)
2629 : {
2630 0 : level1DataSegsSlice.resize(level1RankSize);
2631 0 : u64 level1SliceSize = sliceSize * level2RankSize;
2632 0 : for (u32 i = 0; i < level1RankSize; i++) {
2633 0 : level1DataSegsSlice[i].size = level1SliceSize;
2634 0 : level1DataSegsSlice[i].offset = i * level1SliceSize;
2635 : }
2636 0 : }
2637 :
2638 0 : HcclResult CollCommExecutor::GetCommRankInfoNormal(
2639 : u32& level0Rank, u32& level0RankSize, u32& level1Rank, u32& level1RankSize, u32& level2Rank, u32& level2RankSize,
2640 : bool isAHCAlgo)
2641 : {
2642 : // 获取通信域信息
2643 : // ==> Level0
2644 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
2645 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
2646 0 : level0Rank = level0CommInfo.localRank;
2647 0 : level0RankSize = level0CommInfo.localRankSize;
2648 : // ==> Level1
2649 0 : CommPlane commPlaneLevel1 = isAHCAlgo ? COMM_LEVEL1_AHC : COMM_LEVEL1;
2650 0 : CHK_RET(CheckCommSize(commPlaneLevel1, level0Rank + 1));
2651 0 : SubCommInfo level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0Rank);
2652 0 : level1Rank = level1CommInfo.localRank;
2653 0 : level1RankSize = level1CommInfo.localRankSize;
2654 : // ==> Level2
2655 0 : if (isAHCAlgo) { // bypass level2
2656 0 : level2Rank = 0;
2657 0 : level2RankSize = 1;
2658 : } else {
2659 0 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
2660 0 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
2661 0 : level2Rank = level2CommInfo.localRank;
2662 0 : level2RankSize = level2CommInfo.localRankSize;
2663 0 : }
2664 0 : return HCCL_SUCCESS;
2665 0 : }
2666 : /* ↑↑ ====================== 用于ZerocopyExecutor ======================= ↑↑ */
2667 :
2668 : /* ↓↓ ====================== 用于ExchangeExecutor ====================== ↓↓ */
2669 0 : HcclResult CollCommExecutor::CalExchangeRemoteRankForReduceScatter(u32& remoteRankSend, u32& remoteRankRecv)
2670 : {
2671 0 : u32 userRank = topoAttr_.userRank;
2672 0 : u32 userRankSize = topoAttr_.userRankSize;
2673 0 : u32 l2Size = topoAttr_.superPodNum;
2674 0 : CHK_PRT_RET(
2675 : l2Size == 0, HCCL_ERROR("[CollCommExecutor][CalExchangeRemoteRank] invalid rank size, level2RankSize is 0"),
2676 : HCCL_E_PARA);
2677 0 : u32 l1Size = topoAttr_.serverNum / l2Size;
2678 0 : CHK_PRT_RET(
2679 : l1Size == 0, HCCL_ERROR("[CollCommExecutor][CalExchangeRemoteRank] invalid rank size, level1RankSize is 0"),
2680 : HCCL_E_PARA);
2681 0 : u32 l0Size = userRankSize / l1Size / l2Size;
2682 0 : u32 l0Index = userRank % l0Size;
2683 0 : u32 l1ServerIndex = userRank % (l0Size * l1Size) / l0Size;
2684 0 : u32 l2ServerIndex = userRank / l0Size / l1Size;
2685 :
2686 : // 计算本端将要发送数据的目标rank
2687 0 : remoteRankSend = l0Index * l2Size * l1Size + l1ServerIndex * l2Size + l2ServerIndex;
2688 :
2689 : // 计算本端将要接收数据的目标rank
2690 0 : u32 r0 = userRank / (l1Size * l2Size);
2691 0 : u32 r1 = userRank % (l1Size * l2Size) / l2Size;
2692 0 : u32 r2 = userRank % (l1Size * l2Size) % l2Size;
2693 0 : remoteRankRecv = r2 * l1Size * l0Size + r1 * l0Size + r0;
2694 0 : return HCCL_SUCCESS;
2695 : }
2696 :
2697 0 : HcclResult CollCommExecutor::GetTransportForExchange(u32 remoteUserRank, LINK& targetLink)
2698 : {
2699 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
2700 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
2701 0 : u32 level0RankSize = level0CommInfo.localRankSize;
2702 0 : CommPlane commPlane = IsLevel0Neighbor(remoteUserRank, level0RankSize) ? COMM_LEVEL0 : COMM_COMBINE_ORDER;
2703 :
2704 0 : CHK_PRT_RET(
2705 : COMM_INDEX_0 >= algResResp_->opTransportResponse[commPlane].size(),
2706 : HCCL_ERROR(
2707 : "[%s] commIndex[%u] is larger than opTransportResponse size[%zu]", __func__, COMM_INDEX_0,
2708 : algResResp_->opTransportResponse[commPlane].size()),
2709 : HCCL_E_PARA);
2710 0 : SingleSubCommTransport& commCombined = algResResp_->opTransportResponse[commPlane][COMM_INDEX_0];
2711 :
2712 0 : CHK_PRT_RET(
2713 : commCombined.userRank2subCommRank.count(remoteUserRank) == 0,
2714 : HCCL_ERROR("[%s] remoteUserRank[%u] not found in userRank2subCommRank map.", __func__, remoteUserRank),
2715 : HCCL_E_PARA);
2716 :
2717 0 : u32 remoteRank = commCombined.userRank2subCommRank[remoteUserRank];
2718 0 : CHK_PRT_RET(
2719 : remoteRank >= commCombined.links.size(),
2720 : HCCL_ERROR(
2721 : "[%s] remoteUserRank[%u], get remoteRank[%u], the size of combinedComm links is [%zu]", __func__,
2722 : remoteUserRank, remoteRank, commCombined.links.size()),
2723 : HCCL_E_PARA);
2724 0 : targetLink = commCombined.links[remoteRank];
2725 0 : CHK_PTR_NULL(targetLink);
2726 :
2727 0 : return HCCL_SUCCESS;
2728 0 : }
2729 :
2730 0 : bool CollCommExecutor::IsLevel0Neighbor(u32 remoteRank, u32 level0RankSize)
2731 : {
2732 0 : CHK_PRT_RET(level0RankSize == 0, HCCL_ERROR("[%s] invalid rank size, Level0RankSize is 0", __func__), HCCL_E_PARA);
2733 0 : bool isSameServer = remoteRank / level0RankSize == topoAttr_.userRank / level0RankSize;
2734 0 : bool isLeftNeighbor = (topoAttr_.userRank + 1) % level0RankSize == remoteRank % level0RankSize;
2735 0 : bool isRightNeighbor = (topoAttr_.userRank + level0RankSize - 1) % level0RankSize == remoteRank % level0RankSize;
2736 0 : return isSameServer && (isLeftNeighbor || isRightNeighbor);
2737 : }
2738 : /* ↑↑ ====================== 用于ExchangeExecutor ======================= ↑↑ */
2739 :
2740 0 : HcclResult CollCommExecutor::GetAdjInfo(AlgResourceResponse& algRes, AdjInfo& adjInfo)
2741 : {
2742 0 : HCCL_INFO("[nslbdp] Entry GetAdjInfo.");
2743 0 : algResResp_ = &algRes;
2744 0 : SubCommInfo level1CommInfo = {};
2745 0 : AdjInfo nslbAdjInfo = {};
2746 0 : if (Getlevel1CommRank(level1CommInfo) != HCCL_SUCCESS) {
2747 0 : HCCL_INFO("[nslbdp-GetAdjInfo] Getlevel1CommRank is NULL.");
2748 0 : return HCCL_SUCCESS;
2749 : }
2750 0 : u32 localRank = level1CommInfo.localRank;
2751 0 : u32 localRankSize = level1CommInfo.localRankSize;
2752 0 : HCCL_INFO("[nslbdp-GetAdjInfo] level1CommInfo.localRank = [%u] localRankSize = [%u].", localRank, localRankSize);
2753 :
2754 0 : if (localRankSize == 1) {
2755 0 : return HCCL_SUCCESS;
2756 : }
2757 :
2758 0 : if (level1CommInfo.links.size() < localRankSize) {
2759 0 : return HCCL_SUCCESS;
2760 : }
2761 :
2762 0 : std::unique_ptr<AlgTemplateBase> nslbdp_levelTempAlg;
2763 0 : if (SelectTempAlg(nslbdp_levelTempAlg, localRankSize) != HCCL_SUCCESS) {
2764 0 : HCCL_INFO("[nslbdp-GetAdjInfo] SelectTempAlg is unsuccessful.");
2765 0 : return HCCL_SUCCESS;
2766 : }
2767 0 : if (nslbdp_levelTempAlg == nullptr) {
2768 0 : return HCCL_SUCCESS;
2769 : }
2770 0 : CHK_RET(nslbdp_levelTempAlg->GetNslbAdjInfo(localRank, localRankSize, level1CommInfo.links, nslbAdjInfo));
2771 :
2772 0 : adjInfo.dstRankNum = nslbAdjInfo.dstRankNum;
2773 0 : HCCL_INFO("[nslbdp-GetAdjInfo] adjInfo.dstRankNum[%u].", adjInfo.dstRankNum);
2774 :
2775 0 : for (size_t i = 0; i < nslbAdjInfo.nsAdjInfo.size(); i++) {
2776 0 : NslbDpAdjInfo dpAdjInfo = {};
2777 0 : dpAdjInfo.dstLocalRankId = nslbAdjInfo.nsAdjInfo[i].dstLocalRankId;
2778 0 : dpAdjInfo.phaseId = nslbAdjInfo.nsAdjInfo[i].phaseId;
2779 0 : dpAdjInfo.rev = 0;
2780 0 : adjInfo.nsAdjInfo.push_back(dpAdjInfo);
2781 0 : HCCL_INFO(
2782 : "[nslbdp]GetAdjInfo dstLocalRankId[%u], phaseId[%u].", nslbAdjInfo.nsAdjInfo[i].dstLocalRankId,
2783 : nslbAdjInfo.nsAdjInfo[i].phaseId);
2784 : }
2785 0 : return HCCL_SUCCESS;
2786 0 : }
2787 : } // namespace hccl
|