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