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