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 "log.h"
12 :
13 : #include "ins_temp_all_to_all_mesh.h"
14 :
15 : namespace Hccl {
16 0 : InsTempAlltoAllMesh::InsTempAlltoAllMesh(
17 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
18 0 : const std::map<RankId, u32>& tempVirtRankMap)
19 0 : : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 : {
21 0 : if (tempRankSize_ == 0) {
22 0 : THROW<InvalidParamsException>(StringFormat("[InsTempAlltoAllMesh] Invalid tempRankSize[%u].", tempRankSize_));
23 : }
24 0 : }
25 :
26 0 : InsTempAlltoAllMesh::~InsTempAlltoAllMesh() {}
27 :
28 0 : HcclResult InsTempAlltoAllMesh::CalcRes(AlgTempResReq& tempResReq)
29 : {
30 0 : CHK_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq));
31 :
32 0 : auto& linkReq = tempResReq.links;
33 0 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
34 0 : auto remoteRank = resReqIter->first;
35 0 : if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_.at(remoteRank) == 0) {
36 0 : HCCL_ERROR("No path to remoteRank[%d]", remoteRank);
37 0 : return HcclResult::HCCL_E_INTERNAL;
38 : }
39 0 : resReqIter->second = rank2PathNumMap_.at(remoteRank);
40 0 : maxPathNum = std::max(maxPathNum, rank2PathNumMap_.at(remoteRank)); // 每个rank的远端搬运最多需要pathNum条流
41 : }
42 0 : tempResReq.queNum = 1 + maxPathNum * std::min(ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE, tempRankSize_);
43 0 : tempResReq.streamNum = tempResReq.queNum;
44 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
45 :
46 0 : QId centerQ = 0;
47 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
48 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
49 0 : HCCL_DEBUG(
50 : "[InsTempAlltoAllMesh] Rank[%d], VtopoSize[%zu], requiredQue Num [%u].", myRank_, tempVTopo_[0].size(),
51 : tempResReq.queNum);
52 0 : return HcclResult::HCCL_SUCCESS;
53 : }
54 :
55 0 : void InsTempAlltoAllMesh::SetA2ASendRecvInfo(const A2ASendRecvInfo& sendRecvInfo) { localSendRecvInfo_ = sendRecvInfo; }
56 :
57 0 : HcclResult InsTempAlltoAllMesh::GetScratchBufferInfo(const u64& scratchBufferSize, DataType dataType)
58 : {
59 : // 需要变更为CCU/AICPU均只加载scratchSize
60 0 : u32 concurrentSendRecvNum = (tempRankSize_ > ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE) ?
61 : ALLTOALLV_DIRECT_FULLMESH_CONCURRENT_SIZE :
62 0 : tempRankSize_;
63 : // scratch 不需要分两份,userIn直接一步到对端 scratch buffer
64 0 : u64 maxTmpMemSize = scratchBufferSize;
65 0 : u32 typeSize = DataTypeSizeGet(dataType);
66 0 : CHK_PRT_RET(
67 : typeSize == 0,
68 : HCCL_ERROR("[InsTempAlltoAllMesh] Rank [%d], Invalid dataSizePerVolume [%u].", myRank_, typeSize),
69 : HcclResult::HCCL_E_INTERNAL);
70 0 : u64 scratchInputMemSize = maxTmpMemSize / (concurrentSendRecvNum * typeSize) * typeSize;
71 0 : CHK_RET(SetBuffBlockSize(scratchInputMemSize));
72 0 : CHK_RET(SetConcurrentSendRecvNum(concurrentSendRecvNum));
73 :
74 0 : HCCL_INFO(
75 : "[InsTempAlltoAllMesh][GetScratchBufferInfo] concurrentSendRecvNum[%u], scratchInputMemSize[%llu]",
76 : concurrentSendRecvNum, scratchInputMemSize);
77 0 : return HcclResult::HCCL_SUCCESS;
78 : }
79 :
80 0 : HcclResult InsTempAlltoAllMesh::SetBuffBlockSize(const u64 buffBlockSize)
81 : {
82 0 : CHK_PRT_RET(
83 : buffBlockSize == 0, HCCL_ERROR("[InsTempAlltoAllMesh][SetBuffBlockSize] buffBlockSize should not be zero"),
84 : HcclResult::HCCL_E_PARA);
85 0 : buffBlockSize_ = buffBlockSize;
86 0 : return HcclResult::HCCL_SUCCESS;
87 : }
88 :
89 0 : HcclResult InsTempAlltoAllMesh::SetConcurrentSendRecvNum(const u32 concurrentSendRecvNum)
90 : {
91 0 : CHK_PRT_RET(
92 : concurrentSendRecvNum == 0,
93 : HCCL_ERROR("[InsTempAlltoAllMesh][SetConcurrentSendRecvNum] concurrentSendRecvNum should not be zero"),
94 : HcclResult::HCCL_E_PARA);
95 0 : concurrentSendRecvNum_ = concurrentSendRecvNum;
96 0 : return HcclResult::HCCL_SUCCESS;
97 : }
98 :
99 0 : u32 InsTempAlltoAllMesh::CalcStepNum()
100 : {
101 0 : u32 numSubStep = 0;
102 :
103 0 : for (u32 destRank = 0; destRank < tempRankSize_; destRank++) {
104 0 : if (destRank == static_cast<u32>(myRank_)) {
105 0 : continue;
106 : }
107 0 : u32 currRankSendSubStep = ((localSendRecvInfo_.sendLength[destRank] + buffBlockSize_ - 1) / buffBlockSize_);
108 0 : u32 currRankRecvSubStep = ((localSendRecvInfo_.recvLength[destRank] + buffBlockSize_ - 1) / buffBlockSize_);
109 0 : numSubStep = std::max(numSubStep, std::max(currRankSendSubStep, currRankRecvSubStep));
110 : }
111 0 : return numSubStep;
112 : }
113 :
114 0 : HcclResult InsTempAlltoAllMesh::CalcSendSliceInfo(u32 remoteRank, UsrData& sendSliceInfo)
115 : {
116 : // 判断在打平的视角下,当前的remoteRank 是在本rank的左边还是右边
117 0 : u32 pairNum = concurrentSendRecvNum_ / 2; // 每轮次,会与concurrentSendRecvNum个对端通信,左右各pairNum个
118 0 : u32 gapRight = (tempRankSize_ + remoteRank - myRank_) % tempRankSize_;
119 0 : u32 gapLeft = (tempRankSize_ + myRank_ - remoteRank) % tempRankSize_;
120 0 : u32 sendSrcBuffIdx = 0;
121 0 : u32 sendDstBuffIdx = 0;
122 0 : if (gapLeft < gapRight) {
123 : // 离左边更近
124 0 : u32 gap = gapLeft;
125 0 : sendSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
126 0 : sendDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
127 0 : } else if (gapLeft > gapRight) {
128 : // 离右边更近
129 0 : u32 gap = gapRight;
130 0 : sendSrcBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
131 0 : sendDstBuffIdx = pairNum + ((gap - 1) % pairNum);
132 : } else {
133 0 : sendSrcBuffIdx = 0;
134 0 : sendDstBuffIdx = 0;
135 : }
136 0 : u64 sendDataOffset = 0;
137 0 : u64 remainSendLen = localSendRecvInfo_.sendLength[remoteRank];
138 0 : u64 sendSrcBuffOffset = sendSrcBuffIdx * buffBlockSize_;
139 0 : u64 sendDstBuffOffset = sendDstBuffIdx * buffBlockSize_;
140 0 : while (remainSendLen > 0) {
141 0 : u64 currDataRemainLen = localSendRecvInfo_.sendLength[remoteRank] - sendDataOffset;
142 0 : u64 sendLen = std::min(buffBlockSize_, currDataRemainLen);
143 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[remoteRank] + sendDataOffset;
144 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_] + sendDataOffset;
145 :
146 0 : DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendLen);
147 0 : DataSlice sendScratchSlice = DataSlice(buffInfo_.inBuffType, sendSrcBuffOffset, sendLen);
148 0 : DataSlice sendDstSlice = DataSlice(buffInfo_.outBuffType, sendDstBuffOffset, sendLen);
149 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, sendLen);
150 0 : sendSliceInfo.usrInSlices.emplace_back(userInSlice);
151 0 : sendSliceInfo.scratchInSlices.emplace_back(sendScratchSlice);
152 0 : sendSliceInfo.scratchOutSlices.emplace_back(sendDstSlice);
153 0 : sendSliceInfo.usrOutSlices.emplace_back(userOutSlice);
154 0 : sendDataOffset += sendLen;
155 0 : remainSendLen -= sendLen;
156 : }
157 0 : return HcclResult::HCCL_SUCCESS;
158 : }
159 :
160 0 : HcclResult InsTempAlltoAllMesh::CalcRecvSliceInfo(u32 remoteRank, UsrData& readSliceInfo)
161 : {
162 0 : u32 pairNum = concurrentSendRecvNum_ / 2;
163 0 : u32 gapRight = (remoteRank - myRank_ + tempRankSize_) % tempRankSize_;
164 0 : u32 gapLeft = (myRank_ - remoteRank + tempRankSize_) % tempRankSize_;
165 0 : u32 recvSrcBuffIdx = 0;
166 0 : u32 recvDstBuffIdx = 0;
167 0 : if (gapLeft < gapRight) {
168 0 : u32 gap = gapLeft;
169 0 : recvSrcBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
170 0 : recvDstBuffIdx = pairNum + ((gap - 1) % pairNum);
171 0 : } else if (gapLeft > gapRight) {
172 0 : u32 gap = gapRight;
173 0 : recvSrcBuffIdx = pairNum + ((gap - 1) % pairNum);
174 0 : recvDstBuffIdx = pairNum - 1 - ((gap - 1) % pairNum);
175 : } else {
176 0 : recvSrcBuffIdx = 0;
177 0 : recvDstBuffIdx = 0;
178 : }
179 0 : u64 recvDataOffset = 0;
180 0 : u64 remainRecvLen = localSendRecvInfo_.recvLength[remoteRank];
181 0 : u64 recvSrcBuffOffset = recvSrcBuffIdx * buffBlockSize_;
182 0 : u64 recvDstBuffOffset = recvDstBuffIdx * buffBlockSize_;
183 0 : while (remainRecvLen > 0) {
184 0 : u64 currDataRemainLen = localSendRecvInfo_.recvLength[remoteRank] - recvDataOffset;
185 0 : u64 recvLen = std::min(buffBlockSize_, currDataRemainLen);
186 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[remoteRank] + recvDataOffset;
187 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_] + recvDataOffset;
188 :
189 0 : DataSlice recvInSlice = DataSlice(BufferType::INPUT, userInOffset, recvLen);
190 0 : DataSlice recvScratchSlice = DataSlice(buffInfo_.inBuffType, recvSrcBuffOffset, recvLen);
191 0 : DataSlice recvDstSlice = DataSlice(buffInfo_.outBuffType, recvDstBuffOffset, recvLen);
192 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvLen);
193 0 : readSliceInfo.usrInSlices.emplace_back(recvInSlice);
194 0 : readSliceInfo.scratchInSlices.emplace_back(recvScratchSlice);
195 0 : readSliceInfo.scratchOutSlices.emplace_back(recvDstSlice);
196 0 : readSliceInfo.usrOutSlices.emplace_back(userOutSlice);
197 0 : recvDataOffset += recvLen;
198 0 : remainRecvLen -= recvLen;
199 : }
200 0 : return HcclResult::HCCL_SUCCESS;
201 : }
202 :
203 0 : HcclResult InsTempAlltoAllMesh::CalcSendRecvAllSliceInfo(
204 : std::unordered_map<u32, UsrData>& sendSliceInfoMap, std::unordered_map<u32, UsrData>& recvSliceInfoMap)
205 : {
206 0 : for (u32 remoteRank = 0; remoteRank < tempRankSize_; remoteRank++) {
207 0 : if (remoteRank == static_cast<u32>(myRank_)) {
208 0 : continue;
209 : }
210 0 : UsrData readSliceInfo;
211 0 : UsrData sendSliceInfo;
212 0 : CalcSendSliceInfo(remoteRank, sendSliceInfo);
213 0 : CalcRecvSliceInfo(remoteRank, readSliceInfo);
214 0 : sendSliceInfoMap[remoteRank] = sendSliceInfo;
215 0 : recvSliceInfoMap[remoteRank] = readSliceInfo;
216 0 : }
217 0 : return HcclResult::HCCL_SUCCESS;
218 : }
219 :
220 : HcclResult
221 0 : InsTempAlltoAllMesh::CalcCommRankSetforOneLoop(u32 roundIdx, const u32 groupRankSize, std::vector<u32>& commRanks) const
222 : {
223 0 : commRanks.clear();
224 0 : u32 pairNumPerRound = concurrentSendRecvNum_ / 2;
225 0 : u32 pairSize = (groupRankSize < concurrentSendRecvNum_) ? (groupRankSize + 1) / 2 : pairNumPerRound;
226 0 : for (u32 i = roundIdx * pairNumPerRound + 1; i < (roundIdx * pairNumPerRound + pairSize + 1); i++) {
227 0 : u32 leftRemoteRank = (myRank_ + tempRankSize_ - i) % tempRankSize_;
228 0 : u32 rightRemoteRank = (myRank_ + i) % tempRankSize_;
229 0 : if (leftRemoteRank == rightRemoteRank) {
230 0 : commRanks.push_back(leftRemoteRank);
231 0 : break;
232 : } else {
233 0 : commRanks.push_back(leftRemoteRank);
234 0 : commRanks.push_back(rightRemoteRank);
235 : }
236 : }
237 0 : return HcclResult::HCCL_SUCCESS;
238 : }
239 :
240 0 : HcclResult InsTempAlltoAllMesh::CopySendDataToScratch(
241 : u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfo,
242 : const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
243 : {
244 0 : HCCL_INFO("[InsTempAlltoAllMesh] CopySendDataToScratch");
245 0 : u32 queueId = 0;
246 0 : for (u32 i = 0; i < commRanks.size(); i++) {
247 0 : u32 remoteRank = commRanks[i];
248 0 : HCCL_INFO("remoteRank = %u", remoteRank);
249 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
250 0 : std::vector<LinkData> links = tempLinks.at(remoteRank);
251 0 : std::vector<float> dataSplitRate(linkNum);
252 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
253 0 : UsrData& currSendSliceInfo = sendSliceInfo[remoteRank];
254 0 : for (u32 j = 0; j < linkNum; j++) {
255 0 : InsQuePtr queue = queues[queueId];
256 0 : queueId++;
257 0 : HCCL_INFO("queueId=%u", queueId);
258 0 : if (step < currSendSliceInfo.usrInSlices.size()) {
259 0 : DataSlice& userInSliceAllLinks = currSendSliceInfo.usrInSlices[step];
260 0 : DataSlice& sendScratchSliceAllLinks = currSendSliceInfo.scratchInSlices[step];
261 0 : DataSlice userInSlice = CalcDataSliceForLinks(userInSliceAllLinks, dataSplitRate, j, dataType_);
262 : DataSlice sendScratchSlice
263 0 : = CalcDataSliceForLinks(sendScratchSliceAllLinks, dataSplitRate, j, dataType_);
264 :
265 0 : InsQuePtr queue = queues[i];
266 0 : CHK_RET(LocalCopy(queue, userInSlice, sendScratchSlice));
267 0 : }
268 0 : }
269 0 : }
270 0 : return HcclResult::HCCL_SUCCESS;
271 : }
272 :
273 0 : HcclResult InsTempAlltoAllMesh::SendRecvData(
274 : u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfo,
275 : std::unordered_map<u32, UsrData>& readSliceInfo, const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
276 : {
277 0 : CHK_PRT_RET(
278 : queues.empty(), HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] empty queues"), HcclResult::HCCL_E_INTERNAL);
279 0 : CHK_PTR_NULL(queues[0]);
280 :
281 0 : if (commRanks.size() * maxPathNum > queues.size()) {
282 0 : HCCL_ERROR("[InsTempAlltoAllMesh][SendRecvData] commRanks.size() * maxPathNum > queues.size() is wrong");
283 0 : return HcclResult::HCCL_E_INTERNAL;
284 : }
285 0 : uint64_t queuesId = 0;
286 0 : for (u32 i = 0; i < commRanks.size(); i++) {
287 0 : s32 remoteRank = static_cast<s32>(commRanks[i]);
288 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
289 0 : UsrData& currSendSliceInfo = sendSliceInfo[remoteRank];
290 0 : UsrData& currReadSliceInfo = readSliceInfo[remoteRank];
291 0 : std::vector<LinkData> links = tempLinks.at(remoteRank);
292 0 : std::vector<float> dataSplitRate(linkNum);
293 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
294 : std::vector<DataSlice>& currSendSrcSlices
295 0 : = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.scratchInSlices : currSendSliceInfo.usrInSlices;
296 : std::vector<DataSlice>& currSendDstSlices
297 0 : = (dmaMode_ == DmaMode::GET) ? currSendSliceInfo.usrOutSlices : currSendSliceInfo.scratchOutSlices;
298 : std::vector<DataSlice>& currReadSrcSlices
299 0 : = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.scratchInSlices : currReadSliceInfo.usrInSlices;
300 : std::vector<DataSlice>& currReadDstSlices
301 0 : = (dmaMode_ == DmaMode::GET) ? currReadSliceInfo.usrOutSlices : currReadSliceInfo.scratchOutSlices;
302 0 : for (u32 j = 0; j < linkNum; j++) {
303 0 : InsQuePtr queue = queues[queuesId];
304 0 : queuesId++;
305 :
306 0 : LinkData link = tempLinks.at(remoteRank)[j];
307 0 : if (step < currSendSrcSlices.size() && step < currReadSrcSlices.size()) {
308 0 : DataSlice& sendSrcSliceAllLinks = currSendSrcSlices[step];
309 0 : DataSlice& sendDstSliceAllLinks = currSendDstSlices[step];
310 0 : DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
311 0 : DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
312 0 : std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
313 0 : std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
314 0 : SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
315 0 : DataSlice& recvSrcSliceAllLinks = currReadSrcSlices[step];
316 0 : DataSlice& recvDstSliceAllLinks = currReadDstSlices[step];
317 0 : DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
318 0 : DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
319 0 : std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
320 0 : std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
321 0 : SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
322 0 : TxRxSlicesList sendRecvSlice(sendDataSlice, recvDataSlice);
323 :
324 0 : TxRxLinks sendRecvLinks(link, link);
325 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlice);
326 0 : CHK_RET(SendRecv(sendRecvInfo, queue, 0, true, dmaMode_));
327 0 : HCCL_DEBUG(
328 : "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send and recv.",
329 : step, i, remoteRank);
330 0 : } else if (step < currSendSrcSlices.size()) {
331 0 : DataSlice& sendSrcSliceAllLinks = currSendSrcSlices[step];
332 0 : DataSlice& sendDstSliceAllLinks = currSendDstSlices[step];
333 0 : DataSlice sendSrcSlice = CalcDataSliceForLinks(sendSrcSliceAllLinks, dataSplitRate, j, dataType_);
334 0 : DataSlice sendDstSlice = CalcDataSliceForLinks(sendDstSliceAllLinks, dataSplitRate, j, dataType_);
335 0 : std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
336 0 : std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
337 0 : SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
338 0 : DataInfo sendDataInfo(link, sendDataSlice);
339 0 : CHK_RET(Send(sendDataInfo, queue));
340 0 : HCCL_DEBUG(
341 : "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run send.", step, i,
342 : remoteRank);
343 0 : } else if (step < currReadSrcSlices.size()) {
344 0 : DataSlice& recvSrcSliceAllLinks = currReadSrcSlices[step];
345 0 : DataSlice& recvDstSliceAllLinks = currReadDstSlices[step];
346 0 : DataSlice recvSrcSlice = CalcDataSliceForLinks(recvSrcSliceAllLinks, dataSplitRate, j, dataType_);
347 0 : DataSlice recvDstSlice = CalcDataSliceForLinks(recvDstSliceAllLinks, dataSplitRate, j, dataType_);
348 0 : std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
349 0 : std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
350 0 : SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
351 0 : DataInfo recvDataInfo(link, recvDataSlice);
352 0 : CHK_RET(Recv(recvDataInfo, queue));
353 0 : HCCL_DEBUG(
354 : "[InsTempAlltoAllMesh][SendRecvData] step[%u], commRank[%u], remoteRank[%d] run recv.", step, i,
355 : remoteRank);
356 0 : }
357 0 : }
358 0 : }
359 0 : return HcclResult::HCCL_SUCCESS;
360 : }
361 :
362 0 : HcclResult InsTempAlltoAllMesh::CopyRecvDataFromScratch(
363 : u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& readSliceInfo,
364 : const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
365 : {
366 0 : HCCL_INFO("[InsTempAlltoAllMesh] CopyRecvDataFromScratch");
367 0 : u32 queueId = 0;
368 0 : for (u32 i = 0; i < commRanks.size(); i++) {
369 0 : u32 remoteRank = commRanks[i];
370 0 : HCCL_INFO("remoteRank = %u", remoteRank);
371 0 : u32 linkNum = rank2PathNumMap_.at(remoteRank);
372 0 : std::vector<LinkData> links = tempLinks.at(remoteRank);
373 0 : std::vector<float> dataSplitRate(linkNum);
374 0 : CHK_RET(CalcDataSplitRateForLinks(links, dataSplitRate));
375 0 : UsrData& currReadSliceInfo = readSliceInfo[remoteRank];
376 0 : for (u32 j = 0; j < linkNum; j++) {
377 0 : InsQuePtr queue = queues[queueId];
378 0 : queueId++;
379 0 : HCCL_INFO("queueId=%u", queueId);
380 0 : if (step < currReadSliceInfo.scratchOutSlices.size()) {
381 0 : DataSlice& recvScratchSliceAllLinks = currReadSliceInfo.scratchOutSlices[step];
382 0 : DataSlice& userOutSliceAllLinks = currReadSliceInfo.usrOutSlices[step];
383 : DataSlice recvScratchSlice
384 0 : = CalcDataSliceForLinks(recvScratchSliceAllLinks, dataSplitRate, j, dataType_);
385 0 : DataSlice userOutSlice = CalcDataSliceForLinks(userOutSliceAllLinks, dataSplitRate, j, dataType_);
386 0 : CHK_RET(LocalCopy(queue, recvScratchSlice, userOutSlice));
387 : }
388 0 : }
389 0 : }
390 0 : return HcclResult::HCCL_SUCCESS;
391 : }
392 :
393 0 : HcclResult InsTempAlltoAllMesh::RunSendRecvBufferLoop(
394 : u32 step, const std::vector<u32>& commRanks, std::unordered_map<u32, UsrData>& sendSliceInfoMap,
395 : std::unordered_map<u32, UsrData>& recvSliceInfoMap, const ResLinks& tempLinks, std::vector<InsQuePtr>& queues) const
396 : {
397 0 : if (commRanks.size() > 1) {
398 0 : CHK_RET(PreSyncInterQueues(queues));
399 : }
400 :
401 0 : if (dmaMode_ == DmaMode::PUT) {
402 0 : CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
403 0 : CHK_RET(CopyRecvDataFromScratch(step, commRanks, recvSliceInfoMap, tempLinks, queues));
404 0 : } else if (dmaMode_ == DmaMode::GET) {
405 0 : CHK_RET(CopySendDataToScratch(step, commRanks, sendSliceInfoMap, tempLinks, queues));
406 0 : CHK_RET(SendRecvData(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
407 : }
408 :
409 0 : if (commRanks.size() > 1) {
410 0 : CHK_RET(PostSyncInterQueues(queues));
411 : }
412 0 : return HcclResult::HCCL_SUCCESS;
413 : }
414 :
415 0 : HcclResult InsTempAlltoAllMesh::RunSendRecvForAllRanks(
416 : u32 step, std::unordered_map<u32, UsrData>& sendSliceInfoMap, std::unordered_map<u32, UsrData>& recvSliceInfoMap,
417 : const ResLinks& tempLinks, std::vector<InsQuePtr>& queues)
418 : {
419 0 : u64 commLoops = (tempRankSize_ + concurrentSendRecvNum_ - 1) / concurrentSendRecvNum_;
420 0 : u32 leftRankSize = tempRankSize_ - 1; // leftRankSize中去掉本卡
421 0 : std::vector<u32> commRanks;
422 0 : for (u32 roundIdx = 0; roundIdx < commLoops && leftRankSize > 0; roundIdx++) {
423 0 : u32 groupRankSize = (leftRankSize > concurrentSendRecvNum_) ? concurrentSendRecvNum_ : leftRankSize;
424 0 : CHK_RET(CalcCommRankSetforOneLoop(roundIdx, groupRankSize, commRanks));
425 0 : CHK_RET(RunSendRecvBufferLoop(step, commRanks, sendSliceInfoMap, recvSliceInfoMap, tempLinks, queues));
426 0 : leftRankSize -= groupRankSize;
427 0 : HCCL_DEBUG(
428 : "[InsTempAlltoAllMesh][RunSendRecvForAllRanks] commRanksSize[%zu], roundIdx[%u] finish.", commRanks.size(),
429 : roundIdx);
430 : }
431 0 : return HcclResult::HCCL_SUCCESS;
432 0 : }
433 :
434 0 : HcclResult InsTempAlltoAllMesh::LocalDataCopy(InsQuePtr tempInsQue)
435 : {
436 0 : u64 userInOffset = localSendRecvInfo_.sendOffset[myRank_];
437 0 : u64 sendSize = localSendRecvInfo_.sendLength[myRank_];
438 0 : u64 userOutOffset = localSendRecvInfo_.recvOffset[myRank_];
439 0 : u64 recvSize = localSendRecvInfo_.recvLength[myRank_];
440 0 : DataSlice userInSlice = DataSlice(BufferType::INPUT, userInOffset, sendSize);
441 0 : DataSlice userOutSlice = DataSlice(BufferType::OUTPUT, userOutOffset, recvSize);
442 0 : CHK_RET(LocalCopy(tempInsQue, userInSlice, userOutSlice));
443 0 : return HcclResult::HCCL_SUCCESS;
444 : }
445 :
446 0 : HcclResult InsTempAlltoAllMesh::Run(
447 : const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
448 : std::vector<InsQuePtr>& tempInsQues)
449 : {
450 : (void)tempFuncs;
451 : (void)sliceInfoVec;
452 0 : if (tempInsQues.size() == 0) {
453 0 : HCCL_ERROR("[CcuTempAlltoAllMesh1D] tempInsQues.size() is zero.");
454 0 : return HcclResult::HCCL_E_PARA;
455 : }
456 0 : dmaMode_ = DmaMode::PUT;
457 0 : if (IsPcieLink(tempLinks)) {
458 0 : dmaMode_ = DmaMode::GET;
459 : }
460 0 : buffInfo_ = buffInfo;
461 0 : u32 totalStep = CalcStepNum();
462 0 : HCCL_INFO(
463 : "[InsTempAlltoAllMesh] AllToAll full mesh total step is [%u], tempInsQuesSize[%zu].", totalStep,
464 : tempInsQues.size());
465 0 : std::unordered_map<u32, UsrData> sendSliceInfoMap;
466 0 : std::unordered_map<u32, UsrData> recvSliceInfoMap;
467 0 : CHK_RET(CalcSendRecvAllSliceInfo(sendSliceInfoMap, recvSliceInfoMap));
468 0 : std::vector<InsQuePtr> localCopyQues;
469 0 : if (tempInsQues.size() > 1) {
470 0 : localCopyQues.push_back(tempInsQues[0]);
471 0 : localCopyQues.push_back(tempInsQues.back());
472 : }
473 :
474 0 : if (localCopyQues.size() > 1) {
475 0 : CHK_RET(PreSyncInterQueues(localCopyQues));
476 : }
477 : // 最后一条流做localCopy
478 0 : CHK_RET(LocalDataCopy(tempInsQues.back()));
479 : // 其余的流去做sendRecv
480 0 : std::vector<InsQuePtr> sendRecvQues(tempInsQues.begin(), tempInsQues.begin() + (tempInsQues.size() - 1));
481 0 : for (u32 step = 0; step < totalStep; step++) {
482 0 : CHK_RET(RunSendRecvForAllRanks(step, sendSliceInfoMap, recvSliceInfoMap, tempLinks, sendRecvQues));
483 0 : HCCL_DEBUG("[InsTempAlltoAllMesh] AllToAll full mesh step[%u] execute success", step);
484 : }
485 0 : if (localCopyQues.size() > 1) {
486 0 : CHK_RET(PostSyncInterQueues(localCopyQues));
487 : }
488 0 : HCCL_INFO("[InsTempAlltoAllMesh] AllToAll full mesh rank[%d] finish.", myRank_);
489 0 : return HcclResult::HCCL_SUCCESS;
490 0 : }
491 :
492 : } // namespace Hccl
|