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