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_mesh1d.h"
12 : #include "ccu_instruction_all_to_all_mesh1d.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 :
20 0 : CcuContextAllToAllMesh1D::CcuContextAllToAllMesh1D(
21 0 : const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
22 0 : : CcuContextAlgBase(arg, transports, group)
23 : {
24 0 : const CcuCtxArgAllToAllMesh1D* ctxArg = dynamic_cast<const CcuCtxArgAllToAllMesh1D*>(&arg);
25 0 : if (ctxArg == nullptr) {
26 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::ctxArg ptr is null"));
27 : }
28 0 : rankId_ = ctxArg->rankId;
29 0 : if (ctxArg->dimSize.size() > 0) {
30 0 : rankSize_ = ctxArg->dimSize[0];
31 : }
32 0 : if (transports.size() == 0) {
33 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D transports is empty"));
34 : }
35 0 : loadFromMem_ = ctxArg->loadFromMem;
36 0 : }
37 :
38 0 : void CcuContextAllToAllMesh1D::Algorithm()
39 : {
40 0 : HCCL_INFO("[ccuAllToAllMesh1D_context] AllToAllMesh1D run.");
41 : // 创建Variable,用于交换地址及token
42 0 : u32 transportId = 0;
43 0 : for (u64 id = 0; id < rankSize_; id++) {
44 0 : if (id == rankId_) {
45 0 : input_.push_back(CreateVariable());
46 0 : output_.push_back(CreateVariable());
47 0 : token_.push_back(CreateVariable());
48 : } else { // 非本地,使用远端Variable
49 0 : CHK_PRT_RET(
50 : transports[transportId] == nullptr,
51 : HCCL_ERROR("[CcuContextAllToAllMesh1D] Algorithm transport ptr is null"), );
52 0 : input_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_0));
53 0 : output_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_1));
54 0 : token_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_2));
55 0 : transportId++;
56 : }
57 : }
58 0 : sliceSize_ = CreateVariable();
59 0 : srcStride_ = CreateVariable();
60 0 : srcOffset_ = CreateVariable();
61 0 : dstOffset_ = CreateVariable();
62 0 : groupOpSize_ = CreateGroupOpSize();
63 :
64 : // 从SQE load args,本rank需要的input、output地址等信息
65 : // inputAddr, outputAddr, tokenInfo, srcStride, srcOffset, dstOffset, groupOpSize
66 0 : Load(input_[rankId_]);
67 0 : Load(output_[rankId_]);
68 0 : Load(token_[rankId_]);
69 0 : Load(sliceSize_); // 本轮传输的分片大小
70 0 : Load(srcStride_); // 单片数据大小
71 0 : Load(srcOffset_);
72 0 : Load(dstOffset_);
73 0 : Load(groupOpSize_);
74 :
75 : // 前同步。交换信息,将本Rank load的in\out等地址信息写到所有对端的对应Variable中,并同步
76 0 : uint16_t selfBit = 1 << rankId_; // 本rank的mask
77 0 : uint16_t allBit = ((1 << rankSize_) - 1) & (~(1 << rankId_));
78 :
79 0 : srcOffset_ += input_[rankId_];
80 :
81 0 : for (auto t : transports) {
82 : // (transport, param, paramID, SemID, mask)
83 0 : WriteVariableWithSignal(*t, output_[rankId_], CKE_IDX_1, CKE_IDX_1, selfBit); // index = 1,传递output信息
84 0 : WriteVariableWithSignal(*t, token_[rankId_], CKE_IDX_2, CKE_IDX_2, selfBit); // index = 2,传递token信息
85 : }
86 :
87 0 : GroupWait(*transportGroup, CKE_IDX_1, allBit); // index = 1,传递output信息
88 0 : GroupWait(*transportGroup, CKE_IDX_2, allBit); // index = 2,传递token信息
89 :
90 : // 创建GSA, src为本地的各片HBM地址GSA列表,dst为所有对端的HBM地址GSA列表
91 0 : std::vector<CcuRep::Memory> src;
92 0 : for (uint64_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
93 0 : src.push_back(CreateMemory());
94 : }
95 0 : std::vector<CcuRep::Memory> dst;
96 0 : for (uint64_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
97 0 : dst.push_back(CreateMemory());
98 : }
99 :
100 : // 考虑stride信息
101 0 : for (uint64_t r = 0; r < rankSize_; r++) {
102 0 : src[r].token = token_[r];
103 0 : dst[r].token = token_[r];
104 :
105 : // src[r] = srcOffset + r*srcStride
106 0 : src[r].addr = srcOffset_;
107 0 : for (uint64_t i = 0; i < r; i++) {
108 0 : src[r].addr += srcStride_;
109 : }
110 : // dst[r] = recvBuf[r] + dstOffset
111 0 : dst[r].addr = output_[r];
112 0 : dst[r].addr += dstOffset_;
113 : }
114 :
115 : // 创建CKE,源端保序
116 0 : CcuRep::MaskSignal locMask = CreateMaskSignal();
117 : // all2all 数据搬运
118 0 : transportId = 0;
119 0 : for (uint64_t r = 0; r < rankSize_; r++) {
120 0 : if (r != rankId_) {
121 0 : Write(*transports[transportId], dst[r], src[r], sliceSize_, locMask, 1 << r);
122 0 : transportId++;
123 : }
124 : }
125 0 : GroupCopy(dst[rankId_], src[rankId_], groupOpSize_);
126 0 : LocalWait(locMask, allBit);
127 :
128 : // 后同步
129 0 : for (auto t : transports) {
130 0 : if (t == nullptr) {
131 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::Algorithm transport ptr is null"));
132 : }
133 0 : RemotePost(*t, CKE_IDX_0, selfBit);
134 : }
135 0 : GroupWait(*transportGroup, CKE_IDX_0, allBit);
136 0 : HCCL_INFO("[AllToAllAlgo] AllToAllMesh1D end");
137 :
138 0 : return;
139 0 : }
140 :
141 0 : std::vector<uint64_t> CcuContextAllToAllMesh1D::GeneArgs(const CcuTaskArg& arg)
142 : {
143 0 : const CcuTaskArgAllToAllMesh1D* taskArg = dynamic_cast<const CcuTaskArgAllToAllMesh1D*>(&arg);
144 0 : if (taskArg == nullptr) {
145 0 : THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::taskArg ptr is null"));
146 : }
147 0 : uint64_t inputAddr = taskArg->inputAddr;
148 0 : uint64_t outputAddr = taskArg->outputAddr;
149 0 : uint64_t tokenInfo = taskArg->token;
150 :
151 0 : uint64_t srcStride = taskArg->srcStride;
152 0 : uint64_t srcOffset = taskArg->srcOffset;
153 0 : uint64_t dstOffset = taskArg->dstOffset;
154 :
155 0 : uint64_t sliceSize = taskArg->sliceSize;
156 0 : auto goSize = CalGoSize(sliceSize);
157 0 : HCCL_INFO(
158 : "[AllToAllAlgo] inputAddr[%llu], outputAddr[%llu], sliceSize[%llu], srcStride[%llu], srcOffset[%llu], "
159 : "dstOffset[%llu].",
160 : inputAddr, outputAddr, sliceSize, srcStride, srcOffset, dstOffset);
161 :
162 : return {inputAddr, outputAddr, tokenInfo, sliceSize, srcStride, srcOffset,
163 0 : dstOffset, goSize[0], goSize[1], goSize[2], goSize[3]};
164 0 : }
165 :
166 : } // namespace Hccl
|