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 : #include "alg_data_trans_wrapper.h"
13 : #include "ins_temp_all_reduce_nhr.h"
14 :
15 : namespace Hccl {
16 0 : InsTempAllReduceNHR::InsTempAllReduceNHR(
17 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
18 0 : const std::map<RankId, u32>& tempVirtRankMap)
19 0 : : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 0 : {}
21 :
22 0 : InsTempAllReduceNHR::~InsTempAllReduceNHR() {}
23 :
24 0 : HcclResult InsTempAllReduceNHR::CalcRes(AlgTempResReq& tempResReq)
25 : {
26 : // NHR 需要的 que Num 为 1
27 0 : CHK_PRT_RET(
28 : CalcResLinksNHR(myRank_, tempRankSize_, tempVTopo_, tempResReq) != HcclResult::HCCL_SUCCESS,
29 : HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceNHR] Rank [%d], resLinks calculation error!", myRank_),
30 : HcclResult::HCCL_E_INTERNAL);
31 0 : auto& linkReq = tempResReq.links;
32 0 : u32 pathNum = 0;
33 0 : for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
34 0 : auto remoteRank = resReqIter->first;
35 0 : if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_[remoteRank] == 0) {
36 0 : HCCL_ERROR("[InsTempAllReduceNHR] No path to remoteRank[%d]", remoteRank);
37 0 : return HcclResult::HCCL_E_INTERNAL;
38 : }
39 0 : if (pathNum == 0) {
40 0 : pathNum = rank2PathNumMap_[remoteRank];
41 0 : } else if (rank2PathNumMap_[remoteRank] != pathNum) {
42 0 : HCCL_ERROR(
43 : "[InsTempAllReduceNHR] Inconsistency pathNum to remoteRanks, Previous consistent pathNum=[%u], "
44 : "mismatched "
45 : "remoteRank=[%d], pathNum=[%u]",
46 : pathNum, remoteRank, rank2PathNumMap_[remoteRank]);
47 0 : return HcclResult::HCCL_E_INTERNAL;
48 : }
49 0 : resReqIter->second = pathNum;
50 : }
51 0 : tempResReq.queNum = 1 * pathNum;
52 0 : HCCL_INFO("[InsTempAllReduceNHR] tempResReq.queNum = %u", tempResReq.queNum);
53 0 : tempResReq.streamNum = tempResReq.queNum;
54 0 : tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
55 :
56 0 : return HcclResult::HCCL_SUCCESS;
57 : }
58 :
59 : /*
60 : * Desc: 将数据按照rank切分为chuck 块,给后续的allreduce操作使用
61 : * param: dataSize: 待处理的输入数据大小
62 : * return: sliceInfoVec: 存储数据切分结果
63 : * return: HcclResult
64 : */
65 0 : HcclResult InsTempAllReduceNHR::CalcSlice(const u64 dataSize, const u64 baseOff, RankSliceInfo& sliceInfoVec)
66 : {
67 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
68 0 : sliceInfoVec.resize(tempRankSize_, tmp);
69 :
70 0 : u64 unitAllignSize = DataTypeSizeGet(dataType_);
71 0 : u64 chunkSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
72 :
73 0 : u64 accumOff = 0;
74 0 : for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
75 0 : u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
76 0 : SliceInfo slice = {accumOff + baseOff, currChunkSize};
77 0 : sliceInfoVec[rankIdx][0] = slice;
78 0 : accumOff += currChunkSize;
79 : }
80 :
81 0 : CHK_PRT_RET(
82 : (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != baseOff + dataSize),
83 : HCCL_ERROR(
84 : "[InsTempAllReduceNHR] chunkSize:[%llu], Rank:[%d], SliceInfo calculation error!", chunkSize, myRank_),
85 : HcclResult::HCCL_E_INTERNAL);
86 0 : return HcclResult::HCCL_SUCCESS;
87 0 : }
88 :
89 : /*
90 : * Desc: 返回当前rank能处理的数据量和scratch buffer之间的比例关系
91 : * param: input: 输入数据位置
92 : * param: output 输出数据位置
93 : */
94 0 : u32 InsTempAllReduceNHR::CalcScratchMultiple(BufferType input, BufferType output)
95 : {
96 : (void)input;
97 : (void)output;
98 : // 单算子模式,cclBuffer和usrIn一样大,图模式,不需要cclBuffer
99 0 : u32 multiple = 0;
100 0 : if (op_.opMode == OpMode::OPBASE) {
101 0 : multiple = 1;
102 : }
103 :
104 0 : return multiple;
105 : }
106 :
107 0 : HcclResult InsTempAllReduceNHR::GenExtIns(
108 : const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
109 : std::vector<InsQuePtr>& tempInsQues)
110 : {
111 0 : HCCL_INFO("[InsTempAllReduceNHR][GenExtIns] AllReduceNHR begin: rank[%d] start", myRank_);
112 0 : if (IsPcieLink(tempLinks)) {
113 0 : dmaMode_ = DmaMode::GET;
114 : }
115 :
116 0 : opMode_ = tempFuncs.opMode;
117 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
118 :
119 0 : uint32_t linkNum = tempLinks.begin()->second.size();
120 : // 流的数量不能少于linkNum
121 0 : CHK_PRT_RET(
122 : linkNum > tempInsQues.size(),
123 : HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceNHR] Rank [%d], requiredQue Error.", myRank_),
124 : HcclResult::HCCL_E_INTERNAL);
125 :
126 0 : std::vector<float> dataSplitRate(linkNum);
127 0 : CHK_RET(CalcDataSplitRateForLinks(tempLinks.begin()->second, dataSplitRate));
128 :
129 0 : RankSliceInfo sliceInfoVec;
130 0 : CHK_RET(CalcSlice(tempAlgParams.sliceSize, 0, sliceInfoVec));
131 :
132 : // 将一个RankSliceInfo,拆分成linkNum 个RankSliceInfo
133 0 : u64 typeSize = DataTypeSizeGet(dataType_);
134 0 : std::vector<RankSliceInfo> sliceInfoVecForAllLinks(linkNum);
135 0 : for (auto sliceInfoPerRank : sliceInfoVec) {
136 0 : std::vector<std::vector<SliceInfo>> sliceInfoPerRankForAllLinks(linkNum);
137 0 : for (auto sliceInfo : sliceInfoPerRank) {
138 0 : u64 size = sliceInfo.size;
139 0 : u64 offset = sliceInfo.offset;
140 0 : u64 AccSize = 0;
141 0 : u64 dataCnt = size / typeSize;
142 0 : vector<SliceInfo> sliceInfoForAllLinks(linkNum);
143 0 : for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
144 0 : if (linkIdx != linkNum - 1) {
145 0 : sliceInfoForAllLinks[linkIdx].size
146 0 : = static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize;
147 : } else {
148 0 : sliceInfoForAllLinks[linkIdx].size = size - AccSize;
149 : }
150 0 : sliceInfoForAllLinks[linkIdx].offset = offset + AccSize;
151 0 : AccSize += sliceInfoForAllLinks[linkIdx].size;
152 : }
153 0 : for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
154 0 : sliceInfoPerRankForAllLinks[linkIdx].emplace_back(sliceInfoForAllLinks[linkIdx]);
155 : }
156 0 : }
157 0 : for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
158 0 : sliceInfoVecForAllLinks[linkIdx].emplace_back(sliceInfoPerRankForAllLinks[linkIdx]);
159 : }
160 0 : }
161 :
162 : // 预拷贝
163 0 : CHK_RET(PreCopy(tempAlgParams, tempInsQues));
164 :
165 0 : u32 mainQueIdx = 0;
166 : // 流间前同步,主流通知从流,只有一个流则不做任何事
167 0 : CHK_RET(PreSyncQues(tempInsQues, mainQueIdx));
168 :
169 : // 主从流执行nhr
170 : // 待修改数据切分方式
171 0 : for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
172 0 : CHK_RET(RunReduceScatter(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQues, linkIdx));
173 0 : CHK_RET(PrepareDataForAllGather(sliceInfoVecForAllLinks[linkIdx], tempInsQues, linkIdx));
174 0 : CHK_RET(RunAllGather(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQues, linkIdx));
175 : }
176 : // 流间后同步,从流通知主流
177 0 : CHK_RET(PostSyncQues(tempInsQues, mainQueIdx));
178 : // 结果拷贝
179 0 : CHK_RET(PostCopy(tempAlgParams, tempInsQues));
180 0 : HCCL_INFO("[InsTempAllReduceNHR][GenExtIns] AllReduceNHR finished: rank[%d] end", myRank_);
181 0 : return HcclResult::HCCL_SUCCESS;
182 0 : }
183 :
184 0 : HcclResult InsTempAllReduceNHR::PreCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
185 : {
186 : // 单算子模式,需要先将数据拷贝到cclBuffer
187 0 : if (opMode_ == OpMode::OPBASE) {
188 0 : nhrInBuffType_ = BufferType::SCRATCH;
189 0 : nhrInBuffBaseOff_ = tempAlgParams.buffInfo.inBuffBaseOff;
190 :
191 0 : if (tempAlgParams.buffInfo.inBuffType != BufferType::SCRATCH) {
192 0 : HCCL_INFO("[InsTempAllReduceNHR][PreCopy] Opbase copy from userIn to scratchBuffer");
193 : DataSlice usrInSlices = DataSlice(
194 0 : tempAlgParams.buffInfo.inBuffType, tempAlgParams.buffInfo.inBuffBaseOff, tempAlgParams.sliceSize);
195 : DataSlice scratchSlices
196 0 : = DataSlice(BufferType::SCRATCH, tempAlgParams.buffInfo.scratchBuffBaseOff, tempAlgParams.sliceSize);
197 0 : CHK_RET(LocalCopy(tempInsQues[0], usrInSlices, scratchSlices));
198 :
199 0 : nhrInBuffBaseOff_ = tempAlgParams.buffInfo.scratchBuffBaseOff;
200 : } else {
201 0 : HCCL_INFO("[InsTempAllReduceNHR][PreCopy] skip precopy");
202 : }
203 : } else {
204 0 : HCCL_INFO("[InsTempAllReduceNHR][PreCopy] offload skip precopy");
205 0 : nhrInBuffType_ = tempAlgParams.buffInfo.inBuffType;
206 0 : nhrInBuffBaseOff_ = tempAlgParams.buffInfo.inBuffBaseOff;
207 : }
208 :
209 0 : nhrOutBuffType_ = tempAlgParams.buffInfo.outBuffType;
210 0 : nhrOutBuffBaseOff_ = tempAlgParams.buffInfo.outBuffBaseOff;
211 :
212 0 : return HcclResult::HCCL_SUCCESS;
213 : }
214 :
215 : // 将reduceScatter之后的数据先放到usrOut
216 0 : HcclResult InsTempAllReduceNHR::PrepareDataForAllGather(
217 : const RankSliceInfo& sliceInfoVec, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
218 : {
219 : // 如果是单算子模式,在原来的位置要先做完allGather,然后postCopy把数据放到usrOut
220 : // 如果是图模式,直接把数据放到usrOUt,然后在usrOut上做allGather
221 0 : HCCL_INFO("[InsTempAllReduceNHR][PrepareDataForAllGather] prepare data for allGather");
222 :
223 0 : if (opMode_ == OpMode::OFFLOAD) {
224 0 : u64 size = sliceInfoVec[tempVirtRankMap_[myRank_]][0].size;
225 0 : u64 srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
226 0 : u64 dstOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
227 0 : DataSlice srcSlice = DataSlice(nhrInBuffType_, nhrInBuffBaseOff_ + srcOffset, size);
228 0 : DataSlice dstSlice = DataSlice(nhrOutBuffType_, nhrOutBuffBaseOff_ + dstOffset, size);
229 0 : CHK_RET(LocalCopy(tempInsQues[linkIdx], srcSlice, dstSlice));
230 :
231 0 : nhrInBuffType_ = nhrOutBuffType_;
232 0 : nhrInBuffBaseOff_ = nhrOutBuffBaseOff_;
233 : }
234 :
235 0 : return HcclResult::HCCL_SUCCESS;
236 : }
237 :
238 0 : HcclResult InsTempAllReduceNHR::PostCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
239 : {
240 : // 单算子模式,需要将数据拷贝到usrOut
241 0 : if (opMode_ == OpMode::OPBASE) {
242 0 : HCCL_INFO("[InsTempAllReduceNHR][PostCopy] Opbase copy from scratchBuffer to userOut");
243 0 : DataSlice scratchSlices = DataSlice(nhrInBuffType_, nhrInBuffBaseOff_, tempAlgParams.sliceSize);
244 0 : DataSlice usrOutSlices = DataSlice(nhrOutBuffType_, nhrOutBuffBaseOff_, tempAlgParams.sliceSize);
245 0 : CHK_RET(LocalCopy(tempInsQues[0], scratchSlices, usrOutSlices));
246 : } else {
247 0 : HCCL_INFO("[InsTempAllReduceNHR][PostCopy] offload skip postcopy");
248 : }
249 :
250 0 : return HcclResult::HCCL_SUCCESS;
251 : }
252 :
253 0 : HcclResult InsTempAllReduceNHR::RunReduceScatter(
254 : const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
255 : {
256 0 : std::vector<AicpuNHRStepInfo> stepInfoList;
257 0 : GetStepInfoList(stepInfoList);
258 0 : for (auto& stepInfo : stepInfoList) {
259 0 : HCCL_DEBUG(
260 : "[InsTempAllReduceNHR][RunReduceScatter] step[%u], myRank[%u], toRank[%u], fromRank[%u], nSlices[%u].",
261 : stepInfo.step, stepInfo.myRank, stepInfo.toRank, stepInfo.fromRank, stepInfo.nSlices);
262 :
263 0 : const std::vector<LinkData>& linkRecv = tempLinks.at(GetRankFromMap(stepInfo.fromRank));
264 0 : const std::vector<LinkData>& linkSend = tempLinks.at(GetRankFromMap(stepInfo.toRank));
265 0 : std::vector<DataSlice> txSlices;
266 0 : std::vector<DataSlice> rxSlices;
267 :
268 : // 在 nhrInBuffType_ 上进行 ReduceScatter 操作
269 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
270 0 : u64 txOffset = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
271 0 : u64 txSize = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].size;
272 0 : u64 rxOffset = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
273 0 : u64 rxSize = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].size;
274 0 : DataSlice txSlice = DataSlice(nhrInBuffType_, txOffset, txSize);
275 0 : DataSlice rxSlice = DataSlice(nhrInBuffType_, rxOffset, rxSize);
276 0 : txSlices.push_back(txSlice);
277 0 : rxSlices.push_back(rxSlice);
278 : }
279 : SendRecvReduceInfo sendRecvReduceInfo{
280 0 : {linkSend[linkIdx], linkRecv[linkIdx]}, {{txSlices, txSlices}, {rxSlices, rxSlices}}, dataType_, redOp_};
281 0 : CHK_PRT_RET(
282 : SendRecvReduce(sendRecvReduceInfo, tempInsQues[linkIdx], 0, true, dmaMode_) != HcclResult::HCCL_SUCCESS,
283 : HCCL_ERROR("[InsTempAllReduceNHR] RunReduceScatter SendRecvReduce failed"), HcclResult::HCCL_E_INTERNAL);
284 0 : }
285 0 : return HcclResult::HCCL_SUCCESS;
286 0 : }
287 :
288 0 : HcclResult InsTempAllReduceNHR::RunAllGather(
289 : const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
290 : {
291 0 : u32 nSteps = GetNHRStepNum(tempRankSize_);
292 0 : for (u32 step = 0; step < nSteps; step++) {
293 0 : AicpuNHRStepInfo stepInfo;
294 0 : CHK_RET(GetStepInfo(step, nSteps, stepInfo));
295 :
296 0 : const std::vector<LinkData>& linkRecv = tempLinks.at(GetRankFromMap(stepInfo.fromRank));
297 0 : const std::vector<LinkData>& linkSend = tempLinks.at(GetRankFromMap(stepInfo.toRank));
298 :
299 0 : std::vector<DataSlice> txSlices;
300 0 : std::vector<DataSlice> rxSlices;
301 :
302 0 : HCCL_DEBUG(
303 : "[InsTempAllReduceNHR] rank[%d] rankSize[%u] recvFrom[%u] sendTo[%u] step[%u] nSteps[%u] nSlices[%u]",
304 : myRank_, tempRankSize_, stepInfo.fromRank, stepInfo.toRank, step, nSteps, stepInfo.nSlices);
305 :
306 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
307 0 : u64 txOffset = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
308 0 : u64 txSize = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].size;
309 0 : u64 rxOffset = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
310 0 : u64 rxSize = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].size;
311 0 : DataSlice txSlice = DataSlice(nhrInBuffType_, txOffset, txSize);
312 0 : DataSlice rxSlice = DataSlice(nhrInBuffType_, rxOffset, rxSize);
313 0 : txSlices.push_back(txSlice);
314 0 : rxSlices.push_back(rxSlice);
315 : }
316 :
317 0 : TxRxLinks sendRecvLinks(linkSend[linkIdx], linkRecv[linkIdx]);
318 0 : TxRxSlicesList sendRecvSlicesList({txSlices, txSlices}, {rxSlices, rxSlices});
319 :
320 0 : SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
321 0 : CHK_PRT_RET(
322 : SendRecv(sendRecvInfo, tempInsQues[linkIdx], 0, true, dmaMode_) != HcclResult::HCCL_SUCCESS,
323 : HCCL_ERROR("[InsTempAllReduceNHR] RunAllGather send/recv failed"), HcclResult::HCCL_E_INTERNAL);
324 0 : }
325 0 : return HcclResult::HCCL_SUCCESS;
326 : }
327 :
328 0 : HcclResult InsTempAllReduceNHR::GetStepInfo(u32 step, u32 nSteps, AicpuNHRStepInfo& stepInfo)
329 : {
330 0 : u32 rankIdx = tempVirtRankMap_[myRank_];
331 0 : stepInfo.txSliceIdxs.clear();
332 0 : stepInfo.rxSliceIdxs.clear();
333 0 : stepInfo.step = step;
334 0 : stepInfo.myRank = rankIdx;
335 :
336 : // 计算通信对象
337 0 : u32 deltaRank = 1 << (nSteps - 1 - step);
338 0 : u32 recvFrom = (rankIdx + tempRankSize_ - deltaRank) % tempRankSize_;
339 0 : u32 sendTo = (rankIdx + deltaRank) % tempRankSize_;
340 :
341 : // 数据份数和数据编号增量
342 0 : u32 nSlices = (tempRankSize_ - 1 + (1 << (nSteps - 1 - step))) / (1 << (nSteps - step));
343 0 : u32 deltaSliceIndex = 1 << (nSteps - step);
344 0 : u32 txSliceIdx = rankIdx;
345 0 : u32 rxSliceIdx = (rankIdx - (1 << (nSteps - 1 - step)) + tempRankSize_) % tempRankSize_;
346 :
347 0 : stepInfo.nSlices = nSlices;
348 0 : stepInfo.toRank = sendTo;
349 0 : stepInfo.fromRank = recvFrom;
350 :
351 0 : for (u32 i = 0; i < nSlices; i++) {
352 0 : stepInfo.txSliceIdxs.push_back(txSliceIdx);
353 0 : stepInfo.rxSliceIdxs.push_back(rxSliceIdx);
354 :
355 0 : HCCL_DEBUG("[InsTempAllReduceNHR][GetStepInfo] i[%u] txSliceIdx[%u] rxSliceIdx[%u]", i, txSliceIdx, rxSliceIdx);
356 :
357 0 : txSliceIdx = (txSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
358 0 : rxSliceIdx = (rxSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
359 : }
360 0 : return HcclResult::HCCL_SUCCESS;
361 : }
362 :
363 : // 计算每轮收发的对端以及slice编号
364 0 : HcclResult InsTempAllReduceNHR::GetStepInfoList(std::vector<AicpuNHRStepInfo>& stepInfoList)
365 : {
366 : // 将本 rank 号转换成算法使用的索引号
367 0 : u32 rankIdx = tempVirtRankMap_[myRank_];
368 0 : stepInfoList.clear();
369 :
370 0 : u32 nSteps = GetNHRStepNum(tempRankSize_);
371 0 : stepInfoList.resize(nSteps);
372 0 : for (u32 step = 0; step < nSteps; step++) {
373 : // 计算通信对象
374 0 : u32 deltaRank = 1 << step;
375 0 : u32 sendTo = (rankIdx + tempRankSize_ - deltaRank) % tempRankSize_;
376 0 : u32 recvFrom = (rankIdx + deltaRank) % tempRankSize_;
377 :
378 : // 数据份数和数据编号增量
379 0 : u32 nSlices = (tempRankSize_ - 1 + (1 << step)) / (1 << (step + 1));
380 0 : u32 deltaSliceIndex = 1 << (step + 1);
381 0 : u32 txSliceIdx = sendTo;
382 0 : u32 rxSliceIdx = rankIdx;
383 :
384 0 : AicpuNHRStepInfo& currStepInfo = stepInfoList[step];
385 0 : currStepInfo.step = step;
386 0 : currStepInfo.myRank = rankIdx;
387 0 : currStepInfo.nSlices = nSlices;
388 0 : currStepInfo.toRank = sendTo;
389 0 : currStepInfo.fromRank = recvFrom;
390 :
391 : // 计算本rank在每轮收/发中的slice编号
392 0 : currStepInfo.txSliceIdxs.reserve(nSlices);
393 0 : currStepInfo.rxSliceIdxs.reserve(nSlices);
394 0 : for (u32 i = 0; i < nSlices; i++) {
395 0 : currStepInfo.txSliceIdxs.push_back(txSliceIdx);
396 0 : currStepInfo.rxSliceIdxs.push_back(rxSliceIdx);
397 0 : HCCL_DEBUG(
398 : "[InsTempAllReduceNHR][GetStepInfoList] i[%u] txSliceIdx[%u] rxSliceIdx[%u]", i, txSliceIdx,
399 : rxSliceIdx);
400 0 : txSliceIdx = (txSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
401 0 : rxSliceIdx = (rxSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
402 : }
403 : }
404 0 : return HcclResult::HCCL_SUCCESS;
405 : }
406 :
407 0 : RankId InsTempAllReduceNHR::GetRankFromMap(const u32 rankIdx)
408 : {
409 0 : RankId rank = -1;
410 0 : HCCL_INFO("[InsTempAllReduceNHR] GetRankFromMap");
411 0 : for (auto& pair : tempVirtRankMap_) {
412 0 : if (pair.second == rankIdx) {
413 0 : rank = pair.first;
414 0 : break;
415 : }
416 : }
417 0 : return rank;
418 : }
419 : } // namespace Hccl
|