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