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 "alltoall_symmetric_memory.h"
12 :
13 : namespace hccl {
14 0 : AlltoAllFullMeshSymmetricMemory::AlltoAllFullMeshSymmetricMemory(const HcclDispatcher dispatcher)
15 0 : : AlgTemplateBase(dispatcher)
16 0 : {}
17 :
18 0 : AlltoAllFullMeshSymmetricMemory::~AlltoAllFullMeshSymmetricMemory() {}
19 :
20 0 : HcclResult AlltoAllFullMeshSymmetricMemory::GenerateSubStreamInfo(
21 : const std::vector<Stream>& subStreams, const std::vector<std::shared_ptr<LocalNotify>>& meshSignalMainToSub,
22 : const std::vector<std::shared_ptr<LocalNotify>>& meshSignalSubToMain)
23 : {
24 0 : u32 totalSubstreamSize = sdmaConcurrentNum_;
25 0 : if (subStreams.size() < totalSubstreamSize || meshSignalMainToSub.size() < totalSubstreamSize
26 0 : || meshSignalSubToMain.size() < totalSubstreamSize) {
27 0 : HCCL_ERROR(
28 : "[AlltoAllFullMeshSymmetricMemory][GenerateSubStreamInfo]subStreamsSize[%zu], meshSignalMainToSubSize[%zu]"
29 : "meshSignalSubToMainSize[%zu] is smaller than totalSubstreamSize[%u]",
30 : subStreams.size(), meshSignalMainToSub.size(), meshSignalSubToMain.size(), totalSubstreamSize);
31 0 : return HCCL_E_PARA;
32 : }
33 0 : CHK_PRT_RET(
34 : links_.size() < userRankSize_,
35 : HCCL_ERROR(
36 : "[AlltoAllFullMeshSymmetricMemory][GenerateSubStreamInfo]"
37 : "links_.size()[%zu] is smaller than userRankSize_[%u].",
38 : links_.size(), userRankSize_),
39 : HCCL_E_PARA);
40 0 : HCCL_DEBUG(
41 : "subStreams.size[%zu], meshSignalMainToSub.size[%zu], links_.size[%zu]", subStreams.size(),
42 : meshSignalMainToSub.size(), links_.size());
43 0 : for (u32 sdmaIndex = 0; sdmaIndex < sdmaConcurrentNum_; sdmaIndex++) {
44 0 : sdmaSubStream_.push_back(subStreams[sdmaIndex]);
45 0 : sdmaMeshSignalMainToSub_.push_back(meshSignalMainToSub[sdmaIndex]);
46 0 : sdmaMeshSignalSubToMain_.push_back(meshSignalSubToMain[sdmaIndex]);
47 : }
48 0 : return HCCL_SUCCESS;
49 : }
50 :
51 0 : HcclResult AlltoAllFullMeshSymmetricMemory::Prepare(PrepareData& param)
52 : {
53 0 : mainStream_ = param.stream;
54 0 : userRank_ = param.userRank;
55 0 : userRankSize_ = param.userRankSize;
56 0 : links_ = *param.linksPtr;
57 0 : sendRecvInfoPtr_ = param.sendRecvInfoPtr;
58 0 : devNumInlocalPod_ = param.devNumInlocalPod;
59 0 : rankIdxInPod_ = param.rankIdxInPod;
60 0 : opType_ = param.opType;
61 0 : algOpContext_ = param.algOpContext;
62 :
63 0 : podStartRank_ = userRank_ - rankIdxInPod_;
64 0 : podEndRank_ = podStartRank_ + devNumInlocalPod_ - 1;
65 0 : sdmaConcurrentNum_ = (devNumInlocalPod_ > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
66 : (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) :
67 0 : (devNumInlocalPod_);
68 :
69 0 : HCCL_DEBUG(
70 : "[AlltoAllFullMeshSymmetricMemory]devNumInlocalPod_[%u], userRankSize_[%u] podStartRank_[%u]"
71 : "podEndRank_[%u], sdmaConcurrentNum_[%u]",
72 : devNumInlocalPod_, userRankSize_, podStartRank_, podEndRank_, sdmaConcurrentNum_);
73 :
74 0 : CHK_PRT_RET(
75 : userRankSize_ == 0, HCCL_ERROR("[AlltoAllFullMeshSymmetricMemory][Prepare]userRankSize_ is zero."),
76 : HCCL_E_PARA);
77 :
78 0 : userInput_ = param.inputMem;
79 0 : userOutput_ = param.outputMem;
80 0 : workMode_ = param.workMode;
81 :
82 0 : CHK_RET(GenerateSubStreamInfo(*param.subStreamsPtr, *param.signalPtr, *param.signalAuxPtr));
83 0 : return HCCL_SUCCESS;
84 : }
85 :
86 0 : std::string AlltoAllFullMeshSymmetricMemory::GetStreamIndexString()
87 : {
88 0 : std::string res = "";
89 0 : for (auto& info : subStreamReadInfo_) {
90 0 : u32 destRank = info.first;
91 0 : u32 streamIndex = destRank % sdmaConcurrentNum_;
92 0 : res += std::to_string(streamIndex) + ", ";
93 : }
94 0 : return res;
95 0 : }
96 :
97 0 : void AlltoAllFullMeshSymmetricMemory::UpdateCurrRankRecvInfo(
98 : [[maybe_unused]] u32 roundIdx, [[maybe_unused]] u32 side, u32 destRank, ReadDataBlock& readInfo)
99 : {
100 0 : const ZCopySendRecvInfo& sendRecvInfo = *sendRecvInfoPtr_;
101 0 : u64 recvLen = sendRecvInfo.localRecvLength[destRank];
102 0 : u64 userOutOffset = sendRecvInfo.localRecvOffset[destRank];
103 0 : u64 remoteUserInOffset = sendRecvInfo.remoteSendOffset[destRank];
104 0 : HCCL_DEBUG(
105 : "[AlltoAllFullMeshSymmetricMemory][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
106 : "recvLen[%lu] remoteUserInOffset[%llu] userOutOffset[%llu]",
107 : userRank_, destRank, recvLen, remoteUserInOffset, userOutOffset);
108 0 : readInfo = {recvLen, remoteUserInOffset, userOutOffset};
109 0 : }
110 :
111 0 : void AlltoAllFullMeshSymmetricMemory::UpdateSendRecvInfo(
112 : u32 roundIdx, std::unordered_map<u32, ReadDataBlock>& subStreamReadInfo,
113 : const std::vector<std::vector<std::pair<u32, u32>>>& partialCommRankSet)
114 : {
115 0 : for (u32 side = 0; side < partialCommRankSet.size(); side++) {
116 0 : for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
117 0 : u32 readRemoteRank = partialCommRankSet[side][j].first;
118 0 : if (readRemoteRank == userRank_) {
119 0 : continue;
120 : }
121 : ReadDataBlock readInfo;
122 0 : UpdateCurrRankRecvInfo(roundIdx, side, readRemoteRank, readInfo);
123 :
124 0 : subStreamReadInfo[readRemoteRank] = readInfo;
125 : }
126 : }
127 0 : }
128 :
129 0 : void AlltoAllFullMeshSymmetricMemory::UpdateRemoteRankSet(u32 roundIdx, u32 groupRankSize)
130 : {
131 0 : if (sdmaConcurrentNum_ == 1) {
132 0 : UpdatePartialCommunicationRankSetPairWise(roundIdx, groupRankSize);
133 : } else {
134 0 : UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
135 : }
136 0 : }
137 :
138 0 : void AlltoAllFullMeshSymmetricMemory::UpdatePartialCommunicationRankSetPairWise(u32 roundIdx, u32 groupRankSize)
139 : {
140 0 : partialCommRankSet_.clear();
141 0 : partialCommRankSet_.resize(1);
142 0 : for (u32 i = roundIdx * sdmaConcurrentNum_; i < (roundIdx * sdmaConcurrentNum_ + groupRankSize); i++) {
143 0 : u32 readRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
144 0 : u32 sendRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
145 0 : partialCommRankSet_[0].push_back(std::make_pair(readRemoteRank, sendRemoteRank));
146 0 : HCCL_DEBUG(
147 : "[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSetPairWise] userRank [%u] i[%u]"
148 : "readRemoteRank[%u] writeRemoteRank[%u]",
149 : userRank_, i, readRemoteRank, sendRemoteRank);
150 : }
151 0 : HCCL_DEBUG(
152 : "[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSetPairWise] partialCommRankSet_ size[%zu]",
153 : partialCommRankSet_[0].size());
154 0 : }
155 :
156 0 : void AlltoAllFullMeshSymmetricMemory::UpdatePartialCommunicationRankSet(
157 : u32 roundIdx, u32 groupRankSize, std::vector<std::vector<std::pair<u32, u32>>>& partialCommRankSet)
158 : {
159 0 : partialCommRankSet.clear();
160 0 : partialCommRankSet.resize(RANK_SET_COMPUTE_CONST + 1);
161 0 : u32 pairNumPerRound = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
162 0 : u32 pairSize = (groupRankSize < sdmaConcurrentNum_) ?
163 0 : (groupRankSize + RANK_SET_COMPUTE_CONST - 1) / RANK_SET_COMPUTE_CONST :
164 : pairNumPerRound;
165 0 : for (u32 i = roundIdx * pairNumPerRound + 1; i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
166 0 : u32 leftRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
167 0 : u32 rightRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
168 0 : if (leftRemoteRank == rightRemoteRank) {
169 0 : partialCommRankSet[2].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
170 : } else {
171 0 : partialCommRankSet[0].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
172 0 : partialCommRankSet[1].push_back(std::make_pair(rightRemoteRank, rightRemoteRank));
173 : }
174 0 : HCCL_DEBUG(
175 : "[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSet] round[%u] userRank [%u] i[%u]"
176 : "read/write leftRemoteRank[%u] rightRemoteRank[%u]",
177 : roundIdx, userRank_, i, leftRemoteRank, rightRemoteRank);
178 : }
179 0 : HCCL_DEBUG(
180 : "[AlltoAllFullMeshSymmetricMemory][UpdatePartialCommunicationRankSet] round[%u] partialCommRankSet_ total "
181 : "size[%zu]",
182 : roundIdx, partialCommRankSet[0].size() + partialCommRankSet[1].size() + partialCommRankSet[2].size());
183 0 : }
184 :
185 : // 主流只需要通知当前子步骤需要收发数据的 SDMA 流,减少同步开销
186 0 : HcclResult AlltoAllFullMeshSymmetricMemory::NotifySubStreamStart()
187 : {
188 0 : for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
189 0 : CHK_RET(
190 : LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
191 0 : CHK_RET(LocalNotify::Wait(
192 : sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
193 : }
194 0 : HCCL_DEBUG(
195 : "[AlltoAllFullMeshSymmetricMemory][NotifySubStreamStart] userRank [%u] main stream notify sdma stream [%s]",
196 : userRank_, GetStreamIndexString().c_str());
197 0 : return HCCL_SUCCESS;
198 : }
199 :
200 0 : HcclResult AlltoAllFullMeshSymmetricMemory::WaitSubStreamFinish()
201 : {
202 0 : for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
203 0 : CHK_RET(LocalNotify::Post(
204 : sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
205 0 : CHK_RET(
206 : LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
207 : }
208 0 : HCCL_DEBUG(
209 : "[AlltoAllFullMeshSymmetricMemory][WaitSubStreamFinish] userRank [%u] main stream wait sdma stream [%s]",
210 : userRank_, GetStreamIndexString().c_str());
211 0 : return HCCL_SUCCESS;
212 : }
213 :
214 0 : HcclResult AlltoAllFullMeshSymmetricMemory::NotifyRemoteRankStart()
215 : {
216 0 : u32 streamIndex = 0;
217 0 : for (auto& sendRecvSide : partialCommRankSet_) {
218 0 : for (auto& sendRecvPair : sendRecvSide) {
219 0 : u32 recvRank = sendRecvPair.first;
220 0 : u32 sendRank = sendRecvPair.second;
221 0 : if (sendRank == userRank_) {
222 0 : continue;
223 : }
224 0 : Stream& currStream = sdmaSubStream_[streamIndex];
225 0 : const LINK& readTransport = links_[recvRank];
226 0 : const LINK& sendTransport = links_[sendRank];
227 :
228 0 : CHK_RET(sendTransport->TxAck(currStream));
229 0 : CHK_RET(readTransport->RxAck(currStream));
230 :
231 0 : streamIndex++;
232 : }
233 : }
234 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][NotifyRemoteRankStart] done");
235 0 : return HCCL_SUCCESS;
236 : }
237 :
238 0 : bool AlltoAllFullMeshSymmetricMemory::IsPostSyncEnable(u32 roundIdx)
239 : {
240 0 : bool isPostSyncEnable = false;
241 0 : isPostSyncEnable = (roundIdx == lastRoundIdx_) && algOpContext_.opRetryHandler.retryEnable;
242 0 : return isPostSyncEnable;
243 : }
244 :
245 0 : HcclResult AlltoAllFullMeshSymmetricMemory::SdmaMainStreamWait(u32 roundIdx)
246 : {
247 : // SDMA wait
248 0 : u32 streamIndex = 0;
249 0 : for (auto& sendRecvSide : partialCommRankSet_) {
250 0 : for (auto& sendRecvPair : sendRecvSide) {
251 0 : u32 recvRank = sendRecvPair.first;
252 0 : u32 sendRank = sendRecvPair.second;
253 0 : if (sendRank == userRank_) {
254 0 : continue;
255 : }
256 0 : HCCL_DEBUG(
257 : "[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
258 : "sendRank[%u], sdma stream [%u], "
259 : "post sync info: roundIdx[%u], lastRoundIdx_[%u] main stream wait",
260 : userRank_, recvRank, sendRank, streamIndex, roundIdx, lastRoundIdx_);
261 0 : CHK_RET(LocalNotify::Wait(
262 : mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
263 :
264 0 : streamIndex++;
265 : }
266 : }
267 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamWait] done");
268 0 : return HCCL_SUCCESS;
269 : }
270 :
271 0 : HcclResult AlltoAllFullMeshSymmetricMemory::SdmaMainStreamPost(u32 roundIdx)
272 : {
273 : // SDMA post
274 0 : u32 streamIndex = 0;
275 0 : for (auto& sendRecvSide : partialCommRankSet_) {
276 0 : for (auto& sendRecvPair : sendRecvSide) {
277 0 : u32 recvRank = sendRecvPair.first;
278 0 : u32 sendRank = sendRecvPair.second;
279 0 : if (sendRank == userRank_) {
280 0 : continue;
281 : }
282 0 : HCCL_DEBUG(
283 : "[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
284 : "sendRank[%u], sdma stream [%u], "
285 : "post sync info: roundIdx[%u], lastRoundIdx_[%u] main stream post",
286 : userRank_, recvRank, sendRank, streamIndex, roundIdx, lastRoundIdx_);
287 0 : CHK_RET(LocalNotify::Post(
288 : mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
289 :
290 0 : streamIndex++;
291 : }
292 : }
293 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SdmaMainStreamPost] done");
294 0 : return HCCL_SUCCESS;
295 : }
296 :
297 0 : HcclResult AlltoAllFullMeshSymmetricMemory::SetPostSyncTasks(u32 roundIdx)
298 : {
299 : // SDMA wait
300 0 : CHK_RET(SdmaMainStreamWait(roundIdx));
301 : // SDMA post
302 0 : CHK_RET(SdmaMainStreamPost(roundIdx));
303 0 : HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SetPostSyncTasks] done");
304 0 : return HCCL_SUCCESS;
305 : }
306 :
307 0 : HcclResult AlltoAllFullMeshSymmetricMemory::SDMAwithRemoteRankAndNotifyEnd(u32 roundIdx)
308 : {
309 0 : bool isPostSyncEnable = IsPostSyncEnable(roundIdx);
310 0 : if (isPostSyncEnable) {
311 : // 下发主流上的后同步wait和post
312 0 : CHK_RET(SetPostSyncTasks(roundIdx));
313 : }
314 0 : u32 streamIndex = 0;
315 0 : for (auto& sendRecvSide : partialCommRankSet_) {
316 0 : for (auto& sendRecvPair : sendRecvSide) {
317 0 : u32 recvRank = sendRecvPair.first;
318 0 : u32 sendRank = sendRecvPair.second;
319 0 : if (sendRank == userRank_) {
320 0 : continue;
321 : }
322 0 : const ReadDataBlock& readInfo = subStreamReadInfo_[recvRank];
323 0 : Stream& currStream = sdmaSubStream_[streamIndex];
324 0 : const LINK& readTransport = links_[recvRank];
325 0 : const LINK& sendTransport = links_[sendRank];
326 :
327 0 : const LINK& intraNeighboorTransport = links_[recvRank];
328 0 : CHK_PTR_NULL(intraNeighboorTransport);
329 0 : void* remDMAMemPtr = nullptr;
330 0 : CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
331 0 : DeviceMem remoteUserInMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), userInput_.size());
332 0 : DeviceMem srcMem = remoteUserInMem.range(readInfo.remoteOffset, readInfo.recvLen);
333 0 : DeviceMem dstMem = userOutput_.range(readInfo.recvOffset, readInfo.recvLen);
334 0 : CHK_RET(HcclD2DMemcpyAsync(
335 : dispatcher_, dstMem, srcMem, currStream, readTransport->GetRemoteRank(), readTransport->GetLinkType()));
336 0 : HCCL_DEBUG(
337 : "[AlltoAllFullMeshSymmetricMemory][SendRecvData] userRank [%u], recvRank[%u],"
338 : "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
339 : "post sync info: roundIdx[%u], lastRoundIdx_[%u]",
340 : userRank_, recvRank, streamIndex, readInfo.remoteOffset, readInfo.recvLen, readInfo.recvOffset,
341 : roundIdx, lastRoundIdx_);
342 0 : if (isPostSyncEnable) {
343 0 : HCCL_DEBUG("[AlltoAllFullMeshSymmetricMemory][SendRecvData] post sync begins");
344 0 : CHK_RET(LocalNotify::Post(
345 : currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
346 0 : CHK_RET(LocalNotify::Wait(
347 : currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
348 : }
349 0 : CHK_RET(readTransport->TxDataSignal(currStream));
350 0 : CHK_RET(sendTransport->RxDataSignal(currStream));
351 :
352 0 : streamIndex++;
353 0 : }
354 : }
355 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][SDMAwithRemoteRankAndNotifyEnd] done");
356 0 : return HCCL_SUCCESS;
357 : }
358 :
359 0 : HcclResult AlltoAllFullMeshSymmetricMemory::SendRecvData(u32 roundIdx)
360 : {
361 0 : HCCL_DEBUG(
362 : "[AlltoAllFullMeshSymmetricMemory][SendRecvData] userRank [%u] sdma stream [%s] wait main stream", userRank_,
363 : GetStreamIndexString().c_str());
364 0 : CHK_RET(NotifyRemoteRankStart());
365 0 : CHK_RET(WaitSubStreamFinish());
366 0 : CHK_RET(NotifySubStreamStart());
367 0 : CHK_RET(SDMAwithRemoteRankAndNotifyEnd(roundIdx));
368 :
369 0 : return HCCL_SUCCESS;
370 : }
371 :
372 0 : HcclResult AlltoAllFullMeshSymmetricMemory::LocalCopy()
373 : {
374 0 : const ZCopySendRecvInfo& sendRecvInfo = *sendRecvInfoPtr_;
375 0 : DeviceMem src = userInput_.range(sendRecvInfo.remoteSendOffset[userRank_], sendRecvInfo.localRecvLength[userRank_]);
376 0 : DeviceMem dst = userOutput_.range(sendRecvInfo.localRecvOffset[userRank_], sendRecvInfo.localRecvLength[userRank_]);
377 0 : HCCL_DEBUG(
378 : "[AlltoAllFullMeshSymmetricMemory][LocalCopy]userRank [%u] copy from userInput [%llu]"
379 : "to userOutput [%llu] dstLen[%llu]",
380 : userRank_, sendRecvInfo.remoteSendOffset[userRank_], sendRecvInfo.localRecvOffset,
381 : sendRecvInfo.localRecvLength[userRank_]);
382 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
383 0 : return HCCL_SUCCESS;
384 0 : }
385 :
386 0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunGroupFullMeshAlltoall(u32 roundIdx)
387 : {
388 0 : subStreamReadInfo_.clear();
389 0 : UpdateSendRecvInfo(roundIdx, subStreamReadInfo_, partialCommRankSet_);
390 0 : CHK_RET(NotifySubStreamStart());
391 0 : CHK_RET(SendRecvData(roundIdx));
392 0 : if (!islocalCpyDone_) {
393 0 : CHK_RET(LocalCopy());
394 0 : islocalCpyDone_ = true;
395 : }
396 0 : CHK_RET(WaitSubStreamFinish());
397 0 : return HCCL_SUCCESS;
398 : }
399 :
400 : HcclResult
401 0 : AlltoAllFullMeshSymmetricMemory::RunSDMATasks(u32 roundIdx, u32 groupRankSize, [[maybe_unused]] u32 leftRankSize)
402 : {
403 0 : UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
404 0 : CHK_RET(RunGroupFullMeshAlltoall(roundIdx));
405 0 : return HCCL_SUCCESS;
406 : }
407 :
408 0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunSDMA(HcclOpMetaInfoDef& opMeta)
409 : {
410 : // 计算每个rank分组fullmesh后需要通信的轮次,向上取整
411 0 : commRounds_ = (devNumInlocalPod_ + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_;
412 0 : u32 leftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
413 : lastRoundIdx_
414 0 : = std::min((leftRankSize + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_, static_cast<u32>(commRounds_)) - 1;
415 0 : HCCL_DEBUG(
416 : "[AlltoAllFullMeshSymmetricMemory][RunSDMA] userRank [%u] communication rounds[%llu]"
417 : "post sync info: lastRoundIdx_[%u] devNumInlocalPod_[%u] sdmaConcurrentNum_[%u]",
418 : userRank_, commRounds_, lastRoundIdx_, devNumInlocalPod_, sdmaConcurrentNum_);
419 :
420 0 : u32 currentLeftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
421 0 : for (u32 roundIdx = 0; roundIdx < commRounds_ && currentLeftRankSize > 0; roundIdx++) {
422 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
423 0 : u32 groupRankSize = (currentLeftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : currentLeftRankSize;
424 0 : CHK_RET(RunSDMATasks(roundIdx, groupRankSize, currentLeftRankSize));
425 0 : currentLeftRankSize -= groupRankSize;
426 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
427 : }
428 :
429 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunSDMA] finished.");
430 0 : return HCCL_SUCCESS;
431 : }
432 :
433 0 : HcclResult AlltoAllFullMeshSymmetricMemory::RunAsync()
434 : {
435 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, userInput_.size(), true);
436 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
437 :
438 0 : if (userRankSize_ == 1) {
439 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunAsync] do localcopy with 1 rank");
440 0 : CHK_RET(LocalCopy());
441 0 : return HCCL_SUCCESS;
442 : }
443 :
444 0 : if (devNumInlocalPod_ > 1) {
445 0 : CHK_RET(RunSDMA(opMeta));
446 : }
447 :
448 0 : HCCL_INFO("[AlltoAllFullMeshSymmetricMemory][RunAsync] finished.");
449 0 : return HCCL_SUCCESS;
450 : }
451 :
452 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_FULL_MESH_SYMMETRIC_MEMORY, AlltoAllFullMeshSymmetricMemory);
453 : } // namespace hccl
|