Line data Source code
1 : /**
2 : * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 :
11 : #include "coll_reduce_scatter_ring_for_910_93_executor.h"
12 : #include <numeric>
13 : #include "alg_template_register.h"
14 :
15 : namespace hccl {
16 :
17 3 : CollReduceScatterRingFor91093Executor::CollReduceScatterRingFor91093Executor(const HcclDispatcher dispatcher,
18 3 : std::unique_ptr<TopoMatcher> &topoMatcher)
19 3 : : CollReduceScatterExecutor(dispatcher, topoMatcher)
20 : {
21 3 : DMAReduceFlag_ = (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE);
22 3 : desc_.deterministic = 1;
23 3 : desc_.level1SupportedAlgos = {
24 : AlgTypeLevel1::ALG_LEVEL1_NHR,
25 : AlgTypeLevel1::ALG_LEVEL1_NB,
26 : AlgTypeLevel1::ALG_LEVEL1_RING,
27 : AlgTypeLevel1::ALG_LEVEL1_AHC,
28 : AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE
29 3 : };
30 3 : desc_.level2SupportedAlgos = {
31 : AlgTypeLevel2::ALG_LEVEL2_NHR,
32 : AlgTypeLevel2::ALG_LEVEL2_NB,
33 : AlgTypeLevel2::ALG_LEVEL2_RING
34 3 : };
35 3 : }
36 :
37 16 : bool CollReduceScatterRingFor91093Executor::IsUnifiedMarch(const OpParam ¶m) const
38 : {
39 16 : return IsSupportUnifiedMarch(param, topoType_, topoAttr_.serverNum, topoAttr_.superPodNum);
40 : }
41 :
42 6 : u64 CollReduceScatterRingFor91093Executor::CalcTotalCount(const OpParam ¶m) const
43 : {
44 6 : if (isReduceScatterV_) {
45 0 : const auto *counts = static_cast<const u64 *>(param.VDataDes.counts);
46 0 : return std::accumulate(counts, counts + topoAttr_.userRankSize, 0ULL);
47 : }
48 6 : return param.DataDes.count * topoAttr_.userRankSize;
49 : }
50 :
51 6 : void CollReduceScatterRingFor91093Executor::ParseParam(const OpParam& param)
52 : {
53 6 : tag_ = param.tag;
54 :
55 6 : const HcclDataType dataType = param.GetDataType();
56 : // 是否需要scratch memory
57 16 : if ((workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) &&
58 6 : isSupportSDMAReduce_ && IsSupportRDMAReduce(dataType, param.reduceType)) {
59 4 : scratchMemFlag_ = false;
60 : } else {
61 2 : scratchMemFlag_ = true;
62 : }
63 :
64 6 : HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][ParseParam] tag[%s] isSupportSDMAReduce_[%u] "
65 : "scratchMemFlag_[%u] workflowMode_[%u]", tag_.c_str(), isSupportSDMAReduce_, scratchMemFlag_, workflowMode_);
66 :
67 : // 记录图模式总数据量
68 6 : totalSize_ = CalcTotalCount(param) * SIZE_TABLE[dataType];
69 6 : aicpuUnfoldMode_ = param.aicpuUnfoldMode;
70 6 : isZeroCopy_ = param.isZeroCopy;
71 6 : }
72 :
73 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcScratchMemSize(u64& scratchMemSize)
74 : {
75 3 : if (scratchMemFlag_) {
76 1 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
77 0 : scratchMemSize = inCCLbufferSize_;
78 : } else {
79 1 : scratchMemSize = totalSize_;
80 : }
81 : } else {
82 2 : scratchMemSize = 0U;
83 : }
84 3 : HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcScratchMemSize] tag[%s] scratchMemSize[%llu] "
85 : "scratchMemFlag_[%u] workflowMode_[%u]", tag_.c_str(), scratchMemSize, scratchMemFlag_, workflowMode_);
86 3 : return HCCL_SUCCESS;
87 : }
88 :
89 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcStreamNum(u32& streamNum)
90 : {
91 3 : u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
92 : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
93 3 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
94 2 : totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
95 : }
96 3 : streamNum = totalStreamNum - 1;
97 3 : HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcStreamNum] tag[%s] streamNum[%u]",
98 : tag_.c_str(), streamNum);
99 3 : return HCCL_SUCCESS;
100 : }
101 :
102 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
103 : {
104 3 : TransportMemType inputType = TransportMemType::RESERVED;
105 3 : TransportMemType outputType = TransportMemType::RESERVED;
106 3 : CHK_RET(CalcTransportMemType(inputType, outputType));
107 3 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
108 3 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
109 3 : CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
110 3 : return HCCL_SUCCESS;
111 : }
112 :
113 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcTransportMemType(TransportMemType &inputType,
114 : TransportMemType &outputType)
115 : {
116 3 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
117 2 : inputType = TransportMemType::CCL_INPUT;
118 2 : if (scratchMemFlag_) {
119 0 : outputType = TransportMemType::SCRATCH;
120 : } else {
121 2 : outputType = TransportMemType::CCL_OUTPUT;
122 : }
123 : } else {
124 1 : inputType = TransportMemType::PARAM_INPUT;
125 1 : if (scratchMemFlag_) {
126 1 : outputType = TransportMemType::SCRATCH;
127 : } else {
128 0 : outputType = TransportMemType::PARAM_OUTPUT;
129 : }
130 : }
131 3 : HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
132 : tag_.c_str(), inputType, outputType);
133 3 : return HCCL_SUCCESS;
134 : }
135 :
136 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcLevel0CommInfo(TransportMemType inputType,
137 : TransportMemType outputType,
138 : std::vector<LevelNSubCommTransport>& opTransport)
139 : {
140 3 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
141 3 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
142 3 : return HCCL_SUCCESS;
143 3 : }
144 :
145 3 : HcclResult CollReduceScatterRingFor91093Executor::CalcLevel2CommInfo(TransportMemType inputType,
146 : TransportMemType outputType,
147 : std::vector<LevelNSubCommTransport>& opTransport)
148 : {
149 3 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
150 3 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
151 0 : HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcLevel2CommInfo] select AHC bypass level2 comm calculate");
152 0 : return HCCL_SUCCESS;
153 : }
154 :
155 3 : CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX);
156 3 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
157 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
158 0 : HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
159 3 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
160 0 : commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
161 0 : HCCL_INFO("[%s]Calc NBCommInfo", __func__);
162 : } else {
163 3 : commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
164 3 : HCCL_INFO("[%s]Calc RingCommInfo", __func__);
165 : }
166 3 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
167 3 : return HCCL_SUCCESS;
168 3 : }
169 :
170 2 : u64 CollReduceScatterRingFor91093Executor::CalcLoopMaxCount(const u32 unitSize)
171 : {
172 : // 中转内存单次最多能够接受的output count,放开ranksize限制
173 2 : u64 maxCountPerLoop = inCCLbufferSize_ / topoAttr_.userRankSize / HCCL_MIN_SLICE_ALIGN
174 2 : * HCCL_MIN_SLICE_ALIGN / unitSize;
175 2 : return maxCountPerLoop;
176 : }
177 :
178 16 : bool CollReduceScatterRingFor91093Executor::IsHugeData(const u64 curSize, OpParam *param)
179 : {
180 : u32 level2RankSize;
181 16 : if ((algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
182 16 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE)) {
183 : //AHC非对称场景下没有L2
184 0 : level2RankSize =1;
185 : } else {
186 : // 多QP哈希散列开启且RDMA通信下,强制刷新子图
187 : // 这里如果CheckCommSize返回ERROR,相当于HugeData true,防止GetSubCommInfo越界
188 16 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
189 16 : SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
190 16 : level2RankSize = level2CommInfo.localRankSize;
191 16 : }
192 :
193 16 : const u64 TBE_REDUCE_MAX_COUNT = INT32_MAX;
194 :
195 16 : u64 curCount = curSize / SIZE_TABLE[param->DataDes.dataType];
196 16 : bool issupportRDMAInlineReduce = IsSupportRDMAReduce(param->DataDes.dataType, param->reduceType);
197 16 : bool hugeData =
198 16 : (curSize * level2RankSize / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE) ||
199 16 : (curSize > SDMA_SEND_MAX_SIZE) ||
200 48 : ((!isSupportSDMAReduce_) && (curCount > TBE_REDUCE_MAX_COUNT)) ||
201 16 : ((!issupportRDMAInlineReduce) && (curCount * level2RankSize / HCCL_INTERNODE_MAX_DATA_RATE > TBE_REDUCE_MAX_COUNT));
202 16 : return hugeData;
203 : }
204 :
205 0 : HcclResult CollReduceScatterRingFor91093Executor::RunIntraSeverReduceScatter(
206 : const std::string &tag, DeviceMem &inputMem, DeviceMem &outputMem,
207 : const u64 count, const HcclDataType &dataType, const HcclReduceOp &reductionOp,
208 : const std::vector<std::vector<Slice>> &multRingsSliceZero, const Stream &stream,
209 : s32 profStage, const u64 baseOffset, const HcomCollOpInfo *opInfo,
210 : const std::vector<std::vector<Slice>> &multRingsUserMemSlice, const bool disableDMAReduce)
211 : {
212 0 : CHK_RET(MultiRingReduceScatter(tag, inputMem, outputMem, count, dataType, reductionOp,
213 : multRingsSliceZero, stream, profStage, baseOffset, opInfo, multRingsUserMemSlice, logicalLevel0plane_));
214 0 : return HCCL_SUCCESS;
215 : }
216 :
217 33 : void CollReduceScatterRingFor91093Executor::FillMultiRingSlice(const ExecMem &execMem,
218 : const std::vector<std::vector<Slice>> &multiStreamSlice, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
219 : const u32 ringIndex, std::vector<Slice> &dataSlice)
220 : {
221 99 : for (u32 level0Idx = 0; level0Idx < sliceNum; level0Idx++) {
222 66 : Slice sliceTemp;
223 132 : for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
224 198 : for (u32 level1Idx = 0; level1Idx < level1RankSize; level1Idx++) {
225 132 : sliceTemp.size = multiStreamSlice[ringIndex][level0Idx].size;
226 132 : sliceTemp.offset = multiStreamSlice[ringIndex][level0Idx].offset +
227 264 : level1Idx * sliceNum * execMem.outputMem.size() +
228 132 : level2Idx * sliceNum * level1RankSize * execMem.outputMem.size();
229 132 : dataSlice.push_back(sliceTemp);
230 132 : HCCL_DEBUG("rank[%u] sliceTemp.size[%zu], sliceTemp.offset[%llu]", topoAttr_.userRank,
231 : sliceTemp.size, sliceTemp.offset);
232 : }
233 : }
234 : }
235 33 : }
236 :
237 17 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel0DataSegsSlice(const ExecMem &execMem,
238 : std::vector<std::vector<Slice>> &multiStreamSlice, const OpParam ¶m, u32 ringNum, u32 sliceNum,
239 : u32 level1RankSize, u32 level2RankSize, HcclDataType dataType, std::vector<std::vector<Slice>> &level0DataSegsSlice)
240 : {
241 17 : if (isReduceScatterV_) {
242 0 : return CalLevel0DataSegsSliceV(execMem, multiStreamSlice, param, ringNum, sliceNum, level1RankSize,
243 0 : level2RankSize, dataType, level0DataSegsSlice);
244 : }
245 17 : bool isInlineReduce = IsSupportSDMAReduce(execMem.inputMem.ptr(), execMem.scratchMem.ptr(), dataType,
246 17 : param.reduceType);
247 17 : bool useInlineReduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
248 17 : std::vector<Slice> dataSegsSlice; // 数据分成ranksize份,每份的起始偏移和大小
249 17 : multiStreamSlice = ReduceScatterRingSlicePrepare(ringNum, sliceNum, useInlineReduce, execMem.outputMem,
250 17 : dataSegsSlice, param.tag);
251 :
252 50 : for (u32 ringIndex = 0; ringIndex < multiStreamSlice.size(); ringIndex++) {
253 33 : std::vector<Slice> dataSlice;
254 33 : FillMultiRingSlice(execMem, multiStreamSlice, sliceNum, level1RankSize, level2RankSize, ringIndex, dataSlice);
255 33 : level0DataSegsSlice.push_back(dataSlice);
256 33 : }
257 17 : return HCCL_SUCCESS;
258 17 : }
259 :
260 17 : HcclResult CollReduceScatterRingFor91093Executor::CalUserMemDataSegsSlice(const ExecMem &execMem,
261 : const std::vector<std::vector<Slice>> &level0DataSegsSlice, const std::vector<std::vector<Slice>> &multiStreamSlice,
262 : const OpParam ¶m, u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
263 : u32 perDataSize, HcomCollOpInfo *opInfoPtr, bool disableDMAReduce,
264 : std::vector<std::vector<Slice>> &multRingsUserMemSlice)
265 : {
266 17 : if (isReduceScatterV_) {
267 0 : return CalUserMemDataSegsSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
268 0 : multRingsUserMemSlice);
269 : }
270 17 : CHK_PRT_RET(0 < param.DataDes.strideCount && param.DataDes.strideCount < param.DataDes.count,
271 : HCCL_ERROR("[CollReduceScatterRingFor91093Executor][KernelRun]strideCount[%llu] is smaller than opCount[%llu]",
272 : param.DataDes.strideCount, param.DataDes.count),
273 : HCCL_E_PARA);
274 17 : HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][KernelRun]strideCount[%llu], opCount[%llu]",
275 : param.DataDes.strideCount, param.DataDes.count);
276 :
277 17 : u32 level0RankSize = logicalLevel0CommInfo_.localRankSize;
278 17 : bool ARSFlag = topoMatcher_->GetARSFlag();
279 17 : bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
280 :
281 17 : if (opInfoPtr == nullptr &&
282 1 : (!((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING || ARSDoubleRing) &&
283 0 : (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB || disableDMAReduce)))) {
284 1 : multRingsUserMemSlice = level0DataSegsSlice;
285 : // 图模式,根据strideCount更新slice的offset
286 1 : if (param.DataDes.strideCount != 0) {
287 0 : CHK_RET(UpdateOffsetBasedOnStrideCount(param, multRingsUserMemSlice));
288 : }
289 1 : } else {
290 48 : for (u32 ringIndex = 0; ringIndex < level0DataSegsSlice.size(); ringIndex++) {
291 32 : std::vector<Slice> level1UserMemSlice;
292 160 : for (auto &cclSlice : level0DataSegsSlice[ringIndex]) {
293 128 : Slice tmpSlice;
294 128 : u64 count = (param.DataDes.strideCount == 0) ? param.DataDes.count : param.DataDes.strideCount;
295 128 : tmpSlice.size = cclSlice.size;
296 128 : CHK_PRT_RET(execMem.outputMem.size() == 0,
297 : HCCL_ERROR("[CollReduceScatterRingFor91093Executor][KernelRun]cclout memsize[%llu] is zero",
298 : execMem.outputMem.size()), HCCL_E_PARA);
299 128 : tmpSlice.offset = (cclSlice.offset / execMem.outputMem.size()) * count * perDataSize +
300 128 : multiStreamSlice[ringIndex][0].offset;
301 128 : level1UserMemSlice.push_back(tmpSlice);
302 128 : HCCL_DEBUG("rank[%u], ringIndex[%u], tmpSlice.offset=[%llu], size=[%llu]",
303 : topoAttr_.userRank, ringIndex, tmpSlice.offset, tmpSlice.size);
304 : }
305 32 : multRingsUserMemSlice.push_back(level1UserMemSlice);
306 32 : }
307 : }
308 17 : return HCCL_SUCCESS;
309 : }
310 :
311 16 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel1DataSegsSlice(const ExecMem &execMem, const OpParam ¶m,
312 : CommPlane commPlaneLevel, const u32 &commIndex, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
313 : u32 perDataSize, std::vector<Slice> &level1DataSegsSlice)
314 : {
315 16 : if (isReduceScatterV_) {
316 0 : return CalLevel1DataSegsSliceV(param, commPlaneLevel, commIndex, sliceNum, level1RankSize, level2RankSize,
317 0 : perDataSize, level1DataSegsSlice);
318 : }
319 48 : for (u32 i = 0; i < level1RankSize; i++) {
320 32 : Slice sliceTemp;
321 : u32 level1UserRank;
322 32 : CHK_RET(GetUserRankByRank(commPlaneLevel, commIndex, i, level1UserRank));
323 32 : if (level2RankSize <= 1) {
324 32 : sliceTemp.size = execMem.outputMem.size();
325 32 : sliceTemp.offset = level1UserRank * execMem.outputMem.size();
326 32 : level1DataSegsSlice.push_back(sliceTemp);
327 32 : HCCL_DEBUG("rank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]", topoAttr_.userRank, i,
328 : sliceTemp.offset, sliceTemp.size);
329 : } else {
330 0 : for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
331 0 : sliceTemp.size = execMem.outputMem.size();
332 0 : sliceTemp.offset = (level1UserRank % (level1RankSize * sliceNum)) * execMem.outputMem.size() +
333 0 : level2Idx * sliceNum * level1RankSize * execMem.outputMem.size();
334 0 : level1DataSegsSlice.push_back(sliceTemp);
335 0 : HCCL_DEBUG("rank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]", topoAttr_.userRank, i,
336 : sliceTemp.offset, sliceTemp.size);
337 : }
338 : }
339 : }
340 16 : return HCCL_SUCCESS;
341 : }
342 :
343 17 : HcclResult CollReduceScatterRingFor91093Executor::GetLevelCommInfo()
344 : {
345 17 : logicalLevel0plane_ = COMM_LEVEL0;
346 17 : CHK_RET(CheckCommSize(logicalLevel0plane_, COMM_INDEX_0 + 1));
347 17 : logicalLevel0CommInfo_ = GetSubCommInfo(logicalLevel0plane_, COMM_INDEX_0);
348 17 : u32 commIndex = logicalLevel0CommInfo_.localRank;
349 34 : bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
350 17 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
351 17 : logicalLevel1plane_ = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
352 17 : CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
353 17 : logicalLevel1CommInfo_ = GetSubCommInfo(logicalLevel1plane_, commIndex);
354 17 : return HCCL_SUCCESS;
355 : }
356 :
357 0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel2DataSegsSlice(const ExecMem &execMem, const OpParam ¶m,
358 : u32 level2RankSize, u32 perDataSize, std::vector<Slice> &level2DataSegsSlice)
359 : {
360 0 : if (isReduceScatterV_) {
361 0 : return CalLevel2DataSegsSliceV(param, level2RankSize, perDataSize, level2DataSegsSlice);
362 : }
363 0 : Slice sliceTemp;
364 0 : for (u32 i = 0; i < level2RankSize; i++) {
365 0 : sliceTemp.size = execMem.outputMem.size();
366 : u32 level2UserRank;
367 0 : CHK_RET(GetUserRankByRank(COMM_LEVEL2, COMM_INDEX_0, i, level2UserRank));
368 0 : sliceTemp.offset = level2UserRank * execMem.outputMem.size();
369 0 : level2DataSegsSlice.push_back(sliceTemp);
370 0 : HCCL_DEBUG("rank[%u], level2DataSegsSlice[%u].offset=%llu, size=[%llu], level2RankSize[%u]",
371 : topoAttr_.userRank, i, sliceTemp.offset, sliceTemp.size, level2RankSize);
372 : }
373 0 : return HCCL_SUCCESS;
374 : }
375 :
376 0 : void CollReduceScatterRingFor91093Executor::PrepareLevel0Slices(const OpParam ¶m, u32 sliceNum, u32 level1RankSize,
377 : u32 level1Index, u32 level2Index, u32 perDataSize, std::vector<Slice> &cclSegSlices)
378 : {
379 0 : const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
380 : // 根据counts和displace计算每个rank的数据范围
381 : // cclSlices里的offset是cclBuffer范围内的偏移,就地计算得出,不考虑displs
382 0 : const u32 level1Rank = level2Index * level1RankSize * sliceNum + level1Index * sliceNum;
383 0 : u64 displace = std::accumulate(counts, counts + level1Rank, 0ULL);
384 0 : for (auto rank = 0U; rank < sliceNum; ++rank) {
385 0 : const u32 idx = level1Rank + rank;
386 0 : Slice slice;
387 0 : slice.size = counts[idx] * perDataSize;
388 0 : slice.offset = displace * perDataSize;
389 0 : cclSegSlices.emplace_back(slice);
390 0 : displace += counts[idx];
391 : }
392 0 : }
393 :
394 0 : void CollReduceScatterRingFor91093Executor::PrepareLevel0UserSlices(const OpParam ¶m, u32 sliceNum,
395 : u32 level1RankSize, u32 level1Index, u32 level2Index, u32 perDataSize, std::vector<Slice> &userSegSlices)
396 : {
397 0 : const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
398 0 : const auto *displsPtr = static_cast<const u64*>(param.VDataDes.displs);
399 0 : const u32 level1Rank = level2Index * level1RankSize * sliceNum + level1Index * sliceNum;
400 : // 根据counts和displace计算每个rank的数据范围
401 : // userSlices里的offset是user input的偏移,使用传入的displs算得
402 0 : for (auto rank = 0U; rank < sliceNum; ++rank) {
403 0 : const u32 idx = level1Rank + rank;
404 0 : Slice slice;
405 0 : slice.size = counts[idx] * perDataSize;
406 0 : slice.offset = displsPtr[idx] * perDataSize;
407 0 : userSegSlices.emplace_back(std::move(slice));
408 : }
409 0 : }
410 :
411 0 : bool CollReduceScatterRingFor91093Executor::IsCceReduceAligned(const std::vector<Slice> &dataSlices) const
412 : {
413 0 : for (const auto &slice : dataSlices) {
414 0 : if (slice.size % CCE_REDUCE_ALIGN_SIZE != 0) {
415 0 : return false;
416 : }
417 : }
418 0 : return true;
419 : }
420 :
421 0 : HcclResult CollReduceScatterRingFor91093Executor::FillMultiRingSliceV(const ExecMem &execMem, const OpParam ¶m,
422 : u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
423 : std::vector<std::vector<Slice>> &level0DataSegsSlice, std::vector<std::vector<std::vector<Slice>>> &serverSlices,
424 : const Level0SlicesCalculator &calcLevel0Slices)
425 : {
426 0 : bool isInlineReduce = IsSupportSDMAReduce(execMem.inputMem.ptr(), execMem.scratchMem.ptr(), dataType,
427 0 : param.reduceType);
428 0 : bool useInlineReduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
429 0 : u32 perDataSize = 0;
430 0 : CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
431 0 : for (u32 i = 0; i < level2RankSize; i++) {
432 0 : for (u32 j = 0; j < level1RankSize; j++) {
433 0 : std::vector<Slice> dataSegsSlice; // 数据分成rank size份,每份的起始偏移和大小
434 0 : calcLevel0Slices(param, sliceNum, level1RankSize, j, i, perDataSize, dataSegsSlice);
435 :
436 0 : std::vector<std::vector<Slice>> multiStreamSlices;
437 : // 再将每个 slice 划分为 ringNum 份
438 0 : if (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) {
439 0 : if (useInlineReduce) {
440 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag);
441 0 : } else if (IsCceReduceAligned(dataSegsSlice)) {
442 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag);
443 : } else {
444 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, true);
445 : }
446 0 : } else if (ringNum == LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE) {
447 : // 双环场景,需要传入正确的 niclist (不涉及网口裁剪)
448 0 : if (useInlineReduce) {
449 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
450 0 : } else if (IsCceReduceAligned(dataSegsSlice)) {
451 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
452 : } else {
453 0 : multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, true, topoAttr_.nicList);
454 : }
455 : } else {
456 0 : multiStreamSlices.push_back(dataSegsSlice);
457 : }
458 0 : serverSlices.push_back(multiStreamSlices);
459 0 : }
460 : }
461 0 : level0DataSegsSlice.resize(ringNum);
462 0 : for (u32 level0Idx = 0; level0Idx < sliceNum; level0Idx++) {
463 0 : for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
464 0 : for (u32 level1Idx = 0; level1Idx < level1RankSize; level1Idx++) {
465 0 : u32 serverIdx = level2Idx * level1RankSize + level1Idx;
466 0 : const auto &multiStreamSlices = serverSlices[serverIdx];
467 0 : for (u32 ringIndex = 0; ringIndex < multiStreamSlices.size(); ringIndex++) {
468 0 : const auto &slice = multiStreamSlices[ringIndex][level0Idx];
469 0 : level0DataSegsSlice[ringIndex].push_back(slice);
470 0 : HCCL_DEBUG("[RSV]rank[%u], level0[%u]level2[%u]level1[%u], ringIndex[%u] slice.offset=[%llu], "
471 : "size=[%llu]", topoAttr_.userRank, level0Idx, level2Idx, level1Idx, ringIndex, slice.offset,
472 : slice.size);
473 : }
474 : }
475 : }
476 : }
477 0 : return HCCL_SUCCESS;
478 : }
479 :
480 0 : HcclResult CollReduceScatterRingFor91093Executor::CalUserMemDataSegsSliceV(const ExecMem &execMem,
481 : const OpParam ¶m, u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
482 : std::vector<std::vector<Slice>> &multRingsUserMemSlice)
483 : {
484 0 : std::vector<std::vector<std::vector<Slice>>> serverSlices;
485 0 : CHK_RET(FillMultiRingSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
486 : multRingsUserMemSlice, serverSlices, PrepareLevel0UserSlices));
487 0 : return HCCL_SUCCESS;
488 0 : }
489 :
490 0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel0DataSegsSliceV(const ExecMem &execMem,
491 : std::vector<std::vector<Slice>> &multiStreamSlice, const OpParam ¶m, u32 ringNum, u32 sliceNum,
492 : u32 level1RankSize, u32 level2RankSize, HcclDataType dataType, std::vector<std::vector<Slice>> &level0DataSegsSlice)
493 : {
494 0 : std::vector<std::vector<std::vector<Slice>>> serverSlices;
495 0 : CHK_RET(FillMultiRingSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
496 : level0DataSegsSlice, serverSlices, PrepareLevel0Slices));
497 0 : multiStreamSlice = serverSlices[0];
498 0 : return HCCL_SUCCESS;
499 0 : }
500 :
501 0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel1DataSegsSliceV(const OpParam ¶m,
502 : CommPlane commPlaneLevel, const u32 &commIndex, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
503 : u32 perDataSize, std::vector<Slice> &level1DataSegsSlice)
504 : {
505 0 : const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
506 0 : for (u32 i = 0; i < level1RankSize; i++) {
507 0 : Slice sliceTemp;
508 : u32 level1UserRank;
509 0 : CHK_RET(GetUserRankByRank(commPlaneLevel, commIndex, i, level1UserRank));
510 0 : if (level2RankSize <= 1) {
511 0 : sliceTemp.size = counts[level1UserRank] * perDataSize;
512 0 : sliceTemp.offset = std::accumulate(counts, counts + level1UserRank, 0ULL) * perDataSize;
513 0 : level1DataSegsSlice.push_back(sliceTemp);
514 0 : HCCL_DEBUG("[RSV]rank[%u], level1UserRank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]",
515 : topoAttr_.userRank, level1UserRank, i, sliceTemp.offset, sliceTemp.size);
516 : } else {
517 0 : for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
518 0 : const u32 ranksPerServer = level1RankSize * sliceNum;
519 0 : const u32 level2UserRank = level2Idx * ranksPerServer + level1UserRank % ranksPerServer;
520 0 : sliceTemp.size = counts[level2UserRank] * perDataSize;
521 0 : sliceTemp.offset = std::accumulate(counts, counts + level2UserRank, 0ULL) * perDataSize;
522 0 : level1DataSegsSlice.push_back(sliceTemp);
523 0 : HCCL_DEBUG("[RSV]rank[%u], level2UserRank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]",
524 : topoAttr_.userRank, level2UserRank, i, sliceTemp.offset, sliceTemp.size);
525 : }
526 : }
527 : }
528 0 : return HCCL_SUCCESS;
529 : }
530 :
531 0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel2DataSegsSliceV(const OpParam ¶m, u32 level2RankSize,
532 : u32 perDataSize, std::vector<Slice> &level2DataSegsSlice)
533 : {
534 0 : const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
535 0 : Slice sliceTemp;
536 0 : for (u32 i = 0; i < level2RankSize; i++) {
537 : u32 level2UserRank;
538 0 : CHK_RET(GetUserRankByRank(COMM_LEVEL2, COMM_INDEX_0, i, level2UserRank));
539 0 : sliceTemp.size = counts[level2UserRank] * perDataSize;
540 0 : sliceTemp.offset = std::accumulate(counts, counts + level2UserRank, 0ULL) * perDataSize;
541 0 : level2DataSegsSlice.push_back(sliceTemp);
542 0 : HCCL_DEBUG("[RSV]rank[%u], level2UserRank[%u], level2DataSegsSlice[%u].offset=%llu, size=[%llu]",
543 : topoAttr_.userRank, level2UserRank, i, sliceTemp.offset, sliceTemp.size);
544 : }
545 0 : return HCCL_SUCCESS;
546 : }
547 :
548 17 : HcomCollOpInfo CollReduceScatterRingFor91093Executor::GetHcomCollOpInfo(const OpParam ¶m,
549 : const ExecMem &execMem) const
550 : {
551 17 : const u64 count = param.GetDataCount(topoAttr_.userRank);
552 17 : const HcclDataType dataType = param.GetDataType();
553 17 : const u64 strideCount = param.GetStrideCount();
554 17 : HcomCollOpInfo opInfo = {"", execMem.inputPtr, execMem.outputPtr, count, dataType, param.root, param.reduceType,
555 17 : strideCount};
556 17 : HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][KernelRun] execMem.inputPtr[%p], execMem.outputPtr[%p], "
557 : "execMem.inputMem[%p], execMem.outputMem[%p], strideCount[%llu]", execMem.inputPtr, execMem.outputPtr,
558 : execMem.inputMem.ptr(), execMem.outputMem.ptr(), strideCount);
559 17 : return opInfo;
560 : }
561 :
562 16 : u64 CollReduceScatterRingFor91093Executor::CalcSrcMemOffset(const ExecMem &execMem, const OpParam ¶m,
563 : u32 perDataSize) const
564 : {
565 16 : if (isReduceScatterV_) {
566 0 : const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
567 0 : return std::accumulate(counts, counts + topoAttr_.userRank, 0ULL) * perDataSize;
568 : }
569 16 : return topoAttr_.userRank * execMem.outputMem.size();
570 : }
571 :
572 17 : HcclResult CollReduceScatterRingFor91093Executor::KernelRun(const OpParam ¶m, ExecMem &execMem)
573 : {
574 17 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] executor starts, rsv[%u]", __func__, isReduceScatterV_);
575 17 : CHK_RET(GetLevelCommInfo()); // 获取通信域
576 17 : u32 perDataSize = 0;
577 17 : const HcclDataType dataType = param.GetDataType();
578 17 : CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
579 :
580 : u32 ringNum;
581 17 : u32 level0RankSize = logicalLevel0CommInfo_.localRankSize;
582 17 : bool ARSFlag = topoMatcher_->GetARSFlag();
583 17 : bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
584 17 : u32 sliceNum = logicalLevel0CommInfo_.localRankSize;
585 17 : u32 commIndex = logicalLevel0CommInfo_.localRank;
586 :
587 17 : if ((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING && !IsUnifiedMarch(param) && !ARSFlag) || ARSDoubleRing) {
588 16 : ringNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
589 : } else {
590 1 : ringNum = LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
591 : }
592 :
593 34 : bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
594 17 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
595 :
596 17 : SubCommInfo level2CommInfo;
597 17 : if (isSelectAHC) {
598 0 : level2CommInfo = logicalLevel1CommInfo_;
599 0 : level2CommInfo.localRankSize = 1; // AHC bypass level2
600 : } else {
601 17 : CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
602 17 : level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
603 : }
604 17 : const u32 level2RankSize = level2CommInfo.localRankSize;
605 17 : const u32 level1RankSize = logicalLevel1CommInfo_.localRankSize;
606 :
607 : // 节点内reduce scatter
608 17 : CHK_RET(ActiveSlaveStreams(param.stream));
609 :
610 : // 计算slice
611 17 : std::vector<std::vector<Slice>> multiStreamSlice; // 每个stream使用的数据基于用户buffer的偏移
612 17 : std::vector<std::vector<Slice>> level0DataSegsSlice;
613 17 : CalLevel0DataSegsSlice(execMem, multiStreamSlice, param, ringNum, sliceNum, level1RankSize, level2RankSize,
614 : dataType, level0DataSegsSlice);
615 :
616 17 : HcomCollOpInfo opInfo = GetHcomCollOpInfo(param, execMem);
617 17 : HcomCollOpInfo *opInfoPtr = nullptr;
618 17 : if (DMAReduceFlag_) {
619 16 : opInfoPtr = &opInfo;
620 : }
621 :
622 17 : bool disableDMAReduce = algOpContext_.opRetryHandler.retryEnable &&
623 0 : (algOpContext_.opRetryHandler.inPlaceSupportRetryStatus == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE1 ||
624 0 : algOpContext_.opRetryHandler.inPlaceSupportRetryStatus == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE2);
625 17 : std::vector<std::vector<Slice>> multRingsUserMemSlice;
626 17 : CalUserMemDataSegsSlice(execMem, level0DataSegsSlice, multiStreamSlice, param, ringNum, sliceNum, level1RankSize,
627 : level2RankSize, dataType, perDataSize, opInfoPtr, disableDMAReduce, multRingsUserMemSlice);
628 :
629 : // 区分消减拷贝场景
630 17 : if ((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING || ARSDoubleRing) &&
631 16 : (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB)) {
632 : // 图模式opinfo不为空
633 0 : HcomCollOpInfo graphModeOpInfo = {"", execMem.inputMem.ptr(), nullptr, param.GetDataCount(topoAttr_.userRank),
634 0 : dataType, param.root, param.reduceType, param.GetStrideCount()};
635 0 : CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count, dataType,
636 : param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0, &graphModeOpInfo,
637 : multRingsUserMemSlice, disableDMAReduce));
638 17 : } else if (opInfoPtr != nullptr && (level1RankSize > 1 || level2RankSize > 1)) {
639 16 : HcomCollOpInfo opInfoByReduceScatterDMAreduce = *opInfoPtr;
640 16 : opInfoByReduceScatterDMAreduce.outputAddr = nullptr;
641 16 : CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count,
642 : dataType, param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0,
643 : &opInfoByReduceScatterDMAreduce, multRingsUserMemSlice, disableDMAReduce));
644 16 : } else {
645 1 : CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count,
646 : dataType, param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0, opInfoPtr,
647 : multRingsUserMemSlice, disableDMAReduce));
648 : }
649 : // 对于单server图模式的最后一步需要把数据从ccl input拷贝到ccl output上
650 16 : if (level1RankSize == 1 && level2RankSize == 1 && opInfoPtr == nullptr) {
651 0 : const u64 offset = CalcSrcMemOffset(execMem, param, perDataSize);
652 0 : DeviceMem srcMem = execMem.inputMem.range(offset, execMem.outputMem.size());
653 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
654 0 : }
655 :
656 16 : if (level1RankSize > 1) {
657 : // 节点间做reduce scatter(ring/NHR/NB)
658 16 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, dataType, param.reduceType);
659 16 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
660 :
661 : // 计算slice
662 16 : std::vector<Slice> level1DataSegsSlice;
663 16 : CHK_RET(CalLevel1DataSegsSlice(execMem, param, logicalLevel1plane_, commIndex, sliceNum, level1RankSize,
664 : level2RankSize, perDataSize, level1DataSegsSlice));
665 :
666 16 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
667 32 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
668 16 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
669 16 : CHK_SMART_PTR_NULL(level1TempAlg);
670 16 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
671 16 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
672 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
673 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
674 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
675 0 : CHK_SMART_PTR_NULL(level1TempAlg);
676 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
677 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
678 0 : } else if (isSelectAHC) {
679 : // 获取通信域分组信息
680 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
681 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
682 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(logicalLevel1plane_, globalSubGroups));
683 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
684 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
685 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC, dispatcher_);
686 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC in COMM_LEVEL1", __func__);
687 : } else {
688 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC_BROKE, dispatcher_);
689 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC_BROKE in COMM_LEVEL1", __func__);
690 : }
691 0 : HCCL_DEBUG("[CollReduceScatterRingFor91093Executor]runAsync for COMM_LEVEL1 ends");
692 0 : CHK_SMART_PTR_NULL(level1TempAlg);
693 0 : CHK_RET(level1TempAlg->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
694 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr));
695 0 : } else {
696 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
697 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
698 0 : CHK_SMART_PTR_NULL(level1TempAlg);
699 0 : CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
700 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
701 : }
702 :
703 48 : CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count,
704 : dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSegsSlice));
705 16 : CHK_RET(level1TempAlg->RegisterProfiler(
706 : (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + logicalLevel1CommInfo_.localRank,
707 : PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
708 16 : CHK_RET(RunTemplate(level1TempAlg, logicalLevel1CommInfo_));
709 16 : }
710 :
711 16 : if (level2RankSize > 1) {
712 : /* ****************** 超节点间 reducescatter *******************************/
713 0 : u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, dataType, param.reduceType);
714 :
715 : // 计算slice
716 0 : std::vector<Slice> level2DataSegsSlice;
717 0 : CHK_RET(CalLevel2DataSegsSlice(execMem, param, level2RankSize, perDataSize, level2DataSegsSlice));
718 :
719 0 : std::unique_ptr<AlgTemplateBase> level2TempAlg;
720 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
721 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
722 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
723 0 : CHK_SMART_PTR_NULL(level2TempAlg);
724 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr));
725 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL2", __func__);
726 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
727 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
728 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
729 0 : CHK_SMART_PTR_NULL(level2TempAlg);
730 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr, false));
731 0 : if (algoAttr_.isSupportAtomicWrite) {
732 0 : level2TempAlg->CloseBarrier();
733 : }
734 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL2", __func__);
735 : } else {
736 0 : level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
737 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
738 0 : CHK_SMART_PTR_NULL(level2TempAlg);
739 0 : CHK_RET(level2TempAlg->Prepare(reduceAttr));
740 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL2", __func__);
741 : }
742 :
743 0 : CHK_RET(level2TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count, dataType,
744 : param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, level2DataSegsSlice));
745 0 : CHK_RET(level2TempAlg->RegisterProfiler(
746 : (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank,
747 : PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
748 0 : CHK_RET(RunTemplate(level2TempAlg, level2CommInfo));
749 0 : }
750 :
751 16 : if (level1RankSize > 1 || level2RankSize > 1) {
752 : // 区分消减拷贝场景(消减拷贝数据需要拷贝到user output上)
753 16 : const u64 offset = CalcSrcMemOffset(execMem, param, perDataSize);
754 16 : DeviceMem srcMem = execMem.inputMem.range(offset, execMem.outputMem.size());
755 16 : if (opInfoPtr != nullptr) {
756 16 : DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(opInfoPtr->outputAddr), execMem.outputMem.size());
757 16 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, const_cast<Stream&>(param.stream)));
758 16 : } else {
759 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
760 : }
761 16 : }
762 :
763 16 : HCCL_INFO("ReduceScatter ring run success, rsv[%u]", isReduceScatterV_);
764 16 : return HCCL_SUCCESS;
765 17 : }
766 :
767 0 : HcclResult CollReduceScatterRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
768 : {
769 0 : bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
770 0 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
771 :
772 0 : if (isSelectAHC) {
773 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
774 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
775 :
776 0 : u32 commIndex = level0CommInfo.localRank;
777 :
778 0 : CommPlane commPlaneLevel1 = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
779 0 : CHK_RET(CheckCommSize(commPlaneLevel1, commIndex + 1));
780 0 : level1CommInfo = GetSubCommInfo(commPlaneLevel1, commIndex);
781 0 : return HCCL_SUCCESS;
782 0 : }
783 :
784 0 : if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
785 0 : return HCCL_E_UNAVAIL;
786 : }
787 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
788 :
789 0 : return HCCL_SUCCESS;
790 : }
791 :
792 0 : HcclResult CollReduceScatterRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
793 : {
794 0 : bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
795 0 : algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
796 0 : HCCL_DEBUG("[CollReduceScatterRingFor91093Executor]SelectTempAlg begins");
797 0 : if (isSelectAHC) {
798 0 : CommPlane commPlaneLevel1 = COMM_LEVEL1_AHC;
799 : // 获取通信域分组信息
800 0 : std::vector<std::vector<std::vector<u32>>> globalSubGroups;
801 0 : std::map<AHCConcOpType, TemplateType> ahcAlgOption;
802 0 : CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
803 0 : topoMatcher_->GetAHCAlgOption(ahcAlgOption);
804 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
805 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC, dispatcher_);
806 0 : HCCL_INFO("reducescatter ring: using ahc algo inter-server.");
807 : } else {
808 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC_BROKE, dispatcher_);
809 0 : HCCL_INFO("reducescatter ring: using ahc-broke algo inter-server.");
810 : }
811 0 : CHK_SMART_PTR_NULL(level1TempAlg);
812 0 : CHK_RET(level1TempAlg->Prepare(NSLBDP_MIN_COUNT, globalSubGroups, ahcAlgOption));
813 0 : return HCCL_SUCCESS;
814 0 : }
815 0 : if (level1RankSize > 1) {
816 0 : if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
817 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
818 0 : TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
819 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL2", __func__);
820 0 : CHK_SMART_PTR_NULL(level1TempAlg);
821 0 : } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
822 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
823 0 : TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
824 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL2", __func__);
825 0 : CHK_SMART_PTR_NULL(level1TempAlg);
826 : } else {
827 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
828 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
829 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL2", __func__);
830 0 : CHK_SMART_PTR_NULL(level1TempAlg);
831 : }
832 0 : return HCCL_SUCCESS;
833 : }
834 0 : return HCCL_E_UNAVAIL;
835 : }
836 :
837 :
838 : REGISTER_EXEC("ReduceScatterRingFor91093Executor", ReduceScatterRingFor91093, CollReduceScatterRingFor91093Executor);
839 : }
|