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 : #include "alg_data_trans_wrapper.h"
13 : #include "executor_utils.h"
14 : #include "ins_temp_all_to_all_mesh_2D.h"
15 :
16 : namespace Hccl {
17 0 : InsTempAlltoAllMesh2D::InsTempAlltoAllMesh2D(
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 0 : {}
22 :
23 0 : InsTempAlltoAllMesh2D::~InsTempAlltoAllMesh2D() {}
24 :
25 0 : HcclResult InsTempAlltoAllMesh2D::CalcRes(AlgTempResReq& tempResReq)
26 : {
27 0 : if (tempVTopo_.size() >= TEMPVTOPOSIZE) {
28 0 : rankId_ = myRank_;
29 0 : rankSize_ = tempRankSize_;
30 0 : xRankSize_ = tempVTopo_[0].size();
31 0 : yRankSize_ = tempVTopo_[1].size();
32 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[0], xRankId_));
33 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[1], yRankId_));
34 0 : CHK_PRT_RET((xRankSize_ == 0), HCCL_ERROR("xRankSize_ equals to zero."), HcclResult::HCCL_E_PARA);
35 0 : CHK_PRT_RET((yRankSize_ == 0), HCCL_ERROR("yRankSize_ equals to zero."), HcclResult::HCCL_E_PARA);
36 : } else {
37 0 : HCCL_ERROR("tempVTopo_.size() is [%zu]", tempVTopo_.size());
38 0 : return HcclResult::HCCL_E_INTERNAL;
39 : }
40 :
41 0 : HCCL_DEBUG(
42 : "rankId_ is [%u], rankSize_ is [%u], xRankSize_ is [%u], yRankSize_ is [%u]", rankId_, rankSize_, xRankSize_,
43 : yRankSize_);
44 :
45 0 : tempResReq.queNum = tempVTopo_[0].size() + tempVTopo_[1].size();
46 0 : tempResReq.streamNum = tempResReq.queNum;
47 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
48 :
49 0 : QId centerQ = 0;
50 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
51 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
52 :
53 : uint32_t myAlgRank;
54 0 : for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
55 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
56 0 : for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
57 0 : u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
58 0 : RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
59 0 : HCCL_INFO(
60 : "InsTempAlltoAllMesh2D::CalcRes Rank[%d], Dim[%u], NeighborRank[%d].", myRank_, dim, neighborRank);
61 : // LinkNum
62 0 : tempResReq.links[neighborRank] = 1;
63 : }
64 : }
65 0 : HCCL_INFO(
66 : "[InsTempAlltoAllMesh2D] Calculate resource, stream number is[%u], queNotifys size is[%u]",
67 : tempResReq.streamNum, tempResReq.queNotifys.size());
68 0 : return HcclResult::HCCL_SUCCESS;
69 : }
70 :
71 0 : HcclResult InsTempAlltoAllMesh2D::RunMeshX(
72 : std::vector<u64>& xDataInAddr, std::vector<u64>& xDataOutAddr, u64 xSize, BufferType srcBufferType,
73 : BufferType dstBufferType, DmaMode dmaMode, std::vector<InsQuePtr>& xInsQues, const ResLinks& tempLinks) const
74 : {
75 0 : HCCL_DEBUG("RunMeshX begin, xSize is [%u]", xSize);
76 0 : if (xSize == 0) {
77 0 : HCCL_INFO("[InsTempAlltoAllMesh2D] RunMeshX, xSize is 0");
78 0 : return HcclResult::HCCL_SUCCESS;
79 : }
80 0 : std::vector<DataSlice> txSrcSlices, txDstSlices, rxSrcSlices, rxDstSlices;
81 0 : for (u32 i = 0; i < rankSize_; i++) {
82 : // 计算send
83 0 : u64 xOffset = xRankId_;
84 0 : u64 yOffset = i / xRankSize_;
85 0 : u64 dstOffset = yOffset * xRankSize_ + xOffset;
86 :
87 0 : DataSlice txSrcSlice = DataSlice(srcBufferType, xDataInAddr[i], xSize);
88 0 : DataSlice txDstSlice = DataSlice(dstBufferType, xDataOutAddr[dstOffset], xSize);
89 0 : txSrcSlices.push_back(txSrcSlice);
90 0 : txDstSlices.push_back(txDstSlice);
91 :
92 : // 计算recv,recv侧的 xOffset,yOffset,dstOffset的计算方式和send侧一样
93 0 : DataSlice rxSrcSlice = DataSlice(srcBufferType, xDataInAddr[dstOffset], xSize);
94 0 : DataSlice rxDstSlice = DataSlice(dstBufferType, xDataOutAddr[i], xSize);
95 0 : rxSrcSlices.push_back(rxSrcSlice);
96 0 : rxDstSlices.push_back(rxDstSlice);
97 : }
98 :
99 : // 同一列的用一个队列
100 0 : std::vector<DataSlice> txLocalSrcSlices, txLocalDstSlices;
101 0 : for (u32 i = 0; i < rankSize_; i++) {
102 0 : if (i % xRankSize_ == xRankId_) {
103 0 : txLocalSrcSlices.push_back(txSrcSlices[i]);
104 0 : txLocalDstSlices.push_back(txDstSlices[i]);
105 : }
106 : }
107 : // 本地拷贝
108 0 : CHK_RET(LocalCopySlices(xInsQues[xRankId_], txLocalSrcSlices, txLocalDstSlices));
109 :
110 : // 拷贝到其他卡
111 0 : for (u32 queIdx = 0; queIdx < xRankSize_; queIdx++) {
112 0 : if (queIdx == xRankId_) {
113 0 : continue;
114 : }
115 :
116 0 : std::vector<DataSlice> txRmtSrcSlices, txRmtDstSlices, rxRmtSrcSlices, rxRmtDstSlices;
117 0 : for (u32 i = 0; i < rankSize_; i++) {
118 0 : if (i % xRankSize_ == queIdx) {
119 0 : txRmtSrcSlices.push_back(txSrcSlices[i]);
120 0 : txRmtDstSlices.push_back(txDstSlices[i]);
121 0 : rxRmtSrcSlices.push_back(rxSrcSlices[i]);
122 0 : rxRmtDstSlices.push_back(rxDstSlices[i]);
123 : }
124 : }
125 0 : TxRxSlicesList sendRecvSlicesList({txRmtSrcSlices, txRmtDstSlices}, {rxRmtSrcSlices, rxRmtDstSlices});
126 :
127 0 : RankId rankSendRecv = queIdx + yRankId_ * xRankSize_;
128 0 : const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankSendRecv);
129 0 : TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
130 :
131 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
132 0 : CHK_PRT_RET(
133 : SendRecv(sendRecvInfo, xInsQues[queIdx], 0, true, dmaMode),
134 : HCCL_ERROR("[InsTempAlltoAllMesh2D] RunMeshX SendRecv failed"), HcclResult::HCCL_E_INTERNAL);
135 0 : }
136 :
137 0 : return HcclResult::HCCL_SUCCESS;
138 0 : }
139 :
140 0 : HcclResult InsTempAlltoAllMesh2D::RunMeshY(
141 : std::vector<u64>& yDataInAddr, std::vector<u64>& yDataOutAddr, u64 ySize, BufferType srcBufferType,
142 : BufferType dstBufferType, DmaMode dmaMode, std::vector<InsQuePtr>& yInsQues, const ResLinks& tempLinks) const
143 : {
144 0 : HCCL_DEBUG("RunMeshX begin, ySize is [%u]", ySize);
145 0 : if (ySize == 0) {
146 0 : HCCL_INFO("[InsTempAlltoAllMesh2D] RunMeshY, ySize is 0");
147 0 : return HcclResult::HCCL_SUCCESS;
148 : }
149 :
150 0 : std::vector<DataSlice> txSrcSlices, txDstSlices, rxSrcSlices, rxDstSlices;
151 0 : for (u32 i = 0; i < rankSize_; i++) {
152 : // 计算send
153 0 : u64 xOffset = i % xRankSize_;
154 0 : u64 yOffset = yRankId_;
155 0 : u64 dstOffset = yOffset * xRankSize_ + xOffset;
156 :
157 0 : DataSlice txSrcSlice = DataSlice(srcBufferType, yDataInAddr[i], ySize);
158 0 : DataSlice txDstSlice = DataSlice(dstBufferType, yDataOutAddr[dstOffset], ySize);
159 0 : txSrcSlices.push_back(txSrcSlice);
160 0 : txDstSlices.push_back(txDstSlice);
161 :
162 : // 计算recv,recv侧的 xOffset,yOffset,dstOffset的计算方式和send侧一样
163 0 : DataSlice rxSrcSlice = DataSlice(srcBufferType, yDataInAddr[dstOffset], ySize);
164 0 : DataSlice rxDstSlice = DataSlice(dstBufferType, yDataOutAddr[i], ySize);
165 0 : rxSrcSlices.push_back(rxSrcSlice);
166 0 : rxDstSlices.push_back(rxDstSlice);
167 : }
168 :
169 : // 同一行的用一个队列
170 0 : std::vector<DataSlice> txLocalSrcSlices, txLocalDstSlices;
171 0 : for (u32 i = 0; i < rankSize_; i++) {
172 0 : if (i / xRankSize_ == yRankId_) {
173 0 : txLocalSrcSlices.push_back(txSrcSlices[i]);
174 0 : txLocalDstSlices.push_back(txDstSlices[i]);
175 : }
176 : }
177 : // 本地拷贝
178 0 : CHK_RET(LocalCopySlices(yInsQues[yRankId_], txLocalSrcSlices, txLocalDstSlices));
179 :
180 : // 拷贝到其他卡
181 0 : for (u32 queIdx = 0; queIdx < yRankSize_; queIdx++) {
182 0 : if (queIdx == yRankId_) {
183 0 : continue;
184 : }
185 :
186 0 : std::vector<DataSlice> txRmtSrcSlices, txRmtDstSlices, rxRmtSrcSlices, rxRmtDstSlices;
187 0 : for (u32 i = 0; i < rankSize_; i++) {
188 0 : if (i / xRankSize_ == queIdx) {
189 0 : txRmtSrcSlices.push_back(txSrcSlices[i]);
190 0 : txRmtDstSlices.push_back(txDstSlices[i]);
191 0 : rxRmtSrcSlices.push_back(rxSrcSlices[i]);
192 0 : rxRmtDstSlices.push_back(rxDstSlices[i]);
193 : }
194 : }
195 0 : TxRxSlicesList sendRecvSlicesList({txRmtSrcSlices, txRmtDstSlices}, {rxRmtSrcSlices, rxRmtDstSlices});
196 :
197 0 : RankId rankSendRecv = queIdx * xRankSize_ + xRankId_;
198 0 : const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankSendRecv);
199 0 : TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
200 :
201 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
202 0 : CHK_PRT_RET(
203 : SendRecv(sendRecvInfo, yInsQues[queIdx], 0, true, dmaMode),
204 : HCCL_ERROR("[InsTempAlltoAllMesh2D] RunMeshY SendRecv failed"), HcclResult::HCCL_E_INTERNAL);
205 0 : }
206 :
207 0 : return HcclResult::HCCL_SUCCESS;
208 0 : }
209 :
210 0 : HcclResult InsTempAlltoAllMesh2D::GenExtIns(
211 : const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
212 : std::vector<InsQuePtr>& tempInsQues) const
213 : {
214 : (void)tempFuncs;
215 0 : HCCL_INFO("[InsTempAlltoAllMesh2D] Run algorithm start: rank[%d]", myRank_);
216 :
217 0 : u64 xSize = tempAlgParams.sliceSize / 2;
218 0 : u64 ySize = tempAlgParams.sliceSize - xSize;
219 :
220 : // queue arrangement
221 0 : std::vector<InsQuePtr> xInsQues, yInsQues;
222 0 : for (u32 queIdx = 0; queIdx < xRankSize_ + yRankSize_; queIdx++) {
223 0 : if (queIdx < xRankSize_) {
224 0 : xInsQues.push_back(tempInsQues[queIdx]);
225 : } else {
226 0 : yInsQues.push_back(tempInsQues[queIdx]);
227 : }
228 : }
229 :
230 : // stage1
231 0 : if (rankSize_ > 1) {
232 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
233 : }
234 :
235 0 : std::vector<u64> stage1XDataInAddr, stage1YDataInAddr, stage1XDataOutAddr, stage1YDataOutAddr;
236 0 : for (u32 i = 0; i < rankSize_; i++) {
237 0 : stage1XDataInAddr.push_back(tempAlgParams.inputSliceStride * i + tempAlgParams.buffInfo.inBuffBaseOff);
238 0 : stage1YDataInAddr.push_back(tempAlgParams.inputSliceStride * i + tempAlgParams.buffInfo.inBuffBaseOff + xSize);
239 0 : stage1XDataOutAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i);
240 0 : stage1YDataOutAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i + xSize);
241 : }
242 :
243 0 : CHK_RET(RunMeshX(
244 : stage1XDataInAddr, stage1XDataOutAddr, xSize, BufferType::INPUT, BufferType::SCRATCH, DmaMode::PUT, xInsQues,
245 : tempLinks));
246 0 : CHK_RET(RunMeshY(
247 : stage1YDataInAddr, stage1YDataOutAddr, ySize, BufferType::INPUT, BufferType::SCRATCH, DmaMode::PUT, yInsQues,
248 : tempLinks));
249 :
250 0 : if (rankSize_ > 1) {
251 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
252 : }
253 :
254 : // stage2
255 0 : if (rankSize_ > 1) {
256 0 : CHK_RET(PreSyncInterQueues(tempInsQues));
257 : }
258 :
259 0 : std::vector<u64> stage2XDataInAddr, stage2YDataInAddr, stage2XDataOutAddr, stage2YDataOutAddr;
260 0 : for (u32 i = 0; i < rankSize_; i++) {
261 0 : stage2XDataInAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i);
262 0 : stage2YDataInAddr.push_back(tempAlgParams.buffInfo.scratchBuffBaseOff + tempAlgParams.sliceSize * i + xSize);
263 0 : stage2XDataOutAddr.push_back(tempAlgParams.outputSliceStride * i + tempAlgParams.buffInfo.outBuffBaseOff);
264 0 : stage2YDataOutAddr.push_back(
265 0 : tempAlgParams.outputSliceStride * i + tempAlgParams.buffInfo.outBuffBaseOff + xSize);
266 : }
267 :
268 0 : BufferType outType = !tempFuncs.isBottom ? BufferType::SCRATCH : BufferType::OUTPUT;
269 0 : CHK_RET(RunMeshY(
270 : stage2XDataInAddr, stage2XDataOutAddr, xSize, BufferType::SCRATCH, outType, DmaMode::GET, yInsQues, tempLinks));
271 0 : CHK_RET(RunMeshX(
272 : stage2YDataInAddr, stage2YDataOutAddr, ySize, BufferType::SCRATCH, outType, DmaMode::GET, xInsQues, tempLinks));
273 :
274 0 : if (rankSize_ > 1) {
275 0 : CHK_RET(PostSyncInterQueues(tempInsQues));
276 : }
277 :
278 0 : HCCL_INFO("[InsTempAlltoAllMesh2D] Run algorithm end: rank[%d]", myRank_);
279 :
280 0 : return HcclResult::HCCL_SUCCESS;
281 0 : }
282 :
283 : } // namespace Hccl
|