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