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 :
13 : #include "alg_data_trans_wrapper.h"
14 : #include "ins_alg_template/ins_temp_broadcast_mesh_1D_two_shot.h"
15 :
16 : namespace Hccl {
17 0 : InsTempBroadcastMesh1DTwoShot::InsTempBroadcastMesh1DTwoShot(
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 : InsTempBroadcastMesh1DTwoShot::~InsTempBroadcastMesh1DTwoShot() {}
24 :
25 0 : HcclResult InsTempBroadcastMesh1DTwoShot::CalcRes(AlgTempResReq& tempResReq)
26 : {
27 0 : tempResReq.queNum = (tempVTopo_[0].size() > 1) ? (tempVTopo_[0].size() - 1) : 1;
28 0 : tempResReq.streamNum = tempResReq.queNum;
29 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
30 :
31 0 : QId centerQ = 0;
32 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
33 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
34 :
35 0 : CHK_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq));
36 0 : HCCL_DEBUG(
37 : "[InsTempBroadcastMesh1DTwoShot] Rank[%d], VtopoSize[%lu], requiredQue Num [%u].", myRank_,
38 : tempVTopo_[0].size(), tempResReq.queNum);
39 :
40 0 : return HcclResult::HCCL_SUCCESS;
41 : }
42 :
43 0 : u32 InsTempBroadcastMesh1DTwoShot::CalcScratchMultiple(BufferType inBuffType, BufferType outBuffType)
44 : {
45 : (void)inBuffType;
46 : (void)outBuffType;
47 0 : if (op_.opMode == OpMode::OPBASE) {
48 0 : return 1;
49 : } else {
50 0 : return 0;
51 : }
52 : }
53 :
54 : // 按照mesh的方式计算SliceInfo,例如N张卡,就是N份slice
55 0 : HcclResult InsTempBroadcastMesh1DTwoShot::CalcDataSliceInfo(const u64 dataSize, RankSliceInfo& sliceInfoVec)
56 : {
57 : // 一般情况下,mesh的temp是单级的
58 : u64 unitAllignSize;
59 0 : AllignInfo allignInfo = {false, 0, dataType_};
60 0 : CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
61 0 : sliceInfoVec.resize(tempRankSize_);
62 :
63 0 : u64 chunkSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
64 :
65 0 : u64 accumOff = 0;
66 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
67 0 : u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
68 0 : SliceInfo slice = {accumOff, currChunkSize};
69 0 : sliceInfoVec[rankIdx].push_back(slice);
70 0 : accumOff += currChunkSize;
71 : }
72 :
73 0 : CHK_PRT_RET(
74 : (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != dataSize),
75 : HCCL_ERROR("[InsTempBroadcastMesh1DTwoShot] Rank [%d], SliceInfo calculation error!", myRank_),
76 : HcclResult::HCCL_E_INTERNAL);
77 :
78 0 : return HcclResult::HCCL_SUCCESS;
79 : }
80 :
81 : // 计算scatter的通信rank集合
82 : HcclResult
83 0 : InsTempBroadcastMesh1DTwoShot::CalcCommRankSetforScatter(const u32 groupRankSize, std::vector<u32>& commRanks) const
84 : {
85 : (void)groupRankSize;
86 0 : commRanks.clear();
87 :
88 0 : if (u32(myRank_) != root_) {
89 0 : commRanks.emplace_back(root_);
90 0 : return HcclResult::HCCL_SUCCESS;
91 : }
92 :
93 0 : for (auto& rankIter : tempVirtRankMap_) {
94 0 : if (u32(myRank_) != u32(rankIter.first)) {
95 0 : commRanks.emplace_back(u32(rankIter.first));
96 : }
97 : }
98 :
99 0 : return HcclResult::HCCL_SUCCESS;
100 : }
101 :
102 : // 计算allgather的通信rank集合
103 : HcclResult
104 0 : InsTempBroadcastMesh1DTwoShot::CalcCommRankSetforAllGather(const u32 groupRankSize, std::vector<u32>& commRanks) const
105 : {
106 : (void)groupRankSize;
107 0 : commRanks.clear();
108 :
109 0 : for (auto& rankIter : tempVirtRankMap_) {
110 0 : if (u32(myRank_) != u32(rankIter.first) && root_ != u32(rankIter.first)) {
111 0 : commRanks.emplace_back(u32(rankIter.first));
112 : }
113 : }
114 :
115 0 : return HcclResult::HCCL_SUCCESS;
116 : }
117 :
118 0 : HcclResult InsTempBroadcastMesh1DTwoShot::RootSendData(
119 : const u64 memOffset, const s32 remoteRank, const TemplateDataParams& tempAlgParams, const InsQuePtr& queue,
120 : const LinkData& link, const RankSliceInfo& sliceInfoVec) const
121 : {
122 0 : u32 myRankIdx = tempVirtRankMap_.at(myRank_);
123 0 : u32 remoteRankIdx = tempVirtRankMap_.at(remoteRank);
124 :
125 : // root执行常规scatter发送,将remoteRank的数据分片发送至remoteRank的buf中
126 0 : u64 sendSrcOffset0 = sliceInfoVec[remoteRankIdx][0].offset + memOffset;
127 0 : u64 sendDstOffset0 = sliceInfoVec[remoteRankIdx][0].offset;
128 0 : if (dstBufferType_ == BufferType::SCRATCH) {
129 0 : sendDstOffset0 += tempAlgParams.buffInfo.scratchBuffBaseOff;
130 : } else {
131 0 : sendDstOffset0 += tempAlgParams.buffInfo.outBuffBaseOff;
132 : }
133 :
134 0 : DataSlice sendSrcSlice0 = DataSlice(BufferType::INPUT, sendSrcOffset0, sliceInfoVec[remoteRankIdx][0].size);
135 0 : DataSlice sendDstSlice0 = DataSlice(dstBufferType_, sendDstOffset0, sliceInfoVec[remoteRankIdx][0].size);
136 :
137 0 : std::vector<DataSlice> sendSrcSliceVec0 = {sendSrcSlice0};
138 0 : std::vector<DataSlice> sendDstSliceVec0 = {sendDstSlice0};
139 0 : SlicesList sendDataSlice0(sendSrcSliceVec0, sendDstSliceVec0);
140 0 : DataInfo sendDataInfo0(link, sendDataSlice0);
141 0 : CHK_RET(Send(sendDataInfo0, queue, 0, true, DmaMode::PUT));
142 :
143 : // root将自己数据分片发送至对端
144 0 : u64 sendSrcOffset1 = sliceInfoVec[myRankIdx][0].offset + memOffset;
145 0 : u64 sendDstOffset1 = sliceInfoVec[myRankIdx][0].offset;
146 0 : if (dstBufferType_ == BufferType::SCRATCH) {
147 0 : sendDstOffset1 += tempAlgParams.buffInfo.scratchBuffBaseOff;
148 : } else {
149 0 : sendDstOffset1 += tempAlgParams.buffInfo.outBuffBaseOff;
150 : }
151 :
152 0 : DataSlice sendSrcSlice1 = DataSlice(BufferType::INPUT, sendSrcOffset1, sliceInfoVec[myRankIdx][0].size);
153 0 : DataSlice sendDstSlice1 = DataSlice(dstBufferType_, sendDstOffset1, sliceInfoVec[myRankIdx][0].size);
154 :
155 0 : std::vector<DataSlice> sendSrcSliceVec1 = {sendSrcSlice1};
156 0 : std::vector<DataSlice> sendDstSliceVec1 = {sendDstSlice1};
157 0 : SlicesList sendDataSlice1(sendSrcSliceVec1, sendDstSliceVec1);
158 0 : DataInfo sendDataInfo1(link, sendDataSlice1);
159 0 : CHK_RET(Send(sendDataInfo1, queue, 0, true, DmaMode::PUT));
160 :
161 0 : return HcclResult::HCCL_SUCCESS;
162 0 : }
163 :
164 0 : HcclResult InsTempBroadcastMesh1DTwoShot::RankRecvData(
165 : const u64 memOffset, const TemplateDataParams& tempAlgParams, const InsQuePtr& queue, const LinkData& link,
166 : const RankSliceInfo& sliceInfoVec) const
167 : {
168 0 : u32 myRankIdx = tempVirtRankMap_.at(myRank_);
169 0 : u32 rootIdx = tempVirtRankMap_.at(root_);
170 :
171 : // 非root执行常规scatter接收,从root接收本rank的数据分片
172 0 : u64 sendSrcOffset0 = sliceInfoVec[myRankIdx][0].offset + memOffset;
173 0 : u64 sendDstOffset0 = sliceInfoVec[myRankIdx][0].offset;
174 0 : if (dstBufferType_ == BufferType::SCRATCH) {
175 0 : sendDstOffset0 += tempAlgParams.buffInfo.scratchBuffBaseOff;
176 : } else {
177 0 : sendDstOffset0 += tempAlgParams.buffInfo.outBuffBaseOff;
178 : }
179 :
180 0 : DataSlice recvSrcSlice0 = DataSlice(BufferType::INPUT, sendSrcOffset0, sliceInfoVec[myRankIdx][0].size);
181 0 : DataSlice recvDstSlice0 = DataSlice(dstBufferType_, sendDstOffset0, sliceInfoVec[myRankIdx][0].size);
182 :
183 0 : std::vector<DataSlice> recvSrcSliceVec0 = {recvSrcSlice0};
184 0 : std::vector<DataSlice> recvDstSliceVec0 = {recvDstSlice0};
185 0 : SlicesList recvDataSlice0(recvSrcSliceVec0, recvDstSliceVec0);
186 0 : DataInfo recvDataInfo0(link, recvDataSlice0);
187 0 : CHK_RET(Recv(recvDataInfo0, queue, 0, true, DmaMode::PUT));
188 :
189 : // 非root接收root的数据分片
190 0 : u64 sendSrcOffset1 = sliceInfoVec[rootIdx][0].offset + memOffset;
191 0 : u64 sendDstOffset1 = sliceInfoVec[rootIdx][0].offset;
192 0 : if (dstBufferType_ == BufferType::SCRATCH) {
193 0 : sendDstOffset1 += tempAlgParams.buffInfo.scratchBuffBaseOff;
194 : } else {
195 0 : sendDstOffset1 += tempAlgParams.buffInfo.outBuffBaseOff;
196 : }
197 :
198 0 : DataSlice recvSrcSlice1 = DataSlice(BufferType::INPUT, sendSrcOffset1, sliceInfoVec[rootIdx][0].size);
199 0 : DataSlice recvDstSlice1 = DataSlice(dstBufferType_, sendDstOffset1, sliceInfoVec[rootIdx][0].size);
200 :
201 0 : std::vector<DataSlice> recvSrcSliceVec1 = {recvSrcSlice1};
202 0 : std::vector<DataSlice> recvDstSliceVec1 = {recvDstSlice1};
203 0 : SlicesList recvDataSlice1(recvSrcSliceVec1, recvDstSliceVec1);
204 0 : DataInfo recvDataInfo1(link, recvDataSlice1);
205 0 : CHK_RET(Recv(recvDataInfo1, queue, 0, true, DmaMode::PUT));
206 :
207 0 : return HcclResult::HCCL_SUCCESS;
208 0 : }
209 :
210 0 : HcclResult InsTempBroadcastMesh1DTwoShot::RunScatter(
211 : const std::vector<u32>& commRanks, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
212 : std::vector<InsQuePtr>& queues, const RankSliceInfo& sliceInfoVec) const
213 : {
214 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot: Scatter entry.");
215 :
216 : // 主从流同步
217 0 : if (commRanks.size() > 1) {
218 0 : CHK_RET(PreSyncInterQueues(queues));
219 : }
220 :
221 0 : u64 memOffset = tempAlgParams.buffInfo.inBuffBaseOff;
222 :
223 : // DMA消减,直接从root的inputbuf传输数据至对端buf
224 0 : for (u32 i = 0; i < commRanks.size(); i++) {
225 0 : s32 remoteRank = static_cast<s32>(commRanks[i]);
226 0 : InsQuePtr queue = queues[i];
227 0 : LinkData link = tempLinks.at(remoteRank)[0];
228 0 : if (u32(myRank_) == root_) {
229 : // root只发不收
230 0 : CHK_RET(RootSendData(memOffset, remoteRank, tempAlgParams, queue, link, sliceInfoVec));
231 : } else {
232 : // 非root只收不发
233 0 : CHK_RET(RankRecvData(memOffset, tempAlgParams, queue, link, sliceInfoVec));
234 : }
235 0 : }
236 :
237 : // 主从流同步
238 0 : if (commRanks.size() > 1) {
239 0 : CHK_RET(PostSyncInterQueues(queues));
240 : }
241 :
242 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot: Scatter finish.");
243 :
244 0 : return HcclResult::HCCL_SUCCESS;
245 : }
246 :
247 0 : HcclResult InsTempBroadcastMesh1DTwoShot::RunAllGather(
248 : const std::vector<u32>& commRanks, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
249 : std::vector<InsQuePtr>& queues, const RankSliceInfo& sliceInfoVec) const
250 : {
251 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot: AllGather entry.");
252 :
253 0 : if (commRanks.size() > 1) {
254 0 : CHK_RET(PreSyncInterQueues(queues));
255 : }
256 :
257 0 : for (u32 i = 0; i < commRanks.size(); i++) {
258 0 : s32 remoteRank = static_cast<s32>(commRanks[i]);
259 0 : InsQuePtr queue = queues[i];
260 0 : LinkData link = tempLinks.at(remoteRank)[0];
261 :
262 0 : u32 myRankIdx = tempVirtRankMap_.at(myRank_);
263 0 : u32 remoteRankIdx = tempVirtRankMap_.at(remoteRank);
264 :
265 0 : u64 sendSrcOffset = sliceInfoVec[myRankIdx][0].offset;
266 0 : u64 sendDstOffset = sliceInfoVec[myRankIdx][0].offset;
267 0 : u64 recvSrcOffset = sliceInfoVec[remoteRankIdx][0].offset;
268 0 : u64 recvDstOffset = sliceInfoVec[remoteRankIdx][0].offset;
269 :
270 0 : if (srcBufferType_ == BufferType::SCRATCH) {
271 0 : sendSrcOffset += tempAlgParams.buffInfo.scratchBuffBaseOff;
272 0 : recvSrcOffset += tempAlgParams.buffInfo.scratchBuffBaseOff;
273 : } else {
274 0 : sendSrcOffset += tempAlgParams.buffInfo.inBuffBaseOff;
275 0 : recvSrcOffset += tempAlgParams.buffInfo.inBuffBaseOff;
276 : }
277 :
278 0 : if (dstBufferType_ == BufferType::SCRATCH) {
279 0 : sendDstOffset += tempAlgParams.buffInfo.scratchBuffBaseOff;
280 0 : recvDstOffset += tempAlgParams.buffInfo.scratchBuffBaseOff;
281 : } else {
282 0 : sendDstOffset += tempAlgParams.buffInfo.outBuffBaseOff;
283 0 : recvDstOffset += tempAlgParams.buffInfo.outBuffBaseOff;
284 : }
285 :
286 0 : DataSlice sendSrcSlice = DataSlice(srcBufferType_, sendSrcOffset, sliceInfoVec[myRankIdx][0].size);
287 0 : DataSlice sendDstSlice = DataSlice(dstBufferType_, sendDstOffset, sliceInfoVec[myRankIdx][0].size);
288 0 : std::vector<DataSlice> sendSrcSliceVec = {sendSrcSlice};
289 0 : std::vector<DataSlice> sendDstSliceVec = {sendDstSlice};
290 0 : SlicesList sendDataSlice(sendSrcSliceVec, sendDstSliceVec);
291 :
292 0 : DataSlice recvSrcSlice = DataSlice(srcBufferType_, recvSrcOffset, sliceInfoVec[remoteRankIdx][0].size);
293 0 : DataSlice recvDstSlice = DataSlice(dstBufferType_, recvDstOffset, sliceInfoVec[remoteRankIdx][0].size);
294 0 : std::vector<DataSlice> recvSrcSliceVec = {recvSrcSlice};
295 0 : std::vector<DataSlice> recvDstSliceVec = {recvDstSlice};
296 0 : SlicesList recvDataSlice(recvSrcSliceVec, recvDstSliceVec);
297 :
298 0 : TxRxSlicesList sendRecvSlice(sendDataSlice, recvDataSlice);
299 0 : TxRxLinks sendRecvLinks(link, link);
300 :
301 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlice);
302 0 : CHK_RET(SendRecv(sendRecvInfo, queue, 0, true, DmaMode::PUT));
303 0 : }
304 :
305 0 : if (commRanks.size() > 1) {
306 0 : CHK_RET(PostSyncInterQueues(queues));
307 : }
308 :
309 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot: AllGather finish.");
310 :
311 0 : return HcclResult::HCCL_SUCCESS;
312 : }
313 :
314 0 : HcclResult InsTempBroadcastMesh1DTwoShot::PostCopy(
315 : const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues) const
316 : {
317 0 : u64 inOffset = tempAlgParams.buffInfo.scratchBuffBaseOff;
318 :
319 0 : DataSlice usrInSlice = DataSlice(BufferType::SCRATCH, inOffset, tempAlgParams.sliceSize);
320 : DataSlice usrOutSlice
321 0 : = DataSlice(BufferType::INPUT, tempAlgParams.buffInfo.outBuffBaseOff, tempAlgParams.sliceSize);
322 :
323 0 : HCCL_INFO(
324 : "PostCopy usrInSlice: %s, usrOutSlice: %s", usrInSlice.Describe().c_str(), usrOutSlice.Describe().c_str());
325 :
326 0 : CHK_RET(LocalCopy(tempInsQues[0], usrInSlice, usrOutSlice));
327 :
328 0 : return HcclResult::HCCL_SUCCESS;
329 : }
330 :
331 0 : HcclResult InsTempBroadcastMesh1DTwoShot::GenExtIns(
332 : const TempFuncs& tempFuncs, const TemplateDataParams& templateDataParams, const ResLinks& tempLinks,
333 : std::vector<InsQuePtr>& tempInsQues)
334 : {
335 0 : opMode_ = tempFuncs.opMode;
336 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
337 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot entry.");
338 :
339 0 : if (opMode_ == OpMode::OPBASE) {
340 0 : srcBufferType_ = BufferType::SCRATCH;
341 0 : dstBufferType_ = BufferType::SCRATCH;
342 : }
343 :
344 0 : RankSliceInfo sliceInfoVec{};
345 0 : CHK_RET(CalcDataSliceInfo(templateDataParams.sliceSize, sliceInfoVec));
346 :
347 0 : queNum_ = tempVTopo_[0].size() - 1;
348 0 : CHK_PRT_RET(
349 : queNum_ != tempInsQues.size(),
350 : HCCL_ERROR("[CollAlgFactory] [InsTempBroadcastMesh1DTwoShot] Rank [%d], requiredQue Error.", myRank_),
351 : HcclResult::HCCL_E_INTERNAL);
352 :
353 0 : HCCL_INFO(
354 : "[InsTempBroadcastMesh1DTwoShot Run]RankID:[%d], root:[%u], isForepart:[%d], isBottom:[%d]", myRank_, root_,
355 : tempFuncs.isForepart, tempFuncs.isBottom);
356 :
357 0 : std::vector<u32> scatterCommRanks;
358 0 : CHK_RET(CalcCommRankSetforScatter(tempRankSize_, scatterCommRanks)); // 计算scatter步骤的通信对象
359 0 : CHK_RET(RunScatter(scatterCommRanks, templateDataParams, tempLinks, tempInsQues, sliceInfoVec)); // 运行scatter步骤
360 :
361 0 : if (u32(myRank_) != root_) {
362 0 : std::vector<u32> allgatherCommRanks;
363 0 : CHK_RET(CalcCommRankSetforAllGather(tempRankSize_, allgatherCommRanks)); // 计算allgather步骤的通信对象
364 0 : CHK_RET(RunAllGather(
365 : allgatherCommRanks, templateDataParams, tempLinks, tempInsQues, sliceInfoVec)); // 运行allgather步骤
366 0 : }
367 :
368 : // 单算子模式
369 0 : if (opMode_ == OpMode::OPBASE && (u32(myRank_) != root_)) {
370 0 : CHK_RET(PostCopy(templateDataParams, tempInsQues));
371 : }
372 :
373 0 : HCCL_INFO("[InsTempBroadcastMesh1DTwoShot] BroadcastMesh1DTwoShot finish.");
374 :
375 0 : return HcclResult::HCCL_SUCCESS;
376 0 : }
377 :
378 : } // namespace Hccl
|