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