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