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