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(
27 0 : const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, 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(
100 : transports[transportIdx] == nullptr,
101 : HCCL_ERROR("[CcuContextAllGatherNHR1D] Algorithm transport ptr is null"), );
102 0 : output_.push_back(
103 0 : CreateVariable((*transports[transportIdx]), OUTPUT_XN_ID)); // 获取transport中id=1的Var来传递output
104 0 : token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
105 : }
106 0 : output_.push_back(CreateVariable());
107 0 : token_.push_back(CreateVariable());
108 :
109 0 : srcMem_ = CreateMemory();
110 0 : dstMem_ = CreateMemory();
111 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] InitResources finished");
112 : }
113 :
114 0 : void CcuContextAllGatherNHR1D::PreSync()
115 : {
116 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync start");
117 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
118 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
119 0 : for (auto t : transports) {
120 0 : WriteVariableWithSignal(*t, output_[localSize_], OUTPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
121 0 : WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
122 : }
123 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
124 0 : for (auto& pair : indexMap_) {
125 0 : uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
126 0 : uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
127 0 : waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
128 : }
129 0 : for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
130 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
131 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
132 : }
133 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync end");
134 0 : }
135 :
136 0 : void CcuContextAllGatherNHR1D::PostSync()
137 : {
138 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
139 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
140 0 : for (auto& t : transports) {
141 0 : RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
142 : }
143 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
144 0 : for (auto& pair : indexMap_) {
145 0 : uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
146 0 : uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
147 0 : waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
148 : }
149 0 : for (uint32_t sId = 0; sId < signalNum_; sId++) {
150 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
151 : }
152 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] PostSync run finished");
153 0 : }
154 :
155 0 : void CcuContextAllGatherNHR1D::AxisSync(uint32_t signalIndex)
156 : {
157 0 : const uint32_t DIE_NUM = 2;
158 0 : if (signalIndex > 1) {
159 0 : THROW<InvalidParamsException>(
160 0 : StringFormat("[CcuContextAllGatherNHR1D] Unexpected SignalInex[%u]", signalIndex));
161 : }
162 0 : LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
163 0 : LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
164 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AxisSync run finished");
165 0 : return;
166 : }
167 :
168 0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHR()
169 : {
170 0 : tmpSliceOffset_ = 0;
171 0 : myrankInputSliceOffset_ = 0;
172 0 : for (u64 i = 0; i < rankId_; i++) {
173 0 : myrankInputSliceOffset_ += inputSliceStride_;
174 : }
175 0 : for (u64 i = 0; i < dimSize_; i++) {
176 0 : outputSliceOffset_[i] = tmpSliceOffset_;
177 0 : tmpSliceOffset_ += outputSliceStride_;
178 : }
179 0 : srcMem_.addr = input_;
180 0 : srcMem_.addr += myrankInputSliceOffset_;
181 0 : dstMem_.addr = output_[myRankIdx_];
182 0 : dstMem_.addr += outputSliceOffset_[rankId_];
183 0 : srcMem_.token = token_[myRankIdx_];
184 0 : dstMem_.token = token_[myRankIdx_];
185 0 : tmpCopyRepeatNum_ = repeatNum_;
186 0 : repeatTimeflag_ = 0;
187 0 : CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
188 : {
189 0 : tmpCopyRepeatNum_ += constVar1_;
190 0 : CCU_IF(repeatTimeflag_ != 0)
191 : {
192 0 : srcMem_.addr += inputRepeatStride_;
193 0 : dstMem_.addr += outputRepeatStride_;
194 0 : }
195 0 : CCU_IF(repeatTimeflag_ == 0)
196 : {
197 0 : if (axisId_ == 1) {
198 0 : srcMem_.addr += die0Size_;
199 0 : dstMem_.addr += die0Size_;
200 : }
201 0 : }
202 0 : CCU_IF(isInputOutputEqual_ == 0)
203 : {
204 0 : LocalCopy(dstMem_, srcMem_, axisId_ == 0 ? die0Size_ : die1Size_, localSignal_, 1 << rankId_);
205 0 : }
206 0 : CCU_IF(isInputOutputEqual_ != 0) { LocalPost(localSignal_, 1 << rankId_); }
207 0 : LocalWait(localSignal_, 1 << rankId_);
208 0 : repeatTimeflag_ = 1;
209 0 : }
210 :
211 0 : for (auto& nhrStepInfo : stepInfoVector_) {
212 0 : DoRepeatAllGatherNHRSingleStep(nhrStepInfo);
213 : }
214 0 : }
215 :
216 0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHRSingleStep(const NHRStepInfo& nhrStepInfo)
217 : {
218 0 : u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
219 0 : u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
220 0 : u32 sendSliceIdx = 0;
221 0 : CcuTransport* sendTransport = transports[toRankIdx];
222 0 : CcuTransport* recvTransport = transports[fromRankIdx];
223 0 : const std::vector<u32>& sendSliceIdxList = nhrStepInfo.txSliceIdxs;
224 0 : srcMem_.token = token_[myRankIdx_];
225 0 : dstMem_.token = token_[toRankIdx];
226 0 : for (u32 i = 0; i < sendSliceIdxList.size(); i++) { ////这里写的可能有问题
227 0 : sendSliceIdx = sendSliceIdxList[i];
228 0 : if (i != 0) {
229 0 : if (i % BIT_NUM_PER_CKE == 0) {
230 0 : LocalWait(localSignal_, (1 << BIT_NUM_PER_CKE) - 1);
231 : }
232 : }
233 0 : if (nhrStepInfo.step == 0) {
234 0 : srcMem_.addr = input_;
235 0 : srcMem_.addr += myrankInputSliceOffset_;
236 : } else {
237 0 : srcMem_.addr = output_[myRankIdx_];
238 0 : srcMem_.addr += outputSliceOffset_[sendSliceIdx];
239 : }
240 0 : dstMem_.addr = output_[toRankIdx];
241 0 : dstMem_.addr += outputSliceOffset_[sendSliceIdx];
242 0 : DoRepeatSendRecvSlices(nhrStepInfo.toRank, srcMem_, dstMem_, i % BIT_NUM_PER_CKE);
243 : }
244 :
245 0 : if (nhrStepInfo.step + 1 != stepInfoVector_.size()) {
246 0 : uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
247 0 : uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
248 0 : RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true);
249 0 : uint16_t recvSignalId = nhrStepInfo.fromRank / BIT_NUM_PER_CKE;
250 0 : uint16_t recvBit = 1 << (nhrStepInfo.fromRank % BIT_NUM_PER_CKE);
251 0 : RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_3, recvBit);
252 : }
253 0 : }
254 :
255 0 : void CcuContextAllGatherNHR1D::DoRepeatSendRecvSlices(
256 : const u32& toRank, CcuRep::Memory& src, CcuRep::Memory& dst, u32 signalIndex)
257 : {
258 0 : CcuTransport* sendTransport = transports[indexMap_[toRank]];
259 0 : const CcuRep::Variable& sliceSize = axisId_ == 0 ? die0Size_ : die1Size_;
260 0 : repeatTimeflag_ = 0;
261 0 : tmpCopyRepeatNum_ = repeatNum_;
262 0 : CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
263 : {
264 0 : tmpCopyRepeatNum_ += constVar1_;
265 0 : CCU_IF(repeatTimeflag_ == 1)
266 : {
267 0 : src.addr += inputRepeatStride_;
268 0 : dst.addr += outputRepeatStride_;
269 0 : }
270 0 : CCU_IF(repeatTimeflag_ == 0)
271 : {
272 0 : if (axisId_ == 1) {
273 0 : src.addr += die0Size_;
274 0 : dst.addr += die0Size_;
275 : }
276 0 : }
277 0 : Write(*sendTransport, dst, src, sliceSize, localSignal_, 1 << signalIndex);
278 0 : LocalWait(localSignal_, 1 << signalIndex);
279 0 : repeatTimeflag_ = 1;
280 0 : }
281 0 : }
282 :
283 0 : void CcuContextAllGatherNHR1D::Algorithm()
284 : {
285 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D run");
286 0 : InitResources();
287 0 : LoadArgs();
288 0 : if (axisSize_ > 1) {
289 0 : AxisSync(FST_AXIS_ID);
290 : }
291 0 : PreSync();
292 0 : DoRepeatAllGatherNHR();
293 0 : PostSync();
294 0 : if (axisSize_ > 1) {
295 0 : AxisSync(SEC_AXIS_ID);
296 : }
297 0 : HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D end");
298 0 : return;
299 : }
300 :
301 0 : std::vector<uint64_t> CcuContextAllGatherNHR1D::GeneArgs(const CcuTaskArg& arg)
302 : {
303 0 : const CcuTaskArgAllGatherNHR1D* taskArg = dynamic_cast<const CcuTaskArgAllGatherNHR1D*>(&arg);
304 0 : if (taskArg == nullptr) {
305 0 : THROW<NullPtrException>(StringFormat("CcuContextAllGatherNHR1D::taskArg ptr is null"));
306 : }
307 : // input&output&buffer地址
308 0 : uint64_t inputAddr = taskArg->inputAddr_;
309 0 : uint64_t outputAddr = taskArg->outputAddr_;
310 0 : uint64_t token = taskArg->token_;
311 0 : uint64_t die0Size = taskArg->die0Size_;
312 0 : uint64_t die1Size = taskArg->die1Size_;
313 0 : uint64_t repeatNum = UINT64_MAX - taskArg->repeatNum_;
314 0 : uint64_t inputSliceStride = taskArg->inputSliceStride_;
315 0 : uint64_t outputSliceStride = taskArg->outputSliceStride_;
316 0 : uint64_t inputRepeatStride = taskArg->inputRepeatStride_;
317 0 : uint64_t outputRepeatStride = taskArg->outputRepeatStride_;
318 0 : uint64_t isInputOutputEqual = taskArg->isInputOutputEqual_;
319 :
320 0 : HCCL_INFO(
321 : "[CcuContextAllGatherNHR1D] TaskArgs: inputAddr[%llu], outputAddr[%llu], "
322 : "die0Size[%llu], die1Size[%llu], repeatNum[%llu]"
323 : "inputSliceStride[%llu], outputSliceStride[%llu], inputRepeatStride[%llu], outputRepeatStride[%llu]",
324 : inputAddr, outputAddr, die0Size, die1Size, repeatNum, inputSliceStride, outputSliceStride, inputRepeatStride,
325 : outputRepeatStride);
326 :
327 : return {inputAddr, outputAddr, token,
328 : die0Size, die1Size, repeatNum,
329 : inputSliceStride, outputSliceStride, inputRepeatStride,
330 0 : outputRepeatStride, isInputOutputEqual};
331 : }
332 : } // namespace Hccl
|