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_to_all_v_mesh1d.h"
12 : #include "ccu_instruction_all_to_all_v_mesh1d.h"
13 :
14 : namespace Hccl {
15 : constexpr int OUTPUT_XN_ID = 0;
16 : constexpr int TOKEN_XN_ID = 1;
17 : constexpr int CKE_IDX_0 = 0;
18 : constexpr int CKE_IDX_1 = 1;
19 : constexpr int CKE_IDX_2 = 2;
20 :
21 0 : CcuContextAllToAllVMesh1D::CcuContextAllToAllVMesh1D(const CcuCtxArg &arg, const std::vector<CcuTransport*> &transports,
22 0 : const CcuTransportGroup &group)
23 0 : : CcuContextAlgBase(arg, transports, group)
24 : {
25 0 : const CcuCtxArgAllToAllVMesh1D *ctxArg = dynamic_cast<const CcuCtxArgAllToAllVMesh1D *>(&arg);
26 0 : if (ctxArg == nullptr) {
27 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllVMesh1D::ctxArg ptr is null"));
28 : }
29 0 : rankId_ = ctxArg->rankId;
30 0 : if (ctxArg->dimSize.size() > 0) {
31 0 : rankSize_ = ctxArg->dimSize[0];
32 : }
33 0 : loadFromMem = ctxArg->loadFromMem;
34 :
35 0 : if (transports.size() == 0) {
36 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllVMesh1D transports is empty"));
37 : }
38 0 : }
39 :
40 0 : void CcuContextAllToAllVMesh1D::PreSync()
41 : {
42 0 : CcuRep::Variable tempDst = CreateVariable();
43 0 : u32 transportId = 0;
44 0 : for (u32 id = 0; id < rankSize_; id++) {
45 0 : if (id == rankId_) {
46 0 : continue;
47 : }
48 0 : tempDst = output_[rankId_];
49 0 : tempDst += sendRecvInfo_[id].recvOffset;
50 : // index = 0,传递output信息
51 0 : WriteVariableWithSignal(*transports[transportId], tempDst, OUTPUT_XN_ID, CKE_IDX_1, selfBit_);
52 : // index = 1,传递token信息
53 0 : WriteVariableWithSignal(*transports[transportId], token_[rankId_], TOKEN_XN_ID, CKE_IDX_2, selfBit_);
54 0 : transportId++;
55 : }
56 :
57 0 : GroupWait(*transportGroup, CKE_IDX_1, allOtherBit_); // index = 1,传递output信息
58 0 : GroupWait(*transportGroup, CKE_IDX_2, allOtherBit_); // index = 2,传递token信息
59 0 : }
60 :
61 0 : void CcuContextAllToAllVMesh1D::PostSync()
62 : {
63 0 : for (auto t : transports) {
64 0 : if (t == nullptr) {
65 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllVMesh1D::Algorithm transport ptr is null"));
66 : }
67 0 : RemotePost(*t, CKE_IDX_0, selfBit_);
68 : }
69 0 : GroupWait(*transportGroup, CKE_IDX_0, allOtherBit_);
70 0 : }
71 :
72 0 : void CcuContextAllToAllVMesh1D::CreateVariables()
73 : {
74 0 : u32 transportId = 0;
75 0 : input_.push_back(CreateVariable());
76 0 : output_.reserve(rankSize_);
77 0 : token_.reserve(rankSize_);
78 0 : for (u32 id = 0; id < rankSize_; id++) {
79 0 : if (id == rankId_) {
80 0 : output_.push_back(CreateVariable());
81 0 : token_.push_back(CreateVariable());
82 : }
83 : else { // 非本地,使用远端Variable
84 0 : CHK_PRT_RET(transports[transportId] == nullptr || transportId >= transports.size(),
85 : HCCL_ERROR("[CcuContextAllToAllVMesh1D] Algorithm transport ptr is null or transportIdx is out of bounds"),);
86 0 : output_.push_back(CreateVariable((*transports[transportId]), OUTPUT_XN_ID)); // 与远端交换本卡的接收地址
87 0 : token_.push_back(CreateVariable((*transports[transportId]), TOKEN_XN_ID));
88 0 : transportId++;
89 : }
90 : }
91 :
92 0 : src_.reserve(rankSize_);
93 0 : dst_.reserve(rankSize_);
94 0 : for (uint32_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
95 0 : src_.push_back(CreateMemory());
96 0 : dst_.push_back(CreateMemory());
97 : }
98 :
99 0 : srcOffset_ = CreateVariable();
100 0 : dstOffset_ = CreateVariable();
101 0 : a2avXnAddr_ = CreateVariable();
102 :
103 : // 前同步。交换信息,将本Rank load的in\out等地址信息写到所有对端的对应Variable中,并同步
104 0 : selfBit_ = 1 << rankId_; // 本rank的mask
105 0 : allBit_ = (1 << rankSize_) - 1; // 等待包含自身的全部对端
106 0 : allOtherBit_ = ((1 << rankSize_) - 1) & (~(1 << rankId_)); // 等待其他所有对端
107 :
108 0 : locMask_ = CreateMaskSignal();
109 : // all2allv 数据搬运
110 0 : completedRankCount_ = CreateVariable();
111 0 : xnMaxTransportSize_ = CreateVariable();
112 0 : xnMaxTransportGoSize_ = CreateGroupOpSize();
113 0 : localTailGoSize_ = CreateGroupOpSize();
114 0 : xnConst1_ = CreateVariable();
115 :
116 0 : xnLength_ = CreateVariable();
117 0 : xnLength_ = 8; // xn长度为8byte
118 : }
119 :
120 0 : void CcuContextAllToAllVMesh1D::LoadArgs()
121 : {
122 : // 从SQE load args,本rank需要的input、output地址等信息
123 : // inputAddr, outputAddr, tokenInfo, srcStride, dstStride, srcOffset, dstOffset
124 0 : Load(input_[0]);
125 0 : Load(output_[rankId_]); // load的目的存放寄存器
126 0 : Load(token_[rankId_]);
127 0 : Load(srcOffset_);
128 0 : Load(dstOffset_);
129 0 : Load(localTailGoSize_);
130 0 : if (loadFromMem) {
131 0 : Load(a2avXnAddr_);
132 : } else {
133 0 : Load(xnMaxTransportGoSize_);
134 : }
135 :
136 : // 恢复当前卡对所有卡的收发信息
137 0 : sendRecvInfo_.resize(rankSize_);
138 0 : for(uint32_t i = 0; i < rankSize_; i++){
139 0 : sendRecvInfo_[i].tailSize = CreateVariable();
140 0 : sendRecvInfo_[i].loopNum = CreateVariable();
141 0 : sendRecvInfo_[i].sendOffset = CreateVariable();
142 0 : sendRecvInfo_[i].recvOffset = CreateVariable();
143 : }
144 0 : LoadAll2allSendRecvInfo(sendRecvInfo_);
145 0 : }
146 :
147 0 : void CcuContextAllToAllVMesh1D::CalcGroupSrcDst()
148 : {
149 0 : for (uint32_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
150 0 : src_[rankIdx].token = token_[rankIdx];
151 0 : dst_[rankIdx].token = token_[rankIdx];
152 :
153 : // src_[rankIdx] = usrInAddr + sendoffset + srcOffset_
154 0 : src_[rankIdx].addr = input_[0];
155 0 : src_[rankIdx].addr += sendRecvInfo_[rankIdx].sendOffset;
156 0 : src_[rankIdx].addr += srcOffset_;
157 :
158 : // dst_[r] = recvBuf[r] + recvOffset + dstOffset_
159 0 : if (rankIdx == rankId_) {
160 : // 写目的端为本端时需要特殊处理:使用接收基地址 + 块地址offset + 已发送数据量
161 0 : dst_[rankIdx].addr = output_[rankId_];
162 0 : dst_[rankIdx].addr += sendRecvInfo_[rankIdx].recvOffset;
163 0 : dst_[rankIdx].addr += dstOffset_;
164 : } else {
165 : // 对端交换的接收块起始地址 + 已接收的数据偏移
166 0 : dst_[rankIdx].addr = output_[rankIdx];
167 0 : dst_[rankIdx].addr += dstOffset_;
168 : }
169 : }
170 0 : }
171 :
172 0 : void CcuContextAllToAllVMesh1D::DoAll2AllVMultiLoop()
173 : {
174 0 : HCCL_DEBUG("[CcuContextAllToAllVMesh1D] alltoallv mesh 1d use GroupCopy start");
175 0 : xnMaxTransportSize_ = UB_MAX_TRANS_SIZE;
176 0 : completedRankCount_ = 0;
177 0 : xnConst1_ = 1;
178 0 : u32 transportId = 0;
179 0 : CCU_WHILE(completedRankCount_ != rankSize_) { // 循环发送数据,直到所有对端数据都发送完成
180 0 : for(uint32_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) { // 循环发送所有对端数据
181 0 : if (rankIdx == rankId_) {
182 0 : continue;
183 : }
184 0 : CCU_IF(sendRecvInfo_[rankIdx].loopNum == UINT64_MAX) { // 已经完成,直接置位完成信号
185 0 : LocalPost(locMask_, (1 << rankIdx));
186 0 : }
187 0 : CCU_IF(sendRecvInfo_[rankIdx].loopNum != UINT64_MAX) { // 还没有完成,则继续循环
188 0 : CCU_IF(sendRecvInfo_[rankIdx].loopNum == UINT64_MAX - 1) { // 最后一轮循环, 发送尾块数据
189 0 : CCU_IF(sendRecvInfo_[rankIdx].tailSize == 0) { // 尾块数据量为 0,则不需要发送尾块数据
190 0 : LocalPost(locMask_, (1 << rankIdx));
191 0 : }
192 0 : CCU_IF(sendRecvInfo_[rankIdx].tailSize != 0) { // 尾块数据量不为 0,则需要发送尾块数据
193 0 : Write(*transports[transportId], dst_[rankIdx], src_[rankIdx], sendRecvInfo_[rankIdx].tailSize,
194 0 : locMask_, 1 << rankIdx);
195 0 : }
196 0 : completedRankCount_ += xnConst1_; // 之后一轮循环完成,更新已完成的rank数
197 0 : }
198 0 : CCU_IF(sendRecvInfo_[rankIdx].loopNum != UINT64_MAX - 1) { // 未完成,则继续循环,发送整块数据
199 0 : Write(*transports[transportId], dst_[rankIdx], src_[rankIdx], xnMaxTransportSize_, locMask_,
200 0 : 1 << rankIdx);
201 : // 更新偏移
202 0 : src_[rankIdx].addr += xnMaxTransportSize_;
203 0 : dst_[rankIdx].addr += xnMaxTransportSize_;
204 0 : }
205 0 : sendRecvInfo_[rankIdx].loopNum += xnConst1_;
206 0 : }
207 0 : transportId++;
208 : }
209 0 : CCU_IF(sendRecvInfo_[rankId_].loopNum == UINT64_MAX) { // 已经完成,直接置位完成信号
210 0 : LocalPost(locMask_, (1 << rankId_));
211 0 : }
212 :
213 0 : CCU_IF(sendRecvInfo_[rankId_].loopNum != UINT64_MAX) { // 还没有完成,则继续循环
214 0 : CCU_IF(sendRecvInfo_[rankId_].loopNum == UINT64_MAX - 1) { // 最后一轮循环, 发送尾块数据
215 0 : CCU_IF(sendRecvInfo_[rankId_].tailSize == 0) { // 尾块数据量为 0,则不需要发送尾块数据
216 0 : LocalPost(locMask_, (1 << rankId_));
217 0 : }
218 0 : CCU_IF(sendRecvInfo_[rankId_].tailSize != 0) { // 尾块数据量不为 0,则需要发送尾块数据
219 0 : GroupCopy(dst_[rankId_], src_[rankId_], localTailGoSize_);
220 0 : LocalPost(locMask_, 1 << rankId_);
221 0 : }
222 0 : completedRankCount_ += xnConst1_; // 之后一轮循环完成,更新已完成的rank数
223 0 : }
224 0 : CCU_IF(sendRecvInfo_[rankId_].loopNum != UINT64_MAX - 1) { // 未完成,则继续循环,发送整块数据
225 0 : GroupCopy(dst_[rankId_], src_[rankId_], xnMaxTransportGoSize_);
226 0 : LocalPost(locMask_, 1 << rankId_);
227 : // 更新偏移
228 0 : src_[rankId_].addr += xnMaxTransportSize_;
229 0 : dst_[rankId_].addr += xnMaxTransportSize_;
230 0 : }
231 0 : sendRecvInfo_[rankId_].loopNum += xnConst1_;
232 0 : }
233 : // 等待本轮发送完成
234 0 : LocalWait(locMask_, allBit_);
235 0 : }
236 0 : }
237 :
238 0 : void CcuContextAllToAllVMesh1D::Algorithm()
239 : {
240 0 : HCCL_INFO("[ccuAllToAllVMesh1D_context] AllToAllVMesh1D run");
241 0 : CreateVariables();
242 0 : LoadArgs();
243 0 : PreSync();
244 : // 创建GSA, src为本地的各片HBM地址GSA列表,dst为所有对端的HBM地址GSA列表
245 0 : CalcGroupSrcDst();
246 0 : DoAll2AllVMultiLoop();
247 : // 后同步
248 0 : PostSync();
249 0 : HCCL_INFO("[AllToAllAlgo] AllToAllMesh1D end");
250 0 : return;
251 : }
252 :
253 0 : std::vector<uint64_t> CcuContextAllToAllVMesh1D::GeneArgs(const CcuTaskArg &arg)
254 : {
255 0 : const CcuTaskArgAllToAllVMesh1D *taskArg = dynamic_cast<const CcuTaskArgAllToAllVMesh1D *>(&arg);
256 0 : if (taskArg == nullptr) {
257 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllVMesh1D::taskArg ptr is null"));
258 : }
259 0 : uint64_t inputAddr = taskArg->inputAddr_;
260 0 : uint64_t outputAddr = taskArg->outputAddr_;
261 0 : uint64_t tokenInfo = taskArg->token_;
262 :
263 0 : uint64_t srcOffset = taskArg->srcOffset_;
264 0 : uint64_t dstOffset = taskArg->dstOffset_;
265 :
266 0 : HCCL_INFO("[AllToAllVAlgo] inputAddr[%llu], outputAddr[%llu],"
267 : "srcOffset[%llu], dstOffset[%llu]",
268 : inputAddr, outputAddr, srcOffset, dstOffset);
269 0 : std::vector<uint64_t> processReturn = {inputAddr, outputAddr, tokenInfo, srcOffset, dstOffset};
270 :
271 0 : u64 localTailSize = taskArg->localSendRecvInfo_.sendLength[rankId_] % UB_MAX_TRANS_SIZE;
272 0 : auto localTailGoSize = CalGoSize(localTailSize);
273 0 : for (auto val : localTailGoSize) {
274 0 : processReturn.push_back(val);
275 : }
276 :
277 0 : if (loadFromMem) {
278 0 : processReturn.push_back(0); // 空地址占位,保证参数个数与load个数一致
279 0 : return processReturn;
280 : }
281 :
282 0 : uint64_t xnMaxTransportSize = UB_MAX_TRANS_SIZE;
283 0 : HCCL_INFO("[CcuContextAllToAllVMesh1D][GeneArgs] CalGoSize size[%llu]", xnMaxTransportSize);
284 0 : auto xnMaxTransportGoSize = CalGoSize(xnMaxTransportSize);
285 0 : for (auto val : xnMaxTransportGoSize) {
286 0 : processReturn.push_back(val);
287 : }
288 :
289 0 : uint64_t rankSize = taskArg->sliceSize_.size();
290 0 : for (uint64_t i = 0; i < rankSize; i++) {
291 0 : uint64_t tailSize = taskArg->localSendRecvInfo_.sendLength[i] % UB_MAX_TRANS_SIZE;
292 0 : uint64_t loopNum = UINT64_MAX - 1 - (taskArg->localSendRecvInfo_.sendLength[i] / UB_MAX_TRANS_SIZE);
293 0 : uint64_t sendOffset = taskArg->localSendRecvInfo_.sendOffset[i];
294 0 : uint64_t recvOffset = taskArg->localSendRecvInfo_.recvOffset[i];
295 0 : HCCL_INFO("[CcuContextAllToAllVMesh1D][GeneArgs] CalGoSize size[%llu]", tailSize);
296 0 : processReturn.push_back(tailSize);
297 0 : processReturn.push_back(loopNum);
298 0 : processReturn.push_back(sendOffset);
299 0 : processReturn.push_back(recvOffset);
300 0 : HCCL_INFO("[AllToAllVAlgo] rankIdx[i] taskArg->sliceSize[%llu]," \
301 : "loopNum[%llu]," \
302 : "taskArg->localSendRecvInfo.sendOffset[%llu]," \
303 : "taskArg->localSendRecvInfo.recvOffset[%llu]",
304 : taskArg->sliceSize_[i], loopNum, taskArg->localSendRecvInfo_.sendOffset[i],
305 : taskArg->localSendRecvInfo_.recvOffset[i]);
306 : }
307 0 : return processReturn;
308 0 : }
309 :
310 0 : void CcuContextAllToAllVMesh1D::LoadAll2allSendRecvInfo(std::vector<A2AsingleSendRecvInfo> &sendRecvInfo)
311 : {
312 0 : if (loadFromMem) {
313 : //连续加载ranksize个sendSize,loopNum,sendOffset,receiveOffset
314 0 : u32 argsCount = sendRecvInfo.size() * 4;
315 0 : std::vector<CcuRep::Variable> tempArgs(argsCount);
316 0 : HCCL_INFO("AllToAllVAlgo LoadArgsFromMem, argsCount: [%u]", argsCount);
317 0 : for (uint32_t i = 0; i < tempArgs.size(); ++i) {
318 0 : tempArgs[i] = CreateContinuousVariable();
319 : }
320 0 : LoadVariable(a2avXnAddr_, tempArgs[0], argsCount);
321 :
322 : // 赋值给对应的 XN
323 0 : u32 argIdx = 0;
324 0 : for (uint32_t i = 0; i < sendRecvInfo.size(); i++) {
325 0 : sendRecvInfo[i].tailSize = tempArgs[argIdx];
326 0 : argIdx++;
327 0 : sendRecvInfo[i].loopNum = UINT64_MAX - 1;
328 0 : argIdx++;
329 0 : sendRecvInfo[i].sendOffset = tempArgs[argIdx];
330 0 : argIdx++;
331 0 : sendRecvInfo[i].recvOffset = tempArgs[argIdx];
332 0 : argIdx++;
333 : }
334 0 : } else {
335 0 : for(uint32_t i = 0; i < rankSize_; i++){
336 0 : Load(sendRecvInfo[i].tailSize);
337 0 : Load(sendRecvInfo[i].loopNum);
338 0 : Load(sendRecvInfo[i].sendOffset);
339 0 : Load(sendRecvInfo[i].recvOffset);
340 : }
341 : }
342 0 : }
343 :
344 0 : void CcuContextAllToAllVMesh1D::RefreshArgs(CollOpParams opParams, u32 rankSize, std::vector<uint64_t> &args, const u32 myRank)
345 : {
346 : uint64_t inputAddr;
347 : uint64_t outputAddr;
348 0 : uint64_t token = 0;
349 0 : uint64_t srcOffset = 0;
350 0 : uint64_t dstOffset = 0;
351 :
352 0 : inputAddr = reinterpret_cast<uint64_t>(opParams.sendBuf);
353 0 : outputAddr = reinterpret_cast<uint64_t>(opParams.recvBuf);
354 :
355 0 : args.push_back(inputAddr);
356 0 : args.push_back(outputAddr);
357 0 : args.push_back(token);
358 0 : args.push_back(srcOffset);
359 0 : args.push_back(dstOffset);
360 :
361 : //配置本地拷贝的moConfig参数
362 0 : u32 loopCount = LOCAL_COPY_MS_PER_LOOP;
363 0 : u32 memSlice = CCU_MS_LOCAL_COPY_LOOP_COUNT * CcuRep::CCU_MS_SIZE;
364 0 : GroupOpConfig moConfig{CcuRep::CCU_MS_INTERLEAVE, loopCount, memSlice};
365 :
366 0 : u64 mySendCounts = *(static_cast<const u64 *>(opParams.all2AllVDataDes.sendCounts) + myRank);
367 0 : u64 mySendLength = mySendCounts * DataTypeSizeGet(opParams.all2AllVDataDes.sendType);
368 0 : uint64_t localTailSize = mySendLength % UB_MAX_TRANS_SIZE;
369 0 : auto localTailGoSize = CcuContext::CalGoSizeStatic(localTailSize, moConfig);
370 0 : for (auto val : localTailGoSize) {
371 0 : args.push_back(val);
372 : }
373 :
374 0 : uint64_t xnMaxTransportSize = UB_MAX_TRANS_SIZE;
375 0 : HCCL_INFO("[CcuContextAllToAllVMesh1D][RefreshArgs] CalGoSizeStatic size [%llu]", xnMaxTransportSize);
376 0 : auto xnMaxTransportGoSize = CcuContext::CalGoSizeStatic(xnMaxTransportSize, moConfig);
377 0 : for (auto val : xnMaxTransportGoSize) {
378 0 : args.push_back(val);
379 : }
380 :
381 :
382 0 : for (u32 i = 0; i < rankSize; i++) {
383 0 : u64 curSendCounts = *(static_cast<const u64 *>(opParams.all2AllVDataDes.sendCounts) + i);
384 0 : u64 curSendDispls = *(static_cast<const u64 *>(opParams.all2AllVDataDes.sdispls) + i);
385 0 : u64 sendLength = curSendCounts * DataTypeSizeGet(opParams.all2AllVDataDes.sendType);
386 0 : u64 sendOffset = curSendDispls * DataTypeSizeGet(opParams.all2AllVDataDes.sendType);
387 :
388 0 : u64 curRecvDispls = *(static_cast<const u64 *>(opParams.all2AllVDataDes.rdispls) + i);
389 0 : u64 recvOffset = curRecvDispls * DataTypeSizeGet(opParams.all2AllVDataDes.recvType);
390 :
391 0 : uint64_t tailSize = sendLength % UB_MAX_TRANS_SIZE;
392 0 : uint64_t loopNum = UINT64_MAX - 1 - (sendLength / UB_MAX_TRANS_SIZE);
393 0 : HCCL_INFO("[CcuContextAllToAllVMesh1D][RefreshArgs] CalGoSizeStatic size [%llu]", tailSize);
394 :
395 0 : args.push_back(tailSize);
396 0 : args.push_back(loopNum);
397 0 : args.push_back(sendOffset);
398 0 : args.push_back(recvOffset);
399 : }
400 :
401 0 : for (u32 i = 0; i < args.size(); i++) {
402 0 : HCCL_INFO("[CcuContextAllToAllVMesh1D][RefreshArgs] SFL args[%u] is [%llu]", i, args[i]);
403 : }
404 0 : }
405 : }
|