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 "all_gather_unified_march.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 : static const u32 NEIGHBORS_NUM_TWO = 2; // 2: 邻居数量
16 : static const u32 NEIGHBORS_NUM_ONE = 1; // 1: 邻居数量
17 : static const u32 DIVISOR_NUM_TWO = 2;
18 :
19 0 : AllGatherUnifiedMarch::AllGatherUnifiedMarch(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
20 :
21 0 : AllGatherUnifiedMarch::~AllGatherUnifiedMarch() {}
22 :
23 0 : std::string AllGatherUnifiedMarch::GetStreamIndexString()
24 : {
25 0 : std::string res = "";
26 0 : for (u32 streamIndex = 0; streamIndex < subStreams_.size(); streamIndex++) {
27 0 : res += std::to_string(streamIndex) + ", ";
28 : }
29 0 : return res;
30 0 : }
31 :
32 : // 主流所有从流
33 0 : HcclResult AllGatherUnifiedMarch::NotifySubStreamStart(u32 streamSize)
34 : {
35 0 : CHK_PRT_RET(
36 : streamSize > subStreams_.size() || streamSize > meshSignalSubToMain_.size(),
37 : HCCL_ERROR(
38 : "[AllGatherUnifiedMarch][NotifySubStreamStart] streamSize[%u] is out of range"
39 : "subStreams_ size[%zu] or meshSignalSubToMain_ size[%zu]",
40 : streamSize, subStreams_.size(), meshSignalSubToMain_.size()),
41 : HCCL_E_PARA);
42 0 : for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
43 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, meshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
44 0 : CHK_RET(LocalNotify::Wait(
45 : subStreams_[streamIndex], dispatcher_, meshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
46 : }
47 0 : HCCL_DEBUG(
48 : "[AllGatherUnifiedMarch][NotifySubStreamStart] userRank [%u] main stream notify substream [%s]", intraRank_,
49 : GetStreamIndexString().c_str());
50 0 : return HCCL_SUCCESS;
51 : }
52 :
53 0 : HcclResult AllGatherUnifiedMarch::WaitSubStreamFinish(u32 streamSize)
54 : {
55 0 : CHK_PRT_RET(
56 : streamSize > subStreams_.size() || streamSize > meshSignalMainToSub_.size(),
57 : HCCL_ERROR(
58 : "[AllGatherUnifiedMarch][NotifySubStreamStart] streamSize[%u] is out of range"
59 : "subStreams_ size[%zu] or meshSignalMainToSub_ size[%zu]",
60 : streamSize, subStreams_.size(), meshSignalMainToSub_.size()),
61 : HCCL_E_PARA);
62 0 : for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
63 0 : CHK_RET(LocalNotify::Post(
64 : subStreams_[streamIndex], dispatcher_, meshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
65 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, meshSignalMainToSub_[streamIndex], INVALID_VALUE_STAGE));
66 : }
67 0 : HCCL_DEBUG(
68 : "[AllGatherUnifiedMarch][WaitSubStreamFinish] userRank [%u] main stream wait substream [%s]", intraRank_,
69 : GetStreamIndexString().c_str());
70 0 : return HCCL_SUCCESS;
71 : }
72 :
73 0 : HcclResult AllGatherUnifiedMarch::DoSerialSDMA(
74 : void* remoteSrcAddr, u64 remoteOffsetByte, void* dstAddr, Stream& temStream, LINK& tmpLink, u64 memSize, u32 step)
75 : {
76 : (void)step;
77 0 : for (u32 sliceIdx = 0; sliceIdx < (multRingsUserMemSlice_[0].size() / intraRankSize_); sliceIdx++) {
78 0 : struct hccl::Transport::Buffer remoteBuf;
79 0 : remoteBuf.addr = static_cast<u8*>(remoteSrcAddr) + remoteOffsetByte + multRingsUserMemSlice_[0][sliceIdx].offset
80 0 : + baseOffset_;
81 0 : remoteBuf.size = memSize;
82 0 : struct hccl::Transport::Buffer localBuf;
83 0 : localBuf.addr = static_cast<u8*>(dstAddr) + multRingsUserMemSlice_[0][sliceIdx].offset;
84 0 : localBuf.size = memSize;
85 0 : HCCL_DEBUG(
86 : "intralRank[%u] slice[%u] offset[%llu] do SDMA read with remoteBuf[addr[%p], size[%llu]] and "
87 : "localBuf[addr[%p], size[%llu]]",
88 : intraRank_, sliceIdx, multRingsUserMemSlice_[0][sliceIdx].offset, remoteBuf.addr, remoteBuf.size,
89 : localBuf.addr, localBuf.size);
90 0 : CHK_RET(tmpLink->ReadSync(localBuf, remoteBuf, temStream));
91 : }
92 0 : return HCCL_SUCCESS;
93 : }
94 :
95 0 : HcclResult AllGatherUnifiedMarch::NotifyNeighborsStart(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
96 : {
97 : // 图模式保持使用Post/Wait接口
98 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
99 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
100 0 : HCCL_DEBUG("[AllGatherUnifiedMarch][NotifyNeighborsStart]neighborRankId is %u", neighborRankId);
101 : // notify是否越界由平台侧保证
102 0 : if (neighborRankId == 0) {
103 0 : CHK_RET(nextIntralLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
104 0 : CHK_RET(prevIntraLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait
105 0 : } else if (neighborRankId == 1) {
106 0 : CHK_RET(prevIntraLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
107 0 : CHK_RET(nextIntralLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait
108 : }
109 : }
110 0 : HCCL_DEBUG(
111 : "[AllGatherUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done", intraRank_,
112 : neighbors);
113 0 : return HCCL_SUCCESS;
114 : }
115 :
116 : // 一条流负责一个环
117 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
118 : // 交替使用Ack和DataSignal两种notify
119 0 : const u32 NOTIFY_IDX_TWO = 2;
120 0 : if (neighborRankId == 0) {
121 0 : if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
122 0 : CHK_RET(nextIntralLink->TxAck(subStreams_[neighborRankId])); // AckRecord
123 0 : CHK_RET(prevIntraLink->RxAck(subStreams_[neighborRankId])); // AckWait
124 : } else {
125 0 : CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
126 0 : CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId])); // DataWait
127 : }
128 0 : } else if (neighborRankId == 1) {
129 0 : if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
130 0 : CHK_RET(prevIntraLink->TxAck(subStreams_[neighborRankId])); // AckRecord
131 0 : CHK_RET(nextIntralLink->RxAck(subStreams_[neighborRankId])); // AckWait
132 : } else {
133 0 : CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
134 0 : CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId])); // DataWait
135 : }
136 : }
137 : }
138 0 : HCCL_DEBUG(
139 : "[AllGatherUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done", intraRank_,
140 : neighbors);
141 :
142 0 : return HCCL_SUCCESS;
143 : }
144 :
145 0 : HcclResult AllGatherUnifiedMarch::NotifyNeighborsEnd(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
146 : {
147 0 : HCCL_DEBUG("[AllGatherUnifiedMarch]NotifyNeighborsEnd start.");
148 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
149 0 : if (neighborRankId == 0) {
150 0 : CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
151 0 : CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId]));
152 0 : } else if (neighborRankId == 1) {
153 0 : CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId]));
154 0 : CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId]));
155 : }
156 : }
157 0 : HCCL_DEBUG(
158 : "[AllGatherUnifiedMarch][NotifyNeighborsStart] intraRank[%u] notifys [%u]neigbhbors sdma read done", intraRank_,
159 : neighbors);
160 :
161 0 : return HCCL_SUCCESS;
162 : }
163 :
164 0 : HcclResult AllGatherUnifiedMarch::RunSingleStep(u32 ringPrevRank, u32 ringNextRank, u32 step, u32 totalStep)
165 : {
166 0 : LINK prevIntraLink = links_[ringPrevRank];
167 0 : LINK nextIntralLink = links_[ringNextRank];
168 0 : CHK_SMART_PTR_NULL(prevIntraLink);
169 0 : CHK_SMART_PTR_NULL(nextIntralLink);
170 0 : u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
171 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, neighbors));
172 :
173 : // 拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
174 0 : CHK_RET(WaitSubStreamFinish(neighbors));
175 0 : CHK_RET(NotifySubStreamStart(neighbors));
176 :
177 : // 从前向rank读取数据
178 0 : void* preRemDMAMemPtr = nullptr;
179 0 : CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &preRemDMAMemPtr));
180 0 : u32 preDataIndex = (intraRank_ + intraRankSize_ - step - 1) % intraRankSize_;
181 0 : u64 preOffsetByte = preDataIndex * blockDataByte_;
182 0 : void* preDstAddr = static_cast<u8*>(userOutput_.ptr()) + preOffsetByte;
183 0 : u64 preRemoteOffsetByte = preOffsetByte;
184 0 : CHK_RET(DoSerialSDMA(
185 : preRemDMAMemPtr, preRemoteOffsetByte, preDstAddr, subStreams_[0], prevIntraLink, blockDataByte_, step));
186 0 : HCCL_INFO(
187 : "[AllGatherUnifiedMarch][RunSingleStep] intralRank [%u] read from ringPrevRank [%u] done", intraRank_,
188 : ringPrevRank);
189 :
190 : // 从后向rank读取数据
191 0 : if (neighbors > NEIGHBORS_NUM_ONE) {
192 0 : void* nextRemDMAMemPtr = nullptr;
193 0 : CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &nextRemDMAMemPtr));
194 0 : u32 nextDataIndex = (intraRank_ + 1 + step) % intraRankSize_;
195 0 : u64 nextOffsetByte = nextDataIndex * blockDataByte_;
196 0 : void* nextDstAddr = static_cast<u8*>(userOutput_.ptr()) + nextOffsetByte;
197 0 : u64 nextRemoteOffsetByte = nextOffsetByte;
198 0 : CHK_RET(DoSerialSDMA(
199 : nextRemDMAMemPtr, nextRemoteOffsetByte, nextDstAddr, subStreams_[1], nextIntralLink, blockDataByte_, step));
200 0 : HCCL_INFO(
201 : "[AllGatherUnifiedMarch][RunSingleStep] intralRank [%u] read from ringNextRank [%u] done", intraRank_,
202 : ringNextRank);
203 : }
204 :
205 0 : if (step == totalStep - 1) {
206 0 : CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, neighbors));
207 : }
208 0 : notifyIdx_++;
209 :
210 0 : return HCCL_SUCCESS;
211 0 : }
212 :
213 0 : HcclResult AllGatherUnifiedMarch::RunLastStep(u32 ringPrevRank, u32 ringNextRank, u32 totalStep)
214 : {
215 0 : LINK prevIntraLink = links_[ringPrevRank];
216 0 : CHK_SMART_PTR_NULL(prevIntraLink);
217 0 : LINK nextIntralLink = links_[ringNextRank];
218 0 : CHK_SMART_PTR_NULL(nextIntralLink);
219 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
220 :
221 : // 拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
222 0 : CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
223 0 : CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
224 :
225 : // 从前向rank读取数据
226 0 : void* preRemDMAMemPtr = nullptr;
227 0 : CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &preRemDMAMemPtr));
228 0 : u32 preDataIndex = (intraRank_ + intraRankSize_ - totalStep) % intraRankSize_;
229 0 : u64 preOffsetByte = preDataIndex * blockDataByte_;
230 0 : void* preDstAddr = static_cast<u8*>(userOutput_.ptr()) + preOffsetByte;
231 0 : CHK_RET(DoSerialSDMA(
232 : preRemDMAMemPtr, preOffsetByte, preDstAddr, subStreams_[0], prevIntraLink, blockDataByte_ / DIVISOR_NUM_TWO));
233 0 : HCCL_INFO(
234 : "[AllGatherUnifiedMarch][RunLastStep] intralRank [%u] read from ringPrevRank [%u] done", intraRank_,
235 : ringPrevRank);
236 :
237 : // 从后向rank读取数据
238 0 : void* nextRemDMAMemPtr = nullptr;
239 0 : CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &nextRemDMAMemPtr));
240 0 : u32 nextDataIndex = (intraRank_ + totalStep) % intraRankSize_;
241 0 : u64 nextOffsetByte = nextDataIndex * blockDataByte_ + blockDataByte_ / DIVISOR_NUM_TWO;
242 0 : void* nextDstAddr = static_cast<u8*>(userOutput_.ptr()) + nextOffsetByte;
243 0 : CHK_RET(DoSerialSDMA(
244 : nextRemDMAMemPtr, nextOffsetByte, nextDstAddr, subStreams_[1], nextIntralLink,
245 : (blockDataByte_ - blockDataByte_ / DIVISOR_NUM_TWO))); // 兼容单块儿allgather数据量不是2的倍数场景
246 0 : HCCL_INFO(
247 : "[AllGatherUnifiedMarch][RunLastStep] intralRank [%u] read from ringNextRank [%u] done", intraRank_,
248 : ringNextRank);
249 :
250 : // 单算子使用Ack/Datasignal接口,必须保证两者交替使用
251 0 : const u32 NOTIFY_IDX_TWO = 2;
252 0 : if (notifyIdx_ % NOTIFY_IDX_TWO != 0 && GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
253 0 : notifyIdx_++;
254 0 : CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
255 0 : CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
256 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
257 : }
258 0 : CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
259 :
260 0 : return HCCL_SUCCESS;
261 0 : }
262 :
263 0 : HcclResult AllGatherUnifiedMarch::RunAsync()
264 : {
265 0 : HCCL_INFO("[AllGatherUnifiedMarch][RunAsync] starts.");
266 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForAllGather();
267 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
268 :
269 : // 获取link的收、发
270 0 : u32 ringPrevRank = (intraRank_ + intraRankSize_ - 1) % intraRankSize_;
271 0 : u32 ringNextRank = (intraRank_ + 1) % intraRankSize_;
272 :
273 0 : u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
274 0 : CHK_RET(NotifySubStreamStart(neighbors));
275 :
276 0 : u32 totalStep = intraRankSize_ / DIVISOR_NUM_TWO;
277 0 : HCCL_INFO(
278 : "[AllGatherUnifiedMarch][RunAsync] intraRank [%u] ringPrevRank [%u] ringNextRank [%u] totalStep [%u]",
279 : intraRank_, ringPrevRank, ringNextRank, totalStep);
280 0 : if (totalStep == 1) {
281 0 : CHK_RET(RunSingleStep(ringPrevRank, ringNextRank, 0, totalStep));
282 : } else {
283 0 : for (u32 step = 0; step < totalStep - 1; step++) {
284 0 : CHK_RET(RunSingleStep(ringPrevRank, ringNextRank, step, totalStep));
285 : }
286 0 : CHK_RET(RunLastStep(ringPrevRank, ringNextRank, totalStep));
287 : }
288 0 : CHK_RET(WaitSubStreamFinish(neighbors));
289 :
290 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, subStreams_));
291 :
292 0 : HCCL_INFO("[AllGatherUnifiedMarch][RunAsync] finished.");
293 0 : return HCCL_SUCCESS;
294 : }
295 :
296 0 : HcclResult AllGatherUnifiedMarch::Prepare(
297 : const Stream& mainStream, SubCommInfo& level0CommInfo, DeviceMem& userInput, DeviceMem& userOutput,
298 : DeviceMem& usrInMem, DeviceMem& usrOutMem, u64 blockDataByte, std::vector<Stream>& subStreams,
299 : const std::vector<std::shared_ptr<LocalNotify>>& meshSignalMainToSub,
300 : const std::vector<std::shared_ptr<LocalNotify>>& meshSignalSubToMain,
301 : const std::vector<std::vector<Slice>>& multRingsUserMemSlice, const u64 baseOffset)
302 : {
303 0 : mainStream_ = mainStream;
304 0 : intraRank_ = level0CommInfo.localRank;
305 0 : intraRankSize_ = level0CommInfo.localRankSize;
306 0 : CHK_PRT_RET(
307 : intraRankSize_ == 0 || (intraRankSize_ % DIVISOR_NUM_TWO != 0),
308 : HCCL_ERROR("[AllGatherUnifiedMarch][Prepare]intraRankSize_ is zero."), HCCL_E_PARA);
309 0 : links_ = level0CommInfo.links;
310 :
311 0 : userInput_ = userInput;
312 0 : userOutput_ = userOutput;
313 0 : usrInMem_ = usrInMem;
314 0 : usrOutMem_ = usrOutMem;
315 0 : HCCL_INFO(
316 : "userInput_[%p] size[%llu], userOutput_[%p] size[%llu], usrInMem_[%p] size[%llu], usrOutMem_[%p] size[%llu]",
317 : userInput_.ptr(), userInput_.size(), userOutput_.ptr(), userOutput_.size(), usrInMem_.ptr(), usrInMem_.size(),
318 : usrOutMem_.ptr(), usrOutMem_.size());
319 :
320 0 : subStreams_ = subStreams;
321 0 : meshSignalMainToSub_ = meshSignalMainToSub;
322 0 : meshSignalSubToMain_ = meshSignalSubToMain;
323 0 : CHK_PRT_RET(
324 : subStreams_.size() < NEIGHBORS_NUM_TWO || meshSignalMainToSub_.size() < NEIGHBORS_NUM_TWO
325 : || meshSignalSubToMain_.size() < NEIGHBORS_NUM_TWO,
326 : HCCL_ERROR(
327 : "[AllGatherUnifiedMarch] subStreams_ size[%u] or meshSignalMainToSub_ size[%u] or "
328 : "meshSignalSubToMain_ size[%u] is less than 2",
329 : subStreams_.size(), meshSignalMainToSub_.size(), meshSignalSubToMain_.size()),
330 : HCCL_E_PARA);
331 0 : blockDataByte_ = blockDataByte;
332 0 : multRingsUserMemSlice_ = multRingsUserMemSlice;
333 0 : CHK_PRT_RET(
334 : multRingsUserMemSlice_[0].size() % intraRankSize_ != 0,
335 : HCCL_ERROR(
336 : "[AllGatherUnifiedMarch] multRingsUserMemSlice_[0] size[%u] can not be divided by rank size[%u]",
337 : multRingsUserMemSlice_[0].size(), intraRankSize_),
338 : HCCL_E_PARA);
339 :
340 0 : baseOffset_ = baseOffset;
341 :
342 0 : return HCCL_SUCCESS;
343 : }
344 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_UNIFIED_MARCH, AllGatherUnifiedMarch);
345 : } // namespace hccl
|