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