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