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 "temp_all_gather_ring.h"
14 :
15 : namespace Hccl {
16 0 : TempAllGatherRing::TempAllGatherRing(
17 : const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
18 0 : const std::map<RankId, u32>& tempVirtRankMap)
19 0 : : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 0 : {}
21 :
22 0 : TempAllGatherRing::~TempAllGatherRing() {}
23 :
24 0 : HcclResult TempAllGatherRing::CalcRes(AlgTempResReq& tempResReq)
25 : {
26 0 : tempResReq.queNum = tempVTopo_.size();
27 :
28 0 : CHK_PRT_RET(
29 : CalcResLinksRing(myRank_, tempRankSize_, tempVTopo_, tempResReq) != HcclResult::HCCL_SUCCESS,
30 : HCCL_ERROR("[CollAlgFactory] [TempAllGatherRing] Rank [%d], resLinks calculation error!", myRank_),
31 : HcclResult::HCCL_E_INTERNAL);
32 :
33 0 : return HcclResult::HCCL_SUCCESS;
34 : }
35 :
36 : HcclResult
37 0 : TempAllGatherRing::CalcSliceInfo(const AllignInfo& allignInfo, const u64 dataSize, RankSliceInfo& sliceInfoVec)
38 : {
39 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
40 0 : sliceInfoVec.resize(tempRankSize_, tmp);
41 : // for reduce scatter, dataSize = chunkSize
42 0 : CHK_RET(CalcRsAgSliceInfoRing(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
43 :
44 0 : return HcclResult::HCCL_SUCCESS;
45 0 : }
46 :
47 0 : HcclResult TempAllGatherRing::GenPrimQue(
48 : const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
49 : std::vector<PrimQuePtr>& tempPrimQues)
50 : {
51 0 : opMode_ = tempFuncs.opMode;
52 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
53 0 : buffInfo_ = buffInfo;
54 :
55 0 : queNum_ = tempVTopo_.size();
56 0 : CHK_PRT_RET(
57 : queNum_ != tempPrimQues.size(),
58 : HCCL_ERROR("[CollAlgFactory] [TempAllGatherRing] Rank [%d], requiredQue Error.", myRank_),
59 : HcclResult::HCCL_E_INTERNAL);
60 :
61 : // Local Copy from Input to Scratch Buffer for OPBASE
62 0 : if ((opMode_ == OpMode::OPBASE) && tempFuncs.isForepart && !tempFuncs.forAllReduce) {
63 0 : CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
64 : }
65 :
66 : // Local Copy from Input to Output Buffer for OFFLOAD
67 0 : if ((opMode_ == OpMode::OFFLOAD) && (!tempFuncs.forAlgSeqComb)) {
68 0 : CHK_RET(PreCopyOffload(sliceInfoVec, tempFuncs.forAllReduce, tempPrimQues));
69 : }
70 :
71 0 : stepNum_ = tempRankSize_ - 1;
72 0 : for (u32 queIdx = 0; queIdx < tempVTopo_.size(); queIdx++) {
73 : // semaphore sync for standAlone AllReduce
74 0 : if (!tempFuncs.forAllReduce && (queNum_ > 1)) {
75 0 : CHK_PRT_RET(
76 : PreSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
77 : HCCL_ERROR(
78 : "[CollAlgFactory] [TempAllGatherRing] Rank [%d], Que [%u], Semaphore Synchronization Failed.",
79 : myRank_, tempPrimQues[queIdx]->GetId()),
80 : HcclResult::HCCL_E_INTERNAL);
81 : }
82 :
83 0 : PrimQuePtr currPrimQue = tempPrimQues[queIdx];
84 0 : CHK_RET(RunIndividualRing(queIdx, sliceInfoVec, tempLinks, currPrimQue));
85 :
86 : // semaphore sync
87 0 : if (queNum_ > 1) {
88 0 : CHK_PRT_RET(
89 : PostSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
90 : HCCL_ERROR(
91 : "[CollAlgFactory] [TempAllGatherRing] Rank [%d], Unable to synchronize all queues.", myRank_),
92 : HcclResult::HCCL_E_INTERNAL);
93 : }
94 0 : }
95 :
96 : // LocalCopy: from scratch to output for opbase
97 0 : if ((opMode_ == OpMode::OPBASE) && tempFuncs.isBottom) {
98 0 : CHK_RET(PostCopyOpbase(tempFuncs.usrData, tempPrimQues));
99 : }
100 :
101 0 : return HcclResult::HCCL_SUCCESS;
102 : }
103 :
104 0 : HcclResult TempAllGatherRing::PreCopyOffload(
105 : const RankSliceInfo& sliceInfoVec, const bool forAllReduce, std::vector<PrimQuePtr>& tempPrimQues)
106 : {
107 0 : if (!forAllReduce) {
108 0 : u64 srcOffset = 0;
109 0 : u64 srcSize = sliceInfoVec[tempVirtRankMap_[myRank_]][queNum_ - 1].offset
110 0 : - sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset
111 0 : + sliceInfoVec[tempVirtRankMap_[myRank_]][queNum_ - 1].size;
112 0 : u64 dstOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
113 0 : DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, srcOffset + buffInfo_.inBuffBaseOff, srcSize);
114 0 : DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, dstOffset + buffInfo_.outBuffBaseOff, srcSize);
115 0 : std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
116 0 : tempPrimQues[0]->Append(std::move(primLocalCopy));
117 0 : } else { // allreduce: split local copy per queue to eliminate the notify/wait
118 0 : for (u32 qIdx = 0; qIdx < queNum_; qIdx++) {
119 0 : u64 queSliceOff = sliceInfoVec[tempVirtRankMap_[myRank_]][qIdx].offset;
120 0 : u64 queSliceSize = sliceInfoVec[tempVirtRankMap_[myRank_]][qIdx].size;
121 0 : DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, queSliceOff + buffInfo_.inBuffBaseOff, queSliceSize);
122 0 : DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, queSliceOff + buffInfo_.outBuffBaseOff, queSliceSize);
123 0 : std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
124 0 : tempPrimQues[qIdx]->Append(std::move(primLocalCopy));
125 0 : }
126 : }
127 :
128 0 : return HcclResult::HCCL_SUCCESS;
129 : }
130 :
131 0 : HcclResult TempAllGatherRing::RunIndividualRing(
132 : const u32 queIdx, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, PrimQuePtr currPrimQue)
133 : {
134 : // locate myRank in tempVTopo -> algRank
135 : u32 myAlgRank;
136 0 : CHK_RET(GetAlgRank(myRank_, tempVTopo_[queIdx], myAlgRank));
137 :
138 : // find neighbors -> virtualRank
139 0 : RankId sendToRank = tempVTopo_[queIdx][(myAlgRank + 1) % tempRankSize_];
140 0 : RankId recvFromRank = tempVTopo_[queIdx][(myAlgRank - 1 + tempRankSize_) % tempRankSize_]; // virtualRank
141 :
142 : // Link
143 0 : LinkData sendLinkData = tempLinks.at(sendToRank)[0];
144 0 : LinkData recvLinkData = tempLinks.at(recvFromRank)[0];
145 :
146 : // run stepNum steps to complete the ring
147 0 : for (u32 step = 0; step < stepNum_; step++) {
148 0 : u32 sendChunkIdx = tempVirtRankMap_[tempVTopo_[queIdx][(myAlgRank - step + tempRankSize_) % tempRankSize_]];
149 0 : u64 sendOffset = sliceInfoVec[sendChunkIdx][queIdx].offset;
150 0 : u64 sendSize = sliceInfoVec[sendChunkIdx][queIdx].size;
151 0 : u32 recvChunkIdx = tempVirtRankMap_[tempVTopo_[queIdx][(myAlgRank - 1 - step + tempRankSize_) % tempRankSize_]];
152 :
153 0 : u64 recvOffset = sliceInfoVec[recvChunkIdx][queIdx].offset;
154 0 : u64 recvSize = sliceInfoVec[recvChunkIdx][queIdx].size;
155 :
156 : // PrimGroup
157 0 : std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
158 :
159 : // Send
160 0 : DataSlice sendLocSlice = DataSlice(buffInfo_.outBuffType, sendOffset + buffInfo_.outBuffBaseOff, sendSize);
161 0 : DataSlice sendRemSlice = DataSlice(buffInfo_.outBuffType, sendOffset + buffInfo_.outBuffBaseOff, sendSize);
162 : std::unique_ptr<Primitive> primSend
163 0 : = std::make_unique<PrimSend>(sendToRank, sendLinkData, sendLocSlice, sendRemSlice, dmaMode_);
164 :
165 0 : primGroup->Append(std::move(primSend));
166 :
167 : // Recv
168 0 : DataSlice recvRemSlice = DataSlice(buffInfo_.outBuffType, recvOffset + buffInfo_.outBuffBaseOff, recvSize);
169 0 : DataSlice recvLocSlice = DataSlice(buffInfo_.outBuffType, recvOffset + buffInfo_.outBuffBaseOff, recvSize);
170 : std::unique_ptr<Primitive> primRecv
171 0 : = std::make_unique<PrimRecv>(recvFromRank, recvLinkData, recvLocSlice, recvRemSlice, dmaMode_);
172 :
173 0 : primGroup->Append(std::move(primRecv));
174 :
175 0 : currPrimQue->Append(std::move(primGroup));
176 0 : }
177 0 : return HcclResult::HCCL_SUCCESS;
178 : }
179 :
180 : } // namespace Hccl
|