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 <cmath>
12 :
13 : #include "log.h"
14 :
15 : #include "coll_alg_registry.h"
16 : #include "reduce_scatter_sole_executor.h"
17 :
18 : namespace Hccl {
19 : template <typename AlgTopoMatch, typename AlgTemplate>
20 0 : ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::ReduceScatterSoleExecutor() : CollAlgBase()
21 0 : {}
22 :
23 : template <typename AlgTopoMatch, typename AlgTemplate>
24 0 : ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::~ReduceScatterSoleExecutor()
25 0 : {}
26 :
27 : template <typename AlgTopoMatch, typename AlgTemplate>
28 0 : HcclResult ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::CalcResOffload(
29 : const RankGraph* rankGraph, const u64& dataSize, CollOffloadOpResReq& resReq)
30 : {
31 : // Topo Match
32 0 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
33 0 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
34 :
35 : // instantiate a template
36 0 : AlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
37 :
38 : // calculate required primQues and prepare queue
39 0 : AlgTempResReq tempResReq;
40 0 : u32 requiredScratchMultiplier = 0;
41 0 : if (enableDetour_) {
42 0 : CHK_RET(tempAlg.CalcResDetour(false, rankGraph, tempResReq, requiredScratchMultiplier));
43 : } else {
44 0 : CHK_RET(tempAlg.CalcRes(false, tempResReq, requiredScratchMultiplier));
45 : }
46 0 : resReq.requiredSubQueNum = tempResReq.queNum - 1;
47 :
48 0 : resReq.requiredScratchMemSize = requiredScratchMultiplier * dataSize;
49 :
50 0 : return HcclResult::HCCL_SUCCESS;
51 0 : }
52 :
53 : template <typename AlgTopoMatch, typename AlgTemplate>
54 0 : HcclResult ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::GenPrimQues(
55 : const RankGraph* rankGraph, const CollAlgOperator& op, const CollAlgParams& params, PrimQuePtr primQue)
56 : {
57 : // init and check params
58 0 : CHK_RET(Init(op, params, primQue));
59 :
60 : // Topo Match
61 0 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
62 0 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
63 0 : HCCL_INFO("[CollAlgFactory] Rank[%d], [%s].", myRank_, topoMatch.Describe().c_str());
64 :
65 : // instantiate a template
66 0 : AlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
67 0 : tempAlg.InitReduceInfo(redOp_, dataType_);
68 0 : tempAlg.SetDmaMode(dmaMode_);
69 :
70 : // calculate required primQues and prepare queue
71 0 : AlgTempResReq tempResReq;
72 0 : u32 requiredScratchMultiplier = 0;
73 0 : if (enableDetour_) {
74 0 : CHK_RET(tempAlg.CalcResDetour(false, rankGraph, tempResReq, requiredScratchMultiplier));
75 : } else {
76 0 : CHK_RET(tempAlg.CalcRes(false, tempResReq, requiredScratchMultiplier));
77 : }
78 :
79 0 : CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
80 0 : HCCL_INFO(
81 : "[CollAlgFactory] Rank[%d], template [%s]: requiredQue Num [%u].", myRank_, tempAlg.Describe().c_str(),
82 : tempResReq.queNum);
83 :
84 0 : CHK_RET(PrepResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, tempResLinks_));
85 :
86 0 : u32 dataSizePerVolume = DataTypeSizeGet(dataType_);
87 0 : dataSize_ = dataCount_ * dataSizePerVolume;
88 :
89 0 : if (opMode_ == OpMode::OFFLOAD) {
90 0 : HCCL_INFO("[CollAlgFactory] Rank[%d], Generating Primitive Queues in OFFLOAD Mode for HOST.", myRank_);
91 0 : CHK_RET(GenPrimQues4Offload(tempAlg));
92 : } else { // OPBASE
93 0 : HCCL_INFO("[CollAlgFactory] Rank[%d], Generating Primitive Queues in OPBASE Mode for HOST.", myRank_);
94 0 : CHK_RET(GenPrimQues4Opbase(requiredScratchMultiplier, dataSizePerVolume, tempAlg));
95 : }
96 :
97 0 : return HcclResult::HCCL_SUCCESS;
98 0 : }
99 :
100 : template <typename AlgTopoMatch, typename AlgTemplate>
101 : HcclResult
102 0 : ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::CalcRes(const RankGraph* rankGraph, CollAlgResReq& algResReq)
103 : {
104 : // Topo Match
105 0 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
106 0 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
107 0 : algResReq.topoInfo.UpdateSingleLevelTopo(virtRanks_, virtRankMap_, vTopo_);
108 :
109 : // instantiate a template
110 0 : AlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
111 0 : tempAlg.InitReduceInfo(redOp_, dataType_);
112 :
113 : // calculate required primQues and prepare queue
114 0 : AlgTempResReq tempResReq;
115 0 : u32 requiredScratchMultiplier = 0;
116 0 : if (enableDetour_) {
117 0 : CHK_RET(tempAlg.CalcResDetour(false, rankGraph, tempResReq, requiredScratchMultiplier));
118 : } else {
119 0 : CHK_RET(tempAlg.CalcRes(false, tempResReq, requiredScratchMultiplier));
120 : }
121 :
122 0 : algResReq.primQueueNum = tempResReq.queNum;
123 0 : CHK_RET(CalcResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, algResReq.links));
124 :
125 0 : return HcclResult::HCCL_SUCCESS;
126 0 : }
127 :
128 : template <typename AlgTopoMatch, typename AlgTemplate>
129 0 : HcclResult ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::GenPrimQuesAIC(
130 : const AlgTopoInfo& topoInfo, const CollAlgOperator& op, const CollAlgParams& params, ConnectedLinkMgr* linkMgr,
131 : PrimQuePtr primQue)
132 : {
133 : // init and check params
134 0 : CHK_RET(Init(op, params, primQue));
135 :
136 : // instantiate a template
137 0 : AlgTemplate tempAlg(myRank_, rankSize_, topoInfo.vTopo[0], topoInfo.virtRankMap[0]);
138 0 : tempAlg.InitReduceInfo(redOp_, dataType_);
139 0 : tempAlg.SetDmaMode(dmaMode_);
140 0 : virtRankMap_ = topoInfo.virtRankMap[0];
141 :
142 : // calculate required primQues and prepare queue
143 0 : AlgTempResReq tempResReq;
144 0 : u32 requiredScratchMultiplier = 0;
145 0 : if (enableDetour_) {
146 0 : CHK_RET(tempAlg.CalcResDetour(false, linkMgr, tempResReq, requiredScratchMultiplier));
147 : } else {
148 0 : CHK_RET(tempAlg.CalcRes(false, tempResReq, requiredScratchMultiplier));
149 : }
150 :
151 0 : CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
152 0 : HCCL_INFO(
153 : "[CollAlgFactory] Rank[%d], template [%s]: requiredQue Num [%u].", myRank_, tempAlg.Describe().c_str(),
154 : tempResReq.queNum);
155 :
156 0 : CHK_RET(PrepResLinks(myRank_, tempResReq.links, linkMgr, tempResLinks_));
157 :
158 0 : u32 dataSizePerVolume = DataTypeSizeGet(dataType_);
159 0 : dataSize_ = dataCount_ * dataSizePerVolume;
160 :
161 0 : if (opMode_ == OpMode::OFFLOAD) {
162 0 : HCCL_INFO("[CollAlgFactory] Rank[%d], Generating Primitive Queues in OFFLOAD Mode for AICPU.", myRank_);
163 0 : CHK_RET(GenPrimQues4Offload(tempAlg));
164 : } else { // OPBASE
165 0 : HCCL_INFO("[CollAlgFactory] Rank[%d], Generating Primitive Queues in OPBASE Mode for AICPU.", myRank_);
166 0 : CHK_RET(GenPrimQues4Opbase(requiredScratchMultiplier, dataSizePerVolume, tempAlg));
167 : }
168 :
169 0 : return HcclResult::HCCL_SUCCESS;
170 0 : }
171 :
172 : template <typename AlgTopoMatch, typename AlgTemplate>
173 0 : HcclResult ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::GenPrimQues4Offload(AlgTemplateBase& tempAlg)
174 : {
175 0 : RankSliceInfo sliceInfoVec;
176 0 : AllignInfo allignInfo = {enableAllign_, allignSize_, dataType_};
177 0 : CHK_RET(tempAlg.CalcSliceInfo(allignInfo, false, dataSize_, sliceInfoVec));
178 :
179 0 : BuffInfo buffInfo;
180 0 : buffInfo.inBuffType = BufferType::INPUT;
181 0 : buffInfo.outBuffType = BufferType::OUTPUT;
182 0 : buffInfo.scratBuffType = BufferType::SCRATCH;
183 0 : buffInfo.inBuffBaseOff = 0;
184 0 : buffInfo.outBuffBaseOff = 0;
185 0 : buffInfo.scratchBuffBaseOff = 0;
186 :
187 0 : TempFuncs tempFuncs;
188 0 : tempFuncs.opMode = opMode_;
189 0 : tempFuncs.enableCounterNotify = IsEnableCounterNotify();
190 :
191 0 : CHK_RET(tempAlg.GenPrimQue(tempFuncs, sliceInfoVec, buffInfo, tempResLinks_, requiredQue_));
192 :
193 0 : return HcclResult::HCCL_SUCCESS;
194 0 : }
195 :
196 : template <typename AlgTopoMatch, typename AlgTemplate>
197 0 : HcclResult ReduceScatterSoleExecutor<AlgTopoMatch, AlgTemplate>::GenPrimQues4Opbase(
198 : const u32 requiredScratchMultiplier, const u32 dataSizePerVolume, AlgTemplateBase& tempAlg)
199 : {
200 0 : CHK_PRT_RET(
201 : dataSizePerVolume == 0,
202 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Invalid dataSizePerVolume [%u].", myRank_, dataSizePerVolume),
203 : HcclResult::HCCL_E_INTERNAL);
204 :
205 0 : u32 scratchCCLMultiplier = (requiredScratchMultiplier == 0) ? 1 : requiredScratchMultiplier;
206 0 : u64 scratchInputMemSize = static_cast<int>(
207 0 : ((rankSize_ + scratchCCLMultiplier) % dataSizePerVolume == 0) ?
208 0 : floor(maxTmpMemSize_ / (rankSize_ + scratchCCLMultiplier)) :
209 0 : floor(maxTmpMemSize_ / ((rankSize_ + scratchCCLMultiplier) * dataSizePerVolume)) * dataSizePerVolume);
210 :
211 0 : CHK_PRT_RET(
212 : scratchInputMemSize == 0,
213 : HCCL_ERROR("[CollAlgFactory] Rank [%d], Invalid input maxTmpMemSize [%u].", myRank_, maxTmpMemSize_),
214 : HcclResult::HCCL_E_PARA);
215 :
216 0 : BuffInfo buffInfo;
217 0 : buffInfo.inBuffType = BufferType::SCRATCH;
218 0 : buffInfo.outBuffType = BufferType::SCRATCH;
219 0 : buffInfo.scratBuffType = BufferType::SCRATCH;
220 0 : buffInfo.inBuffBaseOff = 0;
221 :
222 0 : u32 sendRecvTimes = (dataSize_ / scratchInputMemSize) + ((dataSize_ % scratchInputMemSize) == 0 ? 0 : 1);
223 0 : HCCL_INFO("[CollAlgFactory] Rank [%d], sendRecvTimes [%u].", myRank_, sendRecvTimes);
224 :
225 0 : u64 resDataSize = dataSize_;
226 0 : for (u32 idx = 0; idx < sendRecvTimes; idx++) {
227 0 : u64 currDataSize = resDataSize > scratchInputMemSize ? scratchInputMemSize : resDataSize;
228 :
229 0 : buffInfo.scratchBuffBaseOff = currDataSize * rankSize_;
230 0 : buffInfo.outBuffBaseOff = currDataSize * rankSize_;
231 :
232 0 : RankSliceInfo sliceInfoVec;
233 0 : AllignInfo allignInfo = {enableAllign_, allignSize_, dataType_};
234 0 : CHK_RET(tempAlg.CalcSliceInfo(allignInfo, false, currDataSize, sliceInfoVec));
235 :
236 0 : TempFuncs tempFuncs;
237 0 : tempFuncs.opMode = opMode_;
238 0 : tempFuncs.enableCounterNotify = IsEnableCounterNotify();
239 0 : tempFuncs.isForepart = true; // Usr Buff to CCL Buff required
240 0 : tempFuncs.isBottom = true; // CCL Buff to Usr Buff required
241 :
242 0 : UsrData usrData;
243 0 : for (u32 rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
244 0 : DataSlice usrInSlice
245 0 : = DataSlice(BufferType::INPUT, rankIdx * dataSize_ + idx * scratchInputMemSize, currDataSize);
246 0 : DataSlice scratchInSlice = DataSlice(BufferType::SCRATCH, rankIdx * currDataSize, currDataSize);
247 0 : usrData.usrInSlices.push_back(usrInSlice);
248 0 : usrData.scratchInSlices.push_back(scratchInSlice);
249 : }
250 :
251 0 : DataSlice scratchOutSlice = DataSlice(BufferType::SCRATCH, virtRankMap_[myRank_] * currDataSize, currDataSize);
252 0 : DataSlice usrOutSlice = DataSlice(BufferType::OUTPUT, idx * scratchInputMemSize, currDataSize);
253 0 : usrData.scratchOutSlices.push_back(scratchOutSlice);
254 0 : usrData.usrOutSlices.push_back(usrOutSlice);
255 :
256 0 : tempFuncs.usrData = usrData;
257 0 : CHK_RET(tempAlg.GenPrimQue(tempFuncs, sliceInfoVec, buffInfo, tempResLinks_, requiredQue_));
258 0 : resDataSize -= currDataSize;
259 : }
260 :
261 0 : return HcclResult::HCCL_SUCCESS;
262 : }
263 :
264 : REGISTER_IMPL_BY_TEMP(
265 : OpType::REDUCESCATTER, ReduceScatterConcurrMesh, ReduceScatterSoleExecutor, TopoMatchConcurrMesh,
266 : TempReduceScatterConcurrMesh);
267 : REGISTER_IMPL_BY_TEMP(
268 : OpType::REDUCESCATTER, ReduceScatterMesh, ReduceScatterSoleExecutor, TopoMatchMesh, TempReduceScatterMesh);
269 : } // namespace Hccl
|