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 "ins_temp_reduce_scatter_mesh_2D.h"
12 :
13 : #include "log.h"
14 : #include "alg_data_trans_wrapper.h"
15 :
16 : namespace Hccl {
17 0 : InsTempReduceScatterMesh2D::InsTempReduceScatterMesh2D(
18 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
19 0 : const std::map<RankId, u32>& tempVirtRankMap)
20 0 : : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
21 : {
22 0 : xQueNum_ = tempVTopo_[0].size() - 1; // x轴的卡数-1
23 0 : yQueNum_ = tempVTopo_[1].size() - 1; // y轴的卡数-1
24 0 : xRankSize_ = tempVTopo_[0].size(); // x轴的卡数
25 0 : yRankSize_ = tempVTopo_[1].size(); // y轴的卡数
26 0 : }
27 :
28 0 : InsTempReduceScatterMesh2D::~InsTempReduceScatterMesh2D() {}
29 :
30 0 : u64 InsTempReduceScatterMesh2D::CalcScratchMultiple(const BufferType& inBuffType, const BufferType& outBuffType)
31 : {
32 : (void)inBuffType;
33 : (void)outBuffType;
34 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
35 0 : u64 scratchMultiple = xyMaxRankSize * (xRankSize_ + yRankSize_);
36 0 : return scratchMultiple;
37 : }
38 :
39 0 : HcclResult InsTempReduceScatterMesh2D::CalcResLinksMesh2D(const u32 linkNumBtwPeers, AlgTempResReq& tempResReq)
40 : {
41 : u32 myAlgRank;
42 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
43 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
44 0 : for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
45 0 : RankId neighborRank = tempVTopo_[dim][(myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size())];
46 0 : tempResReq.links[neighborRank] = linkNumBtwPeers;
47 : }
48 : }
49 0 : return HcclResult::HCCL_SUCCESS;
50 : }
51 :
52 0 : HcclResult InsTempReduceScatterMesh2D::CalcRes(AlgTempResReq& tempResReq)
53 : {
54 : // Mesh 需要的 que Num 为 tempVTopo_[0].size() + tempVTopo_[1].size() - 2
55 0 : tempResReq.queNum = (xRankSize_ > 1 && yRankSize_ > 1) ? (xQueNum_ + yQueNum_) : 1;
56 0 : tempResReq.streamNum = tempResReq.queNum;
57 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
58 0 : QId centerQ = 0;
59 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
60 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
61 : // linkNumBtwPeers_这个在没有绕路的情况下,是设置成1
62 0 : CHK_PRT_RET(
63 : CalcResLinksMesh2D(linkNumBtwPeers_, tempResReq) != HcclResult::HCCL_SUCCESS,
64 : HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], resLinks calculation error!", myRank_),
65 : HcclResult::HCCL_E_INTERNAL);
66 :
67 0 : return HcclResult::HCCL_SUCCESS;
68 : }
69 :
70 0 : HcclResult InsTempReduceScatterMesh2D::GenExtIns(
71 : const TempFuncs& tempFuncs, TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
72 : std::vector<InsQuePtr>& tempInsQues)
73 : {
74 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[0], xRankId_)); // 得到当前卡在x轴上的编号
75 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[1], yRankId_)); // 得到当前卡在y轴上的编号
76 0 : opMode_ = tempFuncs.opMode;
77 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
78 0 : queNum_ = xQueNum_ + yQueNum_;
79 0 : u64 sliceNum = tempAlgParams.sliceSize / DataTypeSizeGet(dataType_); // 先计算得到本次迭代处理的数据量
80 0 : halfDataSize_ = sliceNum / PARALLEL_SIZE * DataTypeSizeGet(dataType_); // 前一半数据的size
81 0 : HCCL_INFO("[InsTempReduceScatterMesh2D] Run Start");
82 : // 这里不支持绕路的时候,应该就用原始的tempInsQues就行
83 0 : CHK_PRT_RET(
84 : queNum_ != tempInsQues.size(),
85 : HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], requiredQue Error.", myRank_),
86 : HcclResult::HCCL_E_INTERNAL);
87 0 : PreCopy(tempAlgParams, tempInsQues); // stream 0作为主流,负责把本卡的数据拷贝到scratchbuffer上
88 0 : if (queNum_ > 1) {
89 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
90 : }
91 0 : CHK_RET(RunFirstLevel(tempLinks, tempInsQues, tempAlgParams));
92 0 : if (queNum_ > 1) {
93 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
94 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
95 : }
96 0 : CHK_RET(RunFirstReduce(tempInsQues, tempAlgParams));
97 0 : if (queNum_ > 1) {
98 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
99 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
100 : }
101 0 : RunSecondLevel(tempLinks, tempInsQues, tempAlgParams);
102 0 : if (queNum_ > 1) {
103 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
104 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
105 : }
106 0 : RunSecondReduce(tempInsQues, tempAlgParams);
107 0 : if (queNum_ > 1) {
108 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
109 : }
110 0 : return HcclResult::HCCL_SUCCESS;
111 : }
112 :
113 : HcclResult
114 0 : InsTempReduceScatterMesh2D::PreCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
115 : {
116 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
117 0 : u64 remainDataSize = tempAlgParams.sliceSize - halfDataSize_;
118 : // 前一半数据,将本卡数据从input拷贝到scratchbuffer
119 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
120 0 : u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
121 0 : for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
122 0 : u32 rankId = yRankId * xRankSize_ + xRankId_; // 同y轴平面的所有卡,
123 : DataSlice inputRankSlice = DataSlice(
124 : tempAlgParams.buffInfo.inBuffType,
125 0 : tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride
126 0 : + rpt * tempAlgParams.inputRepeatStride,
127 0 : halfDataSize_);
128 : DataSlice scratchRankSlice = DataSlice(
129 : tempAlgParams.buffInfo.scratBuffType,
130 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
131 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_) + scratchRepeatStride,
132 0 : halfDataSize_);
133 0 : CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
134 0 : HCCL_DEBUG(
135 : "[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] top inputRankSlice: %s, scratchRankSlice: %s",
136 : myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
137 : }
138 : }
139 : // 后一半数据,将本卡数据从input拷贝到scratchbuffer
140 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
141 0 : u64 scratchRepeatStride
142 0 : = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
143 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
144 0 : for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
145 0 : u32 rankId = yRankId_ * xRankSize_ + xRankId; // 同x轴平面的所有卡,
146 : DataSlice inputRankSlice = DataSlice(
147 : tempAlgParams.buffInfo.inBuffType,
148 0 : tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride + halfDataSize_
149 0 : + rpt * tempAlgParams.inputRepeatStride,
150 0 : remainDataSize);
151 : DataSlice scratchRankSlice = DataSlice(
152 : tempAlgParams.buffInfo.scratBuffType,
153 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
154 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_) + scratchRepeatStride,
155 0 : remainDataSize);
156 0 : CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
157 0 : HCCL_DEBUG(
158 : "[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] bottom inputRankSlice: %s, scratchRankSlice: %s",
159 : myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
160 : }
161 : }
162 0 : HCCL_INFO("[InsTempReduceScatterMesh2D][PreCopy], copy from userIn to scratch");
163 0 : return HcclResult::HCCL_SUCCESS;
164 : }
165 :
166 0 : HcclResult InsTempReduceScatterMesh2D::SendRecvProcess(
167 : const ResLinks& tempLinks, std::vector<std::vector<DataSlice>> allSliceVec, std::vector<InsQuePtr>& tempInsQues,
168 : u32 remoteRank, u32 queIdx) const
169 : {
170 0 : CHK_PRT_RET(
171 : tempInsQues.empty(), HCCL_ERROR("[InsTempReduceScatterMesh2D][SendRecvProcess] empty queue"),
172 : HcclResult::HCCL_E_INTERNAL);
173 0 : CHK_PTR_NULL(tempInsQues[0]);
174 0 : HCCL_DEBUG("[InsTempReduceScatterMesh2D][SendRecvProcess] SendRecvProcess start");
175 0 : const std::vector<LinkData>& linkRecv = tempLinks.at(remoteRank);
176 0 : const std::vector<LinkData>& linkSend = tempLinks.at(remoteRank);
177 : SendRecvInfo sendRecvInfo{
178 0 : {linkSend[0], linkRecv[0]}, {{allSliceVec[2], allSliceVec[3]}, {allSliceVec[0], allSliceVec[1]}}};
179 :
180 0 : CHK_PRT_THROW(
181 : queIdx >= tempInsQues.size(),
182 : HCCL_ERROR(
183 : "[InsTempReduceScatterMesh2D] queIdx[%u] is bigger than tempInsQues size[%zu].", queIdx,
184 : tempInsQues.size()),
185 : InvalidParamsException, "queIdx is invalid");
186 : // 做了DMA消减之后只支持PUT
187 0 : CHK_PRT_RET(
188 : SendRecv(sendRecvInfo, tempInsQues[queIdx], 0, true, DmaMode::PUT),
189 : HCCL_ERROR("[InsTempReduceScatterMesh2D] RunReduceScatter SendReduce failed"), HcclResult::HCCL_E_INTERNAL);
190 0 : return HcclResult::HCCL_SUCCESS;
191 0 : }
192 :
193 : // 前一半数据的先x轴 和 后一半数据的先y轴
194 0 : HcclResult InsTempReduceScatterMesh2D::RunFirstLevel(
195 : const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
196 : {
197 0 : HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstLevel] myRank[%d]", myRank_);
198 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
199 : u64 processSize;
200 0 : for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
201 : u32 remoteRank;
202 : u32 index;
203 0 : std::vector<DataSlice> rxSrcSlices;
204 0 : std::vector<DataSlice> rxDstSlices;
205 0 : std::vector<DataSlice> txSrcSlices;
206 0 : std::vector<DataSlice> txDstSlices;
207 0 : if (queIdx < xQueNum_) { // 前xRankSize-1个stream,首先拉取前一半数据
208 0 : index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
209 0 : remoteRank = tempVTopo_[0][index];
210 0 : processSize = halfDataSize_;
211 0 : HCCL_DEBUG(
212 : "[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] toRank[%u] fromRank[%u] "
213 : "rpt[%u], index[%u]",
214 : myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
215 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
216 0 : u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
217 0 : for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
218 0 : u32 readRankId = yRankId * xRankSize_ + xRankId_;
219 0 : u32 writeRankId = yRankId * xRankSize_ + index;
220 : // 数据从其他卡,传输到本卡,接收数据
221 0 : rxSrcSlices.emplace_back(
222 0 : tempAlgParams.buffInfo.inBuffType,
223 0 : tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride
224 0 : + rpt * tempAlgParams.inputRepeatStride,
225 : processSize);
226 0 : rxDstSlices.emplace_back(
227 0 : tempAlgParams.buffInfo.scratBuffType,
228 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
229 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + index) + scratchRepeatStride,
230 : processSize);
231 0 : txSrcSlices.emplace_back(
232 0 : tempAlgParams.buffInfo.inBuffType,
233 0 : tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride
234 0 : + rpt * tempAlgParams.inputRepeatStride,
235 : processSize);
236 0 : txDstSlices.emplace_back(
237 0 : tempAlgParams.buffInfo.scratBuffType,
238 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
239 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_)
240 0 : + scratchRepeatStride,
241 : processSize);
242 0 : HCCL_DEBUG(
243 : "[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] *****sendrecv*****, "
244 : "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
245 : myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
246 : txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
247 : }
248 : }
249 : } else { // 后yRankSize-1个stream,首先拉取后一半数据
250 0 : index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
251 0 : remoteRank = tempVTopo_[1][index];
252 0 : processSize = tempAlgParams.sliceSize - halfDataSize_;
253 0 : HCCL_DEBUG(
254 : "[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u], "
255 : "rpt[%u], index[%u]",
256 : myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
257 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
258 0 : u64 scratchRepeatStride
259 0 : = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
260 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
261 0 : for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
262 0 : u32 readRankId = yRankId_ * xRankSize_ + xRankId; // 同x轴平面的所有卡,
263 0 : u32 writeRankId = index * xRankSize_ + xRankId;
264 0 : rxSrcSlices.emplace_back(
265 0 : tempAlgParams.buffInfo.inBuffType,
266 0 : tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride
267 0 : + halfDataSize_ + rpt * tempAlgParams.inputRepeatStride,
268 : processSize);
269 0 : rxDstSlices.emplace_back(
270 0 : tempAlgParams.buffInfo.scratBuffType,
271 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
272 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + index) + scratchRepeatStride,
273 : processSize);
274 0 : txSrcSlices.emplace_back(
275 0 : tempAlgParams.buffInfo.inBuffType,
276 0 : tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride
277 0 : + halfDataSize_ + rpt * tempAlgParams.inputRepeatStride,
278 : processSize);
279 0 : txDstSlices.emplace_back(
280 0 : tempAlgParams.buffInfo.scratBuffType, // tempAlgParams.buffInfo.scratBuffType,
281 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
282 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_)
283 0 : + scratchRepeatStride,
284 : processSize);
285 0 : HCCL_DEBUG(
286 : "[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
287 : "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
288 : myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
289 : txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
290 : }
291 : }
292 : }
293 0 : if (processSize == 0) {
294 0 : continue;
295 : }
296 0 : std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
297 0 : CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
298 0 : }
299 0 : return HcclResult::HCCL_SUCCESS;
300 0 : }
301 :
302 : HcclResult
303 0 : InsTempReduceScatterMesh2D::RunFirstReduce(std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
304 : {
305 0 : HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
306 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
307 0 : u64 processSize = 0;
308 : // 这里的stream 0和stream xRankSize-1分别负责前一半数据与后一半数据的本地reduce
309 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
310 0 : u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
311 0 : for (u32 tmpRank = 0; tmpRank < yRankSize_; tmpRank++) { // 前一半数据做local reduce,由这部分的第一个stream做
312 0 : processSize = halfDataSize_;
313 0 : for (u32 dataIdx = 1; dataIdx < xRankSize_;
314 : dataIdx++) { // 原始这个位置已经有数据了,因此从后一片数据开始累加
315 : DataSlice srcDataSlice = DataSlice(
316 : tempAlgParams.buffInfo.scratBuffType,
317 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
318 0 : + (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
319 0 : processSize);
320 : DataSlice dstDataSlice = DataSlice(
321 : tempAlgParams.buffInfo.scratBuffType,
322 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
323 0 : + (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride,
324 0 : processSize);
325 0 : HCCL_DEBUG(
326 : "[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
327 : "srcDataSlice: %s, dstDataSlice: %s",
328 : myRank_, srcDataSlice.Describe().c_str(), dstDataSlice.Describe().c_str());
329 0 : CHK_RET(LocalReduce(tempInsQues[0], srcDataSlice, dstDataSlice, dataType_, redOp_));
330 : }
331 : }
332 : }
333 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
334 0 : u64 scratchRepeatStride
335 0 : = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
336 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
337 0 : for (u32 tmpRank = 0; tmpRank < xRankSize_; tmpRank++) { // 后一半数据做local reduce,由这部分的第一个stream做
338 0 : processSize = tempAlgParams.sliceSize - halfDataSize_;
339 0 : for (u32 dataIdx = 1; dataIdx < yRankSize_;
340 : dataIdx++) { // 原始这个位置已经有数据了,因此从后一片数据开始累加
341 : DataSlice srcDataSlice = DataSlice(
342 : tempAlgParams.buffInfo.scratBuffType,
343 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
344 0 : + (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
345 0 : processSize);
346 : DataSlice dstDataSlice = DataSlice(
347 : tempAlgParams.buffInfo.scratBuffType,
348 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
349 0 : + (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride,
350 0 : processSize);
351 0 : HCCL_DEBUG(
352 : "[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
353 : "srcDataSlice: %s, dstDataSlice: %s",
354 : myRank_, srcDataSlice.Describe().c_str(), dstDataSlice.Describe().c_str());
355 0 : CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcDataSlice, dstDataSlice, dataType_, redOp_));
356 : }
357 : }
358 : }
359 0 : return HcclResult::HCCL_SUCCESS;
360 : }
361 :
362 : // 后一半数据的后x轴 和 前一半数据的后y轴
363 0 : HcclResult InsTempReduceScatterMesh2D::RunSecondLevel(
364 : const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
365 : {
366 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
367 : u64 processSize;
368 0 : for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
369 : u32 remoteRank;
370 : u32 index;
371 0 : std::vector<DataSlice> rxSrcSlices;
372 0 : std::vector<DataSlice> rxDstSlices;
373 0 : std::vector<DataSlice> txSrcSlices;
374 0 : std::vector<DataSlice> txDstSlices;
375 0 : if (queIdx < xQueNum_) { // 前xRankSize-1个stream,后一半数据
376 0 : index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
377 0 : remoteRank = tempVTopo_[0][index];
378 0 : HCCL_DEBUG(
379 : "[InsTempReduceScatterMesh2D][RunSecondLevel] queIdx < xQueNum myRank[%d] toRank[%u] fromRank[%u]",
380 : myRank_, remoteRank, remoteRank);
381 0 : processSize = tempAlgParams.sliceSize - halfDataSize_;
382 : // 这里过来的数据,直接按照queIdx的顺序放置,不一定是按照rankId顺序排列的
383 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
384 0 : u64 scratchRepeatStride
385 0 : = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
386 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
387 : DataSlice rxSrcSlice = DataSlice(
388 : tempAlgParams.buffInfo.scratBuffType,
389 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
390 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_) + scratchRepeatStride,
391 0 : processSize);
392 : DataSlice rxDstSlice = DataSlice(
393 : tempAlgParams.buffInfo.scratBuffType,
394 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
395 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_ + queIdx + 1)
396 : + scratchRepeatStride,
397 0 : processSize);
398 : DataSlice txSrcSlice = DataSlice(
399 : tempAlgParams.buffInfo.scratBuffType,
400 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
401 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride,
402 0 : processSize);
403 : DataSlice txDstSlice = DataSlice(
404 : tempAlgParams.buffInfo.scratBuffType,
405 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
406 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx + 1) + scratchRepeatStride,
407 0 : processSize);
408 :
409 0 : rxSrcSlices.emplace_back(rxSrcSlice);
410 0 : rxDstSlices.emplace_back(rxDstSlice);
411 0 : txSrcSlices.emplace_back(txSrcSlice);
412 0 : txDstSlices.emplace_back(txDstSlice);
413 :
414 0 : HCCL_DEBUG(
415 : "[InsTempReduceScatterMesh2D][RunSecondLevel] queId < xQueNum myRank[%d] *****sendrecv*****, "
416 : "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
417 : myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
418 : txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
419 : }
420 : } else { // 后yRankSize-1个stream,前一半数据
421 0 : index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
422 0 : remoteRank = tempVTopo_[1][index];
423 0 : HCCL_DEBUG(
424 : "[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u] "
425 : "rpt[%u]",
426 : myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum);
427 0 : processSize = halfDataSize_;
428 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
429 0 : u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
430 : DataSlice rxSrcSlice = DataSlice(
431 : tempAlgParams.buffInfo.scratBuffType,
432 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
433 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_) + scratchRepeatStride,
434 0 : processSize);
435 : DataSlice rxDstSlice = DataSlice(
436 : tempAlgParams.buffInfo.scratBuffType,
437 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
438 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_ + queIdx - xQueNum_ + 1)
439 : + scratchRepeatStride,
440 0 : processSize);
441 : DataSlice txSrcSlice = DataSlice(
442 : tempAlgParams.buffInfo.scratBuffType,
443 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
444 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride,
445 0 : processSize);
446 : DataSlice txDstSlice = DataSlice(
447 : tempAlgParams.buffInfo.scratBuffType,
448 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
449 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx - xQueNum_ + 1)
450 : + scratchRepeatStride,
451 0 : processSize);
452 :
453 0 : rxSrcSlices.emplace_back(rxSrcSlice);
454 0 : rxDstSlices.emplace_back(rxDstSlice);
455 0 : txSrcSlices.emplace_back(txSrcSlice);
456 0 : txDstSlices.emplace_back(txDstSlice);
457 :
458 0 : HCCL_DEBUG(
459 : "[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
460 : "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
461 : myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
462 : txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
463 : }
464 : }
465 0 : if (processSize == 0) {
466 0 : continue;
467 : }
468 0 : std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
469 0 : CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
470 0 : }
471 0 : return HcclResult::HCCL_SUCCESS;
472 0 : }
473 :
474 0 : HcclResult InsTempReduceScatterMesh2D::RunSecondReduce(
475 : std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
476 : {
477 0 : HCCL_INFO("[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
478 0 : u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
479 0 : u64 processSize = 0;
480 : // 这里的stream 0和stream xRankSize-1分别负责后一半数据与前一半数据的本地reduce
481 : // 后一半数据做local reduce,由这部分的第一个stream做
482 0 : processSize = tempAlgParams.sliceSize - halfDataSize_;
483 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
484 0 : u64 scratchRepeatStride
485 0 : = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
486 0 : + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
487 0 : for (u32 dataIdx = 0; dataIdx < xRankSize_; dataIdx++) {
488 : DataSlice srcSecDataSlice = DataSlice(
489 : tempAlgParams.buffInfo.scratBuffType,
490 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
491 0 : + (xyMaxRankSize * xRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
492 0 : processSize);
493 0 : u64 outOffset
494 0 : = tempAlgParams.buffInfo.outBuffBaseOff + halfDataSize_ + rpt * tempAlgParams.outputRepeatStride;
495 : DataSlice dstSecDataSlice = DataSlice(
496 : tempAlgParams.buffInfo.outBuffType, // BufferType::OUTPUT,
497 0 : outOffset, processSize);
498 0 : HCCL_DEBUG(
499 : "[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
500 : "srcDataSlice: %s, dstDataSlice: %s",
501 : myRank_, srcSecDataSlice.Describe().c_str(), dstSecDataSlice.Describe().c_str());
502 0 : if (srcSecDataSlice != dstSecDataSlice) {
503 0 : if (dataIdx == 0) {
504 0 : CHK_RET(LocalCopy(tempInsQues[0], srcSecDataSlice, dstSecDataSlice));
505 : } else {
506 0 : CHK_RET(LocalReduce(tempInsQues[0], srcSecDataSlice, dstSecDataSlice, dataType_, redOp_));
507 : }
508 : }
509 : }
510 : }
511 : // 前一半数据做local reduce,由这部分的第一个stream做
512 0 : processSize = halfDataSize_;
513 0 : for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
514 0 : u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
515 0 : bool hasInplace = false;
516 0 : std::vector<DataSlice> srcFirDataSlices;
517 0 : std::vector<DataSlice> dstFirDataSlices;
518 0 : for (u32 dataIdx = 0; dataIdx < yRankSize_; dataIdx++) {
519 : DataSlice srcFirDataSlice = DataSlice(
520 : tempAlgParams.buffInfo.scratBuffType,
521 0 : tempAlgParams.buffInfo.scratchBuffBaseOff
522 0 : + (xyMaxRankSize * yRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
523 0 : processSize);
524 0 : u64 outOffset = tempAlgParams.buffInfo.outBuffBaseOff + rpt * tempAlgParams.outputRepeatStride;
525 : DataSlice dstFirDataSlice = DataSlice(
526 : tempAlgParams.buffInfo.outBuffType, // BufferType::OUTPUT,
527 0 : outOffset, processSize);
528 0 : HCCL_DEBUG(
529 : "[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
530 : "srcDataSlice: %s, dstDataSlice: %s",
531 : myRank_, srcFirDataSlice.Describe().c_str(), dstFirDataSlice.Describe().c_str());
532 0 : if (srcFirDataSlice != dstFirDataSlice) {
533 : #if DATASLICE_ONE
534 0 : srcFirDataSlices.push_back(srcFirDataSlice);
535 0 : dstFirDataSlices.push_back(dstFirDataSlice);
536 : #else
537 : if (dataIdx == 0) {
538 : CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice));
539 : } else {
540 : CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice, dataType_, redOp_));
541 : }
542 : #endif
543 : } else {
544 0 : hasInplace = true;
545 : }
546 : }
547 0 : for (u32 dataIdx = 0; dataIdx < srcFirDataSlices.size(); dataIdx++) {
548 0 : if (!hasInplace && dataIdx == 0) {
549 0 : CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx]));
550 0 : } else {
551 0 : CHK_RET(LocalReduce(
552 : tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx], dataType_, redOp_));
553 : }
554 : }
555 0 : }
556 0 : return HcclResult::HCCL_SUCCESS;
557 : }
558 :
559 : } // namespace Hccl
|