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