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 "log.h"
12 :
13 : #include "ins_coll_alg_registry.h"
14 : #include "ins_all_to_all_sole_executor.h"
15 :
16 : #include "ccu_temp_all_to_all_mesh_1D.h"
17 : #include "ccu_temp_all_to_all_v_mesh_1D.h"
18 : #include "ccu_temp_all_to_all_mesh2d.h"
19 : #include "ccu_temp_all_to_all_v_mesh_2D.h"
20 : #include "ccu_temp_all_to_all_v_mesh_2Die.h"
21 :
22 : namespace Hccl {
23 : template <typename AlgTopoMatch, typename InsAlgTemplate>
24 1 : InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::InsAlltoAllSoleExecutor() : InsCollAlgBase()
25 : {
26 1 : }
27 :
28 : template <typename AlgTopoMatch, typename InsAlgTemplate>
29 2 : InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::~InsAlltoAllSoleExecutor()
30 : {
31 2 : }
32 :
33 : template <typename AlgTopoMatch, typename InsAlgTemplate>
34 1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::InitParams(const CollAlgOperator &op, const CollAlgParams ¶ms)
35 : {
36 1 : opMode_ = params.opMode;
37 1 : maxTmpMemSize_ = params.maxTmpMemSize;
38 1 : CHK_PRT_RET((maxTmpMemSize_ == 0),
39 : HCCL_ERROR("[InitParams] maxTmpMemSize equals to zero for OPBASE."), HcclResult::HCCL_E_PARA);
40 :
41 1 : CHK_PRT_RET(GetAlltoAllLocalSendRecvInfo(op, myRank_, rankSize_, localSendRecvInfo_), HCCL_ERROR("[InitParams] unable to init DataInfo."),
42 : HcclResult::HCCL_E_PARA);
43 1 : if (op.opType == OpType::ALLTOALL) {
44 1 : sendType_ = op.all2AllDataDes.sendType;
45 1 : recvType_ = op.all2AllDataDes.recvType;
46 0 : } else if (op.opType == OpType::ALLTOALLV) {
47 0 : sendType_ = op.all2AllVDataDes.sendType;
48 0 : recvType_ = op.all2AllVDataDes.recvType;
49 0 : } else if (op.opType == OpType::ALLTOALLVC) {
50 0 : sendType_ = op.all2AllVCDataDes.sendType;
51 0 : recvType_ = op.all2AllVCDataDes.recvType;
52 0 : } else if (op.opType != OpType::HALFALLTOALLV) {
53 0 : HCCL_ERROR("[InsAlltoAllSoleExecutor] opType [%s] is invalid.", op.opType.Describe().c_str());
54 0 : return HcclResult::HCCL_E_PARA;
55 : }
56 1 : CHK_PRT_RET(InitOpInfo(op, opType_, redOp_, root_), HCCL_ERROR("[InitParams] unable to init OpInfo."),
57 : HcclResult::HCCL_E_PARA);
58 1 : return HcclResult::HCCL_SUCCESS;
59 : }
60 :
61 : template <typename AlgTopoMatch, typename InsAlgTemplate>
62 0 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::CalcResOffload(const RankGraph *rankGraph,
63 : const u64 &dataSize,
64 : CollOffloadOpResReq &resReq)
65 : {
66 : (void)dataSize;
67 0 : resReq.requiredScratchMemSize = 200 * 1024 * 1024; // 200 * 1024*1024 = 200M
68 :
69 : // Topo Match
70 0 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
71 0 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
72 :
73 : // instantiate a template
74 0 : InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
75 :
76 0 : std::map<u32, u32>rank2PathNumMap;
77 0 : HCCL_INFO("[InsAlltoAllSoleExecutor] CalcResOffload SetPathNumMap");
78 0 : for(auto rankIdx : virtRanks_){
79 0 : if(rankIdx==myRank_){
80 0 : continue;
81 : }
82 0 : std::vector<NetInstance::Path> tmpPaths0 =
83 : rankGraph->GetPaths(0, myRank_, rankIdx);
84 0 : std::vector<NetInstance::Path> tmpPaths1 =
85 : rankGraph->GetPaths(1, myRank_, rankIdx);
86 0 : HCCL_INFO("[InsAlltoAllSoleExecutor]tmpPaths0.size() = %zu,tmpPaths1.size() = %zu", tmpPaths0.size(), tmpPaths1.size());
87 0 : rank2PathNumMap[rankIdx] = tmpPaths0.size() + tmpPaths1.size();
88 : }
89 :
90 0 : tempAlg.setPathNumMap(rank2PathNumMap);
91 :
92 : // calculate required insQueues and prepare queue
93 0 : AlgTempResReq tempResReq;
94 0 : if (enableDetour_) {
95 0 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor], CalcResOffload with detouring enabled.");
96 0 : CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
97 : } else {
98 0 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor], CalcResOffload with detouring disabled.");
99 0 : CHK_RET(tempAlg.CalcRes(tempResReq));
100 : }
101 :
102 0 : resReq.requiredSubQueNum = tempResReq.streamNum - 1;
103 :
104 0 : HCCL_INFO("[InsAlltoAllSoleExecutor][CalcResOffload] requiredSubQueNum[%llu], requiredScratchMemSize[%llu].",
105 : resReq.requiredSubQueNum, resReq.requiredScratchMemSize);
106 :
107 0 : return HcclResult::HCCL_SUCCESS;
108 0 : }
109 :
110 : template <typename AlgTopoMatch, typename InsAlgTemplate>
111 1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::CalcRes(const RankGraph *rankGraph,
112 : CollAlgResReq &algResReq)
113 : {
114 : // Topo Match
115 1 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
116 1 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
117 1 : algResReq.topoInfo.UpdateSingleLevelTopo(virtRanks_, virtRankMap_, vTopo_);
118 3 : HCCL_DEBUG("[InsAlltoAllSoleExecutor][CalcRes]topoInfo.virtRanks[%u], topoInfo.virtRankMap[%u], topoInfo.vTopo[%u].",
119 : algResReq.topoInfo.virtRanks.size(), algResReq.topoInfo.virtRankMap.size(), algResReq.topoInfo.vTopo.size());
120 : // instantiate a template
121 1 : InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
122 :
123 1 : std::map<u32, u32>rank2PathNumMap;
124 3 : HCCL_INFO("[InsAlltoAllSoleExecutor] CalcRes SetPathNumMap");
125 8 : for(auto rankIdx : virtRanks_){
126 4 : if(rankIdx==myRank_){
127 1 : continue;
128 : }
129 3 : std::vector<NetInstance::Path> tmpPaths0 =
130 : rankGraph->GetPaths(0, myRank_, rankIdx);
131 3 : std::vector<NetInstance::Path> tmpPaths1 =
132 : rankGraph->GetPaths(1, myRank_, rankIdx);
133 3 : rank2PathNumMap[rankIdx] = tmpPaths0.size() + tmpPaths1.size();
134 : }
135 1 : tempAlg.setPathNumMap(rank2PathNumMap);
136 :
137 : // calculate required insQues and prepare queue
138 1 : AlgTempResReq tempResReq;
139 1 : if (enableDetour_) {
140 0 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
141 0 : CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
142 : } else {
143 3 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
144 1 : CHK_RET(tempAlg.CalcRes(tempResReq));
145 : }
146 1 : CHK_RET(CalcLinkInfo(myRank_, rankGraph, tempResReq.links, algResReq.levelRankPairs));
147 1 : algResReq.primQueueNum= tempResReq.streamNum;
148 1 : algResReq.queueNotifys = tempResReq.queNotifys;
149 3 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], requiredQueNum [%u].", myRank_, algResReq.primQueueNum);
150 1 : CHK_RET(CalcResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, algResReq.links));
151 1 : return HcclResult::HCCL_SUCCESS;
152 1 : }
153 :
154 : // dataSize_ as input
155 : template <typename AlgTopoMatch, typename InsAlgTemplate>
156 1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::Orchestrate(const RankGraph *rankGraph,
157 : const CollAlgOperator &op,
158 : const CollAlgParams ¶ms,
159 : InsQuePtr insQue)
160 : {
161 : // init and check params
162 1 : CHK_RET(Init(op, params, insQue));
163 : // Topo Match
164 1 : AlgTopoMatch topoMatch(myRank_, rankSize_, rankGraph, devType_);
165 1 : CHK_RET(topoMatch.MatchTopo(vTopo_, virtRanks_, virtRankMap_));
166 3 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], [%s].", myRank_, topoMatch.Describe().c_str());
167 :
168 : // instantiate a template
169 1 : InsAlgTemplate tempAlg(myRank_, rankSize_, vTopo_, virtRankMap_);
170 1 : tempAlg.SetDmaMode(dmaMode_);
171 1 : tempAlg.SetCollOp(op); // CCU template需要传递op信息
172 1 : tempAlg.SetA2ASendRecvInfo(localSendRecvInfo_);
173 1 : tempAlg.SetLoadInfo(params);
174 :
175 : // calculate required insQues and prepare queue
176 1 : AlgTempResReq tempResReq;
177 1 : if (enableDetour_) {
178 0 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
179 0 : CHK_RET(tempAlg.CalcResDetour(rankGraph, tempResReq));
180 : } else {
181 3 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
182 1 : CHK_RET(tempAlg.CalcRes(tempResReq));
183 : }
184 :
185 1 : CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
186 3 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], template [%s], requiredQue Num [%u].", myRank_,
187 : tempAlg.Describe().c_str(), tempResReq.queNum);
188 :
189 1 : CHK_RET(PrepResLinks(myRank_, rankGraph, linkPriority_, tempResReq.links, tempResLinks_));
190 :
191 1 : CHK_RET(OrchestrateOpbase(tempAlg));
192 :
193 1 : return HcclResult::HCCL_SUCCESS;
194 1 : }
195 :
196 : template <typename AlgTopoMatch, typename InsAlgTemplate>
197 0 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::Orchestrate(const AlgTopoInfo &topoInfo,
198 : const CollAlgOperator &op,
199 : const CollAlgParams ¶ms,
200 : ConnectedLinkMgr *linkMgr,
201 : InsQuePtr insQue)
202 : {
203 0 : HCCL_INFO("[InsAlltoAllSoleExecutor] Begin to orchestrate.");
204 : // init and check params
205 0 : CHK_RET(Init(op, params, insQue));
206 0 : dataType_ = op.dataType;
207 : // instantiate a template
208 0 : if(topoInfo.vTopo.size() == 0) {
209 0 : HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank[%d], vTopo size is zero.", myRank_);
210 0 : return HcclResult::HCCL_E_PARA;
211 : }
212 0 : if(topoInfo.virtRankMap.size() == 0) {
213 0 : HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank[%d], virtRankMap size is zero.", myRank_);
214 0 : return HcclResult::HCCL_E_PARA;
215 : }
216 0 : InsAlgTemplate tempAlg(myRank_, rankSize_, topoInfo.vTopo[0], topoInfo.virtRankMap[0]);
217 0 : tempAlg.SetDmaMode(dmaMode_);
218 0 : tempAlg.SetCollOp(op);
219 0 : tempAlg.SetA2ASendRecvInfo(localSendRecvInfo_);
220 0 : tempAlg.SetLoadInfo(params);
221 0 : tempAlg.SetDataType(dataType_);
222 0 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], Init insAlgTemplate with rankSize [%u] and dmaMode [%s].", myRank_,
223 : rankSize_, dmaMode_.Describe().c_str());
224 0 : virtRankMap_ = topoInfo.virtRankMap[0];
225 0 : virtRanks_ = topoInfo.virtRanks[0];
226 0 : std::map<u32, u32>rank2PathNumMap;
227 0 : for(u32 rankIdx:virtRanks_){
228 0 : auto links0 = linkMgr->GetLinks(0, rankIdx);
229 0 : auto links1 = linkMgr->GetLinks(1, rankIdx);
230 0 : if(links0.size() + links1.size() != 0){
231 0 : rank2PathNumMap[rankIdx] = links0.size() + links1.size();
232 : }
233 : }
234 0 : tempAlg.setPathNumMap(rank2PathNumMap);
235 :
236 : // calculate required insQues and prepare queue
237 0 : AlgTempResReq tempResReq;
238 0 : if (enableDetour_) {
239 0 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring enabled.", myRank_);
240 0 : CHK_RET(tempAlg.CalcResDetour(linkMgr, tempResReq));
241 : } else {
242 0 : HCCL_DEBUG("[InsCollAlgFactory] [InsAlltoAllSoleExecutor] Rank[%d], CalcRes with detouring disabled.", myRank_);
243 0 : CHK_RET(tempAlg.CalcRes(tempResReq));
244 : }
245 :
246 0 : CHK_RET(InitQueue(tempResReq.queNum, requiredQue_));
247 0 : HCCL_DEBUG("[InsAlltoAllSoleExecutor] Rank[%d], template [%s], requiredQue Num [%u].", myRank_,
248 : tempAlg.Describe().c_str(), tempResReq.queNum);
249 :
250 0 : CHK_RET(PrepResLinks(myRank_, tempResReq.links, linkMgr, tempResLinks_));
251 :
252 0 : CHK_RET(OrchestrateOpbase(tempAlg));
253 0 : HCCL_INFO("[InsAlltoAllSoleExecutor] Orchestrate success.");
254 :
255 0 : return HcclResult::HCCL_SUCCESS;
256 0 : }
257 :
258 : template <typename AlgTopoMatch, typename InsAlgTemplate>
259 1 : HcclResult InsAlltoAllSoleExecutor<AlgTopoMatch, InsAlgTemplate>::OrchestrateOpbase(InsAlgTemplate &tempAlg)
260 : {
261 3 : HCCL_DEBUG("[CollAlgFactory][InsAlltoAllSoleExecutor] AlgTemplate is [%s]", tempAlg.Describe().c_str());
262 1 : CHK_PRT_RET(maxTmpMemSize_ == 0,
263 : HCCL_ERROR("[InsAlltoAllSoleExecutor] Rank [%d], Invalid input maxTmpMemSize [%u].", myRank_, maxTmpMemSize_),
264 : HcclResult::HCCL_E_PARA);
265 :
266 1 : CHK_RET(tempAlg.GetScratchBufferInfo(maxTmpMemSize_, sendType_));
267 :
268 1 : BuffInfo buffInfo;
269 1 : buffInfo.inBuffType = BufferType::SCRATCH;
270 1 : buffInfo.outBuffType = BufferType::SCRATCH;
271 1 : buffInfo.inBuffBaseOff = 0;
272 1 : buffInfo.outBuffBaseOff = maxTmpMemSize_ / 2; // 占据scratch memory的后半部分,除以2
273 1 : RankSliceInfo sliceInfoVec;
274 :
275 1 : TempFuncs tempFuncs;
276 1 : tempFuncs.opMode = opMode_;
277 1 : tempFuncs.enableCounterNotify = IsEnableCounterNotifyByDevType(myRank_, devType_);
278 1 : tempFuncs.isForepart = true; // Usr Buff to CCL Buff required
279 1 : tempFuncs.isBottom = true; // CCL Buff to Usr Buff required
280 1 : CHK_RET(tempAlg.Run(tempFuncs, sliceInfoVec, buffInfo, tempResLinks_, requiredQue_));
281 3 : HCCL_INFO("[InsAlltoAllSoleExecutor][OrchestrateOpbase] Run templet success.");
282 :
283 1 : return HcclResult::HCCL_SUCCESS;
284 1 : }
285 :
286 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, InsAlltoAllMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
287 : InsTempAlltoAllMesh);
288 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, InsAlltoAllvMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
289 : InsTempAlltoAllMesh);
290 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLVC, InsAlltoAllvcMesh, InsAlltoAllSoleExecutor, TopoMatchMesh,
291 : InsTempAlltoAllMesh);
292 : #ifndef CCL_KERNEL_AICPU
293 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, CcuAlltoAllMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
294 : CcuTempAllToAllMesh1D);
295 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
296 : CcuTempAlltoAllVMesh1D);
297 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALL, CcuAlltoAllMesh2D, InsAlltoAllSoleExecutor, TopoMatchConcurrMesh,
298 : CcuTempAlltoAllMesh2D);
299 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh2D, InsAlltoAllSoleExecutor, TopoMatchConcurrMesh,
300 : CcuTempAlltoAllVMesh2D);
301 : INS_REGISTER_IMPL_BY_TEMP(OpType::HALFALLTOALLV, CcuHalfAll2AllVMesh1D, InsAlltoAllSoleExecutor, TopoMatchMesh,
302 : CcuTempHalfAllToAllVMesh1D);
303 : INS_REGISTER_IMPL_BY_TEMP(OpType::ALLTOALLV, CcuAlltoAllVMesh2Die, InsAlltoAllSoleExecutor, TopoMatchMesh,
304 : CcuTempAlltoAllVMesh2Die);
305 : #endif
306 :
307 : } // namespace Hccl
|