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 "alltoallv_direct_fullmesh.h"
12 : #include "dispatcher_pub.h"
13 :
14 : namespace hccl {
15 0 : AlltoAllVDirectFullMesh::AlltoAllVDirectFullMesh(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
16 :
17 0 : AlltoAllVDirectFullMesh::~AlltoAllVDirectFullMesh() {}
18 :
19 0 : HcclResult AlltoAllVDirectFullMesh::GenerateSubStreamInfo(
20 : const std::vector<Stream>& subStreams, const std::vector<std::shared_ptr<LocalNotify>>& meshSignalMainToSub,
21 : const std::vector<std::shared_ptr<LocalNotify>>& meshSignalSubToMain)
22 : {
23 0 : u32 totalSubstreamSize
24 0 : = (totalRdmaRankNum_ > 0) ? (sdmaConcurrentNum_ + rdmaConcurrentNum_ + 1) : (sdmaConcurrentNum_);
25 0 : if (subStreams.size() < totalSubstreamSize || meshSignalMainToSub.size() < totalSubstreamSize
26 0 : || meshSignalSubToMain.size() < totalSubstreamSize) {
27 0 : HCCL_ERROR(
28 : "[AlltoAllVDirectFullMesh][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 : "[AlltoAllVDirectFullMesh][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 : u32 index = 0;
44 0 : for (u32 sdmaIndex = 0; sdmaIndex < sdmaConcurrentNum_; sdmaIndex++) {
45 0 : sdmaSubStream_.push_back(subStreams[index]);
46 0 : sdmaMeshSignalMainToSub_.push_back(meshSignalMainToSub[index]);
47 0 : sdmaMeshSignalSubToMain_.push_back(meshSignalSubToMain[index]);
48 0 : index++;
49 : }
50 0 : for (u32 localIndex = 0; localIndex < sdmaConcurrentNum_; localIndex++) {
51 0 : localSubStream_.push_back(subStreams[index]);
52 0 : localSignalMainToSub_.push_back(meshSignalMainToSub[index]);
53 0 : localSignalSubToMain_.push_back(meshSignalSubToMain[index]);
54 0 : index++;
55 : }
56 0 : if (totalRdmaRankNum_ > 0) {
57 0 : rdmaSubStreams_.push_back(subStreams[index]);
58 0 : main2RdmaControlStreamNotify_ = meshSignalMainToSub[index];
59 0 : rdmaControl2MainStreamNotify_ = meshSignalSubToMain[index];
60 0 : index++;
61 0 : for (u32 rdmaIndex = 0; rdmaIndex < rdmaConcurrentNum_; rdmaIndex++) {
62 0 : rdmaSubStreams_.push_back(subStreams[index]);
63 0 : rdmaControl2SubNotifies_.push_back(meshSignalMainToSub[index]);
64 0 : rdmaSub2ControlNotifies_.push_back(meshSignalSubToMain[index]);
65 0 : index++;
66 : }
67 : }
68 0 : return HCCL_SUCCESS;
69 : }
70 :
71 0 : HcclResult AlltoAllVDirectFullMesh::Prepare(PrepareData& param)
72 : {
73 0 : needAlltoallvCache_ = param.needAlltoallvCache;
74 0 : HCCL_INFO(
75 : "[AlltoAllVDirectFullMesh][Prepare] set needAlltoallvCache_[%u] for alltoallv aicpu cache",
76 : needAlltoallvCache_);
77 :
78 0 : mainStream_ = param.stream;
79 0 : userRank_ = param.userRank;
80 0 : userRankSize_ = param.userRankSize;
81 0 : links_ = *param.linksPtr;
82 0 : localSendRecvInfoPtr_ = param.localSendRecvInfoPtr;
83 0 : devNumInlocalPod_ = param.devNumInlocalPod;
84 0 : rankIdxInPod_ = param.rankIdxInPod;
85 0 : opType_ = param.opType;
86 0 : algOpContext_ = param.algOpContext;
87 :
88 0 : podStartRank_ = userRank_ - rankIdxInPod_;
89 0 : podEndRank_ = podStartRank_ + devNumInlocalPod_ - 1;
90 0 : sdmaConcurrentNum_ = (devNumInlocalPod_ > ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) ?
91 : (ALLTOALLV_DIRECT_FULLMESH_SDMA_CONCURRENT_SIZE) :
92 0 : (devNumInlocalPod_);
93 :
94 0 : totalRdmaRankNum_ = userRankSize_ - devNumInlocalPod_;
95 0 : rdmaConcurrentNum_ = (totalRdmaRankNum_ > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
96 : (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) :
97 0 : (totalRdmaRankNum_);
98 0 : HCCL_DEBUG(
99 : "[AlltoAllVDirectFullMesh]devNumInlocalPod_[%u], userRankSize_[%u] podStartRank_[%u]"
100 : "podEndRank_[%u], totalRdmaRankNum_[%u], sdmaConcurrentNum_[%u], rdmaConcurrentNum_[%u]",
101 : devNumInlocalPod_, userRankSize_, podStartRank_, podEndRank_, totalRdmaRankNum_, sdmaConcurrentNum_,
102 : rdmaConcurrentNum_);
103 :
104 0 : CHK_PRT_RET(
105 : userRankSize_ == 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][Prepare]userRankSize_ is zero."), HCCL_E_PARA);
106 :
107 0 : userInput_ = param.inputMem;
108 0 : userOutput_ = param.outputMem;
109 0 : cclInMem_ = param.cclInMem;
110 0 : cclOutMem_ = param.cclOutMem;
111 0 : workMode_ = param.workMode;
112 0 : isSuPodAsym_ = param.isSuPodAsym;
113 :
114 : // 注意: 如果isBigCount的计算逻辑发生变化, 需要同步修改IsBigCountForAlltoallv()中的代码
115 0 : u64 maxSendLen = CalcMaxSendLen();
116 0 : isBigCount_ = (maxSendLen > ALLTOALLV_DIRECT_FULLMESH_BIG_SIZE) ? true : false;
117 0 : CHK_RET(GenerateSubStreamInfo(*param.subStreamsPtr, *param.signalPtr, *param.signalAuxPtr));
118 :
119 0 : if (algOpContext_.mc2Handler.stepSize > 0) {
120 0 : sdmaConcurrentNum_ = (devNumInlocalPod_ > 1) ? 1 : (devNumInlocalPod_);
121 : // MC2细粒度不需要本地并发处理
122 0 : isBigCount_ = false;
123 : }
124 :
125 : /* 考虑当group0 的rank 跟 group 1的所有rank通信时,每次都要收发,所以取sdmaConcurrentNum_块;
126 : 跟group 0内的rank通信有一块儿浪费 */
127 : // 注意: 如果sdmaDataBlockSize_的计算逻辑发生变化,
128 : // 需要同步修改framework下CalcMetadataForFirstAlltoallv()函数中的part 1
129 0 : u32 blockGroup
130 0 : = (isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC) ?
131 : 2 :
132 : 1;
133 0 : sdmaDataBlockSize_ = (cclInMem_.size() / std::max(1u, sdmaConcurrentNum_ * blockGroup));
134 : // 向下对齐到16k Byte
135 0 : if (sdmaDataBlockSize_ > HCCL_MIN_SLICE_ALIGN_910B) {
136 0 : sdmaDataBlockSize_ = (sdmaDataBlockSize_ / HCCL_MIN_SLICE_ALIGN_910B) * HCCL_MIN_SLICE_ALIGN_910B;
137 : }
138 0 : CHK_PRT_RET(
139 : sdmaDataBlockSize_ == 0, HCCL_ERROR("[AlltoAllVDirectFullMesh][Prepare]sdmaDataBlockSize_is zero."),
140 : HCCL_E_INTERNAL);
141 0 : HCCL_DEBUG(
142 : "[AlltoAllVDirectFullMesh][Prepare] userRank [%u] total cclsize[%llu],"
143 : "sdmaDataBlockSize_[%llu], BigCountFlag[%d], stepSize[%u]",
144 : userRank_, cclInMem_.size(), sdmaDataBlockSize_, isBigCount_, algOpContext_.mc2Handler.stepSize);
145 :
146 : // 一半的CCLOut用来发送RDMA数据,另一半用来接收RDMA数据,因此需要除以2
147 0 : rdmaDataBlockSize_ = cclOutMem_.size() / std::max(1u, rdmaConcurrentNum_) / 2;
148 :
149 0 : return HCCL_SUCCESS;
150 : }
151 :
152 0 : std::string AlltoAllVDirectFullMesh::GetStreamIndexString()
153 : {
154 0 : std::string res = "";
155 0 : for (auto& info : subStreamReadInfo_) {
156 0 : u32 destRank = info.first;
157 0 : u32 streamIndex = destRank % sdmaConcurrentNum_;
158 0 : res += std::to_string(streamIndex) + ", ";
159 : }
160 0 : return res;
161 0 : }
162 :
163 0 : u64 AlltoAllVDirectFullMesh::CalcMaxSendLen()
164 : {
165 0 : u64 maxSendLen = 0;
166 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
167 :
168 0 : for (u32 dstRank = 0; dstRank < localSendRecvInfo.sendLength.size(); dstRank++) {
169 0 : maxSendLen = std::max(maxSendLen, localSendRecvInfo.sendLength[dstRank]);
170 : }
171 :
172 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][CalcMaxSendLen] maxSendLen[%llu]", maxSendLen);
173 0 : return maxSendLen;
174 : }
175 :
176 0 : HcclResult AlltoAllVDirectFullMesh::UpdateCurrRankRecvInfo(
177 : u32 step, u32 roundIdx, u32 side, u32 destRank, std::vector<ReadDataBlock>& readInfo,
178 : std::unordered_map<u32, ReadDataBlock>& subStreamZcopyReadInfo, u32 maxRecvStep)
179 : {
180 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
181 0 : u64 remainRecvLen = localSendRecvInfo.recvLength[destRank];
182 0 : u64 scratchOffset = 0;
183 0 : u32 bufferIdx = 0;
184 0 : u32 pairNum = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
185 0 : if (sdmaConcurrentNum_ == 1) { // 保证和当前rank距离一样时,send/recv用的是同一块buff
186 0 : bufferIdx = 0;
187 0 : } else if (side == 0) { // 在curRank左边
188 0 : u32 gap = (userRank_ - destRank + devNumInlocalPod_) % devNumInlocalPod_;
189 0 : bufferIdx = pairNum - (gap - roundIdx * pairNum);
190 0 : } else if (side == 1) { // 在curRank右边
191 0 : u32 gap = (destRank - userRank_ + devNumInlocalPod_) % devNumInlocalPod_;
192 0 : bufferIdx = pairNum - 1 + (gap - roundIdx * pairNum);
193 : } else { // 最后一个中间位置的rank
194 0 : bufferIdx = 0;
195 : }
196 :
197 0 : if ((isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC)
198 0 : && (roundIdx % RANK_SET_COMPUTE_CONST != 0)) { // 奇数轮,用下半Buffer
199 0 : bufferIdx += sdmaConcurrentNum_;
200 : }
201 :
202 0 : scratchOffset = bufferIdx * sdmaDataBlockSize_;
203 :
204 0 : u32 recvStepIdx = 0;
205 0 : u64 dataOffset = 0;
206 0 : HCCL_DEBUG(
207 : "step[%u] round[%u] usrRank[%u] total recv localSendRecvInfo.recvLength[%llu] from dstRank[%u] bufferIdx[%u]",
208 : step, roundIdx, userRank_, remainRecvLen, destRank, bufferIdx);
209 :
210 : // alltoallv类算子的零长拷贝, 需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发
211 : // (cache不使能时为空函数调用)
212 0 : if (needAlltoallvCache_ && remainRecvLen == 0) {
213 : // 获取local user output offset
214 0 : const u64 recvLen = 0;
215 0 : u64 userOutOffset = localSendRecvInfo.recvOffset[destRank];
216 0 : HCCL_DEBUG(
217 : "[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
218 : "recvStepIdx[%u] recvLen[%lu] userOutOffset[%llu] scratchOffset[%llu]",
219 : userRank_, destRank, recvStepIdx, recvLen, userOutOffset, scratchOffset);
220 :
221 : // 更新零长拷贝的read info
222 0 : ReadDataBlock readBlock = {recvLen, scratchOffset, userOutOffset};
223 0 : subStreamZcopyReadInfo[destRank] = readBlock;
224 :
225 : // sendCount为0, step和readInfo.size一定为0
226 0 : CHK_PRT_RET(
227 : maxRecvStep > 0,
228 : HCCL_ERROR(
229 : "[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] maxRecvStep[%u] != 0 for remainRecvLen[%llu]",
230 : maxRecvStep, remainRecvLen),
231 : HCCL_E_INTERNAL);
232 0 : CHK_PRT_RET(
233 : readInfo.size() != 0,
234 : HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] invalid readInfo.size[%u]", readInfo.size()),
235 : HCCL_E_INTERNAL);
236 0 : } else {
237 0 : while (recvStepIdx < maxRecvStep && remainRecvLen > 0) {
238 0 : u64 currDataRemainLen = localSendRecvInfo.recvLength[destRank] - dataOffset;
239 0 : u64 recvLen = std::min(sdmaDataBlockSize_, currDataRemainLen);
240 0 : u64 userOutOffset = localSendRecvInfo.recvOffset[destRank] + dataOffset;
241 0 : HCCL_DEBUG(
242 : "[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] usrRank[%u] recv from destRank [%u]"
243 : "recvStepIdx[%u] recvLen[%lu] userOutOffset[%llu] scratchOffset[%llu]",
244 : userRank_, destRank, recvStepIdx, recvLen, userOutOffset, scratchOffset);
245 0 : readInfo.push_back({recvLen, scratchOffset, userOutOffset});
246 0 : dataOffset += recvLen;
247 0 : recvStepIdx++;
248 0 : remainRecvLen -= recvLen;
249 : }
250 : }
251 :
252 0 : return HCCL_SUCCESS;
253 : }
254 :
255 0 : HcclResult AlltoAllVDirectFullMesh::UpdateCurrRankSendInfo(
256 : u32 step, u32 roundIdx, u32 side, u32 destRank, std::vector<SendDataBlock>& sendInfo,
257 : std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo, u32 maxSendStep)
258 : {
259 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
260 0 : u64 remainSendLen = localSendRecvInfo.sendLength[destRank];
261 :
262 0 : u64 scratchOffset = 0;
263 0 : u32 bufferIdx = 0;
264 0 : u32 pairNum = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
265 0 : if (sdmaConcurrentNum_ == 1) { // 保证和当前rank距离一样时,send/recv用的是同一块buff
266 0 : bufferIdx = 0;
267 0 : } else if (side == 0) { // 在curRank左边
268 0 : u32 gap = (userRank_ - destRank + devNumInlocalPod_) % devNumInlocalPod_;
269 0 : bufferIdx = pairNum - 1 + (gap - roundIdx * pairNum);
270 0 : } else if (side == 1) { // 在curRank右边
271 0 : u32 gap = (destRank - userRank_ + devNumInlocalPod_) % devNumInlocalPod_;
272 0 : bufferIdx = pairNum - (gap - roundIdx * pairNum);
273 : } else { // 最后一个中间位置的rank
274 0 : bufferIdx = 0;
275 : }
276 :
277 0 : if ((isBigCount_ || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLV || opType_ == HcclCMDType::HCCL_CMD_ALLTOALLVC)
278 0 : && (roundIdx % RANK_SET_COMPUTE_CONST != 0)) { // 奇数轮,用下半Buffer
279 0 : bufferIdx += sdmaConcurrentNum_;
280 : }
281 0 : scratchOffset = bufferIdx * sdmaDataBlockSize_;
282 :
283 : // 更新hcclOffset到dstRank的映射, 用于alltoallv算子aicpu展开的SQE缓存
284 0 : if (needAlltoallvCache_) {
285 : // alltoallv cache只针对小数据量, 至多只有1个step
286 0 : CHK_PRT_RET(
287 : step != 0,
288 : HCCL_ERROR(
289 : "[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] needAlltoallvCache_[%u] step[%u]",
290 : needAlltoallvCache_, step),
291 : HCCL_E_INTERNAL);
292 :
293 : std::unordered_map<uint64_t, std::vector<uint32_t>>::iterator mapIter
294 0 : = hcclOffsetDstRanksMap_.find(scratchOffset);
295 0 : if (mapIter == hcclOffsetDstRanksMap_.end()) {
296 0 : constexpr uint32_t singleRankVecSize = 1;
297 : std::pair<std::unordered_map<uint64_t, std::vector<uint32_t>>::iterator, bool> emplaceResult
298 0 : = hcclOffsetDstRanksMap_.emplace(scratchOffset, std::vector<uint32_t>(singleRankVecSize, destRank));
299 0 : CHK_PRT_RET(
300 : !emplaceResult.second,
301 : HCCL_ERROR(
302 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] fail to insert hcclOffset[%llu]-dstRank[%u] "
303 : "pair",
304 : scratchOffset, destRank),
305 : HCCL_E_INTERNAL);
306 0 : mapIter = emplaceResult.first;
307 : } else {
308 : // 虽然同一个dstRank不需要重复计算sendInfo, 但不同dstRanks在multi-round case下可能对应相同的hcclOffset
309 0 : CHK_PRT_RET(
310 : mapIter->second.size() == 0,
311 : HCCL_ERROR(
312 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] empty dstRanks for hcclOffset[%llu] before add "
313 : "destRank[%u]",
314 : mapIter->second, destRank, scratchOffset),
315 : HCCL_E_INTERNAL);
316 0 : mapIter->second.push_back(destRank);
317 : }
318 0 : HCCL_DEBUG(
319 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] mapIter->first[%llu] mapIter->second.size[%u] "
320 : "destRank[%u]",
321 : mapIter->first, mapIter->second.size(), destRank);
322 : }
323 :
324 0 : u32 sendStepIdx = 0;
325 0 : u64 dataOffset = 0;
326 0 : HCCL_DEBUG(
327 : "step[%u] round[%u] usrRank[%u] total send localSendRecvInfo.sendLength[%llu] to dstRank[%u] bufferIdx[%u]",
328 : step, roundIdx, userRank_, remainSendLen, destRank, bufferIdx);
329 :
330 0 : if (needAlltoallvCache_ && remainSendLen == 0) { // alltoallv类算子的零长拷贝, 需要调用MemcpyAsync保证aicpu
331 : // cache使能时placeholder正确下发 (cache不使能时为空函数调用)
332 : // 获取local user input offset
333 0 : const u64 sendLen = 0;
334 0 : u64 userInOffset = localSendRecvInfo.sendOffset[destRank];
335 0 : HCCL_DEBUG(
336 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] usrRank[%u] send to destRank [%u]"
337 : " sendStepIdx[%u] sendLen[%lu] userInOffset[%llu] scratchOffset[%llu]",
338 : userRank_, destRank, sendStepIdx, sendLen, userInOffset, scratchOffset);
339 :
340 : // 更新零长拷贝的send info
341 0 : SendDataBlock sendBlock = {sendLen, userInOffset, scratchOffset};
342 0 : subStreamZcopySendInfo[destRank] = sendBlock;
343 :
344 : // sendCount为0, step和sendInfo.size一定为0
345 0 : CHK_PRT_RET(
346 : maxSendStep > 0,
347 : HCCL_ERROR(
348 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] maxSendStep[%u] != 0 for remainSendLen[%llu]",
349 : maxSendStep, remainSendLen),
350 : HCCL_E_INTERNAL);
351 0 : CHK_PRT_RET(
352 : sendInfo.size() != 0,
353 : HCCL_ERROR("[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] invalid sendInfo.size[%u]", sendInfo.size()),
354 : HCCL_E_INTERNAL);
355 0 : } else {
356 0 : while (sendStepIdx < maxSendStep && remainSendLen > 0) {
357 0 : u64 currDataRemainLen = localSendRecvInfo.sendLength[destRank] - dataOffset;
358 0 : u64 sendLen = std::min(sdmaDataBlockSize_, currDataRemainLen);
359 0 : u64 userInOffset = localSendRecvInfo.sendOffset[destRank] + dataOffset;
360 0 : HCCL_DEBUG(
361 : "[AlltoAllVDirectFullMesh][UpdateCurrRankSendInfo] usrRank[%u] send to destRank [%u]"
362 : " sendStepIdx[%u] sendLen[%lu] userInOffset[%llu] scratchOffset[%llu]",
363 : userRank_, destRank, sendStepIdx, sendLen, userInOffset, scratchOffset);
364 0 : sendInfo.push_back({sendLen, userInOffset, scratchOffset});
365 0 : dataOffset += sendLen;
366 0 : sendStepIdx++;
367 0 : remainSendLen -= sendLen;
368 : }
369 : }
370 :
371 0 : return HCCL_SUCCESS;
372 : }
373 :
374 0 : void AlltoAllVDirectFullMesh::UpdateSendRecvInfo(
375 : u32 step, u32 roundIdx, std::unordered_map<u32, std::vector<ReadDataBlock>>& subStreamReadInfo,
376 : std::unordered_map<u32, std::vector<SendDataBlock>>& subStreamSendInfo,
377 : std::unordered_map<u32, ReadDataBlock>& subStreamZcopyReadInfo,
378 : std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo,
379 : const std::vector<std::vector<std::pair<u32, u32>>>& partialCommRankSet)
380 : {
381 0 : for (u32 side = 0; side < partialCommRankSet.size(); side++) {
382 0 : for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
383 0 : u32 readRemoteRank = partialCommRankSet[side][j].first;
384 0 : if (readRemoteRank == userRank_) {
385 0 : continue;
386 : }
387 0 : u32 currDestRecvStep = recvNumSubStep_[readRemoteRank];
388 0 : std::vector<ReadDataBlock> readInfo;
389 0 : UpdateCurrRankRecvInfo(
390 : step, roundIdx, side, readRemoteRank, readInfo, subStreamZcopyReadInfo, currDestRecvStep);
391 :
392 0 : subStreamReadInfo[readRemoteRank] = readInfo;
393 0 : }
394 : }
395 :
396 0 : for (u32 side = 0; side < partialCommRankSet.size(); side++) {
397 0 : for (u32 j = 0; j < partialCommRankSet[side].size(); j++) {
398 0 : u32 sendRemoteRank = partialCommRankSet[side][j].second;
399 0 : if (sendRemoteRank == userRank_) {
400 0 : continue;
401 : }
402 0 : u32 currDestSendStep = sendNumSubStep_[sendRemoteRank];
403 0 : std::vector<SendDataBlock> sendInfo;
404 0 : UpdateCurrRankSendInfo(
405 : step, roundIdx, side, sendRemoteRank, sendInfo, subStreamZcopySendInfo, currDestSendStep);
406 :
407 0 : subStreamSendInfo[sendRemoteRank] = sendInfo;
408 0 : }
409 : }
410 0 : }
411 :
412 0 : void AlltoAllVDirectFullMesh::UpdateOpBaseSubStreamInfo(u32 step, u32 roundIdx)
413 : {
414 0 : if (roundIdx == 0 || !isBigCount_) {
415 0 : subStreamReadInfo_.clear();
416 0 : subStreamSendInfo_.clear();
417 0 : if (needAlltoallvCache_) {
418 0 : subStreamZcopyReadInfo_.clear();
419 0 : subStreamZcopySendInfo_.clear();
420 : }
421 0 : UpdateSendRecvInfo(
422 0 : step, roundIdx, subStreamReadInfo_, subStreamSendInfo_, subStreamZcopyReadInfo_, subStreamZcopySendInfo_,
423 0 : partialCommRankSet_);
424 : }
425 0 : if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
426 0 : nextSubStreamReadInfo_.clear();
427 0 : nextSubStreamSendInfo_.clear();
428 0 : if (needAlltoallvCache_) {
429 0 : nextSubStreamZcopyReadInfo_.clear();
430 0 : nextSubStreamZcopySendInfo_.clear();
431 : }
432 0 : UpdateSendRecvInfo(
433 0 : step, roundIdx + 1, nextSubStreamReadInfo_, nextSubStreamSendInfo_, nextSubStreamZcopyReadInfo_,
434 0 : nextSubStreamZcopySendInfo_, nextPartialCommRankSet_);
435 : }
436 0 : }
437 :
438 0 : HcclResult AlltoAllVDirectFullMesh::PrepareIntraData(
439 : u32 step, std::unordered_map<u32, std::vector<SendDataBlock>>& subStreamSendInfo,
440 : std::unordered_map<u32, SendDataBlock>& subStreamZcopySendInfo)
441 : {
442 0 : u32 sendDataIndex = 0;
443 0 : for (auto& sdmaInfo : subStreamSendInfo) {
444 0 : const std::vector<SendDataBlock>& sendInfo = sdmaInfo.second;
445 :
446 : // 对于alltoallv类算子, 零长拷贝需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发
447 : // 注意: alltoallv aicpu cache只考虑小数据量 (即max step为1), 所以只需要在step 0时下发一个placeholder SQE即可
448 0 : if (needAlltoallvCache_) {
449 : // alltoallv cache只针对小数据量, 至多只有1个step
450 0 : CHK_PRT_RET(
451 : step != 0,
452 : HCCL_ERROR(
453 : "[AlltoAllVDirectFullMesh][UpdateCurrRankRecvInfo] needAlltoallvCache_[%u] step[%u]",
454 : needAlltoallvCache_, step),
455 : HCCL_E_INTERNAL);
456 :
457 0 : const u32 sendRank = sdmaInfo.first;
458 0 : std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo.find(sendRank);
459 0 : if (mapIter != subStreamZcopySendInfo.end()) { // sendRank的sendCount为0
460 : // 零长拷贝下, sendRank对应的step和sendInfo.size一定为0
461 0 : CHK_PRT_RET(
462 : sendNumSubStep_[sdmaInfo.first] > 0,
463 : HCCL_ERROR("invalid sendNumSubStep_[%u][%u] != 0", sdmaInfo.first, sendNumSubStep_[sdmaInfo.first]),
464 : HCCL_E_INTERNAL);
465 0 : CHK_PRT_RET(
466 : sendInfo.size() > 0,
467 : HCCL_ERROR(
468 : "[AlltoAllVDirectFullMesh][PrepareIntraData] invalid sendInfo.size[%u] != 0", sendInfo.size()),
469 : HCCL_E_INTERNAL);
470 :
471 : // 获取零长拷贝的发送偏移
472 0 : const SendDataBlock& sendBlock = mapIter->second;
473 0 : CHK_PRT_RET(
474 : sendBlock.sendLen != 0,
475 : HCCL_ERROR(
476 : "[AlltoAllVDirectFullMesh][PrepareIntraData] invalid sendBlock.sendLen[%llu] != 0",
477 : sendBlock.sendLen),
478 : HCCL_E_INTERNAL);
479 :
480 : // 强制调用HcclD2DMemcpyAsync下发cache-memcpy placeholder (aicpu cache使能时才会生效,
481 : // 未使能时会直接返回)
482 0 : DeviceMem src = userInput_.range(sendBlock.userInOffset, sendBlock.sendLen);
483 0 : DeviceMem dst = cclInMem_.range(sendBlock.scratchOffset, sendBlock.sendLen);
484 0 : HCCL_DEBUG(
485 : "[AlltoAllVDirectFullMesh][PrepareIntraData]userRank [%u] copy from userInOffset[%llu]"
486 : "len[%u] to scratchOffset [%llu]",
487 : userRank_, sendBlock.userInOffset, sendBlock.sendLen, sendBlock.scratchOffset);
488 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
489 0 : HCCL_INFO(
490 : "[AlltoAllVDirectFullMesh][PrepareIntraData] generate cache-memcpy placeholder for sendRank[%u]",
491 : sendRank);
492 0 : if (isBigCount_) {
493 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, localSubStream_[sendDataIndex]));
494 : } else {
495 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
496 : }
497 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
498 0 : }
499 : }
500 :
501 0 : if (step < sendNumSubStep_[sdmaInfo.first]) {
502 0 : DeviceMem src = userInput_.range(sendInfo[step].userInOffset, sendInfo[step].sendLen);
503 0 : DeviceMem dst = cclInMem_.range(sendInfo[step].scratchOffset, sendInfo[step].sendLen);
504 0 : HCCL_DEBUG(
505 : "[AlltoAllVDirectFullMesh][PrepareIntraData]userRank [%u] copy from userInOffset[%llu]"
506 : "len[%u] to scratchOffset [%llu]",
507 : userRank_, sendInfo[step].userInOffset, sendInfo[step].sendLen, sendInfo[step].scratchOffset);
508 0 : if (isBigCount_) {
509 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, localSubStream_[sendDataIndex]));
510 : } else {
511 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
512 : }
513 0 : }
514 0 : sendDataIndex++;
515 : }
516 0 : return HCCL_SUCCESS;
517 : }
518 :
519 0 : void AlltoAllVDirectFullMesh::UpdateRemoteRankSet(u32 roundIdx, u32 groupRankSize)
520 : {
521 0 : if (sdmaConcurrentNum_ == 1) {
522 0 : UpdatePartialCommunicationRankSetPairWise(roundIdx, groupRankSize);
523 : } else {
524 0 : UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
525 : }
526 0 : }
527 :
528 0 : void AlltoAllVDirectFullMesh::UpdatePartialCommunicationRankSetPairWise(u32 roundIdx, u32 groupRankSize)
529 : {
530 0 : partialCommRankSet_.clear();
531 0 : partialCommRankSet_.resize(1);
532 0 : for (u32 i = roundIdx * sdmaConcurrentNum_; i < (roundIdx * sdmaConcurrentNum_ + groupRankSize); i++) {
533 0 : u32 readRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
534 0 : u32 sendRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
535 0 : partialCommRankSet_[0].push_back(std::make_pair(readRemoteRank, sendRemoteRank));
536 0 : HCCL_DEBUG(
537 : "[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSetPairWise] userRank [%u] i[%u]"
538 : "readRemoteRank[%u] writeRemoteRank[%u]",
539 : userRank_, i, readRemoteRank, sendRemoteRank);
540 : }
541 0 : HCCL_DEBUG(
542 : "[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSetPairWise] partialCommRankSet_ size[%zu]",
543 : partialCommRankSet_[0].size());
544 0 : }
545 :
546 0 : void AlltoAllVDirectFullMesh::UpdatePartialCommunicationRankSet(
547 : u32 roundIdx, u32 groupRankSize, std::vector<std::vector<std::pair<u32, u32>>>& partialCommRankSet)
548 : {
549 0 : partialCommRankSet.clear();
550 0 : partialCommRankSet.resize(RANK_SET_COMPUTE_CONST + 1);
551 0 : u32 pairNumPerRound = sdmaConcurrentNum_ / RANK_SET_COMPUTE_CONST;
552 0 : u32 pairSize = (groupRankSize < sdmaConcurrentNum_) ?
553 0 : (groupRankSize + RANK_SET_COMPUTE_CONST - 1) / RANK_SET_COMPUTE_CONST :
554 : pairNumPerRound;
555 0 : for (u32 i = roundIdx * pairNumPerRound + 1; i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
556 0 : u32 leftRemoteRank = podStartRank_ + (rankIdxInPod_ + devNumInlocalPod_ - i) % devNumInlocalPod_;
557 0 : u32 rightRemoteRank = podStartRank_ + (rankIdxInPod_ + i) % devNumInlocalPod_;
558 0 : if (leftRemoteRank == rightRemoteRank) {
559 0 : partialCommRankSet[2].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
560 : } else {
561 0 : partialCommRankSet[0].push_back(std::make_pair(leftRemoteRank, leftRemoteRank));
562 0 : partialCommRankSet[1].push_back(std::make_pair(rightRemoteRank, rightRemoteRank));
563 : }
564 0 : HCCL_DEBUG(
565 : "[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSet] round[%u] userRank [%u] i[%u]"
566 : "read/write leftRemoteRank[%u] rightRemoteRank[%u]",
567 : roundIdx, userRank_, i, leftRemoteRank, rightRemoteRank);
568 : }
569 0 : HCCL_DEBUG(
570 : "[AlltoAllVDirectFullMesh][UpdatePartialCommunicationRankSet] round[%u] partialCommRankSet_ total size[%zu]",
571 : roundIdx, partialCommRankSet[0].size() + partialCommRankSet[1].size() + partialCommRankSet[2].size());
572 0 : }
573 :
574 : // 主流只需要通知当前子步骤需要收发数据的 SDMA 流,减少同步开销
575 0 : HcclResult AlltoAllVDirectFullMesh::NotifySubStreamStart()
576 : {
577 0 : for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
578 0 : CHK_RET(
579 : LocalNotify::Post(mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
580 0 : CHK_RET(LocalNotify::Wait(
581 : sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
582 : }
583 0 : for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
584 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, sdmaSubStream_[streamIndex], dispatcher_));
585 : }
586 0 : HCCL_DEBUG(
587 : "[AlltoAllVDirectFullMesh][NotifySubStreamStart] userRank [%u] main stream notify sdma stream [%s]", userRank_,
588 : GetStreamIndexString().c_str());
589 0 : return HCCL_SUCCESS;
590 : }
591 :
592 0 : HcclResult AlltoAllVDirectFullMesh::WaitSubStreamFinish()
593 : {
594 0 : for (u32 streamIndex = 0; streamIndex < subStreamReadInfo_.size(); streamIndex++) {
595 0 : CHK_RET(LocalNotify::Post(
596 : sdmaSubStream_[streamIndex], dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
597 0 : CHK_RET(
598 : LocalNotify::Wait(mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
599 : }
600 0 : HCCL_DEBUG(
601 : "[AlltoAllVDirectFullMesh][WaitSubStreamFinish] userRank [%u] main stream wait sdma stream [%s]", userRank_,
602 : GetStreamIndexString().c_str());
603 0 : return HCCL_SUCCESS;
604 : }
605 :
606 0 : HcclResult AlltoAllVDirectFullMesh::NotifyLocalSubStreamStart()
607 : {
608 0 : for (u32 streamIndex = 0; streamIndex < subStreamSendInfo_.size(); streamIndex++) {
609 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, localSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
610 0 : CHK_RET(LocalNotify::Wait(
611 : localSubStream_[streamIndex], dispatcher_, localSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
612 : }
613 0 : return HCCL_SUCCESS;
614 : }
615 :
616 0 : HcclResult AlltoAllVDirectFullMesh::WaitLocalSubStreamFinish()
617 : {
618 0 : for (u32 streamIndex = 0; streamIndex < subStreamSendInfo_.size(); streamIndex++) {
619 0 : CHK_RET(LocalNotify::Post(
620 : localSubStream_[streamIndex], dispatcher_, localSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
621 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, localSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
622 : }
623 0 : return HCCL_SUCCESS;
624 : }
625 :
626 0 : u32 AlltoAllVDirectFullMesh::CalcNumSubStep()
627 : {
628 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
629 :
630 0 : sendNumSubStep_.clear();
631 0 : recvNumSubStep_.clear();
632 0 : u32 numSubStep = 0;
633 :
634 0 : for (u32 destRank = podStartRank_; destRank < podStartRank_ + devNumInlocalPod_; destRank++) {
635 0 : if (destRank == userRank_) {
636 0 : continue;
637 : }
638 :
639 : u32 currRankSendSubStep
640 0 : = ((localSendRecvInfo.sendLength[destRank] + sdmaDataBlockSize_ - 1) / sdmaDataBlockSize_);
641 0 : sendNumSubStep_[destRank] = currRankSendSubStep;
642 :
643 : u32 currRankRecvSubStep
644 0 : = ((localSendRecvInfo.recvLength[destRank] + sdmaDataBlockSize_ - 1) / sdmaDataBlockSize_);
645 0 : recvNumSubStep_[destRank] = currRankRecvSubStep;
646 0 : HCCL_DEBUG(
647 : "[AlltoAllVDirectFullMesh][CalcNumSubStep] userRank [%u] currRankSendSubStep[%u]"
648 : "currRankRecvSubStep[%u]",
649 : userRank_, currRankSendSubStep, currRankRecvSubStep);
650 0 : numSubStep = std::max(numSubStep, std::max(currRankSendSubStep, currRankRecvSubStep));
651 : }
652 0 : HCCL_DEBUG(
653 : "[AlltoAllVDirectFullMesh][CalcNumSubStep] userRank [%u] max communication step[%u]", userRank_, numSubStep);
654 0 : return numSubStep;
655 : }
656 :
657 0 : HcclResult AlltoAllVDirectFullMesh::NotifyRemoteRankStart(u32 step)
658 : {
659 0 : u32 streamIndex = 0;
660 0 : for (auto& sendRecvSide : partialCommRankSet_) {
661 0 : for (auto& sendRecvPair : sendRecvSide) {
662 0 : u32 recvRank = sendRecvPair.first;
663 0 : u32 sendRank = sendRecvPair.second;
664 0 : if (sendRank == userRank_) {
665 0 : continue;
666 : }
667 0 : const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
668 0 : const std::vector<SendDataBlock>& sendInfo = subStreamSendInfo_[sendRank];
669 0 : Stream& currStream = sdmaSubStream_[streamIndex];
670 0 : const LINK& readTransport = links_[recvRank];
671 0 : const LINK& sendTransport = links_[sendRank];
672 :
673 0 : if (needAlltoallvCache_) {
674 : // alltoallv cache只针对小数据量, 至多只有1个step
675 0 : CHK_PRT_RET(
676 : step != 0,
677 : HCCL_ERROR(
678 : "[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] needAlltoallvCache_[%u] step[%u]",
679 : needAlltoallvCache_, step),
680 : HCCL_E_INTERNAL);
681 :
682 0 : std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo_.find(sendRank);
683 0 : if (mapIter != subStreamZcopySendInfo_.end()) { // sendRank的sendCount为0
684 : // 零长拷贝下, sendRank对应的sendInfo.size一定为0
685 0 : CHK_PRT_RET(
686 : sendInfo.size() > 0,
687 : HCCL_ERROR(
688 : "[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] invalid sendInfo.size[%u] != 0",
689 : sendInfo.size()),
690 : HCCL_E_INTERNAL);
691 :
692 : // 生成cache-write placeholder
693 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
694 0 : HCCL_INFO(
695 : "[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] generate cache-write placeholder for "
696 : "sendRank[%u]",
697 : sendRank);
698 0 : CHK_RET(sendTransport->TxAck(currStream));
699 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
700 : }
701 : }
702 0 : if (step < sendInfo.size()) {
703 0 : CHK_RET(sendTransport->TxAck(currStream));
704 : }
705 :
706 0 : if (needAlltoallvCache_) {
707 0 : std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
708 0 : if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
709 : // 零长拷贝下, recvRank对应的readInfo.size一定为0
710 0 : CHK_PRT_RET(
711 : readInfo.size() > 0,
712 : HCCL_ERROR(
713 : "[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] invalid readInfo.size[%u] != 0",
714 : readInfo.size()),
715 : HCCL_E_INTERNAL);
716 :
717 : // 生成cache-write placeholder
718 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
719 0 : HCCL_INFO(
720 : "[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] generate cache-notify placeholder for "
721 : "recvRank[%u]",
722 : recvRank);
723 0 : CHK_RET(readTransport->RxAck(currStream));
724 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
725 : }
726 : }
727 0 : if (step < readInfo.size()) {
728 0 : CHK_RET(readTransport->RxAck(currStream));
729 : }
730 0 : streamIndex++;
731 : }
732 : }
733 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][NotifyRemoteRankStart] done");
734 0 : return HCCL_SUCCESS;
735 : }
736 :
737 0 : bool AlltoAllVDirectFullMesh::IsPostSyncEnable(u32 step, u32 roundIdx)
738 : {
739 0 : bool isPostSyncEnable = false;
740 0 : isPostSyncEnable = (step == lastStep_) && (roundIdx == lastRoundIdx_) && algOpContext_.opRetryHandler.retryEnable;
741 0 : return isPostSyncEnable;
742 : }
743 :
744 0 : HcclResult AlltoAllVDirectFullMesh::SdmaMainStreamWait(u32 step, u32 roundIdx)
745 : {
746 : // SDMA wait
747 0 : u32 streamIndex = 0;
748 0 : for (auto& sendRecvSide : partialCommRankSet_) {
749 0 : for (auto& sendRecvPair : sendRecvSide) {
750 0 : u32 recvRank = sendRecvPair.first;
751 0 : u32 sendRank = sendRecvPair.second;
752 0 : if (sendRank == userRank_) {
753 0 : continue;
754 : }
755 0 : const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
756 :
757 0 : if (needAlltoallvCache_) {
758 : // alltoallv cache只针对小数据量, 至多只有1个step
759 0 : CHK_PRT_RET(
760 : step != 0,
761 : HCCL_ERROR(
762 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] needAlltoallvCache_[%u] step[%u]",
763 : needAlltoallvCache_, step),
764 : HCCL_E_INTERNAL);
765 :
766 0 : std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
767 0 : if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
768 : // 零长拷贝下, recvRank对应的readInfo.size一定为0
769 0 : CHK_PRT_RET(
770 : readInfo.size() > 0,
771 : HCCL_ERROR(
772 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] invalid readInfo.size[%u] != 0",
773 : readInfo.size()),
774 : HCCL_E_INTERNAL);
775 :
776 : // 正常下NotifyWait SQE (本地主从流同步, 由于从流不存在跨卡数据搬运, 主流wait后会立刻wake up)
777 0 : HCCL_DEBUG(
778 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
779 : "sendRank[%u], sdma stream [%u], "
780 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream wait",
781 : userRank_, recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
782 0 : CHK_RET(LocalNotify::Wait(
783 : mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
784 : }
785 : }
786 :
787 0 : if (step < readInfo.size()) {
788 0 : HCCL_DEBUG(
789 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] userRank [%u], recvRank[%u], "
790 : "sendRank[%u], sdma stream [%u], "
791 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream wait",
792 : userRank_, recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
793 0 : CHK_RET(LocalNotify::Wait(
794 : mainStream_, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
795 : }
796 0 : streamIndex++;
797 : }
798 : }
799 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][SdmaMainStreamWait] done");
800 0 : return HCCL_SUCCESS;
801 : }
802 :
803 0 : HcclResult AlltoAllVDirectFullMesh::SdmaMainStreamPost(u32 step, u32 roundIdx)
804 : {
805 : // SDMA post
806 0 : u32 streamIndex = 0;
807 0 : for (auto& sendRecvSide : partialCommRankSet_) {
808 0 : for (auto& sendRecvPair : sendRecvSide) {
809 0 : u32 recvRank = sendRecvPair.first;
810 0 : u32 sendRank = sendRecvPair.second;
811 0 : if (sendRank == userRank_) {
812 0 : continue;
813 : }
814 0 : const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
815 :
816 0 : if (needAlltoallvCache_) {
817 : // alltoallv cache只针对小数据量, 至多只有1个step
818 0 : CHK_PRT_RET(
819 : step != 0,
820 : HCCL_ERROR(
821 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] needAlltoallvCache_[%u] step[%u]",
822 : needAlltoallvCache_, step),
823 : HCCL_E_INTERNAL);
824 :
825 0 : std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
826 0 : if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
827 : // 零长拷贝下, recvRank对应的readInfo.size一定为0
828 0 : CHK_PRT_RET(
829 : readInfo.size() > 0,
830 : HCCL_ERROR(
831 : "[AlltoAllVDirectFullMesh][SdmaMainStreamWait] invalid readInfo.size[%u] != 0",
832 : readInfo.size()),
833 : HCCL_E_INTERNAL);
834 :
835 : // 正常下NotifyRecord SQE
836 0 : HCCL_DEBUG(
837 : "[AlltoAllVDirectFullMesh][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
838 : "sendRank[%u], sdma stream [%u], "
839 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream post",
840 : userRank_, recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
841 0 : CHK_RET(LocalNotify::Post(
842 : mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
843 : }
844 : }
845 :
846 0 : if (step < readInfo.size()) {
847 0 : HCCL_DEBUG(
848 : "[AlltoAllVDirectFullMesh][SdmaMainStreamPost] userRank [%u], recvRank[%u], "
849 : "sendRank[%u], sdma stream [%u], "
850 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u] main stream post",
851 : userRank_, recvRank, sendRank, streamIndex, step, roundIdx, lastStep_, lastRoundIdx_);
852 0 : CHK_RET(LocalNotify::Post(
853 : mainStream_, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
854 : }
855 0 : streamIndex++;
856 : }
857 : }
858 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][SdmaMainStreamPost] done");
859 0 : return HCCL_SUCCESS;
860 : }
861 :
862 0 : HcclResult AlltoAllVDirectFullMesh::SetPostSyncTasks(u32 step, u32 roundIdx)
863 : {
864 : // SDMA wait
865 0 : CHK_RET(SdmaMainStreamWait(step, roundIdx));
866 0 : if (rdmaConcurrentNum_ > 0) {
867 : // RDMA wait
868 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] rdma post sync info: main stream wait");
869 0 : CHK_RET(RdmaControlNotifyMainFinish());
870 : }
871 : // SDMA post
872 0 : CHK_RET(SdmaMainStreamPost(step, roundIdx));
873 0 : if (rdmaConcurrentNum_ > 0) {
874 : // RDMA post
875 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] rdma post sync info: main stream post");
876 0 : CHK_RET(MainNotifyRdmaControlStart());
877 : }
878 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SetPostSyncTasks] done");
879 0 : return HCCL_SUCCESS;
880 : }
881 :
882 0 : HcclResult AlltoAllVDirectFullMesh::SDMAwithRemoteRankAndNotifyEnd(u32 step, u32 roundIdx)
883 : {
884 0 : bool isPostSyncEnable = IsPostSyncEnable(step, roundIdx);
885 0 : if (isPostSyncEnable) {
886 : // 下发主流上的后同步wait和post
887 0 : CHK_RET(SetPostSyncTasks(step, roundIdx));
888 : }
889 0 : u32 streamIndex = 0;
890 0 : for (auto& sendRecvSide : partialCommRankSet_) {
891 0 : for (auto& sendRecvPair : sendRecvSide) {
892 0 : u32 recvRank = sendRecvPair.first;
893 0 : u32 sendRank = sendRecvPair.second;
894 0 : if (sendRank == userRank_) {
895 0 : continue;
896 : }
897 0 : const std::vector<ReadDataBlock>& readInfo = subStreamReadInfo_[recvRank];
898 0 : const std::vector<SendDataBlock>& sendInfo = subStreamSendInfo_[sendRank];
899 0 : Stream& currStream = sdmaSubStream_[streamIndex];
900 0 : const LINK& readTransport = links_[recvRank];
901 0 : const LINK& sendTransport = links_[sendRank];
902 :
903 : // 对于alltoallv类算子, 零长拷贝需要调用MemcpyAsync保证aicpu cache使能时placeholder正确下发
904 : // 注意: alltoallv aicpu cache只考虑小数据量 (即max step为1), 所以只需要在step 0时下发一个placeholder
905 : // SQE即可
906 0 : if (needAlltoallvCache_) {
907 : // alltoallv cache只针对小数据量, 至多只有1个step
908 0 : CHK_PRT_RET(
909 : step != 0,
910 : HCCL_ERROR(
911 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] needAlltoallvCache_[%u] step[%u]",
912 : needAlltoallvCache_, step),
913 : HCCL_E_INTERNAL);
914 :
915 0 : std::unordered_map<u32, ReadDataBlock>::const_iterator mapIter = subStreamZcopyReadInfo_.find(recvRank);
916 0 : if (mapIter != subStreamZcopyReadInfo_.end()) { // recvRank的recvCount为0
917 : // 零长拷贝下, recvRank对应的readInfo.size一定为0
918 0 : CHK_PRT_RET(
919 : readInfo.size() > 0,
920 : HCCL_ERROR(
921 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid readInfo.size[%u] != 0",
922 : readInfo.size()),
923 : HCCL_E_INTERNAL);
924 :
925 : // 获取零长拷贝的接收偏移
926 0 : const ReadDataBlock& readBlock = mapIter->second;
927 0 : CHK_PRT_RET(
928 : readBlock.recvLen != 0,
929 : HCCL_ERROR(
930 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid readBlock.recvLen[%llu] "
931 : "!= 0",
932 : readBlock.recvLen),
933 : HCCL_E_INTERNAL);
934 :
935 : // 强制调用HcclD2DMemcpyAsync下发cache-memcpy/write placeholder (aicpu cache使能时才会下发,
936 : // 未使能时会直接返回)
937 0 : const LINK& intraNeighboorTransport = links_[recvRank];
938 0 : CHK_PTR_NULL(intraNeighboorTransport);
939 0 : void* remDMAMemPtr = nullptr;
940 0 : CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
941 0 : DeviceMem remoteCCLInMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), cclInMem_.size());
942 0 : DeviceMem srcMem = remoteCCLInMem.range(readBlock.remoteOffset, readBlock.recvLen);
943 0 : DeviceMem dstMem = userOutput_.range(readBlock.recvOffset, readBlock.recvLen);
944 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
945 0 : HCCL_INFO(
946 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-memcpy placeholder "
947 : "for recvRank[%u]",
948 : recvRank);
949 0 : CHK_RET(HcclD2DMemcpyAsync(
950 : dispatcher_, dstMem, srcMem, currStream, readTransport->GetRemoteRank(),
951 : readTransport->GetLinkType()));
952 0 : HCCL_INFO(
953 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-write placeholder "
954 : "for recvRank[%u]",
955 : recvRank);
956 0 : CHK_RET(readTransport->TxDataSignal(currStream));
957 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
958 :
959 : // 正常下NotifyRecord/Wait SQE (本地主从流同步, 从流不存在跨卡数据拷贝,
960 : // 下发placeholder后会立刻post主流并进入wait)
961 0 : HCCL_DEBUG(
962 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] userRank [%u], recvRank[%u], "
963 : "sendRank[%u],"
964 : "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
965 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u]",
966 : userRank_, recvRank, sendRank, streamIndex, readBlock.remoteOffset, readBlock.recvLen,
967 : readBlock.recvOffset, step, roundIdx, lastStep_, lastRoundIdx_);
968 0 : if (isPostSyncEnable) {
969 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] post sync begins");
970 0 : CHK_RET(LocalNotify::Post(
971 : currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
972 0 : CHK_RET(LocalNotify::Wait(
973 : currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
974 : }
975 0 : }
976 : }
977 :
978 0 : if (step < readInfo.size()) {
979 0 : const LINK& intraNeighboorTransport = links_[recvRank];
980 0 : CHK_PTR_NULL(intraNeighboorTransport);
981 0 : void* remDMAMemPtr = nullptr;
982 0 : CHK_RET(intraNeighboorTransport->GetRemoteMem(UserMemType::INPUT_MEM, &remDMAMemPtr));
983 0 : DeviceMem remoteCCLInMem = DeviceMem::create(static_cast<u8*>(remDMAMemPtr), cclInMem_.size());
984 0 : DeviceMem srcMem = remoteCCLInMem.range(readInfo[step].remoteOffset, readInfo[step].recvLen);
985 0 : DeviceMem dstMem = userOutput_.range(readInfo[step].recvOffset, readInfo[step].recvLen);
986 0 : CHK_RET(HcclD2DMemcpyAsync(
987 : dispatcher_, dstMem, srcMem, currStream, readTransport->GetRemoteRank(),
988 : readTransport->GetLinkType()));
989 0 : HCCL_DEBUG(
990 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] userRank [%u], recvRank[%u], "
991 : "sendRank[%u],"
992 : "sdma stream [%u] read data from remote offset [%llu] len [%llu] to local [%llu], "
993 : "post sync info: step[%u], roundIdx[%u], lastStep_[%u], lastRoundIdx_[%u]",
994 : userRank_, recvRank, sendRank, streamIndex, readInfo[step].remoteOffset, readInfo[step].recvLen,
995 : readInfo[step].recvOffset, step, roundIdx, lastStep_, lastRoundIdx_);
996 0 : if (isPostSyncEnable) {
997 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] post sync begins");
998 0 : CHK_RET(LocalNotify::Post(
999 : currStream, dispatcher_, sdmaMeshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
1000 0 : CHK_RET(LocalNotify::Wait(
1001 : currStream, dispatcher_, sdmaMeshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
1002 : }
1003 0 : CHK_RET(readTransport->TxDataSignal(currStream));
1004 0 : }
1005 :
1006 0 : if (needAlltoallvCache_) {
1007 0 : std::unordered_map<u32, SendDataBlock>::const_iterator mapIter = subStreamZcopySendInfo_.find(sendRank);
1008 0 : if (mapIter != subStreamZcopySendInfo_.end()) { // sendRank的sendCount为0
1009 : // 零长拷贝下, sendRank对应的sendInfo.size一定为0
1010 0 : CHK_PRT_RET(
1011 : sendInfo.size() > 0,
1012 : HCCL_ERROR(
1013 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] invalid sendInfo.size[%u] != 0",
1014 : sendInfo.size()),
1015 : HCCL_E_INTERNAL);
1016 :
1017 : // 生成cache-notify placeholder
1018 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
1019 0 : HCCL_INFO(
1020 : "[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] generate cache-notify placeholder "
1021 : "for sendRank[%u]",
1022 : sendRank);
1023 0 : CHK_RET(sendTransport->RxDataSignal(currStream));
1024 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
1025 : }
1026 : }
1027 :
1028 0 : if (step < sendInfo.size()) {
1029 0 : CHK_RET(sendTransport->RxDataSignal(currStream));
1030 : }
1031 0 : streamIndex++;
1032 : }
1033 : }
1034 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][SDMAwithRemoteRankAndNotifyEnd] done");
1035 0 : return HCCL_SUCCESS;
1036 : }
1037 :
1038 0 : HcclResult AlltoAllVDirectFullMesh::SendRecvData(u32 step, u32 roundIdx)
1039 : {
1040 0 : HCCL_DEBUG(
1041 : "[AlltoAllVDirectFullMesh][SendRecvData] userRank [%u] sdma stream [%s] wait main stream", userRank_,
1042 : GetStreamIndexString().c_str());
1043 0 : CHK_RET(NotifyRemoteRankStart(step));
1044 0 : CHK_RET(WaitSubStreamFinish());
1045 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1046 0 : CHK_RET(NotifySubStreamStart());
1047 0 : if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
1048 0 : CHK_RET(NotifyLocalSubStreamStart());
1049 0 : CHK_RET(PrepareIntraData(step, nextSubStreamSendInfo_, nextSubStreamZcopySendInfo_));
1050 : }
1051 0 : CHK_RET(SDMAwithRemoteRankAndNotifyEnd(step, roundIdx));
1052 :
1053 0 : return HCCL_SUCCESS;
1054 : }
1055 :
1056 0 : HcclResult AlltoAllVDirectFullMesh::LocalCopy()
1057 : {
1058 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
1059 0 : DeviceMem src = userInput_.range(localSendRecvInfo.sendOffset[userRank_], localSendRecvInfo.sendLength[userRank_]);
1060 0 : DeviceMem dst = userOutput_.range(localSendRecvInfo.recvOffset[userRank_], localSendRecvInfo.recvLength[userRank_]);
1061 0 : HCCL_DEBUG(
1062 : "[AlltoAllVDirectFullMesh][LocalCopy]userRank [%u] copy from userInput [%llu] len [%llu]"
1063 : "to userOutput [%llu] dstLen[%llu]",
1064 : userRank_, localSendRecvInfo.sendOffset[userRank_], localSendRecvInfo.sendLength[userRank_],
1065 : localSendRecvInfo.recvOffset[userRank_], localSendRecvInfo.recvLength[userRank_]);
1066 0 : if (needAlltoallvCache_ && localSendRecvInfo.sendLength[userRank_] == 0) {
1067 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(true);
1068 : }
1069 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, mainStream_));
1070 0 : if (needAlltoallvCache_ && localSendRecvInfo.sendLength[userRank_] == 0) {
1071 0 : reinterpret_cast<DispatcherPub*>(dispatcher_)->SetPlaceholder(false);
1072 : }
1073 :
1074 0 : return HCCL_SUCCESS;
1075 0 : }
1076 :
1077 0 : HcclResult AlltoAllVDirectFullMesh::RunGroupFullMeshAlltoall(u32 roundIdx, u32 step)
1078 : {
1079 0 : UpdateOpBaseSubStreamInfo(step, roundIdx);
1080 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1081 0 : if (isBigCount_ && (roundIdx == 0)) {
1082 0 : CHK_RET(NotifyLocalSubStreamStart());
1083 0 : CHK_RET(PrepareIntraData(step, subStreamSendInfo_, subStreamZcopySendInfo_));
1084 0 : CHK_RET(WaitLocalSubStreamFinish());
1085 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1086 0 : } else if (!isBigCount_) {
1087 0 : CHK_RET(PrepareIntraData(step, subStreamSendInfo_, subStreamZcopySendInfo_));
1088 : }
1089 0 : CHK_RET(NotifySubStreamStart());
1090 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1091 0 : CHK_RET(SendRecvData(step, roundIdx));
1092 0 : if (step == 0 && !islocalCpyDone_) {
1093 0 : CHK_RET(LocalCopy());
1094 0 : islocalCpyDone_ = true;
1095 : }
1096 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1097 0 : CHK_RET(WaitSubStreamFinish());
1098 0 : if (isBigCount_ && (roundIdx < commRounds_ - 1)) {
1099 0 : CHK_RET(WaitLocalSubStreamFinish());
1100 : }
1101 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1102 0 : return HCCL_SUCCESS;
1103 : }
1104 :
1105 : // 主流通知RDMA控制流启动
1106 0 : HcclResult AlltoAllVDirectFullMesh::MainNotifyRdmaControlStart()
1107 : {
1108 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, rdmaControl2MainStreamNotify_, INVALID_VALUE_STAGE));
1109 0 : CHK_RET(LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2MainStreamNotify_, INVALID_VALUE_STAGE));
1110 0 : return HCCL_SUCCESS;
1111 : }
1112 :
1113 : // RDMA控制流通知主流任务完成
1114 0 : HcclResult AlltoAllVDirectFullMesh::RdmaControlNotifyMainFinish()
1115 : {
1116 0 : CHK_RET(LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, main2RdmaControlStreamNotify_, INVALID_VALUE_STAGE));
1117 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, main2RdmaControlStreamNotify_, INVALID_VALUE_STAGE));
1118 0 : return HCCL_SUCCESS;
1119 : }
1120 :
1121 : // RDMA控制流通知从流启动任务
1122 0 : HcclResult AlltoAllVDirectFullMesh::RdmaControlNotifySubStart()
1123 : {
1124 0 : for (u32 i = 1; i < rdmaSubStreams_.size(); i++) {
1125 0 : CHK_RET(
1126 : LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, rdmaSub2ControlNotifies_[i - 1], INVALID_VALUE_STAGE));
1127 0 : CHK_RET(
1128 : LocalNotify::Wait(rdmaSubStreams_[i], dispatcher_, rdmaSub2ControlNotifies_[i - 1], INVALID_VALUE_STAGE));
1129 : }
1130 :
1131 0 : return HCCL_SUCCESS;
1132 : }
1133 :
1134 : // 从流通知RDMA控制流任务结束
1135 0 : HcclResult AlltoAllVDirectFullMesh::SubNotifyRdmaControlFinish()
1136 : {
1137 0 : for (u32 i = 1; i < rdmaSubStreams_.size(); i++) {
1138 0 : CHK_RET(
1139 : LocalNotify::Post(rdmaSubStreams_[i], dispatcher_, rdmaControl2SubNotifies_[i - 1], INVALID_VALUE_STAGE));
1140 0 : CHK_RET(
1141 : LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2SubNotifies_[i - 1], INVALID_VALUE_STAGE));
1142 : }
1143 :
1144 0 : return HCCL_SUCCESS;
1145 : }
1146 :
1147 0 : u32 AlltoAllVDirectFullMesh::GetNextDstRank(u32& curDstRank)
1148 : {
1149 0 : if (curDstRank >= userRankSize_) {
1150 0 : curDstRank = curDstRank % userRankSize_;
1151 : }
1152 0 : if (curDstRank == podStartRank_) {
1153 0 : curDstRank += devNumInlocalPod_;
1154 : }
1155 0 : curDstRank = curDstRank % userRankSize_;
1156 0 : return curDstRank++;
1157 : }
1158 :
1159 0 : u32 AlltoAllVDirectFullMesh::GetPreSrcRank(u32& curDstRank)
1160 : {
1161 0 : if (curDstRank == podStartRank_ + devNumInlocalPod_ - 1) {
1162 0 : curDstRank = (curDstRank + userRankSize_ - devNumInlocalPod_) % userRankSize_;
1163 : }
1164 :
1165 0 : if (curDstRank == 0) {
1166 0 : curDstRank = userRankSize_ - 1;
1167 0 : return 0;
1168 : }
1169 0 : return curDstRank--;
1170 : }
1171 :
1172 0 : void AlltoAllVDirectFullMesh::GenRdmaSendInfo(u32 dstRank, std::vector<SendDataBlock>& sendInfo)
1173 : {
1174 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
1175 0 : u64 sendOffset = localSendRecvInfo.sendOffset[dstRank];
1176 0 : u64 sendLength = localSendRecvInfo.sendLength[dstRank];
1177 0 : while (sendLength > 0) {
1178 0 : u64 curSendLength = std::min(sendLength, rdmaDataBlockSize_);
1179 : SendDataBlock sendData;
1180 0 : sendData.userInOffset = sendOffset;
1181 0 : sendData.sendLen = curSendLength;
1182 0 : u32 index = dstRank % rdmaConcurrentNum_;
1183 0 : sendData.scratchOffset = rdmaDataBlockSize_ * index;
1184 0 : sendInfo.push_back(sendData);
1185 0 : sendOffset += curSendLength;
1186 0 : sendLength -= curSendLength;
1187 0 : HCCL_DEBUG(
1188 : "[GenRdmaSendInfo] userRank[%u], dstRank[%u], sendData.userInOffset[%llu],"
1189 : "sendData.sendLen[%llu], sendData.scratchOffset[%llu]",
1190 : userRank_, dstRank, sendData.userInOffset, sendData.sendLen, sendData.scratchOffset);
1191 : }
1192 0 : return;
1193 : }
1194 :
1195 0 : void AlltoAllVDirectFullMesh::GenRdmaRecvInfo(u32 srcRank, std::vector<RecvDataBlock>& recvInfo)
1196 : {
1197 0 : const SendRecvInfo& localSendRecvInfo = *localSendRecvInfoPtr_;
1198 0 : u64 recvOffset = localSendRecvInfo.recvOffset[srcRank];
1199 0 : u64 recvLength = localSendRecvInfo.recvLength[srcRank];
1200 0 : while (recvLength > 0) {
1201 0 : u64 curRecvLength = std::min(recvLength, rdmaDataBlockSize_);
1202 : RecvDataBlock recvData;
1203 0 : recvData.recvOffset = recvOffset;
1204 0 : recvData.recvLen = curRecvLength;
1205 0 : u32 index = srcRank % rdmaConcurrentNum_;
1206 0 : recvData.scratchOffset = rdmaDataBlockSize_ * index + rdmaDataBlockSize_ * rdmaConcurrentNum_;
1207 0 : recvInfo.push_back(recvData);
1208 0 : recvOffset += curRecvLength;
1209 0 : recvLength -= curRecvLength;
1210 0 : HCCL_DEBUG(
1211 : "[GenRdmaRecvInfo] userRank[%llu], srcRank[%u], recvData.recvOffset[%llu],"
1212 : "recvData.recvLen[%llu], recvData.scratchOffset[%llu]",
1213 : userRank_, srcRank, recvData.recvOffset, recvData.recvLen, recvData.scratchOffset);
1214 : }
1215 0 : return;
1216 : }
1217 :
1218 : // 将数据从userIn拷贝到CCL out
1219 : HcclResult
1220 0 : AlltoAllVDirectFullMesh::CopyDataForSend(u32 dstRank, std::vector<SendDataBlock>& sendInfo, u32 curStep, Stream stream)
1221 : {
1222 0 : if (curStep >= sendInfo.size()) {
1223 0 : return HCCL_SUCCESS;
1224 : }
1225 0 : DeviceMem src = userInput_.range(sendInfo[curStep].userInOffset, sendInfo[curStep].sendLen);
1226 0 : DeviceMem dst = cclOutMem_.range(sendInfo[curStep].scratchOffset, sendInfo[curStep].sendLen);
1227 0 : HCCL_DEBUG(
1228 : "[CopyDataForSend] userRank[%u], dstRank[%u], userInOffset[%llu], sendLen[%llu], scratchOffset[%llu]",
1229 : userRank_, dstRank, sendInfo[curStep].userInOffset, sendInfo[curStep].sendLen, sendInfo[curStep].scratchOffset);
1230 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream));
1231 0 : return HCCL_SUCCESS;
1232 0 : }
1233 :
1234 0 : HcclResult AlltoAllVDirectFullMesh::RdmaPostSync(Stream& stream)
1235 : {
1236 0 : CHK_RET(LocalNotify::Post(stream, dispatcher_, rdmaControl2SubNotifies_[0], INVALID_VALUE_STAGE));
1237 0 : CHK_RET(LocalNotify::Wait(rdmaSubStreams_[0], dispatcher_, rdmaControl2SubNotifies_[0], INVALID_VALUE_STAGE));
1238 :
1239 0 : CHK_RET(LocalNotify::Post(rdmaSubStreams_[0], dispatcher_, rdmaSub2ControlNotifies_[0], INVALID_VALUE_STAGE));
1240 0 : CHK_RET(LocalNotify::Wait(stream, dispatcher_, rdmaSub2ControlNotifies_[0], INVALID_VALUE_STAGE));
1241 0 : return HCCL_SUCCESS;
1242 : }
1243 :
1244 : // 从流完成RDMA数据的收发
1245 0 : HcclResult AlltoAllVDirectFullMesh::SendRecvRdmaData(
1246 : u32 dstRank, u32 srcRank, std::vector<SendDataBlock>& sendInfo, std::vector<RecvDataBlock>& recvInfo, u32 round,
1247 : u32 index, u32 curStep, Stream stream)
1248 : {
1249 0 : const LINK& sendTransport = links_[dstRank];
1250 0 : const LINK& recvTransport = links_[srcRank];
1251 0 : HCCL_DEBUG(
1252 : "[AlltoAllVDirectFullMesh][SendRecvRdmaData] userRank[%u], dstRank[%u], srcRank[%u]", userRank_, dstRank,
1253 : srcRank);
1254 0 : u32 minStep = std::min(sendInfo.size(), recvInfo.size());
1255 0 : CHK_PTR_NULL(sendTransport);
1256 0 : CHK_PTR_NULL(recvTransport);
1257 0 : if (curStep < minStep) {
1258 0 : CHK_RET(recvTransport->TxAck(stream));
1259 0 : CHK_RET(sendTransport->RxAck(stream));
1260 0 : u64 sendSrcOffset = (dstRank % rdmaConcurrentNum_) * rdmaDataBlockSize_;
1261 0 : void* srcPtr = static_cast<u8*>(cclOutMem_.ptr()) + sendSrcOffset;
1262 0 : u32 dstIndex = userRank_ % rdmaConcurrentNum_;
1263 0 : u64 sendDstOffset = (dstIndex + rdmaConcurrentNum_) * rdmaDataBlockSize_;
1264 0 : CHK_RET(
1265 : sendTransport->TxAsync(UserMemType::OUTPUT_MEM, sendDstOffset, srcPtr, sendInfo[curStep].sendLen, stream));
1266 :
1267 0 : u64 recvDstOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
1268 0 : void* dstPtr = static_cast<u8*>(cclOutMem_.ptr()) + recvDstOffset;
1269 0 : u64 recvSrcOffset = (userRank_ % rdmaConcurrentNum_) * rdmaDataBlockSize_;
1270 0 : CHK_RET(
1271 : recvTransport->RxAsync(UserMemType::OUTPUT_MEM, recvSrcOffset, dstPtr, recvInfo[curStep].recvLen, stream));
1272 0 : if ((round == lastRdmaRoundIdx_) && (index == lastRdmaDstRanksIdx_) && (curStep == lastRdmaStep_)
1273 0 : && (sdmaConcurrentNum_ > 1) && algOpContext_.opRetryHandler.retryEnable) {
1274 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] post sync begins");
1275 0 : CHK_RET(RdmaPostSync(stream));
1276 : }
1277 0 : CHK_RET(recvTransport->PostFinAck(stream));
1278 0 : CHK_RET(sendTransport->WaitFinAck(stream));
1279 0 : HCCL_DEBUG(
1280 : "[AlltoAllVDirectFullMesh][SendRecvRdmaData] sendSrcOffset[%llu], sendDstOffset[%llu],"
1281 : "recvDstOffset[%llu], recvSrcOffset[%llu], srcPtr[%p], dstPtr[%p]",
1282 : sendSrcOffset, sendDstOffset, recvDstOffset, recvSrcOffset, srcPtr, dstPtr);
1283 0 : } else if (curStep < sendInfo.size()) {
1284 0 : CHK_RET(sendTransport->RxAck(stream));
1285 0 : u64 sendSrcOffset = (dstRank % rdmaConcurrentNum_) * rdmaDataBlockSize_;
1286 0 : void* srcPtr = static_cast<u8*>(cclOutMem_.ptr()) + sendSrcOffset;
1287 0 : u32 dstIndex = userRank_ % rdmaConcurrentNum_;
1288 0 : u64 sendDstOffset = (dstIndex + rdmaConcurrentNum_) * rdmaDataBlockSize_;
1289 0 : CHK_RET(
1290 : sendTransport->TxAsync(UserMemType::OUTPUT_MEM, sendDstOffset, srcPtr, sendInfo[curStep].sendLen, stream));
1291 0 : CHK_RET(sendTransport->WaitFinAck(stream));
1292 : } else {
1293 0 : CHK_RET(recvTransport->TxAck(stream));
1294 0 : u64 recvDstOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
1295 0 : void* dstPtr = static_cast<u8*>(cclOutMem_.ptr()) + recvDstOffset;
1296 0 : u64 recvSrcOffset = (userRank_ % rdmaConcurrentNum_) * rdmaDataBlockSize_;
1297 0 : CHK_RET(
1298 : recvTransport->RxAsync(UserMemType::OUTPUT_MEM, recvSrcOffset, dstPtr, recvInfo[curStep].recvLen, stream));
1299 0 : if ((round == lastRdmaRoundIdx_) && (index == lastRdmaDstRanksIdx_) && (curStep == lastRdmaStep_)
1300 0 : && (sdmaConcurrentNum_ > 1) && algOpContext_.opRetryHandler.retryEnable) {
1301 0 : HCCL_DEBUG("[AlltoAllVDirectFullMesh][SendRecvRdmaData] post sync begins");
1302 0 : CHK_RET(RdmaPostSync(stream));
1303 : }
1304 0 : CHK_RET(recvTransport->PostFinAck(stream));
1305 : }
1306 0 : return HCCL_SUCCESS;
1307 : }
1308 :
1309 : // 从流将接收到的数据拷贝到输出
1310 0 : HcclResult AlltoAllVDirectFullMesh::CopyRecvDataToOutput(
1311 : u32 srcRank, std::vector<RecvDataBlock>& recvInfo, u32 curStep, Stream stream)
1312 : {
1313 0 : if (curStep >= recvInfo.size()) {
1314 0 : return HCCL_SUCCESS;
1315 : }
1316 0 : u64 srcOffset = (srcRank % rdmaConcurrentNum_ + rdmaConcurrentNum_) * rdmaDataBlockSize_;
1317 0 : DeviceMem src = cclOutMem_.range(srcOffset, recvInfo[curStep].recvLen);
1318 0 : DeviceMem dst = userOutput_.range(recvInfo[curStep].recvOffset, recvInfo[curStep].recvLen);
1319 0 : HCCL_DEBUG(
1320 : "[AlltoAllVDirectFullMesh][CopyRecvDataToOutput] userRank[%u], srcRank[%u], srcOffset[%llu],"
1321 : "recvInfo[curStep].recvOffset[%llu], recvLen[%llu]",
1322 : userRank_, srcRank, srcOffset, recvInfo[curStep].recvOffset, recvInfo[curStep].recvLen);
1323 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream));
1324 0 : return HCCL_SUCCESS;
1325 0 : }
1326 :
1327 : HcclResult
1328 0 : AlltoAllVDirectFullMesh::ProcessSingleGroupRdmaData(std::vector<u32>& dstRanks, std::vector<u32>& srcRanks, u32 round)
1329 : {
1330 0 : lastRdmaDstRanksIdx_ = dstRanks.size() - 1;
1331 0 : for (u32 index = 0; index < dstRanks.size(); index++) {
1332 0 : u32 dstRank = dstRanks[index];
1333 0 : u32 srcRank = srcRanks[index];
1334 0 : Stream stream = rdmaSubStreams_[index + 1];
1335 :
1336 0 : std::vector<SendDataBlock> sendInfo;
1337 0 : std::vector<RecvDataBlock> recvInfo;
1338 0 : GenRdmaSendInfo(dstRank, sendInfo);
1339 0 : GenRdmaRecvInfo(srcRank, recvInfo);
1340 0 : u32 totalStep = std::max(sendInfo.size(), recvInfo.size());
1341 0 : lastRdmaStep_ = totalStep - 1;
1342 0 : for (u32 curStep = 0; curStep < totalStep; curStep++) {
1343 0 : CHK_RET(CopyDataForSend(dstRank, sendInfo, curStep, stream));
1344 0 : CHK_RET(SendRecvRdmaData(dstRank, srcRank, sendInfo, recvInfo, round, index, curStep, stream));
1345 0 : CHK_RET(CopyRecvDataToOutput(srcRank, recvInfo, curStep, stream));
1346 : }
1347 0 : }
1348 :
1349 0 : return HCCL_SUCCESS;
1350 : }
1351 :
1352 0 : HcclResult AlltoAllVDirectFullMesh::ProcessRdmaData()
1353 : {
1354 : // RDMA通信轮次
1355 0 : u32 rdmaRoundNum = (totalRdmaRankNum_ + rdmaConcurrentNum_ - 1) / rdmaConcurrentNum_;
1356 0 : lastRdmaRoundIdx_ = rdmaRoundNum - 1;
1357 :
1358 0 : u32 leftRankNum = totalRdmaRankNum_;
1359 0 : u32 curSrcRank = INVALID_VALUE_RANKID;
1360 0 : u32 curDstRank = INVALID_VALUE_RANKID;
1361 0 : if (isSuPodAsym_) {
1362 0 : for (u32 i = 0; i < userRankSize_; i++) {
1363 0 : if (i < podStartRank_ || i > podEndRank_) {
1364 0 : curSrcRank = i;
1365 0 : curDstRank = i;
1366 0 : break;
1367 : }
1368 : }
1369 : } else {
1370 0 : curDstRank = (userRank_ + devNumInlocalPod_) % userRankSize_;
1371 0 : curSrcRank = (userRank_ + userRankSize_ - devNumInlocalPod_) % userRankSize_;
1372 : }
1373 :
1374 0 : for (u32 round = 0; round < rdmaRoundNum; round++) {
1375 0 : u32 curProcessRankNum = leftRankNum >= rdmaConcurrentNum_ ? rdmaConcurrentNum_ : leftRankNum;
1376 0 : leftRankNum -= curProcessRankNum;
1377 :
1378 0 : std::vector<u32> dstRanks;
1379 0 : std::vector<u32> srcRanks;
1380 0 : for (u32 i = 0; i < curProcessRankNum; i++) {
1381 0 : dstRanks.push_back(GetNextDstRank(curDstRank));
1382 0 : if (isSuPodAsym_) {
1383 0 : srcRanks.push_back(dstRanks.back());
1384 : } else {
1385 0 : srcRanks.push_back(GetPreSrcRank(curSrcRank));
1386 : }
1387 : }
1388 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
1389 0 : CHK_RET(RdmaControlNotifySubStart());
1390 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
1391 0 : CHK_RET(ProcessSingleGroupRdmaData(dstRanks, srcRanks, round));
1392 0 : CHK_RET(SubNotifyRdmaControlFinish());
1393 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, rdmaSubStreams_[0], dispatcher_));
1394 0 : }
1395 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][ProcessRdmaData] done");
1396 0 : return HCCL_SUCCESS;
1397 : }
1398 :
1399 0 : HcclResult AlltoAllVDirectFullMesh::RunRDMA()
1400 : {
1401 : // 先启动RDMA通信
1402 0 : CHK_RET(MainNotifyRdmaControlStart());
1403 0 : CHK_RET(ProcessRdmaData());
1404 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1405 0 : CHK_RET(LocalCopy());
1406 0 : islocalCpyDone_ = true;
1407 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunRDMA] finished.");
1408 0 : return HCCL_SUCCESS;
1409 : }
1410 :
1411 0 : HcclResult AlltoAllVDirectFullMesh::RunSDMATasks(u32 roundIdx, u32 step, u32 groupRankSize, u32 leftRankSize)
1412 : {
1413 0 : if (isBigCount_) {
1414 0 : if (roundIdx == 0) {
1415 0 : UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
1416 : }
1417 0 : if (roundIdx < commRounds_ - 1) {
1418 0 : u32 nextgroupRankSize = (leftRankSize - groupRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ :
1419 : leftRankSize - groupRankSize;
1420 0 : UpdatePartialCommunicationRankSet(roundIdx + 1, nextgroupRankSize, nextPartialCommRankSet_);
1421 : }
1422 0 : CHK_RET(RunGroupFullMeshAlltoall(roundIdx, step));
1423 :
1424 0 : if (roundIdx < commRounds_ - 1) {
1425 0 : partialCommRankSet_ = nextPartialCommRankSet_;
1426 0 : subStreamSendInfo_ = nextSubStreamSendInfo_;
1427 0 : subStreamReadInfo_ = nextSubStreamReadInfo_;
1428 0 : if (needAlltoallvCache_) {
1429 0 : subStreamZcopySendInfo_ = nextSubStreamZcopySendInfo_;
1430 0 : subStreamZcopyReadInfo_ = nextSubStreamZcopyReadInfo_;
1431 : }
1432 : }
1433 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, localSubStream_));
1434 : } else {
1435 0 : UpdatePartialCommunicationRankSet(roundIdx, groupRankSize, partialCommRankSet_);
1436 0 : CHK_RET(RunGroupFullMeshAlltoall(roundIdx, step));
1437 : }
1438 0 : return HCCL_SUCCESS;
1439 : }
1440 :
1441 0 : HcclResult AlltoAllVDirectFullMesh::RunSDMAFineGrained(u32 totalStep, HcclOpMetaInfoDef& opMeta)
1442 : {
1443 0 : if (totalStep > 1) {
1444 : // 细粒度场景不支持切分
1445 0 : HCCL_ERROR(
1446 : "[AlltoAllVDirectFullMesh][RunSDMAFineGrained] AlltoAllV is not supported when totalStep[%u] > 1, "
1447 : "HCCL buffer is insufficient. stepSize : %u ",
1448 : totalStep, algOpContext_.mc2Handler.stepSize);
1449 0 : return HCCL_E_NOT_SUPPORT;
1450 0 : } else if (totalStep == 0) {
1451 : // totalStep不需要通信,但是需要适配高阶API wait/write
1452 0 : for (u32 roundIdx = 0; roundIdx < commRounds_; roundIdx++) {
1453 0 : CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), roundIdx));
1454 0 : CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
1455 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] step is 0 finished.");
1456 : }
1457 : } else {
1458 : // totalStep == 1 细粒度修改
1459 0 : u32 leftRankSize = devNumInlocalPod_; // leftRankSize中去掉本卡
1460 0 : for (u32 roundIdx = 0; roundIdx < commRounds_ && leftRankSize > 0; roundIdx++) {
1461 0 : CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), roundIdx));
1462 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1463 0 : u32 groupRankSize = (leftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : leftRankSize;
1464 0 : UpdateRemoteRankSet(roundIdx, groupRankSize);
1465 0 : CHK_RET(RunGroupFullMeshAlltoall(roundIdx, 0));
1466 0 : leftRankSize -= groupRankSize;
1467 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
1468 0 : CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
1469 : }
1470 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] fine-grained finished.");
1471 0 : return HCCL_SUCCESS;
1472 : }
1473 :
1474 0 : if (totalStep == 0 && !islocalCpyDone_) {
1475 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1476 0 : CHK_RET(LocalCopy());
1477 0 : islocalCpyDone_ = true;
1478 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
1479 0 : return HCCL_SUCCESS;
1480 : }
1481 :
1482 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMAFineGrained] finished.");
1483 0 : return HCCL_SUCCESS;
1484 : }
1485 :
1486 0 : HcclResult AlltoAllVDirectFullMesh::RunSDMA(HcclOpMetaInfoDef& opMeta)
1487 : {
1488 0 : u32 totalStep = CalcNumSubStep();
1489 0 : lastStep_ = totalStep - 1;
1490 : // 计算每个rank分组fullmesh后需要通信的轮次,向上取整
1491 0 : commRounds_ = (devNumInlocalPod_ + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_;
1492 0 : u32 leftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
1493 : lastRoundIdx_
1494 0 : = std::min((leftRankSize + sdmaConcurrentNum_ - 1) / sdmaConcurrentNum_, static_cast<u32>(commRounds_)) - 1;
1495 0 : HCCL_DEBUG(
1496 : "[AlltoAllVDirectFullMesh][RunSDMA] userRank [%u] communication rounds[%llu] totalStep [%u] "
1497 : "stepSize [%u], post sync info: lastStep_[%u] lastRoundIdx_[%u] devNumInlocalPod_[%u] sdmaConcurrentNum_[%u]",
1498 : userRank_, commRounds_, totalStep, algOpContext_.mc2Handler.stepSize, lastStep_, lastRoundIdx_,
1499 : devNumInlocalPod_, sdmaConcurrentNum_);
1500 :
1501 0 : if (UNLIKELY(algOpContext_.mc2Handler.stepSize > 0)) {
1502 0 : CHK_RET(RunSDMAFineGrained(totalStep, opMeta));
1503 : } else {
1504 0 : if (totalStep == 0 && !islocalCpyDone_) {
1505 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1506 0 : CHK_RET(LocalCopy());
1507 0 : islocalCpyDone_ = true;
1508 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
1509 0 : return HCCL_SUCCESS;
1510 : }
1511 :
1512 0 : for (u32 step = 0; step < totalStep; step++) {
1513 0 : u32 currentLeftRankSize = devNumInlocalPod_ - 1; // leftRankSize中去掉本卡
1514 0 : for (u32 roundIdx = 0; roundIdx < commRounds_ && currentLeftRankSize > 0; roundIdx++) {
1515 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1516 0 : u32 groupRankSize
1517 0 : = (currentLeftRankSize > sdmaConcurrentNum_) ? sdmaConcurrentNum_ : currentLeftRankSize;
1518 0 : CHK_RET(RunSDMATasks(roundIdx, step, groupRankSize, currentLeftRankSize));
1519 0 : currentLeftRankSize -= groupRankSize;
1520 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, sdmaSubStream_));
1521 : }
1522 : }
1523 : }
1524 :
1525 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunSDMA] finished.");
1526 0 : return HCCL_SUCCESS;
1527 : }
1528 :
1529 0 : HcclResult AlltoAllVDirectFullMesh::RunAsync()
1530 : {
1531 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllToAllV(CopyPattern::ZCOPY, cclInMem_.size(), true);
1532 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1533 :
1534 0 : if (algOpContext_.mc2Handler.stepSize > 0) {
1535 0 : if (algOpContext_.mc2Handler.stepSize > userRankSize_
1536 0 : || userRankSize_ % algOpContext_.mc2Handler.stepSize != 0) {
1537 0 : HCCL_ERROR(
1538 : "[AlltoAllVDirectFullMesh][RunAsync] Step size should be less than or equal to the rank size, "
1539 : "and the rank size should be a multiple of the step size, but the step size is [%u] and the rank size "
1540 : "is [%u].",
1541 : algOpContext_.mc2Handler.stepSize, userRankSize_);
1542 0 : return HCCL_E_PARA;
1543 : }
1544 0 : if (userRankSize_ == 1) {
1545 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] AlltoAllV do localcopy with 1 rank");
1546 0 : CHK_RET(mc2HandlerPub.Mc2WaitValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler), 0));
1547 0 : CHK_RET(LocalCopy());
1548 0 : CHK_RET(mc2HandlerPub.Mc2WriteValue(dispatcher_, mainStream_, &(algOpContext_.mc2Handler)));
1549 0 : return HCCL_SUCCESS;
1550 : }
1551 : }
1552 :
1553 0 : if (userRankSize_ == 1) {
1554 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] do localcopy with 1 rank");
1555 0 : CHK_RET(LocalCopy());
1556 0 : return HCCL_SUCCESS;
1557 : }
1558 :
1559 0 : CHK_RET(ExecEmptyTask(userInput_, userOutput_, mainStream_, dispatcher_));
1560 0 : if (totalRdmaRankNum_ > 0) {
1561 0 : CHK_RET(RunRDMA());
1562 : }
1563 :
1564 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, rdmaSubStreams_));
1565 :
1566 0 : if (devNumInlocalPod_ > 1) {
1567 0 : CHK_RET(RunSDMA(opMeta));
1568 : }
1569 :
1570 0 : if (totalRdmaRankNum_ > 0) {
1571 : // 等待RDMA通信结束
1572 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
1573 0 : CHK_RET(RdmaControlNotifyMainFinish());
1574 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, rdmaSubStreams_));
1575 : }
1576 :
1577 0 : HCCL_INFO("[AlltoAllVDirectFullMesh][RunAsync] finished.");
1578 0 : return HCCL_SUCCESS;
1579 : }
1580 :
1581 0 : HcclResult AlltoAllVDirectFullMesh::GetNslbAdjInfo(
1582 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
1583 : {
1584 : (void)links;
1585 0 : if (rankSize == 1) {
1586 0 : return HCCL_SUCCESS;
1587 : }
1588 :
1589 0 : u32 devNumInlocalPod = nslbAdjInfo.dstRankNum;
1590 0 : u32 totalRdmaRankNum = rankSize - devNumInlocalPod;
1591 :
1592 0 : u32 rdmaConcurrentNum = (totalRdmaRankNum > ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) ?
1593 : (ALLTOALLV_DIRECT_FULLMESH_RDMA_CONCURRENT_SIZE) :
1594 0 : (totalRdmaRankNum);
1595 0 : if (rdmaConcurrentNum == 0) {
1596 0 : return HCCL_SUCCESS;
1597 : }
1598 : // RDMA通信轮次
1599 0 : u32 rdmaRoundNum = (totalRdmaRankNum + rdmaConcurrentNum - 1) / rdmaConcurrentNum;
1600 0 : if (rdmaRoundNum == 0) {
1601 0 : return HCCL_SUCCESS;
1602 : }
1603 0 : u32 currStage = rank / devNumInlocalPod;
1604 :
1605 0 : for (u32 step = 0; step < rdmaRoundNum; step++) {
1606 0 : u32 sendTo = (rank + devNumInlocalPod + step) % rankSize;
1607 0 : u32 sendToStag = sendTo / devNumInlocalPod;
1608 0 : if (currStage == sendToStag) {
1609 : // 此时认为时同一个超节点内通讯
1610 0 : sendTo = (sendTo + devNumInlocalPod) % rankSize;
1611 : }
1612 0 : NslbDpAdjInfo adjInfoStep = {};
1613 0 : adjInfoStep.dstLocalRankId = sendTo;
1614 0 : adjInfoStep.phaseId = step + 1;
1615 0 : adjInfoStep.rev = 0;
1616 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
1617 : }
1618 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
1619 0 : return HCCL_SUCCESS;
1620 : }
1621 :
1622 0 : HcclResult AlltoAllVDirectFullMesh::GetHcclOffsetDstRanksMap(
1623 : std::unordered_map<uint64_t, std::vector<uint32_t>>& hcclOffsetDstRanksMap) const
1624 : {
1625 0 : hcclOffsetDstRanksMap.clear();
1626 0 : hcclOffsetDstRanksMap = hcclOffsetDstRanksMap_; // Deep copy
1627 :
1628 0 : return HCCL_SUCCESS;
1629 : }
1630 :
1631 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_2_ALL_V_DIRECT_FULL_MESH, AlltoAllVDirectFullMesh);
1632 : } // namespace hccl
|