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