Line data Source code
1 : /**
2 : * Copyright (c) 2026 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_to_all_mesh1d_2Die.h"
12 : #include "ccu_instruction_all_to_all_mesh1d_2Die.h"
13 :
14 : namespace Hccl {
15 :
16 : constexpr int CKE_IDX_0 = 0;
17 : constexpr int CKE_IDX_1 = 1;
18 : constexpr int CKE_IDX_2 = 2;
19 : constexpr int INPUT_XN_ID = 0;
20 : constexpr int OUPUT_XN_ID = 1;
21 : constexpr int TOKEN_XN_ID = 2;
22 :
23 : constexpr uint64_t CCU_MS_SIZE = 4096;
24 : constexpr uint64_t LOCAL_COPY_MS = 8;
25 :
26 0 : CcuContextAllToAllMesh1D2Die::CcuContextAllToAllMesh1D2Die(const CcuCtxArg &arg,
27 : const std::vector<CcuTransport *> &transports,
28 0 : const CcuTransportGroup &group)
29 0 : : CcuContextAlgBase(arg, transports, group)
30 : {
31 0 : const CcuCtxArgAllToAllMesh1D2Die *ctxArg = dynamic_cast<const CcuCtxArgAllToAllMesh1D2Die *>(&arg);
32 0 : if (ctxArg == nullptr) {
33 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::ctxArg ptr is null"));
34 : }
35 :
36 0 : rankId_ = ctxArg->rankId_;
37 0 : withMyRank_ = ctxArg->withMyRank_;
38 0 : rankGroup_ = ctxArg->rankGroup;
39 0 : if (ctxArg->dimSize_.size() > 0) {
40 0 : rankSize_ = ctxArg->dimSize_[0];
41 : }
42 0 : bitNumPerCKE_ = ctxArg->bitNum_;
43 0 : }
44 :
45 0 : void CcuContextAllToAllMesh1D2Die::InitResource()
46 : {
47 : // 创建Variable,用于交换地址及token
48 0 : u32 transportId = 0;
49 0 : if (transports.size() == 0) {
50 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die transports is empty"));
51 : }
52 0 : virRankSize = transports.size() + 1;
53 :
54 0 : for (u64 id = 0; id < transports.size(); id++) {
55 : // 非本地,使用远端Variable
56 0 : CHK_PRT_RET(transports[transportId] == nullptr,
57 : HCCL_ERROR("[CcuContextAllToAllMesh1D2Die] Algorithm transport ptr is null"), );
58 0 : input_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_0));
59 0 : output_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_1));
60 0 : token_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_2));
61 0 : transportId++;
62 : }
63 : // 最后一个位置放自己地址
64 0 : input_.push_back(CreateVariable());
65 0 : output_.push_back(CreateVariable());
66 0 : token_.push_back(CreateVariable());
67 :
68 0 : sliceSize_ = CreateVariable();
69 0 : inputSliceStride_ = CreateVariable();
70 0 : outputoffset_ = CreateVariable();
71 0 : outBuffBaseOff_ = CreateVariable();
72 0 : groupOpSize_ = CreateGroupOpSize();
73 :
74 0 : moConfig.loopCount = 8; // loop展开8次、16次
75 0 : moConfig.msInterleave = LOCAL_COPY_MS; // 一个loop 8个MS
76 0 : moConfig.memSlice = LOCAL_COPY_MS * CCU_MS_SIZE; // 32k
77 0 : if (moRes.executor.size() == 0) {
78 0 : moRes.executor = CreateBlockExecutor(moConfig.loopCount);
79 0 : moRes.maskSignal = CreateBlockMaskSignal(moConfig.loopCount);
80 0 : moRes.ccuBuffer = CreateBlockCcuBuffer(moConfig.loopCount * moConfig.msInterleave);
81 : }
82 :
83 0 : logicRankSize = withMyRank_ ? transports.size() + 1 : transports.size();
84 0 : signalNum_ = (rankSize_ + bitNumPerCKE_ - 1) / bitNumPerCKE_;
85 0 : HCCL_INFO("[CcuContextAlltoAll2Die] CtxArg: rankId_[%u], rankSize_[%u], signalNum_[%u]", rankId_, rankSize_, signalNum_);
86 0 : return;
87 : }
88 :
89 0 : void CcuContextAllToAllMesh1D2Die::LoadArgs()
90 : {
91 : // 从SQE load args,本rank需要的input、output地址等信息
92 : // inputAddr, outputAddr, tokenInfo, srcStride, srcOffset, dstOffset, groupOpSize
93 0 : Load(input_[virRankSize - 1]);
94 0 : Load(output_[virRankSize - 1]);
95 0 : Load(token_[virRankSize - 1]);
96 0 : Load(sliceSize_); // 本轮传输的分片大小
97 0 : Load(inputSliceStride_);
98 0 : Load(outputoffset_);
99 0 : Load(outBuffBaseOff_);
100 0 : Load(groupOpSize_);
101 0 : return;
102 : }
103 :
104 0 : void CcuContextAllToAllMesh1D2Die::PreSync()
105 : {
106 0 : if (withMyRank_) {
107 0 : uint16_t logicId = rankId_ % logicRankSize;
108 0 : selfBit = 1 << logicId;
109 0 : allBit = ((1 << logicRankSize) - 1) & (~(1 << logicId));
110 0 : for (auto t : transports) {
111 : // (transport, param, paramID, SemID, mask)
112 0 : WriteVariableWithSignal(*t, output_[virRankSize - 1], OUPUT_XN_ID, CKE_IDX_1,
113 0 : selfBit); // index = 1,传递output信息
114 0 : WriteVariableWithSignal(*t, token_[virRankSize - 1], TOKEN_XN_ID, CKE_IDX_2,
115 0 : selfBit); // index = 2,传递token信息
116 : }
117 0 : GroupWait(*transportGroup, CKE_IDX_1, allBit); // index = 1,传递output信息
118 0 : GroupWait(*transportGroup, CKE_IDX_2, allBit); // index = 2,传递token信息
119 : } else {
120 0 : uint16_t selfSignalId = rankId_ / bitNumPerCKE_;
121 0 : uint16_t selfBit = 1 << (rankId_ % bitNumPerCKE_);
122 0 : for (auto t : transports) {
123 : // (transport, param, paramID, SemID, mask)
124 0 : WriteVariableWithSignal(*t, output_[virRankSize - 1], OUPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1,
125 : selfBit); // index = 1,传递output信息
126 0 : WriteVariableWithSignal(*t, token_[virRankSize - 1], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2,
127 : selfBit); // index = 2,传递token信息
128 : }
129 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
130 0 : for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
131 0 : waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
132 0 : if (sId == selfSignalId) {
133 0 : waitBitVector[sId] = 0;
134 : }
135 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]); // index = 1,传递output信息
136 0 : GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]); // index = 2,传递token信息
137 : }
138 0 : }
139 0 : return;
140 : }
141 :
142 0 : void CcuContextAllToAllMesh1D2Die::PostSync()
143 : {
144 0 : if (withMyRank_) {
145 0 : uint16_t logicId = rankId_ % logicRankSize;
146 0 : selfBit = 1 << logicId;
147 0 : allBit = ((1 << logicRankSize) - 1) & (~(1 << logicId));
148 0 : for (auto t : transports) {
149 0 : if (t == nullptr) {
150 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::Algorithm transport ptr is null"));
151 : }
152 0 : RemotePost(*t, CKE_IDX_0, selfBit);
153 : }
154 0 : GroupWait(*transportGroup, CKE_IDX_0, allBit);
155 : } else {
156 0 : uint16_t selfSignalId = rankId_ / bitNumPerCKE_;
157 0 : uint16_t selfBit = 1 << (rankId_ % bitNumPerCKE_);
158 0 : for (auto t : transports) {
159 0 : if (t == nullptr) {
160 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::Algorithm transport ptr is null"));
161 : }
162 0 : RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
163 : }
164 0 : std::vector<uint16_t> waitBitVector(signalNum_, 0);
165 0 : for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
166 0 : waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
167 0 : if (sId == selfSignalId) {
168 0 : waitBitVector[sId] = 0;
169 : }
170 0 : GroupWait(*transportGroup, CKE_IDX_0, allBit);
171 : }
172 0 : }
173 :
174 0 : return;
175 : }
176 :
177 0 : uint32_t CcuContextAllToAllMesh1D2Die::CalcDstRank(uint32_t peerId) const
178 : {
179 0 : if (peerId > rankGroup_.size()) {
180 0 : THROW<InvalidParamsException>(
181 0 : StringFormat("[CcuContextAllToAllMesh1D2Die][CalcDstRank] Unexpected peerId[%u]", peerId));
182 : }
183 0 : return rankGroup_[peerId];
184 : }
185 :
186 0 : void CcuContextAllToAllMesh1D2Die::DoRepeatAllToAll()
187 : {
188 : // 创建GSA, src为本地的各片HBM地址GSA列表,dst为所有对端的HBM地址GSA列表
189 0 : std::vector<CcuRep::Memory> src;
190 0 : for (uint64_t rankIdx = 0; rankIdx < logicRankSize; rankIdx++) {
191 0 : src.push_back(CreateMemory());
192 : }
193 0 : std::vector<CcuRep::Memory> dst;
194 0 : for (uint64_t rankIdx = 0; rankIdx < logicRankSize; rankIdx++) {
195 0 : dst.push_back(CreateMemory());
196 : }
197 :
198 : // 考虑stride信息
199 0 : for (uint64_t r = 0; r < logicRankSize; r++) {
200 0 : const u32 dstRank = CalcDstRank(r);
201 :
202 0 : src[r].token = token_[r];
203 0 : dst[r].token = token_[r];
204 :
205 0 : src[r].addr = input_[virRankSize - 1];
206 0 : dst[r].addr = output_[r];
207 0 : dst[r].addr += outputoffset_;
208 0 : for(uint64_t i = 0; i < dstRank; i++){
209 0 : src[r].addr += inputSliceStride_;
210 : }
211 : }
212 :
213 : // all2all 数据搬运
214 0 : u32 transportIdx = 0;
215 0 : if (withMyRank_) {
216 0 : uint64_t allBit_ = withMyRank_ ? ((1 << logicRankSize) - 1) & (~(1 << transports.size())) : (1 << logicRankSize) - 1;
217 0 : CcuRep::MaskSignal locMask = CreateMaskSignal();
218 0 : for (uint64_t r = 0; r < logicRankSize; r++) {
219 0 : if (withMyRank_ && r == logicRankSize - 1) {
220 0 : LocalCopyByLoopGroup(dst[r], src[r]);
221 0 : continue;
222 : }
223 0 : Write(*transports[transportIdx], dst[r], src[r], sliceSize_, locMask, 1 << r);
224 0 : transportIdx++;
225 : }
226 0 : LocalWait(locMask, allBit_);
227 0 : } else {
228 0 : vector<CcuRep::MaskSignal> locMask;
229 0 : uint16_t signalNum = logicRankSize / bitNumPerCKE_;
230 0 : std::vector<uint16_t> waitBitVector(signalNum, 0);
231 0 : for (uint16_t sId = 0; sId < signalNum; sId++) {
232 0 : locMask.push_back(CreateMaskSignal());
233 : }
234 0 : for (uint16_t r = 0; r < logicRankSize; r++) {
235 0 : uint16_t rmtSignalId = r / bitNumPerCKE_;
236 0 : uint16_t rmtSignalBit = 1 << (r % bitNumPerCKE_);
237 0 : Write(*transports[transportIdx], dst[r], src[r], sliceSize_, locMask[rmtSignalId], rmtSignalBit);
238 0 : transportIdx++;
239 : }
240 0 : for (uint16_t sId = 0; sId < signalNum; sId++) {
241 0 : waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
242 0 : LocalWait(locMask[sId], waitBitVector[sId]);
243 : }
244 0 : }
245 0 : }
246 :
247 0 : void CcuContextAllToAllMesh1D2Die::CreateLocalCopyLoop()
248 : {
249 0 : std::string loopType = "all_to_all";
250 0 : if (registeredLoop.find(loopType) != registeredLoop.end()) {
251 0 : return;
252 : }
253 :
254 0 : for (uint32_t index = 0; index < 2; index++) { // 需要2个Loop
255 0 : CcuRep::Variable len = CreateVariable();
256 0 : CcuRep::Memory src = CreateMemory();
257 0 : CcuRep::Memory dst = CreateMemory();
258 0 : CcuRep::LoopBlock lb(this, loopType + "_localcopy_loop_" + std::to_string(index));
259 0 : lb(src, dst, len);
260 :
261 0 : std::vector<CcuRep::CcuBuffer> bufs;
262 0 : CcuRep::MaskSignal sem = moRes.maskSignal[index];
263 0 : for (uint32_t i = 0; i < LOCAL_COPY_MS; i++) {
264 0 : bufs.push_back(moRes.ccuBuffer[i]);
265 : }
266 :
267 0 : LocalCopy(bufs[0], src, len, sem);
268 0 : LocalWait(sem);
269 0 : LocalCopy(dst, bufs[0], len, sem);
270 0 : LocalWait(sem);
271 0 : }
272 0 : registeredLoop.insert(loopType);
273 0 : return;
274 0 : }
275 :
276 0 : void CcuContextAllToAllMesh1D2Die::LocalCopyByLoopGroup(CcuRep::Memory dst, CcuRep::Memory src)
277 : {
278 0 : CreateLocalCopyLoop();
279 :
280 0 : CCU_IF(groupOpSize_.addrOffset != 0)
281 : {
282 0 : CcuRep::Variable loopParam = CreateVariable();
283 0 : loopParam = CcuRep::GetLoopParam(0, moConfig.memSlice * moConfig.loopCount, 0);
284 0 : loopParam += groupOpSize_.loopParam;
285 :
286 0 : CcuRep::Variable sliceSize = CreateVariable();
287 0 : sliceSize = moConfig.memSlice;
288 0 : auto lc = Loop("all_to_all_localcopy_loop_0")(src, dst, sliceSize);
289 :
290 0 : CcuRep::Variable paraCfg = CreateVariable();
291 0 : paraCfg = CcuRep::GetParallelParam(moConfig.loopCount - 1, 0, 1);
292 0 : CcuRep::Variable offsetCfg = CreateVariable();
293 0 : offsetCfg = CcuRep::GetOffsetParam(moConfig.memSlice, moConfig.msInterleave, 1);
294 0 : LoopGroup({lc}, {loopParam}, paraCfg, offsetCfg);
295 0 : }
296 :
297 0 : CCU_IF(groupOpSize_.parallelParam != 0)
298 : {
299 0 : CcuRep::Condition cond(this, groupOpSize_.parallelParam != 0);
300 :
301 0 : src.addr += groupOpSize_.addrOffset;
302 0 : dst.addr += groupOpSize_.addrOffset;
303 0 : auto lc0 = Loop("all_to_all_localcopy_loop_0")(src, dst, groupOpSize_.residual);
304 :
305 0 : src.addr += groupOpSize_.residual;
306 0 : dst.addr += groupOpSize_.residual;
307 0 : CcuRep::Variable sliceSize = CreateVariable();
308 0 : sliceSize = moConfig.memSlice;
309 0 : auto lc1 = Loop("all_to_all_localcopy_loop_1")(src, dst, sliceSize);
310 :
311 0 : CcuRep::Variable loopCfg0 = CreateVariable();
312 0 : loopCfg0 = CcuRep::GetLoopParam(0, 0, 1);
313 0 : CcuRep::Variable loopCfg1 = CreateVariable();
314 0 : loopCfg1 = CcuRep::GetLoopParam(0, 0, 1);
315 0 : CcuRep::Variable offsetCfg = CreateVariable();
316 0 : offsetCfg = CcuRep::GetOffsetParam(moConfig.memSlice, moConfig.msInterleave, 1);
317 0 : LoopGroup({lc0, lc1}, {loopCfg0, loopCfg1}, groupOpSize_.parallelParam, offsetCfg);
318 0 : }
319 0 : }
320 :
321 0 : void CcuContextAllToAllMesh1D2Die::Algorithm()
322 : {
323 0 : HCCL_INFO("[ccuAllToAllMesh1D2Die_context] AllToAllMesh1D2Die run.");
324 0 : InitResource();
325 :
326 0 : LoadArgs();
327 :
328 0 : PreSync();
329 :
330 0 : DoRepeatAllToAll();
331 :
332 0 : PostSync();
333 0 : HCCL_INFO("[ccuAllToAllMesh1D2Die_context] AllToAllMesh1D2Die end.");
334 0 : return;
335 : }
336 :
337 0 : std::vector<uint64_t> CcuContextAllToAllMesh1D2Die::GeneArgs(const CcuTaskArg &arg)
338 : {
339 0 : const CcuTaskArgAllToAllMesh1D2Die *taskArg = dynamic_cast<const CcuTaskArgAllToAllMesh1D2Die *>(&arg);
340 0 : if (taskArg == nullptr) {
341 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::taskArg ptr is null"));
342 : }
343 0 : uint64_t inputAddr = taskArg->inputAddr_;
344 0 : uint64_t outputAddr = taskArg->outputAddr_;
345 0 : uint64_t tokenInfo = taskArg->token_;
346 0 : uint64_t sliceSize = taskArg->sliceSize_;
347 0 : uint64_t inputSliceStride = taskArg->inputSliceStride_;
348 0 : uint64_t outputSliceStride = taskArg->outputSliceStride_ * rankId_;
349 0 : uint64_t outBuffBaseOff = taskArg->outBuffBaseOff_;
350 :
351 0 : auto goSize = CalGoSize(sliceSize);
352 0 : HCCL_INFO("[CcuContextAllToAllMesh1D2Die] inputAddr[%llu], outputAddr[%llu], sliceSize[%llu], "
353 : "inputSliceStride[%llu], outputSliceStride[%llu], outBuffBaseOff[%llu].",
354 : inputAddr, outputAddr, sliceSize, inputSliceStride, outputSliceStride, outBuffBaseOff);
355 :
356 : return {inputAddr, outputAddr, tokenInfo, sliceSize, inputSliceStride, outputSliceStride,
357 0 : outBuffBaseOff, goSize[0], goSize[1], goSize[2], goSize[3]};
358 0 : }
359 :
360 : } // namespace Hccl
|