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 <numeric>
12 :
13 : #include "log.h"
14 :
15 : #include "temp_reduce_scatter_ring.h"
16 :
17 : namespace Hccl {
18 0 : TempReduceScatterRing::TempReduceScatterRing(
19 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
20 0 : const std::map<RankId, u32>& tempVirtRankMap)
21 0 : : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
22 0 : {}
23 :
24 0 : TempReduceScatterRing::~TempReduceScatterRing() {}
25 :
26 : HcclResult
27 0 : TempReduceScatterRing::CalcRes(const bool forAllReduce, AlgTempResReq& tempResReq, u32& requiredScratchMultiplier)
28 : {
29 0 : tempResReq.queNum = tempVTopo_.size();
30 0 : requiredScratchMultiplier = forAllReduce ? tempRankSize_ : 0;
31 :
32 0 : CHK_PRT_RET(
33 : CalcResLinksRing(myRank_, tempRankSize_, tempVTopo_, tempResReq) != HcclResult::HCCL_SUCCESS,
34 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterRing] Rank [%d], resLinks calculation error!", myRank_),
35 : HcclResult::HCCL_E_INTERNAL);
36 :
37 0 : return HcclResult::HCCL_SUCCESS;
38 : }
39 :
40 : /*
41 : dataSize / (rankSize) --> chunkSize
42 : dataSize / (rankSize * queNum) --> sliceSize
43 :
44 : SliceInfoVecforRing: [1st chunk: [1st Slice, 2nd Slice, ...], 2nd chunk: [1st Slice, 2nd Slice, ...], ...]
45 : */
46 0 : HcclResult TempReduceScatterRing::CalcSliceInfo(
47 : const AllignInfo& allignInfo, const bool forAllReduce, const u64 dataSize, RankSliceInfo& sliceInfoVec)
48 : {
49 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
50 0 : sliceInfoVec.resize(tempRankSize_, tmp);
51 :
52 0 : if (forAllReduce) {
53 : // for allreduce, dataSize = total dataSize
54 0 : CHK_RET(CalcSliceInfoAllReduce(allignInfo, dataSize, sliceInfoVec));
55 : } else {
56 : // for reduce scatter, dataSize = chunkSize
57 0 : CHK_RET(CalcRsAgSliceInfoRing(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
58 : }
59 0 : return HcclResult::HCCL_SUCCESS;
60 0 : }
61 :
62 0 : HcclResult TempReduceScatterRing::CalcSliceInfoAllReduce(
63 : const AllignInfo& allignInfo, const u64 dataSize, RankSliceInfo& sliceInfoVec)
64 : {
65 0 : u32 queNum = tempVTopo_.size();
66 : u64 unitAllignSize;
67 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
68 :
69 0 : u64 queDataSize = RoundUp(dataSize, (queNum * unitAllignSize)) * unitAllignSize;
70 :
71 0 : u64 resDataSize = dataSize;
72 0 : std::vector<u64> resQueData;
73 0 : std::vector<u64> queChunkSize;
74 0 : for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
75 : // split data on queues
76 0 : u64 currQueDataSize = (resDataSize > queDataSize) ? queDataSize : resDataSize;
77 0 : resQueData.push_back(currQueDataSize);
78 0 : resDataSize -= currQueDataSize;
79 :
80 : // support ReduceScatterV and AllGatherV for better data alignment when enable Data Align
81 0 : u64 currQueChunkSize = RoundUp(currQueDataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
82 0 : queChunkSize.push_back(currQueChunkSize);
83 : }
84 0 : CHK_PRT_RET(
85 : resDataSize != 0,
86 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterRing] Rank [%d], SliceInfo calculation error!", myRank_),
87 : HcclResult::HCCL_E_INTERNAL);
88 :
89 0 : u64 accumOff = 0;
90 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
91 0 : for (u32 queIdx = 0; queIdx < queNum; queIdx++) {
92 0 : u64 currSliceSize = (resQueData[queIdx] > queChunkSize[queIdx]) ? queChunkSize[queIdx] : resQueData[queIdx];
93 0 : SliceInfo currSlice = {accumOff, currSliceSize};
94 0 : resQueData[queIdx] -= currSliceSize;
95 0 : accumOff += currSliceSize;
96 0 : sliceInfoVec[rankIdx][queIdx] = currSlice;
97 : }
98 : }
99 :
100 0 : CHK_PRT_RET(
101 : ((sliceInfoVec[tempRankSize_ - 1][queNum - 1].offset + sliceInfoVec[tempRankSize_ - 1][queNum - 1].size
102 : != dataSize)
103 : || (accumulate(resQueData.begin(), resQueData.end(), 0) != 0)),
104 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterRing] Rank [%d], SliceInfo calculation error!", myRank_),
105 : HcclResult::HCCL_E_INTERNAL);
106 :
107 0 : return HcclResult::HCCL_SUCCESS;
108 0 : }
109 :
110 0 : HcclResult TempReduceScatterRing::GenPrimQue(
111 : const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
112 : std::vector<PrimQuePtr>& tempPrimQues)
113 : {
114 0 : opMode_ = tempFuncs.opMode;
115 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
116 0 : buffInfo_ = buffInfo;
117 :
118 0 : queNum_ = tempVTopo_.size();
119 0 : CHK_PRT_RET(
120 : queNum_ != tempPrimQues.size(),
121 : HCCL_ERROR("[CollAlgFactory] [TempReduceScatterRing] Rank [%d], requiredQue Error.", myRank_),
122 : HcclResult::HCCL_E_INTERNAL);
123 :
124 : // LocalCopy: from input to scratch In Buffer for OPBASE
125 0 : if (tempFuncs.isForepart) {
126 0 : CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
127 : }
128 :
129 0 : stepNum_ = tempRankSize_ - 1;
130 0 : for (u32 queIdx = 0; queIdx < tempVTopo_.size(); queIdx++) {
131 : // semaphore sync
132 0 : if (queNum_ > 1) {
133 0 : CHK_PRT_RET(
134 : PreSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
135 : HCCL_ERROR(
136 : "[CollAlgFactory] [TempReduceScatterRing] Rank [%d], Unable to synchronize all queues.", myRank_),
137 : HcclResult::HCCL_E_INTERNAL);
138 : }
139 :
140 0 : PrimQuePtr currPrimQue = tempPrimQues[queIdx];
141 0 : CHK_RET(RunIndividualRing(queIdx, tempFuncs.forAllReduce, sliceInfoVec, tempLinks, currPrimQue));
142 :
143 0 : if ((!tempFuncs.forAllReduce) && (queNum_ > 1)) {
144 : // semaphore sync for standalone reducescatter
145 0 : CHK_PRT_RET(
146 : PostSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
147 : HCCL_ERROR(
148 : "[CollAlgFactory] [TempReduceScatterRing] Rank [%d], Unable to synchronize all queues.", myRank_),
149 : HcclResult::HCCL_E_INTERNAL);
150 : }
151 0 : }
152 :
153 : // LocalCopy for standalone reducescatter in Offload Mode
154 0 : if ((opMode_ == OpMode::OFFLOAD) && !tempFuncs.forAllReduce && !tempFuncs.forAlgSeqComb) {
155 0 : CHK_RET(PostCopyOffload(sliceInfoVec, tempPrimQues));
156 : }
157 :
158 : // LocalCopy from scratch to output for Opbase
159 0 : if (tempFuncs.isBottom && !tempFuncs.forAllReduce) {
160 0 : CHK_RET(PostCopyOpbase(tempFuncs.usrData, tempPrimQues));
161 : }
162 :
163 0 : return HcclResult::HCCL_SUCCESS;
164 : }
165 :
166 0 : HcclResult TempReduceScatterRing::RunIndividualRing(
167 : const u32 queIdx, const bool& forAllReduce, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
168 : PrimQuePtr currPrimQue)
169 : {
170 : // locate myRank in tempVTopo -> algRank
171 : u32 myAlgRank;
172 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[queIdx], myAlgRank));
173 :
174 : // find neighbors -> virtualRank
175 0 : RankId sendToRank = tempVTopo_[queIdx][(myAlgRank + 1) % tempRankSize_];
176 0 : RankId recvFromRank = tempVTopo_[queIdx][(myAlgRank - 1 + tempRankSize_) % tempRankSize_]; // virtualRank
177 :
178 : // Link
179 0 : LinkData sendLinkData = tempLinks.at(sendToRank)[0];
180 0 : LinkData recvLinkData = tempLinks.at(recvFromRank)[0];
181 :
182 : // run stepNum steps to complete the ring
183 0 : for (u32 step = 0; step < stepNum_; step++) {
184 0 : u32 sendChunkIdx = tempVirtRankMap_[tempVTopo_[queIdx][(myAlgRank - 1 - step + tempRankSize_) % tempRankSize_]];
185 0 : u64 sendOffset = sliceInfoVec[sendChunkIdx][queIdx].offset;
186 0 : u64 sendSize = sliceInfoVec[sendChunkIdx][queIdx].size;
187 0 : u64 tmpScratchSendOff = forAllReduce ? sendOffset : (sendOffset - sliceInfoVec[sendChunkIdx][0].offset);
188 :
189 0 : u32 recvChunkIdx = tempVirtRankMap_[tempVTopo_[queIdx][(myAlgRank - 2 - step + tempRankSize_) % tempRankSize_]];
190 0 : u64 recvOffset = sliceInfoVec[recvChunkIdx][queIdx].offset;
191 0 : u64 recvSize = sliceInfoVec[recvChunkIdx][queIdx].size;
192 0 : u64 tmpScratchRecvOff = forAllReduce ? recvOffset : (recvOffset - sliceInfoVec[recvChunkIdx][0].offset);
193 :
194 : // PrimGroup
195 0 : std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
196 :
197 : // SendReduce
198 0 : DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
199 : DataSlice sendRemSrcSlice
200 0 : = DataSlice(buffInfo_.outBuffType, tmpScratchSendOff + buffInfo_.outBuffBaseOff, sendSize);
201 0 : DataSlice sendRemDstSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
202 0 : std::unique_ptr<Primitive> primSendReduce = std::make_unique<PrimSendReduce>(
203 0 : sendToRank, sendLinkData, sendLocSlice, sendRemSrcSlice, sendRemDstSlice, dataType_, redOp_, dmaMode_);
204 :
205 0 : primGroup->Append(std::move(primSendReduce));
206 :
207 : // RecvReduce
208 0 : DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
209 : DataSlice recvLocSrcSlice
210 0 : = DataSlice(buffInfo_.outBuffType, tmpScratchRecvOff + buffInfo_.outBuffBaseOff, recvSize);
211 0 : DataSlice recvLocDstSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
212 0 : std::unique_ptr<Primitive> primRecvReduce = std::make_unique<PrimRecvReduce>(
213 0 : recvFromRank, recvLinkData, recvRemSlice, recvLocSrcSlice, recvLocDstSlice, dataType_, redOp_, dmaMode_);
214 :
215 0 : primGroup->Append(std::move(primRecvReduce));
216 :
217 0 : currPrimQue->Append(std::move(primGroup));
218 0 : }
219 :
220 0 : return HcclResult::HCCL_SUCCESS;
221 : }
222 :
223 : HcclResult
224 0 : TempReduceScatterRing::PostCopyOffload(const RankSliceInfo& sliceInfoVec, std::vector<PrimQuePtr>& tempPrimQues)
225 : {
226 0 : u64 srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
227 0 : u64 srcSize = sliceInfoVec[tempVirtRankMap_[myRank_]][queNum_ - 1].offset
228 0 : - sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset
229 0 : + sliceInfoVec[tempVirtRankMap_[myRank_]][queNum_ - 1].size;
230 0 : u64 dstOffset = 0;
231 0 : DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, srcOffset + buffInfo_.inBuffBaseOff, srcSize);
232 0 : DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, dstOffset + buffInfo_.outBuffBaseOff, srcSize);
233 0 : std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
234 0 : tempPrimQues[0]->Append(std::move(primLocalCopy));
235 :
236 0 : return HcclResult::HCCL_SUCCESS;
237 0 : }
238 :
239 : } // namespace Hccl
|