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 "temp_reduce_scatter_concurr_mesh.h"
14 :
15 : namespace Hccl {
16 0 : TempReduceScatterConcurrMesh::TempReduceScatterConcurrMesh(const RankId virtualRank, const u32 tempRankSize,
17 : const std::vector<std::vector<RankId>> &tempVTopo,
18 0 : const std::map<RankId, u32> &tempVirtRankMap)
19 0 : : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 : {
21 0 : }
22 :
23 0 : TempReduceScatterConcurrMesh::~TempReduceScatterConcurrMesh()
24 : {
25 0 : }
26 :
27 0 : HcclResult TempReduceScatterConcurrMesh::CalcRes(const bool forAllReduce, AlgTempResReq &tempResReq,
28 : u32 &requiredScratchMultiplier)
29 : {
30 : (void)forAllReduce;
31 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
32 0 : tempResReq.queNum += tempVTopo_[dim].size() - 1;
33 : }
34 0 : requiredScratchMultiplier = tempRankSize_;
35 :
36 : u32 myAlgRank;
37 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
38 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
39 0 : for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
40 : // find neighbors -> virtualRank
41 0 : u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
42 0 : RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
43 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Dim [%u], NeighborRank [%d].",
44 : myRank_, dim, neighborRank);
45 :
46 : // LinkNum
47 0 : tempResReq.links[neighborRank] = 1;
48 : }
49 : }
50 :
51 0 : return HcclResult::HCCL_SUCCESS;
52 : }
53 :
54 : /*
55 : dataSize / (rankSize) --> chunkSize
56 : dataSize / (rankSize * dimNum) --> sliceSize
57 :
58 : SliceInfoVecforConcurrMesh: [1st chunk: [1st Slice, 2nd Slice], 2nd chunk: [1st Slice, 2nd Slice], ...]
59 : */
60 0 : HcclResult TempReduceScatterConcurrMesh::CalcSliceInfo(const AllignInfo &allignInfo, const bool forAllReduce,
61 : const u64 dataSize, RankSliceInfo &sliceInfoVec)
62 : {
63 0 : u32 dimSize = 0;
64 0 : for (u32 dimIdx = 0; dimIdx < tempVTopo_.size(); dimIdx++) {
65 0 : if (tempVTopo_[dimIdx].size() != 1) {
66 0 : dimSize += 1;
67 : }
68 : }
69 0 : std::vector<SliceInfo> tmp(dimSize);
70 0 : sliceInfoVec.resize(tempRankSize_, tmp);
71 :
72 0 : if (forAllReduce) {
73 : // for allreduce, dataSize = total dataSize
74 0 : CHK_RET(CalcSliceInfoAllReduce(allignInfo, dataSize, sliceInfoVec));
75 : } else {
76 : // for reduce scatter, dataSize = chunkSize
77 0 : if (sliceInfoVec[0].size() == 1) {
78 : // one-dimensional mesh
79 0 : CHK_RET(CalcRsAgSliceInfoMesh(myRank_, tempRankSize_, allignInfo, dataSize, sliceInfoVec));
80 : } else {
81 : // multi-dimensional mesh
82 0 : CHK_RET(CalcRsAgSliceInfoConcurrMesh(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
83 : }
84 : }
85 :
86 0 : return HcclResult::HCCL_SUCCESS;
87 0 : }
88 :
89 0 : HcclResult TempReduceScatterConcurrMesh::CalcSliceInfoAllReduce(const AllignInfo &allignInfo, const u64 dataSize,
90 : RankSliceInfo &sliceInfoVec)
91 : {
92 : u64 unitAllignSize;
93 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
94 :
95 0 : u64 rankDataSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
96 :
97 0 : if (sliceInfoVec[0].size() == 1) {
98 : // one dimensional mesh
99 0 : u64 resDataSize = dataSize;
100 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
101 0 : u64 currChunkSize = (resDataSize > rankDataSize) ? rankDataSize : resDataSize;
102 0 : SliceInfo slice = {dataSize - resDataSize, currChunkSize};
103 0 : sliceInfoVec[rankIdx][0] = slice;
104 0 : resDataSize -= currChunkSize;
105 : }
106 :
107 0 : CHK_PRT_RET(
108 : (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != dataSize),
109 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], SliceInfo calculation error for "
110 : "AllReduce ConcurrMesh!",
111 : myRank_),
112 : HcclResult::HCCL_E_INTERNAL);
113 : } else {
114 0 : u32 dimSize0 = tempVTopo_[0].size();
115 0 : u32 dimSize1 = tempVTopo_[1].size();
116 :
117 0 : u64 resDataSize = dataSize;
118 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
119 0 : u64 currChunkSize = (resDataSize > rankDataSize) ? rankDataSize : resDataSize;
120 0 : u64 sliceSize0 = min(currChunkSize, RoundUp(currChunkSize, (dimSize0 + dimSize1) * unitAllignSize)
121 0 : * dimSize0 * unitAllignSize);
122 0 : SliceInfo slice0 = {dataSize - resDataSize, sliceSize0};
123 0 : sliceInfoVec[rankIdx][0] = slice0;
124 0 : resDataSize -= sliceSize0;
125 :
126 0 : u64 sliceSize1 = currChunkSize - sliceSize0;
127 0 : SliceInfo slice1 = {dataSize - resDataSize, sliceSize1};
128 0 : sliceInfoVec[rankIdx][1] = slice1;
129 0 : resDataSize -= sliceSize1;
130 : }
131 :
132 0 : CHK_PRT_RET((sliceInfoVec[tempRankSize_ - 1][1].offset + sliceInfoVec[tempRankSize_ - 1][1].size != dataSize),
133 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], SliceInfo calculation error "
134 : "for AllReduce ConcurrMesh!",
135 : myRank_),
136 : HcclResult::HCCL_E_INTERNAL);
137 : }
138 :
139 0 : return HcclResult::HCCL_SUCCESS;
140 : }
141 :
142 0 : HcclResult TempReduceScatterConcurrMesh::GenPrimQue(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
143 : const BuffInfo &buffInfo, const ResLinks &tempLinks,
144 : std::vector<PrimQuePtr> &tempPrimQues)
145 : {
146 0 : opMode_ = tempFuncs.opMode;
147 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
148 0 : buffInfo_ = buffInfo;
149 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], EnableCounterNotify [%d].", myRank_,
150 : enableCounterNotify_);
151 :
152 0 : queNum_ = 0;
153 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
154 0 : queNum_ += tempVTopo_[dim].size() - 1;
155 : }
156 0 : CHK_PRT_RET(queNum_ != tempPrimQues.size(),
157 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], requiredQue Error.", myRank_),
158 : HcclResult::HCCL_E_INTERNAL);
159 :
160 : // LocalCopy: from input to scratch In Buffer for OPBASE
161 0 : if ((opMode_ == OpMode::OPBASE) && tempFuncs.isForepart) {
162 0 : CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
163 : }
164 :
165 0 : if (sliceInfoVec[0].size() == 1) {
166 0 : CHK_RET(RunOneDimMesh(sliceInfoVec, tempLinks, tempPrimQues));
167 : } else {
168 0 : CHK_RET(RunConcurrMesh(sliceInfoVec, tempLinks, tempPrimQues));
169 : }
170 :
171 : // LocalCopy for standalone reducescatter in Offload Mode
172 0 : if ((opMode_ == OpMode::OFFLOAD) && !tempFuncs.forAllReduce && !tempFuncs.forAlgSeqComb) {
173 0 : CHK_RET(PostCopyOffload(sliceInfoVec, tempPrimQues));
174 : }
175 :
176 : // LocalCopy from scratch to output for Opbase
177 0 : if ((opMode_ == OpMode::OPBASE) && tempFuncs.isBottom && !tempFuncs.forAllReduce) {
178 0 : CHK_RET(PostCopyOpbase(tempFuncs.usrData, tempPrimQues));
179 : }
180 :
181 0 : return HcclResult::HCCL_SUCCESS;
182 : }
183 :
184 0 : HcclResult TempReduceScatterConcurrMesh::RunOneDimMesh(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
185 : std::vector<PrimQuePtr> &tempPrimQues)
186 : {
187 : // semaphore sync
188 0 : if (queNum_ > 1) {
189 0 : CHK_RET(PreSyncInterQueues(tempPrimQues));
190 : }
191 :
192 : // locate myRank in tempVTopo -> algRank
193 : u32 myAlgRank;
194 0 : u32 validDim = (tempVTopo_[0].size() == 1) ? 1 : 0;
195 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], valid Dim [%u].", myRank_, validDim);
196 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[validDim], myAlgRank));
197 :
198 : // runMesh
199 0 : CHK_PRT_RET(
200 : RunMesh(myAlgRank, tempVTopo_[validDim], sliceInfoVec, tempLinks, tempPrimQues) != HcclResult::HCCL_SUCCESS,
201 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], unable to run the mesh algorithm.",
202 : myRank_),
203 : HcclResult::HCCL_E_INTERNAL);
204 :
205 : // semaphore sync
206 0 : if (queNum_ > 1) {
207 0 : CHK_RET(PostSyncInterQueues(tempPrimQues));
208 : }
209 :
210 0 : return HcclResult::HCCL_SUCCESS;
211 : }
212 :
213 0 : HcclResult TempReduceScatterConcurrMesh::RunMesh(const u32 myAlgRank, const std::vector<RankId> &vTopo,
214 : const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
215 : std::vector<PrimQuePtr> &tempPrimQues)
216 : {
217 0 : for (u32 queIdx = 0; queIdx < tempPrimQues.size(); queIdx++) {
218 : // find neighbors -> virtualRank
219 0 : RankId neighborRank = vTopo[(myAlgRank + 1 + queIdx) % tempRankSize_];
220 : // Link
221 0 : LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
222 :
223 0 : u32 sendChunkIdx = tempVirtRankMap_[neighborRank];
224 0 : u64 sendOffset = sliceInfoVec[sendChunkIdx][0].offset;
225 0 : u64 sendSize = sliceInfoVec[sendChunkIdx][0].size;
226 0 : u32 recvChunkIdx = tempVirtRankMap_[myRank_];
227 0 : u64 recvOffset = sliceInfoVec[recvChunkIdx][0].offset;
228 0 : u64 recvSize = sliceInfoVec[recvChunkIdx][0].size;
229 :
230 : // PrimGroup
231 0 : std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
232 :
233 : // SendReduce
234 0 : DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
235 : DataSlice sendRemSrcSlice
236 0 : = DataSlice(buffInfo_.scratBuffType, sendOffset + buffInfo_.scratchBuffBaseOff, sendSize);
237 0 : DataSlice sendRemDstSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
238 : std::unique_ptr<Primitive> primSendReduce
239 0 : = std::make_unique<PrimSendReduce>(neighborRank, neighborLinkData, sendLocSlice, sendRemSrcSlice,
240 0 : sendRemDstSlice, dataType_, redOp_, dmaMode_);
241 :
242 0 : primGroup->Append(std::move(primSendReduce));
243 :
244 : // RecvReduce
245 0 : DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
246 : DataSlice recvLocSrcSlice
247 0 : = DataSlice(buffInfo_.scratBuffType, recvOffset + buffInfo_.scratchBuffBaseOff, recvSize);
248 0 : DataSlice recvLocDstSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
249 : std::unique_ptr<Primitive> primRecvReduce
250 0 : = std::make_unique<PrimRecvReduce>(neighborRank, neighborLinkData, recvRemSlice, recvLocSrcSlice,
251 0 : recvLocDstSlice, dataType_, redOp_, dmaMode_);
252 :
253 0 : primGroup->Append(std::move(primRecvReduce));
254 :
255 0 : tempPrimQues[queIdx]->Append(std::move(primGroup));
256 0 : }
257 0 : return HcclResult::HCCL_SUCCESS;
258 : }
259 :
260 0 : HcclResult TempReduceScatterConcurrMesh::RunConcurrMesh(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
261 : std::vector<PrimQuePtr> &tempPrimQues)
262 : {
263 0 : std::vector<std::vector<PrimQuePtr>> dimQues;
264 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
265 : // assign queues
266 0 : std::vector<PrimQuePtr> tmpQue;
267 0 : for (u32 idx = 0; idx < tempVTopo_[dim].size() - 1; idx++) {
268 0 : if (dim == 0) {
269 0 : tmpQue.push_back(tempPrimQues[idx]);
270 : } else {
271 0 : tmpQue.push_back(tempPrimQues[tempVTopo_[0].size() - 1 + idx]);
272 : }
273 : }
274 0 : dimQues.push_back(tmpQue);
275 0 : }
276 :
277 0 : std::vector<PrimQuePtr> majorDimQue = {tempPrimQues[0], tempPrimQues[tempVTopo_[0].size() - 1]};
278 :
279 : // semaphore sync inter dimensions
280 0 : CHK_RET(PreSyncInterQueues(majorDimQue));
281 :
282 : // run concurrent Mesh Step 0
283 0 : u32 step = 0;
284 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
285 0 : CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
286 : }
287 :
288 : // semaphore sync
289 0 : CHK_RET(PostSyncInterQueues(majorDimQue));
290 :
291 : // semaphore sync inter dimensions
292 0 : CHK_RET(PreSyncInterQueues(majorDimQue));
293 :
294 : // run concurrent Mesh Step 1
295 0 : step = 1;
296 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
297 0 : CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
298 : }
299 :
300 : // semaphore sync
301 0 : CHK_RET(PostSyncInterQueues(majorDimQue));
302 :
303 0 : return HcclResult::HCCL_SUCCESS;
304 0 : }
305 :
306 0 : HcclResult TempReduceScatterConcurrMesh::RunSingleDimension(const u32 &step, const u32 &dim,
307 : const RankSliceInfo &sliceInfoVec,
308 : const ResLinks &tempLinks,
309 : std::vector<PrimQuePtr> &dimPrimQues)
310 : {
311 0 : CHK_PRT_RET(
312 : dim > 1,
313 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], invalid dim [%u].", myRank_, dim),
314 : HcclResult::HCCL_E_INTERNAL);
315 :
316 : // locate myRank in tempVTopo -> algRank
317 : u32 myAlgRank;
318 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
319 :
320 0 : for (u32 queIdx = 0; queIdx < dimPrimQues.size(); queIdx++) {
321 : // semaphore sync
322 0 : if (dimPrimQues.size() > 1) {
323 0 : CHK_PRT_RET(PreSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
324 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Que [%u], Semaphore "
325 : "Synchronization Failed.",
326 : myRank_, dimPrimQues[queIdx]->GetId()),
327 : HcclResult::HCCL_E_INTERNAL);
328 : }
329 :
330 : // find neighbors -> virtualRank
331 0 : u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
332 0 : RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
333 :
334 : // link
335 0 : LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
336 0 : HCCL_INFO(
337 : "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], [%u]-th Que, queId [%u], neighborRank [%d].",
338 : myRank_, queIdx, dimPrimQues[queIdx]->GetId(), neighborRank);
339 :
340 : // PrimGroup
341 0 : std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
342 :
343 0 : std::vector<u32> sendChunkIdxs;
344 0 : std::vector<u32> recvChunkIdxs;
345 :
346 0 : if (step == 0) {
347 0 : for (u32 chunkIdx = 0; chunkIdx < tempVTopo_[1 - dim].size(); chunkIdx++) {
348 0 : u32 sendChunkIdx = (dim == 0) ? (neighborAlgRank + chunkIdx * tempVTopo_[0].size())
349 0 : : (neighborAlgRank * tempVTopo_[0].size() + chunkIdx);
350 0 : sendChunkIdxs.push_back(sendChunkIdx);
351 0 : u32 recvChunkIdx = (dim == 0) ? (myAlgRank + chunkIdx * tempVTopo_[0].size())
352 0 : : (myAlgRank * tempVTopo_[0].size() + chunkIdx);
353 0 : recvChunkIdxs.push_back(recvChunkIdx);
354 : }
355 : } else {
356 0 : sendChunkIdxs.push_back(tempVirtRankMap_[neighborRank]);
357 0 : recvChunkIdxs.push_back(tempVirtRankMap_[myRank_]);
358 : }
359 :
360 : // SendReduce
361 0 : u32 sliceIdx = (step == 0) ? dim : (1 - dim);
362 : std::unique_ptr<PrimSendReduce> primSendReduce
363 0 : = RunSendReduce(sliceInfoVec, sendChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
364 0 : primGroup->Append(std::move(primSendReduce));
365 :
366 : // RecvReduce
367 : std::unique_ptr<PrimRecvReduce> primRecvReduce
368 0 : = RunRecvReduce(sliceInfoVec, recvChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
369 0 : primGroup->Append(std::move(primRecvReduce));
370 :
371 0 : dimPrimQues[queIdx]->Append(std::move(primGroup));
372 :
373 : // semaphore sync
374 0 : if (dimPrimQues.size() > 1) {
375 0 : CHK_PRT_RET(PostSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
376 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Que [%u], Semaphore "
377 : "Synchronization Failed.",
378 : myRank_, dimPrimQues[queIdx]->GetId()),
379 : HcclResult::HCCL_E_INTERNAL);
380 : }
381 0 : }
382 :
383 0 : return HcclResult::HCCL_SUCCESS;
384 : }
385 :
386 0 : std::unique_ptr<PrimSendReduce> TempReduceScatterConcurrMesh::RunSendReduce(const RankSliceInfo &sliceInfoVec,
387 : const std::vector<u32> &sendChunkIdxs,
388 : const u32 &sliceIdx,
389 : const RankId &neighborRank,
390 : const LinkData &priorLinkData)
391 : {
392 0 : std::unique_ptr<PrimSendReduce> primSendReduce;
393 : u64 tmpSendOff;
394 : u64 tmpSendSize;
395 0 : for (u32 chunkIdx = 0; chunkIdx < sendChunkIdxs.size(); chunkIdx++) {
396 0 : if (chunkIdx == 0) {
397 : // first slice
398 0 : tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
399 0 : tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
400 0 : } else if (tmpSendOff + tmpSendSize == sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset) {
401 : // consequent slice
402 0 : tmpSendSize += sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
403 : } else {
404 0 : DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
405 : DataSlice sendRemSrcSlice
406 0 : = DataSlice(buffInfo_.scratBuffType, tmpSendOff + buffInfo_.scratchBuffBaseOff, tmpSendSize);
407 : DataSlice sendRemDstSlice
408 0 : = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
409 0 : if (!primSendReduce) {
410 0 : primSendReduce.reset(new PrimSendReduce(neighborRank, priorLinkData, sendLocSlice, sendRemSrcSlice,
411 0 : sendRemDstSlice, dataType_, redOp_, dmaMode_));
412 : } else {
413 0 : primSendReduce->Append(sendLocSlice, sendRemSrcSlice, sendRemDstSlice);
414 : }
415 0 : tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
416 0 : tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
417 : }
418 :
419 0 : if (chunkIdx == (sendChunkIdxs.size() - 1)) {
420 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk.", myRank_);
421 0 : DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
422 : DataSlice sendRemSrcSlice
423 0 : = DataSlice(buffInfo_.scratBuffType, tmpSendOff + buffInfo_.scratchBuffBaseOff, tmpSendSize);
424 : DataSlice sendRemDstSlice
425 0 : = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
426 0 : if (!primSendReduce) {
427 0 : HCCL_INFO(
428 : "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is a non-consecutive chunk.",
429 : myRank_);
430 0 : primSendReduce.reset(new PrimSendReduce(neighborRank, priorLinkData, sendLocSlice, sendRemSrcSlice,
431 0 : sendRemDstSlice, dataType_, redOp_, dmaMode_));
432 : } else {
433 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is consecutive.",
434 : myRank_);
435 0 : primSendReduce->Append(sendLocSlice, sendRemSrcSlice, sendRemDstSlice);
436 : }
437 : }
438 : }
439 :
440 0 : return primSendReduce;
441 0 : }
442 :
443 0 : std::unique_ptr<PrimRecvReduce> TempReduceScatterConcurrMesh::RunRecvReduce(const RankSliceInfo &sliceInfoVec,
444 : const std::vector<u32> &recvChunkIdxs,
445 : const u32 &sliceIdx,
446 : const RankId &neighborRank,
447 : const LinkData &priorLinkData)
448 : {
449 0 : std::unique_ptr<PrimRecvReduce> primRecvReduce;
450 : u64 tmpRecvOff;
451 : u64 tmpRecvSize;
452 0 : for (u32 chunkIdx = 0; chunkIdx < recvChunkIdxs.size(); chunkIdx++) {
453 0 : if (chunkIdx == 0) {
454 : // first slice
455 0 : tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
456 0 : tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
457 0 : } else if (tmpRecvOff + tmpRecvSize == sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset) {
458 : // consequent slice
459 0 : tmpRecvSize += sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
460 : } else {
461 0 : DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
462 : DataSlice recvLocSrcSlice
463 0 : = DataSlice(buffInfo_.scratBuffType, tmpRecvOff + buffInfo_.scratchBuffBaseOff, tmpRecvSize);
464 : DataSlice recvLocDstSlice
465 0 : = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
466 :
467 0 : if (!primRecvReduce) {
468 0 : HCCL_INFO(
469 : "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is a non-consecutive chunk.",
470 : myRank_);
471 0 : primRecvReduce.reset(new PrimRecvReduce(neighborRank, priorLinkData, recvRemSlice, recvLocSrcSlice,
472 0 : recvLocDstSlice, dataType_, redOp_, dmaMode_));
473 : } else {
474 0 : HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is consecutive.",
475 : myRank_);
476 0 : primRecvReduce->Append(recvRemSlice, recvLocSrcSlice, recvLocDstSlice);
477 : }
478 0 : tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
479 0 : tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
480 : }
481 :
482 0 : if (chunkIdx == (recvChunkIdxs.size() - 1)) {
483 0 : DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
484 : DataSlice recvLocSrcSlice
485 0 : = DataSlice(buffInfo_.scratBuffType, tmpRecvOff + buffInfo_.scratchBuffBaseOff, tmpRecvSize);
486 : DataSlice recvLocDstSlice
487 0 : = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
488 :
489 0 : if (!primRecvReduce) {
490 0 : primRecvReduce.reset(new PrimRecvReduce(neighborRank, priorLinkData, recvRemSlice, recvLocSrcSlice,
491 0 : recvLocDstSlice, dataType_, redOp_, dmaMode_));
492 : } else {
493 0 : primRecvReduce->Append(recvRemSlice, recvLocSrcSlice, recvLocDstSlice);
494 : }
495 : }
496 : }
497 0 : return primRecvReduce;
498 0 : }
499 :
500 0 : HcclResult TempReduceScatterConcurrMesh::PostCopyOffload(const RankSliceInfo &sliceInfoVec,
501 : std::vector<PrimQuePtr> &tempPrimQues)
502 : {
503 0 : u64 srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
504 0 : u64 srcSize = 0;
505 0 : for (u32 dimIdx = 0; dimIdx < sliceInfoVec[0].size(); dimIdx++) {
506 0 : srcSize += sliceInfoVec[tempVirtRankMap_[myRank_]][dimIdx].size;
507 : }
508 0 : u64 dstOffset = 0;
509 0 : DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, srcOffset + buffInfo_.inBuffBaseOff, srcSize);
510 0 : DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, dstOffset + buffInfo_.outBuffBaseOff, srcSize);
511 0 : std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
512 0 : tempPrimQues[0]->Append(std::move(primLocalCopy));
513 :
514 0 : return HcclResult::HCCL_SUCCESS;
515 0 : }
516 :
517 : } // namespace Hccl
|