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(const RankId virtualRank, const u32 tempRankSize,
17 : const std::vector<std::vector<RankId>> &tempVTopo,
18 0 : const std::map<RankId, u32> &tempVirtRankMap)
19 0 : : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
20 : {
21 0 : }
22 :
23 0 : TempAllGatherRing::~TempAllGatherRing()
24 : {
25 0 : }
26 :
27 0 : HcclResult TempAllGatherRing::CalcRes(AlgTempResReq &tempResReq)
28 : {
29 0 : tempResReq.queNum = tempVTopo_.size();
30 :
31 0 : CHK_PRT_RET(CalcResLinksRing(myRank_, tempRankSize_, tempVTopo_, tempResReq) != HcclResult::HCCL_SUCCESS,
32 : HCCL_ERROR("[CollAlgFactory] [TempAllGatherRing] Rank [%d], resLinks calculation error!", myRank_),
33 : HcclResult::HCCL_E_INTERNAL);
34 :
35 0 : return HcclResult::HCCL_SUCCESS;
36 : }
37 :
38 0 : HcclResult TempAllGatherRing::CalcSliceInfo(const AllignInfo &allignInfo, const u64 dataSize,
39 : RankSliceInfo &sliceInfoVec)
40 : {
41 0 : std::vector<SliceInfo> tmp(tempVTopo_.size());
42 0 : sliceInfoVec.resize(tempRankSize_, tmp);
43 : // for reduce scatter, dataSize = chunkSize
44 0 : CHK_RET(CalcRsAgSliceInfoRing(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
45 :
46 0 : return HcclResult::HCCL_SUCCESS;
47 0 : }
48 :
49 0 : HcclResult TempAllGatherRing::GenPrimQue(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
50 : const BuffInfo &buffInfo, const ResLinks &tempLinks,
51 : std::vector<PrimQuePtr> &tempPrimQues)
52 : {
53 0 : opMode_ = tempFuncs.opMode;
54 0 : enableCounterNotify_ = tempFuncs.enableCounterNotify;
55 0 : buffInfo_ = buffInfo;
56 :
57 0 : queNum_ = tempVTopo_.size();
58 0 : CHK_PRT_RET(queNum_ != tempPrimQues.size(),
59 : HCCL_ERROR("[CollAlgFactory] [TempAllGatherRing] Rank [%d], requiredQue Error.", myRank_),
60 : HcclResult::HCCL_E_INTERNAL);
61 :
62 : // Local Copy from Input to Scratch Buffer for OPBASE
63 0 : if ((opMode_ == OpMode::OPBASE) && tempFuncs.isForepart && !tempFuncs.forAllReduce) {
64 0 : CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
65 : }
66 :
67 : // Local Copy from Input to Output Buffer for OFFLOAD
68 0 : if ((opMode_ == OpMode::OFFLOAD) && (!tempFuncs.forAlgSeqComb)) {
69 0 : CHK_RET(PreCopyOffload(sliceInfoVec, tempFuncs.forAllReduce, tempPrimQues));
70 : }
71 :
72 0 : stepNum_ = tempRankSize_ - 1;
73 0 : for (u32 queIdx = 0; queIdx < tempVTopo_.size(); queIdx++) {
74 : // semaphore sync for standAlone AllReduce
75 0 : if (!tempFuncs.forAllReduce && (queNum_ > 1)) {
76 0 : CHK_PRT_RET(
77 : PreSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
78 : HCCL_ERROR(
79 : "[CollAlgFactory] [TempAllGatherRing] Rank [%d], Que [%u], Semaphore Synchronization Failed.",
80 : myRank_, tempPrimQues[queIdx]->GetId()),
81 : HcclResult::HCCL_E_INTERNAL);
82 : }
83 :
84 0 : PrimQuePtr currPrimQue = tempPrimQues[queIdx];
85 0 : CHK_RET(RunIndividualRing(queIdx, sliceInfoVec, tempLinks, currPrimQue));
86 :
87 : // semaphore sync
88 0 : if (queNum_ > 1) {
89 0 : CHK_PRT_RET(PostSync(queIdx, tempPrimQues) != HcclResult::HCCL_SUCCESS,
90 : HCCL_ERROR("[CollAlgFactory] [TempAllGatherRing] Rank [%d], Unable to synchronize all queues.",
91 : 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(const RankSliceInfo &sliceInfoVec, const bool forAllReduce,
105 : 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(const u32 queIdx, const RankSliceInfo &sliceInfoVec,
132 : 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
|