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_executor.h"
12 : #include "device_capacity.h"
13 :
14 : namespace hccl {
15 :
16 5 : CollAlltoAllExecutor::CollAlltoAllExecutor(const HcclDispatcher dispatcher,
17 5 : std::unique_ptr<TopoMatcher> &topoMatcher)
18 5 : : CollNativeExecutorBase(dispatcher, topoMatcher)
19 : {
20 5 : }
21 :
22 0 : HcclResult CollAlltoAllExecutor::Orchestrate(OpParam& param, AlgResourceResponse& algRes)
23 : {
24 0 : HcclUs startut = TIME_NOW();
25 0 : tag_ = param.tag;
26 0 : algResResp_ = &algRes;
27 0 : AlltoAllVParam_ = param;
28 0 : ExecMem execMem;
29 0 : execMem.count = 0;
30 0 : execMem.inputPtr = param.inputPtr;
31 0 : execMem.outputPtr = param.outputPtr;
32 :
33 0 : HcclResult ret = HCCL_SUCCESS;
34 0 : if (GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
35 0 : execMem.inputMem = algRes.cclInputMem;
36 0 : execMem.outputMem = algRes.cclOutputMem;
37 0 : execMem.scratchMem = algRes.scratchMem;
38 :
39 0 : auto opMeta = GetOpMeta(param.opType, algRes.paramInputMem.size()); // override
40 0 : CHK_RET(InitTask(dispatcher_, param.stream, opMeta.isEnableCache, opMeta.GetCacheKey()));
41 0 : bool massTasks = HasMassTasks(allMeshAggregationSendRecvInfo_);
42 0 : if (massTasks) {
43 0 : CHK_RET(SetNormalMode(dispatcher_));
44 : }
45 0 : ret = KernelRun(param, execMem);
46 : } else {
47 0 : execMem.inputMem = algRes.paramInputMem;
48 0 : execMem.outputMem = algRes.paramOutputMem;
49 0 : execMem.scratchMem = algRes.scratchMem;
50 0 : ret = KernelRun(param, execMem);
51 : }
52 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
53 : HCCL_ERROR("[CollAlltoAllExecutor][Orchestrate]errNo[0x%016llx]executor run failed",
54 : HCCL_ERROR_CODE(ret)), ret);
55 :
56 : // Enforce task launch at the end of Orchestrate
57 : // 注意: 不要删除这里的强制launch, 否则会导致aicpu cache功能问题
58 0 : HCCL_INFO("%s: enforce task launch at the end of Orchestrate", __func__);
59 0 : CHK_RET(LaunchTaskExtend(dispatcher_, param.stream, algResResp_->slaveStreams));
60 :
61 0 : HCCL_INFO("tag[%s], AlltoAll executor orchestrate success, take time [%lld]us.",
62 : param.tag.c_str(), DURATION_US(TIME_NOW() - startut));
63 0 : return HCCL_SUCCESS;
64 0 : }
65 :
66 0 : HcclResult CollAlltoAllExecutor::GetAdjInfo(AlgResourceResponse& algRes, AdjInfo& adjInfo)
67 : {
68 0 : algResResp_ = &algRes;
69 0 : SubCommInfo levelCommInfo = {0};
70 0 : AdjInfo nslbAdjInfo = {0};
71 0 : u32 devNumInlocalPod = INVALID_VALUE_RANKSIZE;
72 :
73 0 : if (Getlevel1CommRank(levelCommInfo) != HCCL_SUCCESS) {
74 0 : return HCCL_SUCCESS;
75 : }
76 0 : u32 localRank= levelCommInfo.localRank;
77 0 : u32 localRankSize = levelCommInfo.localRankSize;
78 :
79 0 : std::unique_ptr<AlgTemplateBase> levelTempAlg;
80 0 : if (SelectTempAlg(levelTempAlg, localRankSize) != HCCL_SUCCESS) {
81 0 : return HCCL_SUCCESS;
82 : }
83 0 : GetDevNumInlocalPod(devNumInlocalPod);
84 0 : if (devNumInlocalPod == INVALID_VALUE_RANKSIZE) {
85 0 : HCCL_INFO("[GetAdjInfo-nslbdp] devNumInlocalPod == INVALID_VALUE_RANKSIZE.");
86 0 : return HCCL_SUCCESS;
87 : }
88 :
89 0 : nslbAdjInfo.dstRankNum = devNumInlocalPod;
90 0 : CHK_RET(levelTempAlg->GetNslbAdjInfo(localRank, localRankSize, levelCommInfo.links, nslbAdjInfo));
91 :
92 0 : adjInfo.dstRankNum = nslbAdjInfo.dstRankNum;
93 0 : HCCL_INFO("[GetAdjInfo-nslbdp] adjInfo.dstRankNum[%u].", adjInfo.dstRankNum);
94 :
95 0 : for (size_t i = 0; i < nslbAdjInfo.nsAdjInfo.size(); i++) {
96 0 : NslbDpAdjInfo dpAdjInfo = {0};
97 0 : dpAdjInfo.dstLocalRankId = nslbAdjInfo.nsAdjInfo[i].dstLocalRankId;
98 0 : dpAdjInfo.phaseId = nslbAdjInfo.nsAdjInfo[i].phaseId;
99 0 : dpAdjInfo.rev = 0;
100 0 : adjInfo.nsAdjInfo.push_back(dpAdjInfo);
101 0 : HCCL_INFO("[nslbdp]GetAdjInfo dstLocalRankId[%u], phaseId[%u].",
102 : nslbAdjInfo.nsAdjInfo[i].dstLocalRankId, nslbAdjInfo.nsAdjInfo[i].phaseId);
103 : }
104 0 : return HCCL_SUCCESS;
105 0 : }
106 :
107 : // override----------------------资源计算接口----------------------
108 5 : HcclResult CollAlltoAllExecutor::CalcResRequest(const OpParam& param, AlgResourceRequest& resourceRequest)
109 : {
110 5 : (void)ParseParam(param);
111 :
112 5 : u64 scratchMemSize = 0U;
113 5 : u32 streamNum = 0U;
114 5 : u32 notifyNum = 0U;
115 5 : u64 aivBufferRequest = 0U;
116 : std::vector<LevelNSubCommTransport> opTransport {
117 0 : std::vector<LevelNSubCommTransport>(static_cast<u32>(COMM_LEVEL_RESERVED))
118 5 : };
119 :
120 : // AICPU aicpuUnfold展开模式下临时强制OP_BASE,使整个资源计算路径与AICPU侧一致
121 : // CalcScratchMemSize走OP_BASE分支正确计算scratch
122 : // CalcCommInfo走OP_BASE分支设outputMemType为CCL_OUTPUT而非SCRATCH
123 : // 避免transport因scratch未分配而拿到nullptr
124 : // AIV executor有独立的资源计算逻辑,不需要force OP_BASE
125 5 : const bool needForceOpBase = param.aicpuUnfoldMode && !param.isZeroCopy && !desc_.isAivMode;
126 5 : HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest] aicpuUnfoldMode[%d] isZeroCopy[%d] "
127 : "needForceOpBase[%d] workflowMode[%d] tag[%s]",
128 : param.aicpuUnfoldMode, param.isZeroCopy, needForceOpBase,
129 : workflowMode_, param.tag.c_str());
130 5 : const HcclWorkflowMode savedWorkflowMode = workflowMode_;
131 5 : if (needForceOpBase) {
132 0 : HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest] aicpuUnfoldMode force OpBase, "
133 : "originalWorkflowMode[%d], tag[%s]",
134 : savedWorkflowMode, param.tag.c_str());
135 0 : workflowMode_ = HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
136 : }
137 :
138 5 : CHK_RET(CalcScratchMemSize(scratchMemSize));
139 5 : CHK_RET(CalcStreamNum(streamNum));
140 5 : CHK_RET(CalcNotifyNum(streamNum, notifyNum));
141 5 : CHK_RET(CalcAivBufferRequest(aivBufferRequest));
142 5 : CHK_RET(CalcCommInfo(opTransport));
143 :
144 5 : if (needForceOpBase) {
145 0 : HCCL_DEBUG("[CollAlltoAllExecutor][CalcResRequest] restore workflowMode "
146 : "after resource calc, scratchMemSize[%llu]", scratchMemSize);
147 0 : workflowMode_ = savedWorkflowMode;
148 : }
149 :
150 5 : CHK_RET(BuildResourceRequest(scratchMemSize, streamNum, notifyNum, aivBufferRequest, opTransport, resourceRequest));
151 5 : HCCL_INFO("[CollAlltoAllExecutor][%s] streamNum[%u], notifyNum[%u], sctrachMemSize[%llu], aivBufferRequest[%llu]",
152 : __func__, resourceRequest.streamNum, resourceRequest.notifyNum, resourceRequest.scratchMemSize,
153 : resourceRequest.aivBufferRequest);
154 : // 打印建链诉求
155 85 : for (u32 levelIndex = 0; levelIndex < COMM_LEVEL_RESERVED; levelIndex++) {
156 80 : LevelNSubCommTransport &levelTransport = resourceRequest.opTransport[levelIndex];
157 80 : u32 ringSize = levelTransport.size();
158 85 : for (u32 ringIndex = 0; ringIndex < ringSize; ringIndex++) {
159 5 : SingleSubCommTransport &subCommTransport = levelTransport[ringIndex];
160 5 : u32 rankSize = subCommTransport.transportRequests.size();
161 10 : for (u32 rankIndex = 0; rankIndex < rankSize; rankIndex++) {
162 5 : if (subCommTransport.transportRequests[rankIndex].isValid == true) {
163 0 : HCCL_INFO("[CollAlltoAllExecutor][CalcResRequest]" \
164 : "levelIndex[%u], ringIndex[%u], rankIndex[%u], userRank[%u], remoteRank[%u]" \
165 : "isUsedRdma[%d]",
166 : levelIndex, ringIndex, rankIndex, subCommTransport.transportRequests[rankIndex].localUserRank,
167 : subCommTransport.transportRequests[rankIndex].remoteUserRank,
168 : subCommTransport.transportRequests[rankIndex].isUsedRdma);
169 : }
170 : }
171 : }
172 : }
173 5 : CHK_RET(CheckNeedCreateVirtualLinks(resourceRequest));
174 5 : HCCL_DEBUG("[%s] process success", __func__);
175 5 : return HCCL_SUCCESS;
176 5 : }
177 :
178 5 : HcclResult CollAlltoAllExecutor::CheckNeedCreateVirtualLinks(AlgResourceRequest &resourceRequest)
179 : {
180 5 : return HCCL_SUCCESS;
181 : }
182 :
183 0 : HcclResult CollAlltoAllExecutor::SetExcutorExtraInfo(const std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, u64 cclbufferSize)
184 : {
185 0 : allMeshAggregationSendRecvInfo_.clear();
186 0 : allMeshAggregationSendRecvInfo_ = allMeshAggregationSendRecvInfo;
187 0 : UpdateAlltoAllZCopyMode(allMeshAggregationSendRecvInfo_, cclbufferSize);
188 0 : HCCL_DEBUG("[%s] allMeshAggregationSendRecvInfo_ size[%u]", __func__, allMeshAggregationSendRecvInfo_.size());
189 :
190 0 : return HCCL_SUCCESS;
191 : }
192 :
193 0 : void CollAlltoAllExecutor::UpdateAlltoAllZCopyMode(std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo, u64 cclbufferSize)
194 : {
195 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
196 0 : u64 maxSendSize = 0;
197 0 : u64 maxRecvSize = 0;
198 0 : for (auto &sendRecvInfo : allMeshAggregationSendRecvInfo) {
199 0 : for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
200 0 : u64 curSendSize = sendRecvInfo.sendLength[i] + sendRecvInfo.sendOffset[i];
201 0 : maxSendSize = std::max(maxSendSize, curSendSize);
202 0 : u64 curRecvSize = sendRecvInfo.recvLength[i] + sendRecvInfo.recvOffset[i];
203 0 : maxRecvSize = std::max(maxRecvSize, curRecvSize);
204 : }
205 : }
206 0 : bool isAlltoAllZCopyMode = (maxSendSize <= cclbufferSize) &&
207 0 : (maxRecvSize <= cclbufferSize);
208 0 : if (isAlltoAllZCopyMode) {
209 0 : isAlltoAllZCopyMode_ = true;
210 : }
211 0 : HCCL_INFO("[CollAlltoAllExecutor][UpdateAlltoAllZCopyMode] maxSendSize[%llu], maxRecvSize[%llu], "\
212 : "cclBufferSize[%llu]", maxSendSize, maxRecvSize, cclbufferSize);
213 : } else {
214 : // 图模式走ZCopy实现
215 0 : isAlltoAllZCopyMode_ = true;
216 : }
217 0 : HCCL_DEBUG("UpdateAlltoAllZCopyMode isAlltoAllZCopyMode_[%d]", isAlltoAllZCopyMode_);
218 0 : }
219 :
220 0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationSendInfo(const AlltoAllUserRankInfo &userRankInfo,
221 : const SendRecvInfo &mySendRecvInfo, const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo,
222 : u32 rankInMeshAggregation, u32 infoIndex, OneSendRecvAddrInfo &curSendInfo, u32 meshAggregationRankSize,
223 : const bool &isSingleMesh)
224 : {
225 0 : if (infoIndex >= mySendRecvInfo.sendOffset.size() || infoIndex >= mySendRecvInfo.sendLength.size()) {
226 0 : HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid infoIndex[%u]", infoIndex);
227 0 : return;
228 : }
229 0 : curSendInfo.localOffset = mySendRecvInfo.sendOffset[infoIndex];
230 0 : curSendInfo.localLength = mySendRecvInfo.sendLength[infoIndex];
231 0 : u64 remoteOffset = 0;
232 :
233 0 : if (isSingleMesh) {
234 0 : remoteOffset = myMeshAggregationSendRecvInfo[infoIndex].recvOffset[userRankInfo.userRank];
235 : } else {
236 0 : for (u32 j = infoIndex % meshAggregationRankSize; j <= infoIndex; j += meshAggregationRankSize) {
237 0 : for (u32 k = 0; k < meshAggregationRankSize; k++) {
238 0 : if (j == infoIndex && k == rankInMeshAggregation) {
239 0 : break;
240 : }
241 0 : if (k < myMeshAggregationSendRecvInfo.size() && j <
242 0 : myMeshAggregationSendRecvInfo[k].sendLength.size()) {
243 0 : remoteOffset += myMeshAggregationSendRecvInfo[k].sendLength[j];
244 : } else {
245 0 : HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] invalid MeshAggregationSendRecvInfo size[%zu]",
246 : myMeshAggregationSendRecvInfo.size());
247 0 : return;
248 : }
249 : }
250 : }
251 : }
252 :
253 0 : curSendInfo.remoteOffset = remoteOffset;
254 0 : curSendInfo.remoteLength = curSendInfo.localLength;
255 0 : HCCL_DEBUG("[CalcIntraMeshAggregationSendInfo] localOffset[%llu], localLength[%llu], "\
256 : "remoteOffset[%llu], remoteLength[%llu]", curSendInfo.localOffset,
257 : curSendInfo.localLength, curSendInfo.remoteOffset, curSendInfo.remoteLength);
258 : }
259 :
260 0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationRecvInfoInMeshAggregation(u32 rankIndex, u32 infoIndex,
261 : const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo, u64 &localOffset, u32 &offsetCounter,
262 : u64 &localLength, u64 &remoteOffset, u32 meshAggregationRankSize)
263 : {
264 : // 这里的判断在外部已经保证了,为了应对coverity sc
265 0 : if (myMeshAggregationSendRecvInfo.size() < meshAggregationRankSize) {
266 0 : HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid myMeshAggregationSendRecvInfo[%zu]",
267 : myMeshAggregationSendRecvInfo.size());
268 0 : return;
269 : }
270 0 : if (myMeshAggregationSendRecvInfo[0].sendLength.size() == 0 ||
271 0 : myMeshAggregationSendRecvInfo[0].sendOffset.size() == 0) {
272 0 : HCCL_ERROR("[CalcIntraMeshAggregationSendInfo] Invalid sendLength size[%zu] or sendOffset size[%zu]",
273 : myMeshAggregationSendRecvInfo[0].sendLength.size(), myMeshAggregationSendRecvInfo[0].sendOffset.size());
274 0 : return;
275 : }
276 0 : for (u32 k = 0; k < meshAggregationRankSize; k++) {
277 0 : if (infoIndex == 0) {
278 0 : localOffset = 0;
279 0 : localLength = myMeshAggregationSendRecvInfo[k].sendLength[rankIndex];
280 0 : remoteOffset = myMeshAggregationSendRecvInfo[k].sendOffset[rankIndex];
281 0 : break;
282 : }
283 :
284 0 : localOffset += myMeshAggregationSendRecvInfo[k].sendLength[rankIndex];
285 0 : offsetCounter++;
286 0 : if (offsetCounter == infoIndex) {
287 0 : if (k == meshAggregationRankSize - 1) {
288 0 : localLength = myMeshAggregationSendRecvInfo[0].sendLength[rankIndex + meshAggregationRankSize];
289 0 : remoteOffset = myMeshAggregationSendRecvInfo[0].sendOffset[rankIndex + meshAggregationRankSize];
290 : } else {
291 0 : localLength = myMeshAggregationSendRecvInfo[k + 1].sendLength[rankIndex];
292 0 : remoteOffset = myMeshAggregationSendRecvInfo[k + 1].sendOffset[rankIndex];
293 : }
294 0 : break;
295 : }
296 : }
297 0 : HCCL_DEBUG("[%s] process success", __func__);
298 : }
299 :
300 0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationRecvInfo(const AlltoAllUserRankInfo &userRankInfo,
301 : const std::vector<SendRecvInfo> &myMeshAggregationSendRecvInfo, u32 infoIndex, OneSendRecvAddrInfo &curRecvInfo,
302 : u32 meshAggregationRankSize, const bool &isSingleMesh)
303 : {
304 0 : u64 localOffset = 0, localLength = 0, remoteLength = 0, remoteOffset = 0;
305 0 : u32 offsetCounter = 0;
306 :
307 0 : if (isSingleMesh) {
308 0 : localOffset = myMeshAggregationSendRecvInfo[userRankInfo.userRank].recvOffset[infoIndex];
309 0 : localLength = myMeshAggregationSendRecvInfo[userRankInfo.userRank].recvLength[infoIndex];
310 0 : remoteLength = myMeshAggregationSendRecvInfo[infoIndex].sendLength[userRankInfo.userRank];
311 0 : remoteOffset = myMeshAggregationSendRecvInfo[infoIndex].sendOffset[userRankInfo.userRank];
312 : } else {
313 0 : for (u32 j = userRankInfo.userRank % meshAggregationRankSize; j < userRankInfo.userRankSize;
314 0 : j += meshAggregationRankSize) {
315 0 : CalcIntraMeshAggregationRecvInfoInMeshAggregation(j, infoIndex, myMeshAggregationSendRecvInfo, localOffset,
316 : offsetCounter, localLength, remoteOffset, meshAggregationRankSize);
317 0 : if (offsetCounter == infoIndex || infoIndex == 0) {
318 : break;
319 : }
320 : }
321 0 : remoteLength = localLength;
322 : }
323 0 : curRecvInfo.localOffset = localOffset;
324 0 : curRecvInfo.localLength = localLength;
325 :
326 0 : curRecvInfo.remoteOffset = remoteOffset;
327 0 : curRecvInfo.remoteLength = remoteLength;
328 0 : HCCL_DEBUG("[CalcIntraMeshAggregationRecvInfo] localOffset[%llu], localLength[%llu], "\
329 : "remoteOffset[%llu], remoteLength[%llu]", localOffset, localLength, remoteOffset, remoteLength);
330 0 : }
331 :
332 0 : void CollAlltoAllExecutor::CalcIntraMeshAggregationAlltoAllMemInfo(const AlltoAllUserRankInfo &userRankInfo,
333 : const std::vector<SendRecvInfo> &allSendRecvInfo,
334 : std::map<u32, std::list<OneSendRecvAddrInfo>> &sendAddrInfosIntra,
335 : std::map<u32, std::list<OneSendRecvAddrInfo>> &recvAddrInfosIntra, u32 meshAggregationRankSize,
336 : const bool &isSingleMesh)
337 : {
338 0 : sendAddrInfosIntra.clear();
339 0 : recvAddrInfosIntra.clear();
340 0 : if (allSendRecvInfo.size() != userRankInfo.userRankSize) {
341 0 : HCCL_ERROR("Invalid All send recv info size[%zu], should be[%u]", allSendRecvInfo.size(),
342 : userRankInfo.userRankSize);
343 0 : return;
344 : }
345 0 : SendRecvInfo mySendRecvInfo = allSendRecvInfo[userRankInfo.userRank];
346 0 : u32 rankInMeshAggregation = userRankInfo.userRank % meshAggregationRankSize;
347 0 : u32 cluserIndex = userRankInfo.userRank / meshAggregationRankSize;
348 0 : auto itBegin = allSendRecvInfo.begin();
349 0 : auto itEnd = allSendRecvInfo.begin();
350 0 : std::advance(itBegin, cluserIndex * meshAggregationRankSize);
351 0 : std::advance(itEnd, (cluserIndex + 1) * meshAggregationRankSize);
352 0 : std::vector<SendRecvInfo> myMeshAggregationSendRecvInfo(itBegin, itEnd);
353 :
354 0 : for (u32 i = 0; i < userRankInfo.userRankSize; i++) {
355 : // sendInfo 的计算
356 : OneSendRecvAddrInfo curSendInfo;
357 0 : u32 remoteRankInMeshAggregation = i % meshAggregationRankSize;
358 0 : CalcIntraMeshAggregationSendInfo(userRankInfo, mySendRecvInfo, myMeshAggregationSendRecvInfo,
359 : rankInMeshAggregation, i, curSendInfo, meshAggregationRankSize, isSingleMesh);
360 0 : sendAddrInfosIntra[remoteRankInMeshAggregation].push_back(curSendInfo);
361 :
362 : // recvInfo 的计算
363 : OneSendRecvAddrInfo curRecvInfo;
364 0 : CalcIntraMeshAggregationRecvInfo(userRankInfo, myMeshAggregationSendRecvInfo, i,
365 : curRecvInfo, meshAggregationRankSize, isSingleMesh);
366 0 : recvAddrInfosIntra[remoteRankInMeshAggregation].push_back(curRecvInfo);
367 : }
368 0 : }
369 :
370 0 : HcclOpMetaInfo CollAlltoAllExecutor::GetOpMeta(HcclCMDType opType, const u64 size)
371 : {
372 0 : bool hugeData = size > SDMA_SEND_MAX_SIZE;
373 0 : HcclOpMetaInfoDef opMeta;
374 :
375 0 : if (isAlltoAllZCopyMode_) {
376 : /* zcopy拆分4GB以上SDMA任务前,准备好子图不复用标志 */
377 0 : if (opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
378 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, size, hugeData);
379 : } else {
380 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllVC(CopyPattern::ZCOPY, size, hugeData);
381 : }
382 : } else {
383 : /* bcopy每次重新生成子图 */
384 0 : if (opType == HcclCMDType::HCCL_CMD_ALLTOALLV) {
385 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::BCOPY, size, false);
386 : } else {
387 0 : opMeta = HcclOpMetaInfo::GetOneForAllToAllVC(CopyPattern::BCOPY, size, false);
388 : }
389 : }
390 :
391 0 : return opMeta;
392 : }
393 :
394 0 : u64 CollAlltoAllExecutor::CalAlltoAllVScratchMemSize(u64 &workSpaceMemSize)
395 : {
396 0 : u64 scratchMemSize = 0U;
397 0 : if (workSpaceMemSize == 0) {
398 0 : scratchMemSize = TINY_MEM_SIZE;
399 0 : HCCL_DEBUG("[CalAlltoAllVScratchMemSize] workSpaceMemSize==0, use TINY_MEM_SIZE[%llu]",
400 : TINY_MEM_SIZE);
401 : } else {
402 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
403 0 : scratchMemSize = std::max(std::max(workSpaceMemSize, inCCLbufferSize_), TINY_MEM_SIZE);
404 0 : HCCL_DEBUG("[CalAlltoAllVScratchMemSize] OpBase mode, workSpaceMemSize[%llu], "
405 : "inCCLbufferSize_[%llu], scratchMemSize[%llu]",
406 : workSpaceMemSize, inCCLbufferSize_, scratchMemSize);
407 : } else {
408 0 : scratchMemSize = workSpaceMemSize;
409 0 : HCCL_DEBUG("[CalAlltoAllVScratchMemSize] non-OpBase mode, workSpaceMemSize[%llu], "
410 : "scratchMemSize[%llu]", workSpaceMemSize, scratchMemSize);
411 : }
412 : }
413 0 : return scratchMemSize;
414 : }
415 :
416 0 : bool CollAlltoAllExecutor::HasMassTasks(std::vector<SendRecvInfo> &allMeshAggregationSendRecvInfo)
417 : {
418 0 : if (isAlltoAllZCopyMode_) {
419 0 : return false;
420 : }
421 :
422 0 : u64 maxSendTimes = 0;
423 0 : u64 maxRecvTimes = 0;
424 0 : const u64 cclBufferSize = algResResp_->cclInputMem.size();
425 0 : for (auto &sendRecvInfo : allMeshAggregationSendRecvInfo) {
426 0 : u64 sendTimes = 0;
427 0 : u64 recvTimes = 0;
428 0 : for (u32 i = 0; i < topoAttr_.userRankSize; i++) {
429 0 : sendTimes += (sendRecvInfo.sendLength[i] + cclBufferSize - 1) / cclBufferSize;
430 0 : recvTimes += (sendRecvInfo.recvLength[i] + cclBufferSize - 1) / cclBufferSize;
431 : }
432 0 : maxSendTimes = (maxSendTimes > sendTimes) ? maxSendTimes : sendTimes;
433 0 : maxRecvTimes = (maxRecvTimes > recvTimes) ? maxRecvTimes : recvTimes;
434 : }
435 0 : const u64 massThreshold = 65535; // 65535: 单个ffts+任务中,最多承载64K个task
436 0 : const u64 maxTasksPerStep = 10; // BCOPY中每次和远端通信最多消耗task数
437 0 : const u64 maxTasksBaseCost = 50; // BCOPY中除每步和远端通信外,最多消耗的task数
438 0 : u64 maxTasks = (maxSendTimes + maxRecvTimes) * maxTasksPerStep + maxTasksBaseCost;
439 0 : HCCL_DEBUG("[AlltoAll] bcopy maxSendTimes[%llu], maxRecvTimes[%llu], maxTasks[%llu], hasMassTask[%u]",
440 : maxSendTimes, maxRecvTimes, maxTasks, (maxTasks > massThreshold));
441 0 : return (maxTasks > massThreshold);
442 : }
443 :
444 5 : HcclResult CollAlltoAllExecutor::SetVirtualDispatcher(const HcclDispatcher virtualDispatcher)
445 : {
446 5 : vDispatcher_ = virtualDispatcher;
447 5 : return HCCL_SUCCESS;
448 : }
449 :
450 1 : HcclResult CollAlltoAllExecutor::CheckNeedRecreateComm(u64 lastScratchMemSize, bool& needRecreateAlltoallComm)
451 : {
452 1 : needRecreateAlltoallComm = false;
453 1 : return HCCL_SUCCESS;
454 : }
455 :
456 0 : HcclResult CollAlltoAllExecutor::RunAlltoAllTemplate(const std::unique_ptr<AlgTemplateBase> &executor,
457 : const SubCommInfo &commInfo)
458 : {
459 0 : HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.links);
460 0 : CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunAlltoAllTemplate]" \
461 : "group has been destroyed. Break!"), ret);
462 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
463 : HCCL_ERROR("[CollAlltoAllExecutor][RunAlltoAllTemplate]run executor rank[%u] rank size[%u] failed",
464 : commInfo.localRank, commInfo.localRankSize), ret);
465 0 : return HCCL_SUCCESS;
466 : }
467 :
468 0 : HcclResult CollAlltoAllExecutor::RunAlltoAllVTemplateStaged(const std::unique_ptr<AlgTemplateBase> &executor,
469 : const SubCommInfo &commInfo)
470 : {
471 0 : HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.links);
472 0 : CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunAlltoAllVTemplateStaged]" \
473 : "group has been destroyed. Break!"), ret);
474 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
475 : HCCL_ERROR("[CollAlltoAllExecutor][RunAlltoAllVTemplateStaged]run executor rank[%u] rank size[%u] failed",
476 : commInfo.localRank, commInfo.localRankSize), ret);
477 0 : return HCCL_SUCCESS;
478 : }
479 :
480 : // deprecated
481 0 : HcclResult CollAlltoAllExecutor::RunTemplateWithVirtualLink(const std::unique_ptr<AlgTemplateBase> &executor,
482 : const SubCommInfo &commInfo)
483 : {
484 0 : HcclResult ret = executor->RunAsync(commInfo.localRank, commInfo.localRankSize, commInfo.virtualLinks);
485 0 : CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[CollAlltoAllExecutor][RunTemplateWithVirtualLink]" \
486 : "group has been destroyed. Break!"), ret);
487 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
488 : HCCL_ERROR("[CollAlltoAllExecutor][RunTemplateWithVirtualLink]run executor rank[%u] rank size[%u] failed",
489 : commInfo.localRank, commInfo.localRankSize), ret);
490 0 : return HCCL_SUCCESS;
491 : }
492 :
493 : } // namespace hccl
|