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_all_to_all_v_direct_fullmesh_executor.h"
12 :
13 : namespace hccl {
14 :
15 0 : CollRunAlltoAllDirectFullmesh::CollRunAlltoAllDirectFullmesh(const HcclDispatcher dispatcher,
16 0 : std::unique_ptr<TopoMatcher> &topoMatcher)
17 0 : : CollAlltoAllExecutor(dispatcher, topoMatcher)
18 : {
19 0 : }
20 :
21 0 : HcclResult CollRunAlltoAllDirectFullmesh::Orchestrate(OpParam& param, AlgResourceResponse& algRes)
22 : {
23 0 : HcclUs startut = TIME_NOW();
24 0 : HcclResult ret = HCCL_SUCCESS;
25 0 : tag_ = param.tag;
26 0 : algResResp_ = &algRes;
27 0 : AlltoAllVParam_ = param;
28 :
29 0 : ExecMem execMem;
30 0 : execMem.count = 0;
31 0 : execMem.inputPtr = param.inputPtr;
32 0 : execMem.outputPtr = param.outputPtr;
33 0 : execMem.inputMem = algRes.cclInputMem;
34 0 : execMem.outputMem = algRes.cclOutputMem;
35 0 : ret = KernelRun(param, execMem);
36 :
37 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
38 : HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][Orchestrate]errNo[0x%016llx]executor run failed",
39 : HCCL_ERROR_CODE(ret)), ret);
40 :
41 : // Enforce task launch at the end of Orchestrate
42 : // 注意: 不要删除这里的强制launch, 否则会导致aicpu cache功能问题
43 0 : HCCL_INFO("%s: enforce task launch at the end of Orchestrate", __func__);
44 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
45 :
46 0 : HCCL_INFO("tag[%s], AlltoAllDirectFullmesh tempAlg orchestrate success, take time [%lld]us.",
47 : param.tag.c_str(), DURATION_US(TIME_NOW() - startut));
48 0 : return HCCL_SUCCESS;
49 0 : }
50 :
51 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetAdjInfo(AlgResourceResponse& algRes, AdjInfo& adjInfo)
52 : {
53 0 : HCCL_INFO("[GetAdjInfo-nslbdp] GetAdjInfo.");
54 0 : algResResp_ = &algRes;
55 0 : SubCommInfo levelCommInfo = {0};
56 0 : AdjInfo nslbAdjInfo = {0};
57 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
58 :
59 0 : u32 localRank= topoAttr_.userRank;
60 0 : u32 localRankSize = topoAttr_.userRankSize;
61 :
62 0 : std::unique_ptr<AlgTemplateBase> levelTempAlg;
63 :
64 0 : HCCL_INFO("[GetAdjInfo-nslbdp] SelectTempAlg.");
65 0 : levelTempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
66 0 : TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
67 0 : CHK_SMART_PTR_NULL(levelTempAlg);
68 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
69 0 : CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
70 :
71 0 : if (devNumInlocalPod == INVALID_VALUE_RANKSIZE) {
72 0 : HCCL_INFO("[GetAdjInfo-nslbdp] devNumInlocalPod == INVALID_VALUE_RANKSIZE.");
73 0 : return HCCL_SUCCESS;
74 : }
75 :
76 0 : nslbAdjInfo.dstRankNum = devNumInlocalPod;
77 0 : CHK_RET(levelTempAlg->GetNslbAdjInfo(localRank, localRankSize, levelCommInfo.links, nslbAdjInfo));
78 :
79 0 : adjInfo.dstRankNum = nslbAdjInfo.dstRankNum;
80 0 : HCCL_INFO("[GetAdjInfo-nslbdp] adjInfo.dstRankNum[%u].", adjInfo.dstRankNum);
81 :
82 0 : for (size_t i = 0; i < nslbAdjInfo.nsAdjInfo.size(); i++) {
83 0 : NslbDpAdjInfo dpAdjInfo = {0};
84 0 : dpAdjInfo.dstLocalRankId = nslbAdjInfo.nsAdjInfo[i].dstLocalRankId;
85 0 : dpAdjInfo.phaseId = nslbAdjInfo.nsAdjInfo[i].phaseId;
86 0 : dpAdjInfo.rev = 0;
87 0 : adjInfo.nsAdjInfo.push_back(dpAdjInfo);
88 0 : HCCL_INFO("[nslbdp]GetAdjInfo dstLocalRankId[%u], phaseId[%u].",
89 : nslbAdjInfo.nsAdjInfo[i].dstLocalRankId, nslbAdjInfo.nsAdjInfo[i].phaseId);
90 : }
91 0 : return HCCL_SUCCESS;
92 0 : }
93 :
94 0 : HcclResult CollRunAlltoAllDirectFullmesh::MarkNeedAlltoallvCache() {
95 0 : needAlltoallvCache_ = true;
96 0 : HCCL_INFO("[CollRunAlltoAllDirectFullmesh][MarkNeedAlltoallvCache] set needAlltoallvCache_[%u]"\
97 : "for alltoallv aicpu cache", needAlltoallvCache_);
98 0 : return HCCL_SUCCESS;
99 : }
100 :
101 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetHcclOffsetDstRanksMap(std::unordered_map<uint64_t,
102 : std::vector<uint32_t>>& hcclOffsetDstRanksMap) const {
103 0 : hcclOffsetDstRanksMap.clear();
104 0 : hcclOffsetDstRanksMap = hcclOffsetDstRanksMap_; // Deep copy
105 :
106 0 : return HCCL_SUCCESS;
107 : }
108 :
109 0 : HcclOpMetaInfo CollRunAlltoAllDirectFullmesh::GetOpMeta(HcclCMDType opType, const u64 size)
110 : {
111 : (void)opType;
112 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, size, true);
113 0 : return opMeta;
114 : }
115 :
116 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSDMAGroupInfo(const u32 userRank,
117 : u32& devNumInlocalPod, u32& rankIdxInPod)
118 : {
119 : (void) userRank;
120 0 : bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B &&
121 0 : !topoAttr_.isSingleMeshAggregation;
122 0 : if (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) {
123 0 : CHK_RET(topoMatcher_->GetLocalServerRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
124 : } else {
125 0 : CHK_RET(topoMatcher_->GetLocalSuperPodRankSize(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
126 : }
127 0 : CHK_PRT_RET(devNumInlocalPod == INVALID_VALUE_RANKSIZE,
128 : HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][GetLocalSDMAGroupInfo]get local superPod total ranksize failed."),
129 : HCCL_E_PARA);
130 0 : return HCCL_SUCCESS;
131 : }
132 :
133 0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcStreamNum(u32& streamNum)
134 : {
135 : // 每个超节点内的卡数
136 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
137 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
138 0 : CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
139 :
140 : // 单超节点场景需要的从流数量
141 0 : streamNum = (devNumInlocalPod > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
142 0 : (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE * RANK_SET_COMPUTE_CONST) : (devNumInlocalPod * RANK_SET_COMPUTE_CONST);
143 :
144 : // 多超节点场景下,RDMA会设置独立的并发度
145 0 : if ((topoAttr_.userRankSize - devNumInlocalPod) > 0) {
146 0 : streamNum += 1; // 一条从流专门用来管理超节点间的RDMA通信
147 0 : u32 totalRdmaRankNum = topoAttr_.userRankSize - devNumInlocalPod;
148 0 : streamNum += (totalRdmaRankNum > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
149 0 : (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) : (totalRdmaRankNum);
150 : }
151 :
152 0 : HCCL_INFO("[CollRunAlltoAllDirectFullmesh][CalcStreamNum] tag[%s] streamNum[%u]",
153 : tag_.c_str(), streamNum);
154 0 : return HCCL_SUCCESS;
155 : }
156 :
157 : // level0-level1 打平fullmesh
158 : // 超节点内建SDMA链路;超节点间建RDMA链路
159 0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcLevel0CommInfo(TransportMemType inputType, TransportMemType outputType,
160 : std::vector<LevelNSubCommTransport>& opTransport)
161 : {
162 0 : CommParaInfo commCombinePara(COMM_COMBINE_ORDER, CommType::COMM_TAG_MESH);
163 0 : CHK_RET(CalcCommPlaneInfo(tag_, commCombinePara, opTransport[COMM_COMBINE_ORDER], inputType, outputType));
164 :
165 0 : LevelNSubCommTransport &commTransportLevel0 = opTransport[COMM_COMBINE_ORDER];
166 0 : for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
167 0 : for (auto &transportRequest : commTransportLevel0[subCommIndex].transportRequests) {
168 0 : transportRequest.isUsedRdma = topoAttr_.isUsedRdmaMap.at(transportRequest.remoteUserRank);
169 : }
170 : }
171 0 : return HCCL_SUCCESS;
172 0 : }
173 :
174 0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcTransportMemType(TransportMemType &inputType, TransportMemType &outputType)
175 : {
176 0 : inputType = TransportMemType::CCL_INPUT;
177 0 : outputType = TransportMemType::CCL_OUTPUT;
178 :
179 0 : HCCL_INFO("[CollRunAlltoAllDirectFullmesh][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
180 : tag_.c_str(), inputType, outputType);
181 0 : return HCCL_SUCCESS;
182 : }
183 :
184 0 : HcclResult CollRunAlltoAllDirectFullmesh::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
185 : {
186 0 : TransportMemType inputType = TransportMemType::RESERVED;
187 0 : TransportMemType outputType = TransportMemType::RESERVED;
188 :
189 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
190 : // level0 - level1 全连接通信域
191 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
192 0 : return HCCL_SUCCESS;
193 : }
194 :
195 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoallV(const OpParam ¶m)
196 : {
197 : // 注意: 如果send/recv info的计算逻辑发生变化, 需要同步修改framework下的IsSmallDataAlltoallv()函数
198 0 : for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
199 0 : u64 curSendCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCounts) + j);
200 0 : u64 curSendDispls = *(static_cast<const u64 *>(param.All2AllDataDes.sdispls) + j);
201 0 : localSendRecvInfo_.sendCounts[j] = curSendCounts;
202 0 : localSendRecvInfo_.sendDispls[j] = curSendDispls;
203 0 : localSendRecvInfo_.sendLength[j] = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
204 0 : localSendRecvInfo_.sendOffset[j] = curSendDispls * SIZE_TABLE[param.All2AllDataDes.sendType];
205 :
206 0 : u64 curRecvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.recvCounts) + j);
207 0 : u64 curRecvDispls = *(static_cast<const u64 *>(param.All2AllDataDes.rdispls) + j);
208 0 : localSendRecvInfo_.recvCounts[j] = curRecvCounts;
209 0 : localSendRecvInfo_.recvDispls[j] = curRecvDispls;
210 0 : localSendRecvInfo_.recvLength[j] = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
211 0 : localSendRecvInfo_.recvOffset[j] = curRecvDispls * SIZE_TABLE[param.All2AllDataDes.recvType];
212 :
213 0 : HCCL_DEBUG("GetLocalSendRecvInfoforAlltoallV rank[%u], sendCounts[%llu], sendDispls[%llu] "\
214 : "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
215 : localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
216 : localSendRecvInfo_.recvDispls[j]);
217 : }
218 0 : return HCCL_SUCCESS;
219 : }
220 :
221 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoall(const OpParam ¶m)
222 : {
223 0 : u64 curSendDispls = 0;
224 0 : u64 curSendOffset = 0;
225 0 : u64 curRecvDispls = 0;
226 0 : u64 curRecvOffset = 0;
227 0 : for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
228 0 : u64 curSendCounts = param.All2AllDataDes.sendCount;
229 0 : u64 curSendLength = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
230 0 : localSendRecvInfo_.sendCounts[j] = curSendCounts;
231 0 : localSendRecvInfo_.sendDispls[j] = curSendDispls;
232 0 : localSendRecvInfo_.sendLength[j] = curSendLength;
233 0 : localSendRecvInfo_.sendOffset[j] = curSendOffset;
234 0 : curSendDispls += curSendCounts;
235 0 : curSendOffset += curSendLength;
236 :
237 0 : u64 curRecvCounts = param.All2AllDataDes.sendCount;
238 0 : u64 curRecvLength = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
239 0 : localSendRecvInfo_.recvCounts[j] = curRecvCounts;
240 0 : localSendRecvInfo_.recvDispls[j] = curRecvDispls;
241 0 : localSendRecvInfo_.recvLength[j] = curRecvLength;
242 0 : localSendRecvInfo_.recvOffset[j] = curRecvOffset;
243 0 : curRecvDispls += curRecvCounts;
244 0 : curRecvOffset += curRecvLength;
245 0 : HCCL_DEBUG("GetLocalSendRecvInfoforAlltoAll rank[%u], sendCounts[%llu], sendDispls[%llu] "\
246 : "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
247 : localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
248 : localSendRecvInfo_.recvDispls[j]);
249 : }
250 0 : return HCCL_SUCCESS;
251 : }
252 :
253 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetLocalSendRecvInfoforAlltoallVC(const OpParam ¶m)
254 : {
255 0 : u64 curSendDispls = 0;
256 0 : u64 curSendOffset = 0;
257 0 : u64 curRecvDispls = 0;
258 0 : u64 curRecvOffset = 0;
259 0 : u64 rankSize = topoAttr_.userRankSize;
260 0 : u64 usrRank = topoAttr_.userRank;
261 0 : for (u32 j = 0; j < topoAttr_.userRankSize; j++) {
262 0 : u64 curSendCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank * rankSize + j);
263 0 : u64 curSendLength = curSendCounts * SIZE_TABLE[param.All2AllDataDes.sendType];
264 0 : localSendRecvInfo_.sendCounts[j] = curSendCounts;
265 0 : localSendRecvInfo_.sendDispls[j] = curSendDispls;
266 0 : localSendRecvInfo_.sendLength[j] = curSendLength;
267 0 : localSendRecvInfo_.sendOffset[j] = curSendOffset;
268 0 : curSendDispls += curSendCounts;
269 0 : curSendOffset += curSendLength;
270 :
271 0 : u64 curRecvCounts = *(static_cast<const u64 *>(param.All2AllDataDes.sendCountMatrix) + usrRank + rankSize * j);
272 0 : u64 curRecvLength = curRecvCounts * SIZE_TABLE[param.All2AllDataDes.recvType];
273 0 : localSendRecvInfo_.recvCounts[j] = curRecvCounts;
274 0 : localSendRecvInfo_.recvDispls[j] = curRecvDispls;
275 0 : localSendRecvInfo_.recvLength[j] = curRecvLength;
276 0 : localSendRecvInfo_.recvOffset[j] = curRecvOffset;
277 0 : curRecvDispls += curRecvCounts;
278 0 : curRecvOffset += curRecvLength;
279 0 : HCCL_DEBUG("GetLocalSendRecvInfoforAlltoallVC rank[%u], sendCounts[%llu], sendDispls[%llu] "\
280 : "recvCounts[%llu], recvDispls[%llu]", topoAttr_.userRank, localSendRecvInfo_.sendCounts[j],
281 : localSendRecvInfo_.sendDispls[j], localSendRecvInfo_.recvCounts[j],
282 : localSendRecvInfo_.recvDispls[j]);
283 : }
284 0 : return HCCL_SUCCESS;
285 : }
286 :
287 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetAlltoAllvTmpRankSendRecvInfo(const OpParam ¶m)
288 : {
289 0 : localSendRecvInfo_.sendCounts.resize(topoAttr_.userRankSize, 0);
290 0 : localSendRecvInfo_.sendDispls.resize(topoAttr_.userRankSize, 0);
291 0 : localSendRecvInfo_.sendLength.resize(topoAttr_.userRankSize, 0);
292 0 : localSendRecvInfo_.sendOffset.resize(topoAttr_.userRankSize, 0);
293 :
294 0 : localSendRecvInfo_.recvCounts.resize(topoAttr_.userRankSize, 0);
295 0 : localSendRecvInfo_.recvDispls.resize(topoAttr_.userRankSize, 0);
296 0 : localSendRecvInfo_.recvLength.resize(topoAttr_.userRankSize, 0);
297 0 : localSendRecvInfo_.recvOffset.resize(topoAttr_.userRankSize, 0);
298 0 : if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
299 0 : CHK_RET(GetLocalSendRecvInfoforAlltoallV(param));
300 0 : } else if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALL) {
301 0 : CHK_RET(GetLocalSendRecvInfoforAlltoall(param));
302 0 : } else if (param.opType == HcclCMDType::HCCL_CMD_ALLTOALLVC) {
303 0 : CHK_RET(GetLocalSendRecvInfoforAlltoallVC(param));
304 : } else {
305 0 : HCCL_ERROR("Only support optype AllToAll , AllToAllV and AllToAllVC !");
306 : }
307 0 : return HCCL_SUCCESS;
308 : }
309 :
310 0 : HcclResult CollRunAlltoAllDirectFullmesh::KernelRun(const OpParam ¶m, ExecMem &execMem)
311 : {
312 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] AllToAll fullmesh start.", __func__);
313 :
314 : // 准备数据
315 0 : CHK_RET(ActiveSlaveStreams(param.stream));
316 0 : CHK_RET(GetAlltoAllvTmpRankSendRecvInfo(param));
317 :
318 : // 获取当前超节点内总卡数
319 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
320 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
321 0 : CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
322 :
323 : // 获取通信域
324 0 : CHK_RET(CheckCommSize(COMM_COMBINE_ORDER, COMM_INDEX_0 + 1));
325 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_COMBINE_ORDER, COMM_INDEX_0);
326 0 : bool isA2MultiModule = topoAttr_.deviceType == DevType::DEV_TYPE_910B &&
327 0 : !topoAttr_.isSingleMeshAggregation;
328 : // isSuPodAsym 表示A2A3卡数不一致场景或者A3多超节点server数不同场景
329 0 : bool isSuPodAsym = false;
330 0 : if (topoAttr_.superPodNum > 1) {
331 0 : isSuPodAsym = (static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode) || static_cast<bool>(topoAttr_.multiSuperPodDiffServerNumMode));
332 : } else {
333 0 : isSuPodAsym = (static_cast<bool>(topoMatcher_->GetExternalInputInterHccsDisable()) || isA2MultiModule) && static_cast<bool>(topoAttr_.multiModuleDiffDeviceNumMode);
334 : }
335 :
336 : // 执行
337 : // 注意: 如果使用了非AlltoAllVDirectFullMesh的算法模板, 需要同步修改framework中的NeedOpUnfoldCache()函数
338 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
339 0 : TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
340 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH in COMM_COMBINE_ORDER", __func__);
341 0 : CHK_SMART_PTR_NULL(tempAlg);
342 :
343 0 : PrepareData prepareData;
344 0 : prepareData.stream = param.stream;
345 0 : prepareData.userRank = topoAttr_.userRank;
346 0 : prepareData.userRankSize = topoAttr_.userRankSize;
347 0 : prepareData.linksPtr = &level0CommInfo.links;
348 0 : prepareData.localSendRecvInfoPtr = &localSendRecvInfo_;
349 0 : prepareData.devNumInlocalPod = devNumInlocalPod;
350 0 : prepareData.rankIdxInPod = rankIdxInPod;
351 :
352 0 : prepareData.inputMem = algResResp_->paramInputMem;
353 0 : prepareData.outputMem = algResResp_->paramOutputMem;
354 0 : prepareData.cclInMem = execMem.inputMem;
355 0 : prepareData.cclOutMem = execMem.outputMem;
356 0 : prepareData.workMode = workflowMode_;
357 0 : prepareData.subStreamsPtr = &algResResp_->slaveStreams;
358 0 : prepareData.signalPtr = &algResResp_->notifiesMain;
359 0 : prepareData.signalAuxPtr = &algResResp_->notifiesAux;
360 0 : prepareData.isSuPodAsym = isSuPodAsym;
361 0 : prepareData.opType = param.opType;
362 0 : prepareData.algOpContext = algOpContext_;
363 :
364 : // 如果使能alltoallv aicpu cache
365 0 : if (needAlltoallvCache_) {
366 : // 注意: 一定是alltoallv类算子才有可能设置needAlltoallvCache_, 让alltoallv temp alg感知cache并保存算法中间结果
367 0 : CHK_PRT_RET(!(param.opType == HcclCMDType::HCCL_CMD_ALLTOALLV || param.opType == HcclCMDType::HCCL_CMD_ALLTOALLVC),
368 : HCCL_ERROR("[CollRunAlltoAllDirectFullmesh][KernelRun] needAlltoallvCache_[%u] opType[%u]",
369 : needAlltoallvCache_, param.opType),
370 : HCCL_E_INTERNAL);
371 :
372 : // 使能alltoallv temp alg感知alltoallv aicpu cache
373 0 : prepareData.needAlltoallvCache = true;
374 : } else {
375 0 : prepareData.needAlltoallvCache = false;
376 : }
377 :
378 0 : CHK_RET(tempAlg->Prepare(prepareData));
379 :
380 0 : CHK_RET(tempAlg->RunAsync());
381 :
382 0 : if (needAlltoallvCache_) {
383 : // 在tempAlg被销毁前保存hcclOffset-dstRank之间的mapping信息
384 : // 注意: CollRunAlltoAllDirectFullmesh executor使用的一定是AlltoAllVDirectFullMesh template
385 0 : hcclOffsetDstRanksMap_.clear();
386 0 : HCCL_INFO("[CollRunAlltoAllDirectFullmesh][KernelRun] get hcclOffset-dstRanks mapping for AlltoAllVDirectFullMesh");
387 0 : CHK_RET(tempAlg->GetHcclOffsetDstRanksMap(hcclOffsetDstRanksMap_));
388 : }
389 :
390 0 : HCCL_INFO("[CollRunAlltoAllDirectFullmesh] executor run success.");
391 0 : if (algOpContext_.opRetryHandler.isPostSync == true) {
392 0 : OpParam postSyncParam = param;
393 0 : if ((*prepareData.subStreamsPtr).size() == 0) {
394 0 : CHK_RET(PostSyncWithoutSubstream(postSyncParam, execMem));
395 : } else {
396 0 : PrepareData postSyncPrepareData = prepareData;
397 0 : CHK_RET(PostSyncWithSubstream(postSyncParam, execMem, postSyncPrepareData));
398 0 : }
399 0 : }
400 0 : return HCCL_SUCCESS;
401 0 : }
402 0 : HcclResult CollRunAlltoAllDirectFullmesh::Getlevel1CommRank(SubCommInfo& level1CommInfo)
403 : {
404 0 : HCCL_INFO("[GetAdjInfo-nslbdp] Getlevel1CommRank userRank[%u]--userRankSize[%u].",topoAttr_.userRank, topoAttr_.userRankSize);
405 0 : level1CommInfo.localRank = topoAttr_.userRank;
406 0 : level1CommInfo.localRankSize = topoAttr_.userRankSize;
407 0 : return HCCL_SUCCESS;
408 : }
409 :
410 0 : HcclResult CollRunAlltoAllDirectFullmesh::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
411 : {
412 : (void) level1RankSize;
413 0 : HCCL_INFO("[GetAdjInfo-nslbdp] SelectTempAlg.");
414 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
415 0 : TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, dispatcher_);
416 0 : CHK_SMART_PTR_NULL(level1TempAlg);
417 :
418 0 : return HCCL_SUCCESS;
419 : }
420 :
421 0 : HcclResult CollRunAlltoAllDirectFullmesh::GetDevNumInlocalPod(u32& devNumInlocalPod)
422 : {
423 0 : HCCL_INFO("[GetAdjInfo-nslbdp] GetDevNumInlocalPod.");
424 : // 获取当前超节点内总卡数
425 0 : u32 rankIdxInPod = INVALID_VALUE_RANKID;
426 0 : CHK_RET(GetLocalSDMAGroupInfo(topoAttr_.userRank, devNumInlocalPod, rankIdxInPod));
427 :
428 0 : return HCCL_SUCCESS;
429 : }
430 : REGISTER_EXEC("RunAlltoAllDirectFullmesh", AlltoAllVDirectFullMesh, CollRunAlltoAllDirectFullmesh);
431 : } // namespace hccl
|