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_all_reduce_mesh_1D_two_shot.h"
12 :
13 : #include "log.h"
14 : #include "alg_data_trans_wrapper.h"
15 :
16 : namespace Hccl {
17 0 : InsTempAllReduceMesh1DTwoShot::InsTempAllReduceMesh1DTwoShot(
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 : InsTempAllReduceMesh1DTwoShot::~InsTempAllReduceMesh1DTwoShot() {}
24 :
25 : /*
26 : * Desc: 计算资源需求
27 : * return: tempResReq: 资源计算结果存储,包括notify信息,links信息等
28 : * return: HcclResult
29 : */
30 0 : HcclResult InsTempAllReduceMesh1DTwoShot::CalcRes(AlgTempResReq& tempResReq)
31 : {
32 0 : CHK_PRT_RET(
33 : CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq) != HcclResult::HCCL_SUCCESS,
34 : HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceMesh1DTwoShot] Rank [%d], resLinks calculation error!", myRank_),
35 : HcclResult::HCCL_E_INTERNAL);
36 0 : auto& linkReq = tempResReq.links;
37 0 : u32 pathNum = 0;
38 0 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
39 0 : auto remoteRank = resReqIter->first;
40 0 : if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_[remoteRank] == 0) {
41 0 : HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot] No path to remoteRank[%d]", remoteRank);
42 0 : return HcclResult::HCCL_E_INTERNAL;
43 : }
44 0 : if (pathNum == 0) {
45 0 : pathNum = rank2PathNumMap_[remoteRank];
46 0 : } else if (rank2PathNumMap_[remoteRank] != pathNum) {
47 0 : HCCL_ERROR(
48 : "[InsTempAllReduceMesh1DTwoShot] Inconsistency pathNum to remoteRanks, Previous consistent "
49 : "pathNum=[%u], mismatched "
50 : "remoteRank=[%d], pathNum=[%u]",
51 : pathNum, remoteRank, rank2PathNumMap_[remoteRank]);
52 0 : return HcclResult::HCCL_E_INTERNAL;
53 : }
54 0 : resReqIter->second = pathNum;
55 : }
56 :
57 : // 1D Mesh 需要的 que Num 为 ranksize
58 0 : tempResReq.queNum = tempVTopo_[0].size() * pathNum;
59 0 : HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] tempResReq.queNum = %u", tempResReq.queNum);
60 0 : tempResReq.streamNum = tempResReq.queNum;
61 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
62 :
63 0 : QId centerQ = 0;
64 0 : tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
65 0 : tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
66 :
67 0 : return HcclResult::HCCL_SUCCESS;
68 : }
69 :
70 : /*
71 : * Desc: 将数据按照rank切分为chuck 块,给后续的allreduce操作使用
72 : * param: dataSize: 待处理的输入数据大小
73 : * return: sliceInfoVec: 存储数据切分结果
74 : * return: HcclResult
75 : */
76 0 : HcclResult InsTempAllReduceMesh1DTwoShot::CalcSlice(const u64 dataSize, const u64 baseOff, RankSliceInfo& sliceInfoVec)
77 : {
78 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
79 0 : sliceInfoVec.resize(tempRankSize_, tmp);
80 :
81 0 : u64 unitAllignSize = DataTypeSizeGet(dataType_);
82 0 : u64 chunkSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
83 :
84 0 : u64 accumOff = 0;
85 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
86 0 : u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
87 0 : SliceInfo slice = {accumOff + baseOff, currChunkSize};
88 0 : sliceInfoVec[rankIdx][0] = slice;
89 0 : accumOff += currChunkSize;
90 : }
91 :
92 0 : CHK_PRT_RET(
93 : (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != baseOff + dataSize),
94 : HCCL_ERROR(
95 : "[InsAllReduceCombExecutor] chunkSize:[%llu], Rank:[%d], SliceInfo calculation error!", chunkSize, myRank_),
96 : HcclResult::HCCL_E_INTERNAL);
97 0 : return HcclResult::HCCL_SUCCESS;
98 0 : }
99 :
100 : /*
101 : * Desc: 返回当前rank能处理的数据量和scratch buffer之间的比例关系
102 : * param: input: 输入数据位置
103 : * param: output 输出数据位置
104 : */
105 0 : u32 InsTempAllReduceMesh1DTwoShot::CalcScratchMultiple(BufferType input, BufferType output) const
106 : {
107 : (void)input;
108 : (void)output;
109 0 : u32 multiple = 2;
110 0 : return multiple;
111 : }
112 :
113 : /*
114 : * Desc: GenExtIns 算子执行入口
115 : * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
116 : * param: tempAlgParams: 每个rank的数据切片信息
117 : * param: tempLinks: 当前rank通信链接信息
118 : * param: tempInsQues: 通信队列
119 : * return: HcclResult
120 : */
121 0 : HcclResult InsTempAllReduceMesh1DTwoShot::GenExtIns(
122 : const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
123 : std::vector<InsQuePtr>& tempInsQues)
124 : {
125 0 : HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] start.");
126 0 : opMode_ = tempFuncs.opMode;
127 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
128 :
129 0 : uint32_t linkNum = tempLinks.begin()->second.size();
130 0 : CHK_PRT_RET(
131 : tempInsQues.size() != tempVTopo_[0].size() * linkNum,
132 : HCCL_ERROR(
133 : "[InsTempAllReduceMesh1DTwoShot] RankSize [%lu], linkNum_:[%u], tempInsQues size:[%zu],requiredQue Error.",
134 : tempVTopo_[0].size(), linkNum, tempInsQues.size()),
135 : HcclResult::HCCL_E_INTERNAL);
136 :
137 0 : u64 dataSizePerVolume = DataTypeSizeGet(dataType_);
138 0 : CHK_PRT_RET(
139 : (tempRankSize_ * dataSizePerVolume) + tempAlgParams.sliceSize > tempAlgParams.buffInfo.scratchBuffSize,
140 : HCCL_ERROR(
141 : "[InsTempAllReduceMesh1DTwoShot]Rank [%d], Input size:[%llu], BfSize:[%llu] Insufficient buffer!", myRank_,
142 : tempAlgParams.sliceSize, tempAlgParams.buffInfo.scratchBuffSize),
143 : HcclResult::HCCL_E_INTERNAL);
144 :
145 0 : u64 inBuffSize = tempAlgParams.sliceSize;
146 0 : std::vector<RankSliceInfo> sliceInfoVecForAllLinks(linkNum);
147 0 : std::vector<float> dataSplitRate(linkNum);
148 0 : CHK_RET(CalcDataSplitRateForLinks(tempLinks.begin()->second, dataSplitRate));
149 0 : u64 typeSize = DataTypeSizeGet(dataType_);
150 0 : u64 dataCnt = inBuffSize / typeSize;
151 0 : std::vector<u64> inBuffSizeForLinks(linkNum);
152 0 : std::vector<u64> baseOff(linkNum);
153 0 : u64 offset = 0;
154 0 : for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
155 0 : if (linkIdx != linkNum - 1) {
156 0 : inBuffSizeForLinks[linkIdx]
157 0 : = static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize;
158 : } else {
159 0 : inBuffSizeForLinks[linkIdx] = inBuffSize - offset;
160 : }
161 : // 对每块输入都CalcSlice
162 0 : CHK_RET(CalcSlice(inBuffSizeForLinks[linkIdx], offset, sliceInfoVecForAllLinks[linkIdx]));
163 0 : offset += inBuffSizeForLinks[linkIdx];
164 : }
165 :
166 0 : std::vector<std::vector<InsQuePtr>> tempInsQuesVec;
167 0 : u32 queNumPerlink = tempInsQues.size() / linkNum;
168 0 : HCCL_INFO("tempInsQues.size()=%zu,queNumPerlink=%u", tempInsQues.size(), queNumPerlink);
169 0 : for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
170 0 : tempInsQuesVec.emplace_back(
171 0 : tempInsQues.begin() + linkIdx * queNumPerlink, tempInsQues.begin() + (linkIdx + 1) * queNumPerlink);
172 : }
173 0 : u32 mainQueIdx = 0;
174 0 : CHK_RET(PreSyncQues(tempInsQues, mainQueIdx));
175 0 : HCCL_INFO("[InsTempAllReduce1DMeshTwoShot][PreCopy] Rank [%d].", myRank_);
176 0 : for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
177 0 : CHK_RET(RunAllReduceScatter(
178 : sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQuesVec[linkIdx], tempAlgParams, linkIdx));
179 0 : CHK_RET(RunAllReduceAllgather(
180 : sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQuesVec[linkIdx], tempAlgParams, linkIdx));
181 : }
182 0 : HCCL_INFO("[InsTempAllReduce1DMeshTwoShot][PostCopy] Rank [%d].", myRank_);
183 : // 流间后同步,从流通知主流
184 0 : CHK_RET(PostSyncQues(tempInsQues, mainQueIdx));
185 :
186 0 : return HcclResult::HCCL_SUCCESS;
187 0 : }
188 :
189 : /*
190 : * Desc: 1D Mesh twoshot AllReduce: Scatter+reduce
191 : * param: sliceInfoVec: 每个rank的数据切片信息
192 : * param: tempLinks: 当前rank通信链接信息
193 : * param: tempInsQues: 通信队列
194 : * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
195 : * return: HcclResult
196 : */
197 0 : HcclResult InsTempAllReduceMesh1DTwoShot::RunAllReduceScatter(
198 : const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues,
199 : const TemplateDataParams& tempAlgParams, u32 linkIdx)
200 : {
201 0 : u64 inBuffBaseOff = tempAlgParams.buffInfo.inBuffBaseOff;
202 0 : u32 MyAlgRank = tempVirtRankMap_[myRank_];
203 : // scatter
204 0 : for (u32 rankId = 0; rankId < tempRankSize_; rankId++) {
205 : DataSlice ssrcSlice = DataSlice(
206 0 : tempAlgParams.buffInfo.inBuffType, sliceInfoVec[rankId][0].offset + inBuffBaseOff,
207 0 : sliceInfoVec[rankId][0].size);
208 : DataSlice sdestSlice = DataSlice(
209 : tempAlgParams.buffInfo.scratBuffType,
210 0 : MyAlgRank * sliceInfoVec[rankId][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff,
211 0 : sliceInfoVec[rankId][0].size);
212 0 : if (rankId == u32(MyAlgRank)) {
213 0 : if (sliceInfoVec[rankId][0].size != 0) {
214 : // 如果是本地rank,直接拷贝到scratch对应位置
215 0 : CHK_PRT_RET(
216 : LocalCopy(tempInsQues[rankId], ssrcSlice, sdestSlice),
217 : HCCL_ERROR(
218 : "[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce scatter LocalCopy failed"),
219 : HcclResult::HCCL_E_INTERNAL);
220 : }
221 : } else {
222 0 : const std::vector<LinkData>& linkSendRecv = tempLinks.at(GetRankFromMap(rankId));
223 : // 发送, 未过滤size为0的情况
224 0 : std::vector<DataSlice> sendSrcSlices{ssrcSlice};
225 0 : std::vector<DataSlice> sendDestSlices{sdestSlice};
226 :
227 : // 接收,未过滤size为0的情况
228 : DataSlice rsrcSlice = DataSlice(
229 0 : tempAlgParams.buffInfo.inBuffType, sliceInfoVec[MyAlgRank][0].offset + inBuffBaseOff,
230 0 : sliceInfoVec[MyAlgRank][0].size);
231 : DataSlice rdestSlice = DataSlice(
232 : tempAlgParams.buffInfo.scratBuffType,
233 0 : rankId * sliceInfoVec[MyAlgRank][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff,
234 0 : sliceInfoVec[MyAlgRank][0].size);
235 0 : std::vector<DataSlice> recvSrcSlices{rsrcSlice};
236 0 : std::vector<DataSlice> recvDestSlices{rdestSlice};
237 :
238 0 : TxRxLinks sendRecvLinks(linkSendRecv[linkIdx], linkSendRecv[linkIdx]);
239 0 : TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
240 :
241 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
242 0 : CHK_PRT_RET(
243 : SendRecv(sendRecvInfo, tempInsQues[rankId], 0, true, DmaMode::PUT),
244 : HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce scatter failed"),
245 : HcclResult::HCCL_E_INTERNAL);
246 0 : }
247 : }
248 : // 从流同步,等待所有并发的send和copy完成
249 0 : PostSyncInterQueues(tempInsQues);
250 :
251 : // local reduce
252 0 : if (sliceInfoVec[MyAlgRank][0].size != 0) {
253 : DataSlice ldestSlice = DataSlice(
254 0 : tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff,
255 0 : sliceInfoVec[MyAlgRank][0].size);
256 0 : for (u32 rankId = 1; rankId < tempRankSize_; rankId++) {
257 : DataSlice lsrcSlice = DataSlice(
258 : tempAlgParams.buffInfo.scratBuffType,
259 0 : rankId * sliceInfoVec[MyAlgRank][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff,
260 0 : sliceInfoVec[MyAlgRank][0].size);
261 : // 所有reduce操作在同一个insque中才能保序;
262 0 : CHK_PRT_RET(
263 : LocalReduce(tempInsQues[0], lsrcSlice, ldestSlice, dataType_, redOp_),
264 : HCCL_ERROR(
265 : "[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce reduce LocalReduce failed"),
266 : HcclResult::HCCL_E_INTERNAL);
267 : }
268 : }
269 :
270 0 : return HcclResult::HCCL_SUCCESS;
271 : }
272 :
273 : /*
274 : * Desc: 1D Mesh twoshot AllReduce: Allgather
275 : * param: sliceInfoVec: 每个rank的数据切片信息
276 : * param: tempLinks: 当前rank通信链接信息
277 : * param: tempInsQues: 通信队列
278 : * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
279 : * return: HcclResult
280 : */
281 0 : HcclResult InsTempAllReduceMesh1DTwoShot::RunAllReduceAllgather(
282 : const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues,
283 : const TemplateDataParams& tempAlgParams, u32 linkIdx)
284 : {
285 0 : u64 outBuffBaseOff = tempAlgParams.buffInfo.outBuffBaseOff;
286 : // sync:前同步
287 0 : PreSyncInterQueues(tempInsQues);
288 0 : u32 MyAlgRank = tempVirtRankMap_[myRank_];
289 : // allgather
290 0 : for (u32 rankId = 0; rankId < tempRankSize_; rankId++) {
291 : DataSlice rsrcSlice = DataSlice(
292 0 : tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff,
293 0 : sliceInfoVec[rankId][0].size);
294 : DataSlice rdestSlice = DataSlice(
295 0 : tempAlgParams.buffInfo.outBuffType, sliceInfoVec[rankId][0].offset + outBuffBaseOff,
296 0 : sliceInfoVec[rankId][0].size);
297 0 : if (u32(MyAlgRank) == rankId) {
298 0 : if (sliceInfoVec[rankId][0].size != 0) {
299 : // copy本端计算的结果到user output
300 0 : CHK_PRT_RET(
301 : LocalCopy(tempInsQues[rankId], rsrcSlice, rdestSlice),
302 : HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceAllgather] RunAllReduce AllGather "
303 : "LocalCopy failed"),
304 : HcclResult::HCCL_E_INTERNAL);
305 : }
306 : } else {
307 0 : const std::vector<LinkData>& linkSendRecv = tempLinks.at(GetRankFromMap(rankId));
308 : // 接收, 未过滤size为0的情况
309 0 : std::vector<DataSlice> recvSrcSlices{rsrcSlice};
310 0 : std::vector<DataSlice> recvDestSlices{rdestSlice};
311 :
312 : // 发送,未过滤size为0的情况
313 : DataSlice ssrcSlice = DataSlice(
314 0 : tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff,
315 0 : sliceInfoVec[MyAlgRank][0].size);
316 : DataSlice sdestSlice = DataSlice(
317 0 : tempAlgParams.buffInfo.outBuffType, sliceInfoVec[MyAlgRank][0].offset + outBuffBaseOff,
318 0 : sliceInfoVec[MyAlgRank][0].size);
319 0 : std::vector<DataSlice> sendSrcSlices{ssrcSlice};
320 0 : std::vector<DataSlice> sendDestSlices{sdestSlice};
321 :
322 0 : TxRxLinks sendRecvLinks(linkSendRecv[linkIdx], linkSendRecv[linkIdx]);
323 0 : TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
324 :
325 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
326 0 : CHK_PRT_RET(
327 : SendRecv(sendRecvInfo, tempInsQues[rankId], 0, true, DmaMode::GET),
328 : HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceAllgather] RunAllReduce AllGather failed"),
329 : HcclResult::HCCL_E_INTERNAL);
330 0 : }
331 : }
332 0 : return HcclResult::HCCL_SUCCESS;
333 : }
334 :
335 0 : RankId InsTempAllReduceMesh1DTwoShot::GetRankFromMap(const u32 rankIdx)
336 : {
337 0 : RankId rank = -1;
338 0 : HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] GetRankFromMap");
339 0 : for (auto& pair : tempVirtRankMap_) {
340 0 : if (pair.second == rankIdx) {
341 0 : rank = pair.first;
342 0 : break;
343 : }
344 : }
345 0 : return rank;
346 : }
347 : } // namespace Hccl
|