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