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