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