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 "ccu_context_all_gather_nhr1d_mem2mem.h"
12 :
13 : namespace Hccl {
14 :
15 : constexpr uint16_t OUTPUT_XN_ID = 1;
16 : constexpr uint16_t TOKEN_XN_ID = 2;
17 : constexpr uint16_t FST_AXIS_ID = 0;
18 : constexpr uint16_t SEC_AXIS_ID = 1;
19 : constexpr uint16_t CKE_IDX_0 = 0;
20 : constexpr uint16_t CKE_IDX_1 = 1;
21 : constexpr uint16_t CKE_IDX_2 = 2;
22 : constexpr uint16_t CKE_IDX_3 = 3;
23 : constexpr uint16_t CKE_IDX_4 = 4;
24 : constexpr uint16_t BIT_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
25 :
26 0 : CcuContextAllGatherNHR1D::CcuContextAllGatherNHR1D(const CcuCtxArg &arg, const std::vector<CcuTransport *> &transports,
27 0 : const CcuTransportGroup &group)
28 0 : : CcuContextAlgBase(arg, transports, group)
29 : {
30 0 : const CcuCtxArgAllGatherNHR1D *ctxArg = dynamic_cast<const CcuCtxArgAllGatherNHR1D *>(&arg);
31 0 : if (ctxArg == nullptr) {
32 0 : THROW<NullPtrException>(StringFormat("CcuContextAllGatherNHR1D::ctxArg ptr is null"));
33 : }
34 0 : rankId_ = ctxArg->rankId_;
35 0 : axisId_ = ctxArg->axisId_;
36 0 : axisSize_ = ctxArg->axisSize_;
37 0 : dimSize_ = ctxArg->dimSize_[0];
38 0 : localAxisSignalName_ = "CcuContextAllGatherNHR1DDieSync_" + std::to_string(axisId_);
39 0 : anotherAxisSignalName_ = "CcuContextAllGatherNHR1DDieSync_" + std::to_string(1 - axisId_);
40 0 : stepInfoVector_ = ctxArg->stepInfoVector_;
41 0 : indexMap_ = ctxArg->indexMap_;
42 0 : localSize_ = indexMap_.size();
43 0 : myRankIdx_ = indexMap_.size();
44 0 : signalNum_ = (dimSize_ + BIT_NUM_PER_CKE - 1) / BIT_NUM_PER_CKE; // 每个CKE有16个bit
45 0 : HCCL_INFO(
46 : "[CcuContextAllGatherNHR1D] CtxArg: rankId_[%u], axisId_[%u], axisSize_[%u], dimSize_[%u], localSize_[%u], "
47 : "signalNum_[%u]",
48 : rankId_, axisId_, axisSize_, dimSize_, localSize_, signalNum_);
49 0 : }
50 :
51 0 : void CcuContextAllGatherNHR1D::LoadArgs()
52 : {
53 0 : Load(input_);
54 0 : Load(output_[myRankIdx_]);
55 0 : Load(token_[myRankIdx_]);
56 0 : Load(die0Size_);
57 0 : Load(die1Size_);
58 0 : Load(repeatNum_);
59 0 : Load(inputSliceStride_);
60 0 : Load(outputSliceStride_);
61 0 : Load(inputRepeatStride_);
62 0 : Load(outputRepeatStride_);
63 0 : Load(isInputOutputEqual_);
64 :
65 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] LoadArgs run finished");
66 0 : }
67 :
68 0 : void CcuContextAllGatherNHR1D::InitResources()
69 : {
70 0 : die0Size_ = CreateVariable();
71 0 : die1Size_ = CreateVariable();
72 0 : inputSliceStride_ = CreateVariable();
73 0 : outputSliceStride_ = CreateVariable();
74 0 : inputRepeatStride_ = CreateVariable();
75 0 : outputRepeatStride_ = CreateVariable();
76 0 : repeatNum_ = CreateVariable();
77 0 : tmpCopyRepeatNum_ = CreateVariable();
78 0 : repeatTimeflag_ = CreateVariable();
79 0 : isInputOutputEqual_ = CreateVariable();
80 0 : myrankInputSliceOffset_ = CreateVariable();
81 0 : tmpSliceOffset_ = CreateVariable();
82 0 : for (u64 i = 0; i < dimSize_; i++) {
83 0 : outputSliceOffset_.push_back(CreateVariable());
84 : }
85 0 : constVar1_ = CreateVariable();
86 0 : constVar1_ = 1;
87 :
88 0 : localSignal_ = CreateMaskSignal();
89 0 : localAxisSignal_ = CreateMaskSignal();
90 :
91 0 : if (axisSize_ > 1) {
92 0 : ExportMaskSignal(localAxisSignal_, localAxisSignalName_);
93 0 : anotherAxisSignal_ = ImportMaskSignal(anotherAxisSignalName_);
94 : }
95 :
96 0 : input_ = CreateVariable();
97 0 : for (uint32_t transportIdx = 0; transportIdx < localSize_; transportIdx++) {
98 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] MyRank[%u], TransportId[%u]", rankId_, transportIdx);
99 0 : CHK_PRT_RET(transports[transportIdx] == nullptr,
100 : HCCL_ERROR("[CcuContextAllGatherNHR1D] Algorithm transport ptr is null"), );
101 0 : output_.push_back(
102 0 : CreateVariable((*transports[transportIdx]), OUTPUT_XN_ID)); // 获取transport中id=1的Var来传递output
103 0 : token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
104 : }
105 0 : output_.push_back(CreateVariable());
106 0 : token_.push_back(CreateVariable());
107 :
108 0 : srcMem_ = CreateMemory();
109 0 : dstMem_ = CreateMemory();
110 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] InitResources finished");
111 : }
112 :
113 0 : void CcuContextAllGatherNHR1D::PreSync()
114 : {
115 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync start");
116 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
117 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
118 0 : for (auto t : transports) {
119 0 : WriteVariableWithSignal(*t, output_[localSize_], OUTPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
120 0 : WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
121 : }
122 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
123 0 : for (auto &pair : indexMap_) {
124 0 : uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
125 0 : uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
126 0 : waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
127 : }
128 0 : for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
129 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
130 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
131 : }
132 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync end");
133 0 : }
134 :
135 0 : void CcuContextAllGatherNHR1D::PostSync()
136 : {
137 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
138 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
139 0 : for (auto &t : transports) {
140 0 : RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
141 : }
142 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
143 0 : for (auto &pair : indexMap_) {
144 0 : uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
145 0 : uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
146 0 : waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
147 : }
148 0 : for (uint32_t sId = 0; sId < signalNum_; sId++) {
149 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
150 : }
151 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PostSync run finished");
152 0 : }
153 :
154 0 : void CcuContextAllGatherNHR1D::AxisSync(uint32_t signalIndex)
155 : {
156 0 : const uint32_t DIE_NUM = 2;
157 0 : if (signalIndex > 1) {
158 0 : THROW<InvalidParamsException>(
159 0 : StringFormat("[CcuContextAllGatherNHR1D] Unexpected SignalInex[%u]", signalIndex));
160 : }
161 0 : LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
162 0 : LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
163 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AxisSync run finished");
164 0 : return;
165 : }
166 :
167 0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHR()
168 : {
169 0 : tmpSliceOffset_ = 0;
170 0 : myrankInputSliceOffset_ = 0;
171 0 : for (u64 i = 0; i < rankId_; i++) {
172 0 : myrankInputSliceOffset_ += inputSliceStride_;
173 : }
174 0 : for (u64 i = 0; i < dimSize_; i++) {
175 0 : outputSliceOffset_[i] = tmpSliceOffset_;
176 0 : tmpSliceOffset_ += outputSliceStride_;
177 : }
178 0 : srcMem_.addr = input_;
179 0 : srcMem_.addr += myrankInputSliceOffset_;
180 0 : dstMem_.addr = output_[myRankIdx_];
181 0 : dstMem_.addr += outputSliceOffset_[rankId_];
182 0 : srcMem_.token = token_[myRankIdx_];
183 0 : dstMem_.token = token_[myRankIdx_];
184 0 : tmpCopyRepeatNum_ = repeatNum_;
185 0 : repeatTimeflag_ = 0;
186 0 : CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
187 : {
188 0 : tmpCopyRepeatNum_ += constVar1_;
189 0 : CCU_IF(repeatTimeflag_ != 0)
190 : {
191 0 : srcMem_.addr += inputRepeatStride_;
192 0 : dstMem_.addr += outputRepeatStride_;
193 0 : }
194 0 : CCU_IF(repeatTimeflag_ == 0)
195 : {
196 0 : if (axisId_ == 1) {
197 0 : srcMem_.addr += die0Size_;
198 0 : dstMem_.addr += die0Size_;
199 : }
200 0 : }
201 0 : CCU_IF(isInputOutputEqual_ == 0)
202 : {
203 0 : LocalCopy(dstMem_, srcMem_, axisId_ == 0 ? die0Size_ : die1Size_, localSignal_, 1 << rankId_);
204 0 : }
205 0 : CCU_IF(isInputOutputEqual_ != 0)
206 : {
207 0 : LocalPost(localSignal_, 1 << rankId_);
208 0 : }
209 0 : LocalWait(localSignal_, 1 << rankId_);
210 0 : repeatTimeflag_ = 1;
211 0 : }
212 :
213 0 : for (auto &nhrStepInfo : stepInfoVector_) {
214 0 : DoRepeatAllGatherNHRSingleStep(nhrStepInfo);
215 : }
216 0 : }
217 :
218 0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHRSingleStep(const NHRStepInfo &nhrStepInfo)
219 : {
220 0 : u32 &toRankIdx = indexMap_[nhrStepInfo.toRank];
221 0 : u32 &fromRankIdx = indexMap_[nhrStepInfo.fromRank];
222 0 : u32 sendSliceIdx = 0;
223 0 : CcuTransport *sendTransport = transports[toRankIdx];
224 0 : CcuTransport *recvTransport = transports[fromRankIdx];
225 0 : const std::vector<u32> &sendSliceIdxList = nhrStepInfo.txSliceIdxs;
226 0 : srcMem_.token = token_[myRankIdx_];
227 0 : dstMem_.token = token_[toRankIdx];
228 0 : for (u32 i = 0; i < sendSliceIdxList.size(); i++) { ////这里写的可能有问题
229 0 : sendSliceIdx = sendSliceIdxList[i];
230 0 : if (i != 0) {
231 0 : if (i % BIT_NUM_PER_CKE == 0) {
232 0 : LocalWait(localSignal_, (1 << BIT_NUM_PER_CKE) - 1);
233 : }
234 : }
235 0 : if (nhrStepInfo.step == 0) {
236 0 : srcMem_.addr = input_;
237 0 : srcMem_.addr += myrankInputSliceOffset_;
238 : } else {
239 0 : srcMem_.addr = output_[myRankIdx_];
240 0 : srcMem_.addr += outputSliceOffset_[sendSliceIdx];
241 : }
242 0 : dstMem_.addr = output_[toRankIdx];
243 0 : dstMem_.addr += outputSliceOffset_[sendSliceIdx];
244 0 : DoRepeatSendRecvSlices(nhrStepInfo.toRank, srcMem_, dstMem_, i % BIT_NUM_PER_CKE);
245 : }
246 :
247 0 : if (nhrStepInfo.step + 1 != stepInfoVector_.size()){
248 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
249 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
250 0 : RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true);
251 0 : uint16_t recvSignalId = nhrStepInfo.fromRank / BIT_NUM_PER_CKE;
252 0 : uint16_t recvBit = 1 << (nhrStepInfo.fromRank % BIT_NUM_PER_CKE);
253 0 : RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_3, recvBit);
254 : }
255 0 : }
256 :
257 0 : void CcuContextAllGatherNHR1D::DoRepeatSendRecvSlices(const u32 &toRank, CcuRep::Memory &src, CcuRep::Memory &dst,
258 : u32 signalIndex)
259 : {
260 0 : CcuTransport *sendTransport = transports[indexMap_[toRank]];
261 0 : const CcuRep::Variable &sliceSize = axisId_ == 0 ? die0Size_ : die1Size_;
262 0 : repeatTimeflag_ = 0;
263 0 : tmpCopyRepeatNum_ = repeatNum_;
264 0 : CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
265 : {
266 0 : tmpCopyRepeatNum_ += constVar1_;
267 0 : CCU_IF(repeatTimeflag_ == 1)
268 : {
269 0 : src.addr += inputRepeatStride_;
270 0 : dst.addr += outputRepeatStride_;
271 0 : }
272 0 : CCU_IF(repeatTimeflag_ == 0)
273 : {
274 0 : if (axisId_ == 1) {
275 0 : src.addr += die0Size_;
276 0 : dst.addr += die0Size_;
277 : }
278 0 : }
279 0 : Write(*sendTransport, dst, src, sliceSize, localSignal_, 1 << signalIndex);
280 0 : LocalWait(localSignal_, 1 << signalIndex);
281 0 : repeatTimeflag_ = 1;
282 0 : }
283 0 : }
284 :
285 0 : void CcuContextAllGatherNHR1D::Algorithm()
286 : {
287 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D run");
288 0 : InitResources();
289 0 : LoadArgs();
290 0 : if (axisSize_ > 1) {
291 0 : AxisSync(FST_AXIS_ID);
292 : }
293 0 : PreSync();
294 0 : DoRepeatAllGatherNHR();
295 0 : PostSync();
296 0 : if (axisSize_ > 1) {
297 0 : AxisSync(SEC_AXIS_ID);
298 : }
299 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D end");
300 0 : return;
301 : }
302 :
303 0 : std::vector<uint64_t> CcuContextAllGatherNHR1D::GeneArgs(const CcuTaskArg &arg)
304 : {
305 0 : const CcuTaskArgAllGatherNHR1D *taskArg = dynamic_cast<const CcuTaskArgAllGatherNHR1D *>(&arg);
306 0 : if (taskArg == nullptr) {
307 0 : THROW<NullPtrException>(StringFormat("CcuContextAllGatherNHR1D::taskArg ptr is null"));
308 : }
309 : // input&output&buffer地址
310 0 : uint64_t inputAddr = taskArg->inputAddr_;
311 0 : uint64_t outputAddr = taskArg->outputAddr_;
312 0 : uint64_t token = taskArg->token_;
313 0 : uint64_t die0Size = taskArg->die0Size_;
314 0 : uint64_t die1Size = taskArg->die1Size_;
315 0 : uint64_t repeatNum = UINT64_MAX - taskArg->repeatNum_;
316 0 : uint64_t inputSliceStride = taskArg->inputSliceStride_;
317 0 : uint64_t outputSliceStride = taskArg->outputSliceStride_;
318 0 : uint64_t inputRepeatStride = taskArg->inputRepeatStride_;
319 0 : uint64_t outputRepeatStride = taskArg->outputRepeatStride_;
320 0 : uint64_t isInputOutputEqual = taskArg->isInputOutputEqual_;
321 :
322 0 : HCCL_INFO("[CcuContextAllGatherNHR1D] TaskArgs: inputAddr[%llu], outputAddr[%llu], "
323 : "die0Size[%llu], die1Size[%llu], repeatNum[%llu]"
324 : "inputSliceStride[%llu], outputSliceStride[%llu], inputRepeatStride[%llu], outputRepeatStride[%llu]",
325 : inputAddr, outputAddr, die0Size, die1Size, repeatNum, inputSliceStride, outputSliceStride,
326 : inputRepeatStride, outputRepeatStride);
327 :
328 : return {inputAddr, outputAddr, token,
329 : die0Size, die1Size, repeatNum,
330 : inputSliceStride, outputSliceStride, inputRepeatStride,
331 0 : outputRepeatStride, isInputOutputEqual};
332 : }
333 : } // namespace Hccl
|