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