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 <random>
12 : #include <algorithm>
13 : #include "ccu_context_half_alltoallv_mesh1d.h"
14 : #include "ccu_instruction_half_alltoallv_mesh1d.h"
15 :
16 : namespace Hccl {
17 :
18 : constexpr uint16_t RANK_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
19 :
20 0 : void CcuContextHalfAllToAllVMesh1D::ExchangeCtxResource()
21 : {
22 0 : if (missionId_ == 1) {
23 0 : ExportVariable(userInAddr_, ctxName_ + "_UserInAddr_" + std::to_string(missionId_));
24 0 : ExportVariable(sendSizeAddr_, ctxName_ + "_SendSizeAddr_" + std::to_string(missionId_));
25 0 : ExportVariable(sendOffsetAddr_, ctxName_ + "_SendOffsetAddr__" + std::to_string(missionId_));
26 0 : ExportVariable(recvOffset_, ctxName_ + "_RecvOffset_" + std::to_string(missionId_));
27 : } else {
28 0 : anoUserInAddr_ = ImportVariable(ctxName_ + "_UserInAddr_" + std::to_string(1 - missionId_));
29 0 : anoSendSizeAddr_ = ImportVariable(ctxName_ + "_SendSizeAddr_" + std::to_string(1 - missionId_));
30 0 : anoSendOffsetAddr_ = ImportVariable(ctxName_ + "_SendOffsetAddr__" + std::to_string(1 - missionId_));
31 0 : anoRecvOffset_ = ImportVariable(ctxName_ + "_RecvOffset_" + std::to_string(1 - missionId_));
32 : }
33 0 : ExportMaskSignal(locMiSignal0_, ctxName_ + "_MiSync0_" + std::to_string(missionId_));
34 0 : anoMiSignal0_ = ImportMaskSignal(ctxName_ + "_MiSync0_" + std::to_string(1 - missionId_));
35 0 : ExportMaskSignal(locMiSignal1_, ctxName_ + "_MiSync1_" + std::to_string(missionId_));
36 0 : anoMiSignal1_ = ImportMaskSignal(ctxName_ + "_MiSync1_" + std::to_string(1 - missionId_));
37 :
38 0 : return;
39 : }
40 :
41 0 : CcuContextHalfAllToAllVMesh1D::CcuContextHalfAllToAllVMesh1D(
42 0 : const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
43 0 : : CcuContextAlgBase(arg, transports, group)
44 : {
45 0 : const CcuCtxArgHalfAllToAllVMesh1D* ctxArg = dynamic_cast<const CcuCtxArgHalfAllToAllVMesh1D*>(&arg);
46 0 : if (ctxArg == nullptr) {
47 0 : THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D::ctxArg ptr is null"));
48 : }
49 0 : ctxName_ = ctxArg->GetCtxSignature().Describe();
50 0 : rankId_ = ctxArg->rankId;
51 0 : if (ctxArg->dimSize.size() > 0) {
52 0 : rankSize_ = ctxArg->dimSize[0];
53 : }
54 0 : missionId_ = ctxArg->missionId;
55 0 : signalNum_ = (rankSize_ + RANK_NUM_PER_CKE - 1) / RANK_NUM_PER_CKE; // 每个CKE有16个bit
56 0 : myCclBufferAddr_ = ctxArg->cclBufferAddr;
57 :
58 0 : userInAddr_ = CreateVariable();
59 0 : sendSizeAddr_ = CreateVariable();
60 0 : sendOffsetAddr_ = CreateVariable();
61 0 : recvOffset_ = CreateVariable();
62 0 : goSize_ = CreateGroupOpSize();
63 :
64 0 : curSrc_ = CreateMemory();
65 0 : if (transports.size() == 0 || transports.size() < rankSize_ - 1) {
66 0 : THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D transports is empty or size is less"));
67 : }
68 0 : for (uint32_t i = 0; i < rankSize_; i++) {
69 : // curDst在初始化时即赋值各个rank的cclBuffer地址
70 0 : if (i == rankId_) {
71 0 : curDst_.emplace_back(CreateMemory());
72 : } else {
73 0 : uint16_t transIdx = (i < rankId_) ? i : i - 1;
74 0 : curDst_.emplace_back(GetRmtBuffer(*transports[transIdx], 0));
75 : }
76 0 : token_.emplace_back(CreateVariable());
77 0 : sendSizeA_.emplace_back(CreateVariable());
78 0 : sendSizeB_.emplace_back(CreateVariable());
79 0 : sendOffsetA_.emplace_back(CreateVariable());
80 0 : sendOffsetB_.emplace_back(CreateVariable());
81 : }
82 0 : ccuStartSignal_ = CreateMaskSignal();
83 0 : ccuEndSignal_ = CreateMaskSignal();
84 0 : for (uint32_t i = 0; i < signalNum_; i++) {
85 0 : writeDoneSignal_.emplace_back(CreateMaskSignal());
86 : }
87 :
88 : // mission间交互的资源
89 0 : locMiSignal0_ = CreateMaskSignal();
90 0 : locMiSignal1_ = CreateMaskSignal();
91 0 : ExchangeCtxResource();
92 :
93 0 : return;
94 0 : }
95 :
96 0 : void CcuContextHalfAllToAllVMesh1D::LoadArgs()
97 : {
98 0 : if (missionId_ == 0) {
99 0 : Load(userInAddr_);
100 0 : Load(sendSizeAddr_);
101 0 : Load(token_[rankId_]);
102 0 : Load(sendOffsetAddr_);
103 0 : Load(goSize_);
104 0 : Load(recvOffset_);
105 :
106 : // Mi0将参数同步给Mi1
107 0 : LocalCtxPostVar(userInAddr_, anoUserInAddr_, anoMiSignal0_, 1 << 0); // 用第1个bit标记
108 0 : LocalCtxPostVar(sendSizeAddr_, anoSendSizeAddr_, anoMiSignal0_, 1 << 1); // 用第2个
109 0 : LocalCtxPostVar(sendOffsetAddr_, anoSendOffsetAddr_, anoMiSignal0_, 1 << 2); // 用第3个
110 0 : LocalCtxPostVar(recvOffset_, anoRecvOffset_, anoMiSignal0_, 1 << 3); // 用第4个
111 : } else {
112 0 : LocalWait(locMiSignal0_, (1 << 4) - 1); // 共同步4个参数
113 : }
114 0 : return;
115 : }
116 :
117 0 : void CcuContextHalfAllToAllVMesh1D::LoadArgsFromMem()
118 : {
119 : // 暂不支持用单条指令加载多个参数
120 0 : CcuRep::Variable dataLength = CreateVariable();
121 0 : CcuRep::Variable tempAddr = CreateVariable();
122 0 : LoadArgs();
123 :
124 : // 加载本端的cclbuffer地址
125 0 : curDst_[rankId_].addr = myCclBufferAddr_;
126 :
127 : // 连续加载rankSize * 2个sendSize
128 0 : dataLength = 8; // 每个Xn占8个byte
129 :
130 0 : u32 argsCount = sendSizeA_.size() + sendSizeB_.size() + sendOffsetA_.size() + sendOffsetB_.size();
131 0 : std::vector<CcuRep::Variable> tempArgs(argsCount);
132 0 : HCCL_INFO("CcuContextHalfAllToAllVMesh1D LoadArgsFromMem, argsCount:[%u]", argsCount);
133 :
134 0 : for (uint32_t i = 0; i < tempArgs.size(); ++i) {
135 0 : tempArgs[i] = CreateContinuousVariable();
136 : }
137 0 : LoadVariable(sendSizeAddr_, tempArgs[0], argsCount);
138 :
139 0 : u32 argIdx = 0;
140 0 : for (uint32_t i = 0; i < sendSizeA_.size(); i++) {
141 0 : sendSizeA_[i] = tempArgs[argIdx];
142 0 : argIdx++;
143 : }
144 0 : for (uint32_t i = 0; i < sendSizeB_.size(); i++) {
145 0 : sendSizeB_[i] = tempArgs[argIdx];
146 0 : argIdx++;
147 : }
148 0 : for (uint32_t i = 0; i < sendOffsetA_.size(); i++) {
149 0 : sendOffsetA_[i] = tempArgs[argIdx];
150 0 : argIdx++;
151 : }
152 0 : for (uint32_t i = 0; i < sendOffsetB_.size(); i++) {
153 0 : sendOffsetB_[i] = tempArgs[argIdx];
154 0 : argIdx++;
155 : }
156 0 : return;
157 0 : }
158 :
159 0 : void CcuContextHalfAllToAllVMesh1D::MissionSync(uint32_t signalIndex)
160 : {
161 0 : const uint32_t MISSION_NUM = 2;
162 0 : if (signalIndex > 1) {
163 0 : THROW<InvalidParamsException>(
164 0 : StringFormat("[CcuContextHalfAllToAllVMesh1D] Unexpected SignalInex[%u]", signalIndex));
165 : }
166 0 : LocalCtxPost(anoMiSignal1_, 1 << (missionId_ + signalIndex * MISSION_NUM));
167 0 : LocalWait(locMiSignal1_, 1 << (1 - missionId_ + signalIndex * MISSION_NUM));
168 0 : return;
169 : }
170 :
171 0 : void CcuContextHalfAllToAllVMesh1D::PostSync()
172 : {
173 0 : if (missionId_ == 0) {
174 0 : uint16_t signalId = rankId_ / RANK_NUM_PER_CKE;
175 0 : uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
176 0 : for (auto t : transports) {
177 0 : if (t == nullptr) {
178 0 : THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D::Algorithm transport ptr is null"));
179 : }
180 0 : RemotePost(*t, signalId, selfBit);
181 : }
182 :
183 0 : for (uint16_t sId = 0; sId < signalNum_; sId++) {
184 : uint32_t waitBit;
185 0 : if (sId != signalNum_ - 1) {
186 0 : waitBit = (1 << RANK_NUM_PER_CKE) - 1; // 等待全部16个peer
187 : } else {
188 0 : waitBit = ((1 << (rankSize_ - (signalNum_ - 1) * RANK_NUM_PER_CKE)) - 1);
189 : }
190 0 : if (sId == signalId) {
191 0 : waitBit &= ~selfBit; // 如果这个CKE上有自己对应的bit,设为0
192 : }
193 0 : GroupWait(*transportGroup, sId, waitBit);
194 : }
195 : }
196 0 : return;
197 : }
198 :
199 0 : void CcuContextHalfAllToAllVMesh1D::Algorithm()
200 : {
201 0 : HCCL_INFO("[CcuContextHalfAllToAllVMesh1D] AllgatherMesh1D Algorithm Begins.");
202 0 : LoadArgsFromMem();
203 :
204 : // 向每个对端发送数据
205 0 : CcuRep::Memory lgSrc = CreateMemory();
206 0 : CcuRep::Memory lgDst = CreateMemory();
207 0 : CcuRep::Variable tempCount = CreateVariable();
208 0 : lgSrc.token = token_[rankId_];
209 0 : lgDst.token = token_[rankId_];
210 0 : curSrc_.token = token_[rankId_];
211 :
212 0 : for (uint32_t peerId = 0; peerId < rankSize_; peerId++) {
213 0 : CcuRep::Variable& curCount = missionId_ == 0 ? sendSizeA_[peerId] : sendSizeB_[peerId];
214 0 : CcuRep::Variable& curOffset = missionId_ == 0 ? sendOffsetA_[peerId] : sendOffsetB_[peerId];
215 0 : uint16_t peerSignalId = peerId / RANK_NUM_PER_CKE;
216 0 : uint16_t peerBit = 1 << (peerId % RANK_NUM_PER_CKE);
217 :
218 0 : tempCount = curCount;
219 0 : curSrc_.addr = userInAddr_;
220 0 : curSrc_.addr += curOffset;
221 0 : curDst_[peerId].addr += recvOffset_;
222 0 : if (missionId_ == 1) {
223 0 : curDst_[peerId].addr += sendSizeA_[peerId]; // Mi1的dst需要加chunkOffset
224 : }
225 0 : if (peerId == rankId_) {
226 0 : lgSrc.addr = curSrc_.addr;
227 0 : lgDst.addr = curDst_[peerId].addr;
228 0 : LocalPost(writeDoneSignal_[peerSignalId], peerBit);
229 : } else {
230 0 : uint16_t transIdx = (peerId < rankId_) ? peerId : peerId - 1;
231 0 : CCU_IF(tempCount != 0)
232 : {
233 0 : Write(
234 0 : *(transports[transIdx]), curDst_[peerId], curSrc_, tempCount, writeDoneSignal_[peerSignalId],
235 : peerBit);
236 0 : }
237 0 : CCU_IF(tempCount == 0) { LocalPost(writeDoneSignal_[peerSignalId], peerBit); }
238 : }
239 : }
240 :
241 0 : if (missionId_ == 0) {
242 0 : GroupCopy(lgDst, lgSrc, goSize_);
243 : }
244 0 : for (uint16_t sId = 0; sId < signalNum_; sId++) {
245 : uint32_t waitBit;
246 0 : if (sId != signalNum_ - 1) {
247 0 : waitBit = (1 << RANK_NUM_PER_CKE) - 1; // 等待全部16个peer
248 : } else {
249 0 : waitBit = ((1 << (rankSize_ - (signalNum_ - 1) * RANK_NUM_PER_CKE)) - 1);
250 : }
251 0 : LocalWait(writeDoneSignal_[sId], waitBit);
252 : }
253 :
254 0 : MissionSync(0);
255 0 : PostSync();
256 0 : HCCL_INFO("[CcuContextHalfA2AUnions] Algorithm Ends.");
257 0 : return;
258 0 : }
259 :
260 0 : std::vector<uint64_t> CcuContextHalfAllToAllVMesh1D::GeneArgs(const CcuTaskArg& arg)
261 : {
262 : (void)arg;
263 0 : std::vector<uint64_t> args = {};
264 0 : uint64_t argNum = missionId_ == 0 ? 9 : 0; // Mi0有9个Load,Mi1不做Load
265 0 : for (uint32_t i = 0; i < argNum; i++) {
266 0 : args.emplace_back(0);
267 : }
268 0 : return args;
269 0 : }
270 : } // namespace Hccl
|