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_reduce_multi_deter_pipeline.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : AllReduceMultiDeterPipeline::AllReduceMultiDeterPipeline(const HcclDispatcher dispatcher)
16 0 : : MultiDeterPipeline(dispatcher)
17 0 : {}
18 :
19 0 : AllReduceMultiDeterPipeline::~AllReduceMultiDeterPipeline() {}
20 :
21 0 : HcclResult AllReduceMultiDeterPipeline::GetRemoteCclbufferDeviceMem(
22 : u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
23 : {
24 0 : void* remoteMemPtr = nullptr;
25 0 : CHK_RET(link->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
26 0 : u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
27 0 : u64 size = slices_[inputSliceIndex].size;
28 0 : u64 offset = slices_[outputSliceIndex].offset;
29 0 : u8* intraSrcAddr = beginAddrU8 + offset;
30 0 : remoteMem = DeviceMem::create(intraSrcAddr, size);
31 0 : if (remoteMem.ptr() == nullptr) {
32 0 : HCCL_ERROR("[%s] offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size, outCclBuffer_.size());
33 0 : return HCCL_E_MEMORY;
34 : }
35 0 : HCCL_DEBUG(
36 : "[%s] rank[%u], beginAddr[%p], offset[%llu](slices_[outputSliceIndex].offset), "
37 : "curSize[%llu], totalBufferSize[%llu]",
38 : __func__, outputSliceIndex, remoteMem.ptr(), offset, size, outCclBuffer_.size());
39 0 : return HCCL_SUCCESS;
40 : }
41 :
42 : HcclResult
43 0 : AllReduceMultiDeterPipeline::GetLocalInCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
44 : {
45 0 : u64 size = ifUseLastSize ? lastSize_ : slices_[rankIdInAllRanks].size;
46 0 : u64 offset = slices_[rankIdInAllRanks].offset;
47 0 : localMem = inCclBuffer_.range(offset, size);
48 0 : if (localMem.ptr() == nullptr) {
49 0 : HCCL_ERROR(
50 : "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
51 : inCclBuffer_.size());
52 0 : return HCCL_E_MEMORY;
53 : }
54 0 : HCCL_DEBUG(
55 : "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu]", __func__, rankIdInAllRanks,
56 : localMem.ptr(), offset, size, inCclBuffer_.size());
57 0 : return HCCL_SUCCESS;
58 : }
59 :
60 : // reduce scatter RDMA、SDMA 发送后 取本地 sliceOffset = slices_[rankIdInAllRanks].offset偏移处存放地址
61 : // 什么时候取小内存额外进行判断
62 : // allgather 都是发到相同内存块,所以不需要额外判断是否为小内存
63 0 : HcclResult AllReduceMultiDeterPipeline::GetLocalOutCclbufferDeviceMem(
64 : u32 rankIdInAllRanks, DeviceMem& localMem, bool ifUseLastSize)
65 : {
66 0 : u64 size = ifUseLastSize ? lastSize_ : slices_[rankIdInAllRanks].size;
67 0 : u64 offset = slices_[rankIdInAllRanks].offset;
68 0 : localMem = outCclBuffer_.range(offset, size);
69 0 : if (localMem.ptr() == nullptr) {
70 0 : HCCL_ERROR(
71 : "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
72 : outCclBuffer_.size());
73 0 : return HCCL_E_MEMORY;
74 : }
75 0 : HCCL_DEBUG(
76 : "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu]", __func__, rankIdInAllRanks,
77 : localMem.ptr(), offset, size, outCclBuffer_.size());
78 0 : return HCCL_SUCCESS;
79 : }
80 :
81 0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, bool isUserIn)
82 : {
83 0 : u8* beginAddrU8 = isUserIn ? static_cast<u8*>(usrInMemPtr_) : static_cast<u8*>(usrOutMemPtr_);
84 0 : u64 offset = slices_[rankIdInAllRanks].offset;
85 0 : u64 size = slices_[rankIdInAllRanks].size;
86 0 : u8* intraSrcAddr = beginAddrU8 + offset; // 不用 + offset_,因为usrInMem_已经加过了
87 0 : localMem = DeviceMem::create(intraSrcAddr, size);
88 0 : if (localMem.ptr() == nullptr) {
89 0 : HCCL_ERROR(
90 : "[%s] get localMem failed, offset + size = [%llu] > cclBufferSize[%llu]", __func__, offset + size,
91 : outCclBuffer_.size());
92 0 : return HCCL_E_MEMORY;
93 : }
94 0 : HCCL_DEBUG(
95 : "[%s] rank[%u], beginAddr[%p], offset[%llu], curSize[%llu], totalBufferSize[%llu] isUserIn[%u]", __func__,
96 : rankIdInAllRanks, localMem.ptr(), offset, size, outCclBuffer_.size(), isUserIn);
97 0 : return HCCL_SUCCESS;
98 : }
99 :
100 0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
101 : {
102 0 : CHK_RET(GetLocalUserDeviceMem(rankIdInAllRanks, localMem, true));
103 0 : return HCCL_SUCCESS;
104 : }
105 :
106 0 : HcclResult AllReduceMultiDeterPipeline::GetLocalUserOutDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
107 : {
108 0 : CHK_RET(GetLocalUserDeviceMem(rankIdInAllRanks, localMem, false));
109 0 : return HCCL_SUCCESS;
110 : }
111 :
112 0 : HcclResult AllReduceMultiDeterPipeline::RunLocalCopy()
113 : {
114 0 : if (intraRankId_ != intraRankSize_ - 1) {
115 0 : HCCL_DEBUG("[%s] intra-card no need to copy userRank[%u], intraRankId_[%u]", __func__, userRank_, intraRankId_);
116 0 : return HCCL_SUCCESS;
117 : }
118 :
119 0 : DeviceMem userIn;
120 0 : CHK_RET(GetLocalUserInDeviceMem(userRank_, userIn));
121 0 : DeviceMem cclbuffer;
122 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(userRank_, cclbuffer, false));
123 : // 使用主流搬迁卡内数据
124 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, cclbuffer, userIn, mainStream_));
125 0 : HCCL_DEBUG(
126 : "[%s] intra-card copy data from userInMem to number[%u] cclbuffer[%p] size[%llu]", __func__, userRank_,
127 : cclbuffer.ptr(), curSize_);
128 0 : return HCCL_SUCCESS;
129 0 : }
130 :
131 : // 机内alltoall full mesh收集数据, #step表示pairwise的第step步
132 0 : HcclResult AllReduceMultiDeterPipeline::RunIntraAlltoallPreSync(u32 step)
133 : {
134 0 : HCCL_DEBUG("[%s] intra-server alltoall begin, step[%u]", __func__, step);
135 : // alltoall需要准备跨机要的reduce数据 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块
136 : // 输入是input的第(R1+1)%S1组数据(往后1)
137 : // 每个rank机内只需拷贝intraRankSize_ - 1次
138 0 : HCCL_DEBUG(
139 : "[%s] intra-server SDMA send begin, [serverId, intraRankId] = [%u, %u]", __func__, serverId_, intraRankId_);
140 0 : for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
141 0 : HCCL_DEBUG(
142 : "[%s] intra-server SDMA send begin, userRank[%u] step[%u] pro[%u/%u]", __func__, userRank_, i, i + 1,
143 : intraRankSize_ - 1);
144 : // 从机内rankId为recvIntraRankId收集数据,也发给机内rankId为sendIntraRankId数据
145 0 : u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
146 0 : LINK sendIntraLink = intraLinks_[sendIntraRankId];
147 0 : CHK_RET(sendIntraLink->TxAck(subStreams_[i]));
148 0 : CHK_RET(sendIntraLink->RxAck(subStreams_[i]));
149 0 : }
150 : // 增加主从流同步,目的是让SDMA同时进行
151 0 : CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
152 0 : CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
153 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
154 0 : CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
155 0 : CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
156 0 : return HCCL_SUCCESS;
157 : }
158 :
159 0 : HcclResult AllReduceMultiDeterPipeline::BatchPostNotifyForStreams(
160 : const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
161 : {
162 0 : if (useMainStream) {
163 0 : HCCL_DEBUG("[%s] use mainStream, skip notify wait", __func__);
164 0 : return HCCL_SUCCESS;
165 : }
166 0 : for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
167 0 : if (streamTasks[s].empty())
168 0 : continue; // 无任务的流跳过
169 0 : u32 streamIdx = reduceStreamBegin_ + s;
170 0 : if (reduceMainStreamIdx_ == streamIdx) {
171 0 : continue;
172 : }
173 0 : if (isStartPhase) {
174 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(
175 : inCclBuffer_, outCclBuffer_, subStreams_[reduceMainStreamIdx_], dispatcher_));
176 0 : CHK_RET(LocalNotify::Post(
177 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
178 0 : CHK_RET(LocalNotify::Wait(
179 : subStreams_[streamIdx], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
180 0 : HCCL_DEBUG("[%s] stream[%u] start phase notify done", __func__, streamIdx);
181 : } else {
182 0 : CHK_RET(LocalNotify::Post(
183 : subStreams_[streamIdx], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
184 0 : CHK_RET(LocalNotify::Wait(
185 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
186 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(
187 : inCclBuffer_, outCclBuffer_, subStreams_[reduceMainStreamIdx_], dispatcher_));
188 0 : HCCL_DEBUG("[%s] stream[%u] sync phase notify done", __func__, streamIdx);
189 : }
190 : }
191 0 : return HCCL_SUCCESS;
192 : }
193 :
194 0 : bool AllReduceMultiDeterPipeline::IfUseLastSize(u32 step, u32 sendServerId)
195 : {
196 : // 第0步的最后一块rank,使用小块内存
197 0 : if (step == 0 && userRank_ == userRankSize_ - 1) {
198 0 : return true;
199 : }
200 : // 其他步骤,接收数据的rank是机内最后一个且allreduce后的结果是发给最后一个server
201 0 : if (step != 0 && (sendServerId == serverSize_ - 1) && (intraRankId_ == intraRankSize_ - 1)) {
202 0 : return true;
203 : }
204 0 : return false;
205 : }
206 :
207 : // 机内localreduce首先按序收集所有内存块,接着二分归并reduce,最多使用4条流并行
208 0 : HcclResult AllReduceMultiDeterPipeline::RunIntraLocalReduce(u32 step)
209 : {
210 0 : HCCL_DEBUG("[%s] inter-server local reduce begin, step[%u]", __func__, step);
211 0 : u32 recvServerId = GetPreServerIdByStep(step); // 从上一个收
212 0 : u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个
213 0 : std::vector<DeviceMem> reduceMem;
214 0 : std::vector<bool> isReduceBlock;
215 0 : isReduceBlock.resize(intraRankSize_);
216 0 : reduceMem.resize(intraRankSize_);
217 0 : u32 retIndex = 0;
218 : // 机内,最后rank的规约结果放在倒数第2块,其他放在倒数第1块
219 : // 机内第0步,规约结果放在第userRank块cclbuffer
220 0 : if (intraRankId_ == intraRankSize_ - 1) {
221 0 : retIndex = intraRankSize_ - SECOND_TO_LAST;
222 : } else {
223 0 : retIndex = intraRankSize_ - 1;
224 : }
225 0 : bool ifUseLastSize = IfUseLastSize(step, sendServerId);
226 0 : u32 userInIdx = GetRankIdx(sendServerId, intraRankId_);
227 : // 最后一块留给allreduce,idx为第sendServerId大块内存的第idx小块
228 0 : u32 idx = 0;
229 0 : for (u32 i = 0; i < intraRankSize_; ++i) {
230 0 : u32 outCCLbufferIdx = 0;
231 : // i == intraRankId_时,需要取userIn内存,
232 0 : if (i == intraRankId_) {
233 : // 第0步,机内最后一个rank取cclbuffer,因为localcopy时将该内存搬到了第userRank_块cclbuffer
234 0 : if (step == 0 && i == intraRankSize_ - 1) {
235 0 : isReduceBlock[i] = true;
236 0 : outCCLbufferIdx = userRank_;
237 0 : DeviceMem cclbufferIntraMem;
238 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, ifUseLastSize));
239 0 : reduceMem[i] = std::move(cclbufferIntraMem);
240 0 : HCCL_DEBUG(
241 : "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i,
242 : outCCLbufferIdx);
243 0 : retIndex = i;
244 0 : } else {
245 0 : isReduceBlock[i] = false;
246 0 : DeviceMem usrInIntraMem;
247 0 : CHK_RET(GetLocalUserInDeviceMem(userInIdx, usrInIntraMem));
248 0 : reduceMem[i] = std::move(usrInIntraMem);
249 0 : HCCL_DEBUG(
250 : "[%s] inter-server local reduce, NO.%u reduceMem stores userInIdx[%u]", __func__, i, userInIdx);
251 0 : }
252 0 : continue;
253 0 : }
254 : // 其他情况:统一处理CCLBuffer,
255 0 : isReduceBlock[i] = true;
256 0 : outCCLbufferIdx = GetRankIdx(recvServerId, idx);
257 0 : DeviceMem cclbufferIntraMem;
258 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, ifUseLastSize));
259 0 : reduceMem[i] = std::move(cclbufferIntraMem);
260 0 : HCCL_DEBUG(
261 : "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i, outCCLbufferIdx);
262 0 : idx++;
263 : // step 0, localreduce到userRank_块内存上
264 0 : if (step == 0 && outCCLbufferIdx == userRank_) {
265 0 : retIndex = i;
266 : }
267 0 : }
268 0 : HCCL_DEBUG(
269 : "[%s] intra-server local reduce, retIndex[%u], intraRankId[%u], intraRankSize[%u]", __func__, retIndex,
270 : intraRankId_, intraRankSize_);
271 0 : CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, false));
272 0 : HCCL_INFO("[%s] intra-server step[%u] run local reduce success", __func__, step);
273 0 : return HCCL_SUCCESS;
274 0 : }
275 :
276 0 : HcclResult AllReduceMultiDeterPipeline::RunInterSend(u32 step)
277 : {
278 0 : HCCL_DEBUG("[%s] inter-server RDMA write begin, step[%u]", __func__, step);
279 : // 使用主流进行rdma
280 0 : u32 recvServerId = GetPreServerIdByStep(step); // 从上一个收
281 0 : u32 sendServerId = GetNextServerIdByStep(step); // 发给下一个
282 0 : LINK recvInterLink = serverLinks_[recvServerId];
283 0 : LINK sendInterLink = serverLinks_[sendServerId];
284 : // 跨机 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块的第1块
285 : // 跨机 输入是input的第(R1+1)%S1组数据(往后1)的第2块
286 0 : u32 reduceRetIndex = intraRankSize_ - SECOND_TO_LAST;
287 0 : u32 recvCCLbufferIdx = GetRankIdx(recvServerId, 0);
288 0 : u32 sendCCLbufferIdx = GetRankIdx(recvServerId, reduceRetIndex);
289 : // 跨机写对端
290 0 : u32 sendRankId = GetRankIdx(sendServerId, intraRankId_); // 发送到sendRankId
291 0 : u32 recvRankId = GetRankIdx(recvServerId, intraRankId_); // 从recvRankId接收
292 0 : u32 remoteRecvCCLbufferIdx = GetRankIdx(serverId_, 0); // 接收端:收的位置
293 : // 2机又从serverId_收也从serverId_发
294 0 : u32 remoteSendCCLbufferIdx = serverSize_ == MIN_SERVER_NUM ?
295 0 : GetRankIdx(serverId_, reduceRetIndex) :
296 0 : GetRankIdx(sendServerId, reduceRetIndex); // 发送端:发的位置
297 0 : DeviceMem recvMem;
298 0 : DeviceMem sendMem;
299 0 : bool ifSendToLastServer = IfUseLastSize(step, sendServerId);
300 : // 最后一个rank则只收小块内存
301 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(recvCCLbufferIdx, recvMem, isLastRank_));
302 : // 如果是发给最后一个rank则只发小块内存
303 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(sendCCLbufferIdx, sendMem, ifSendToLastServer));
304 :
305 0 : HCCL_DEBUG(
306 : "[%s] inter-server RDMA write begin, cclbufferIdx: [%u] send to [%u], [%u] recv from [%u]", __func__,
307 : sendCCLbufferIdx, remoteRecvCCLbufferIdx, recvCCLbufferIdx, remoteSendCCLbufferIdx);
308 0 : HCCL_DEBUG(
309 : "[%s] inter-server RDMA write begin, rankId: [%u] send to [%u], [%u] recv from [%u]", __func__, userRank_,
310 : sendRankId, userRank_, recvRankId);
311 0 : HCCL_DEBUG(
312 : "[%s] inter-server RDMA write begin, if use last small mem? : isLastRank[%u], ifSendToLastServer[%u]", __func__,
313 : isLastRank_, ifSendToLastServer);
314 0 : if (recvInterLink->IsSpInlineReduce() && sendInterLink->IsSpInlineReduce()) {
315 0 : CHK_RET(sendInterLink->TxAck(mainStream_));
316 0 : CHK_RET(recvInterLink->RxAck(mainStream_));
317 0 : DeviceMem dstMem = std::move(recvMem);
318 0 : void* remoteMemPtr = nullptr;
319 0 : CHK_RET(recvInterLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
320 0 : u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
321 0 : u64 size = slices_[recvCCLbufferIdx].size;
322 0 : u64 offset = slices_[remoteSendCCLbufferIdx].offset;
323 0 : u8* intraSrcAddr = beginAddrU8 + offset;
324 0 : DeviceMem srcMem = DeviceMem::create(intraSrcAddr, isLastRank_ ? lastSize_ : size);
325 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_, recvRankId, recvInterLink->GetLinkType()));
326 0 : CHK_RET(sendInterLink->TxDataSignal(mainStream_));
327 0 : CHK_RET(recvInterLink->RxDataSignal(mainStream_));
328 0 : } else {
329 0 : CHK_RET(recvInterLink->TxAck(mainStream_));
330 0 : CHK_RET(sendInterLink->RxAck(mainStream_));
331 :
332 0 : u64 size = slices_[sendCCLbufferIdx].size; // 没有特殊情况发送接收内存大小都是一样
333 : // 发送的size取本地发的数据大小,如果是发给最后一个rank则只发小块内存
334 0 : CHK_RET(sendInterLink->TxAsync(
335 : UserMemType::OUTPUT_MEM, slices_[remoteRecvCCLbufferIdx].offset, sendMem.ptr(),
336 : ifSendToLastServer ? lastSize_ : size, mainStream_));
337 : // 接收的size取远端发的数据大小, 如果是最后一个rank则只收小块内存
338 0 : CHK_RET(recvInterLink->RxAsync(
339 : UserMemType::OUTPUT_MEM, slices_[remoteSendCCLbufferIdx].offset, recvMem.ptr(),
340 : isLastRank_ ? lastSize_ : size, mainStream_));
341 0 : CHK_RET(recvInterLink->PostFinAck(mainStream_));
342 0 : CHK_RET(sendInterLink->WaitFinAck(mainStream_));
343 : }
344 0 : HCCL_INFO("[%s] inter-server step[%u] run RDMA send success", __func__, step);
345 0 : return HCCL_SUCCESS;
346 0 : }
347 :
348 0 : HcclResult AllReduceMultiDeterPipeline::RunFinalReduce()
349 : {
350 0 : HCCL_DEBUG("[%s] intra-server final reduce begin", __func__);
351 0 : std::vector<DeviceMem> reduceMem;
352 0 : std::vector<bool> isReduceBlock;
353 0 : isReduceBlock.resize(serverSize_);
354 0 : reduceMem.resize(serverSize_);
355 :
356 0 : u32 retIndex = serverId_;
357 : // userRank_ == userRankSize_ - 1时取小内存进行最后一次reduce
358 0 : bool ifUseLastSize = isLastRank_;
359 0 : HCCL_DEBUG(
360 : "[%s] intra-server retIndex[%u], interRankSize[%u], ifUseLastSize[%u]", __func__, retIndex, serverSize_,
361 : ifUseLastSize);
362 : // 收集每个机子的数据进行最后的reduce
363 0 : for (u32 i = 0; i < serverSize_; ++i) {
364 0 : DeviceMem cclbufferIntraMem;
365 0 : u32 cclbufferIdx = 0;
366 0 : if (i == serverId_) {
367 0 : cclbufferIdx = userRank_;
368 : } else {
369 0 : cclbufferIdx = GetRankIdx(i, 0);
370 : }
371 : // 所有local reduce数据为第serverId大块的第0块
372 0 : isReduceBlock[i] = true;
373 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(cclbufferIdx, cclbufferIntraMem, ifUseLastSize));
374 0 : reduceMem[i] = std::move(cclbufferIntraMem);
375 0 : HCCL_DEBUG(
376 : "[%s] inter-server final local reduce, NO.%u reduceMem stores cclbufferIdx[%u]", __func__, i, cclbufferIdx);
377 0 : }
378 0 : CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, true)); // final reduce使用主流进行操作
379 0 : HCCL_INFO("[%s] intra-server run final local reduce success", __func__);
380 0 : return HCCL_SUCCESS;
381 0 : }
382 :
383 0 : HcclResult AllReduceMultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase)
384 : {
385 0 : if (isStartPhase) {
386 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
387 0 : CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
388 0 : CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
389 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
390 : } else {
391 0 : CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
392 0 : CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
393 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
394 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
395 : }
396 0 : return HCCL_SUCCESS;
397 : }
398 :
399 0 : HcclResult AllReduceMultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase)
400 : {
401 0 : if (isStartPhase) {
402 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
403 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[reduceMainStreamIdx_], -1));
404 0 : CHK_RET(LocalNotify::Wait(
405 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[reduceMainStreamIdx_],
406 : INVALID_VALUE_STAGE));
407 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
408 : } else {
409 0 : CHK_RET(LocalNotify::Post(
410 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], -1));
411 0 : CHK_RET(
412 : LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], INVALID_VALUE_STAGE));
413 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inCclBuffer_, outCclBuffer_, mainStream_, dispatcher_));
414 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
415 : }
416 0 : return HCCL_SUCCESS;
417 : }
418 :
419 : HcclResult
420 0 : AllReduceMultiDeterPipeline::RunAllGatherInterServer(u32 step, const LINK& prevInterLink, const LINK& nextInterLink)
421 : {
422 0 : HCCL_INFO(
423 : "[%s] inter-server allgather run, userRank[%u], step[%u/%u]", __func__, userRank_, step,
424 : serverSize_ - STEP_OFFSET_TWO);
425 0 : CHK_RET(prevInterLink->TxAck(mainStream_));
426 0 : CHK_RET(nextInterLink->RxAck(mainStream_));
427 0 : u32 rxDMAMemSliceId = (serverSize_ + step) % PARITY_BASE;
428 0 : u32 txDMAMemSliceId = (serverSize_ + step - 1) % PARITY_BASE;
429 0 : UserMemType srcMemType = txDMAMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM;
430 0 : UserMemType dstMemType = rxDMAMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM;
431 0 : u32 txSliceId = ((serverId_ + step) % serverSize_) * intraRankSize_ + intraRankId_;
432 0 : u32 txDataSize = slices_[txSliceId].size;
433 0 : DeviceMem txlocalMem;
434 0 : if (txDMAMemSliceId == serverSizeParity_) {
435 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(txSliceId, txlocalMem, false));
436 : } else {
437 0 : CHK_RET(GetLocalInCclbufferDeviceMem(txSliceId, txlocalMem, false));
438 : }
439 0 : CHK_RET(nextInterLink->TxAsync(dstMemType, slices_[txSliceId].offset, txlocalMem.ptr(), txDataSize, mainStream_));
440 :
441 0 : u32 rxSliceId = ((serverId_ + step + 1) % serverSize_) * intraRankSize_ + intraRankId_;
442 0 : DeviceMem rxLocalMem;
443 0 : if (rxDMAMemSliceId == serverSizeParity_) {
444 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(rxSliceId, rxLocalMem, false));
445 : } else {
446 0 : CHK_RET(GetLocalInCclbufferDeviceMem(rxSliceId, rxLocalMem, false));
447 : }
448 0 : u64 rxDataSize = slices_[rxSliceId].size;
449 0 : CHK_RET(prevInterLink->RxAsync(srcMemType, slices_[rxSliceId].offset, rxLocalMem.ptr(), rxDataSize, mainStream_));
450 0 : HCCL_DEBUG(
451 : "[%s] step[%u], txId[%u], rxId[%u], srcMemType[%u], dstMemType[%u]", __func__, step, txDMAMemSliceId,
452 : rxDMAMemSliceId, srcMemType, dstMemType);
453 0 : HCCL_DEBUG(
454 : "[%s] txlocalMem: ptr[%p], size[%llu], rxLocalMem: ptr[%p], size[%llu]", __func__, txlocalMem.ptr(),
455 : txlocalMem.size(), rxLocalMem.ptr(), rxLocalMem.size());
456 0 : HCCL_DEBUG(
457 : "[%s] send txlocalMem to txSliceId[%llu], recv rxLocalMem from rxSliceId[%llu]", __func__, txSliceId,
458 : rxSliceId);
459 0 : HCCL_INFO("[%s] inter-server allgather success", __func__);
460 0 : return HCCL_SUCCESS;
461 0 : }
462 :
463 0 : HcclResult AllReduceMultiDeterPipeline::RunAllGatherIntraServer(u32 step)
464 : {
465 0 : HCCL_INFO("[%s] intra-server allgather run, userRank[%u], step[%u/%u]", __func__, userRank_, step, serverSize_ - 1);
466 0 : u32 dmaMemSliceId = (serverSize_ + step - 1) % PARITY_BASE;
467 0 : for (u32 i = 1; i < intraRankSize_; i++) {
468 0 : u32 remIntraRankId = (intraRankId_ + i) % intraRankSize_;
469 0 : CHK_RET(intraLinks_[remIntraRankId]->TxAck(subStreams_[i - 1]));
470 0 : CHK_RET(intraLinks_[remIntraRankId]->RxAck(subStreams_[i - 1]));
471 0 : void* remoteMemPtr = nullptr;
472 0 : CHK_RET(intraLinks_[remIntraRankId]->GetRemoteMem(
473 : dmaMemSliceId == serverSizeParity_ ? UserMemType::OUTPUT_MEM : UserMemType::INPUT_MEM, &remoteMemPtr));
474 0 : u32 remoteCclbufferId = ((serverId_ + step) % serverSize_) * intraRankSize_ + remIntraRankId;
475 : DeviceMem src = DeviceMem::create(
476 0 : static_cast<u8*>(remoteMemPtr) + slices_[remoteCclbufferId].offset, slices_[remoteCclbufferId].size);
477 : DeviceMem dst = DeviceMem::create(
478 0 : static_cast<u8*>(usrOutMemPtr_) + slices_[remoteCclbufferId].offset, slices_[remoteCclbufferId].size);
479 0 : CHK_RET(HcclD2DMemcpyAsync(
480 : dispatcher_, dst, src, subStreams_[i - 1], intraLinks_[remIntraRankId]->GetRemoteRank(),
481 : intraLinks_[remIntraRankId]->GetLinkType()));
482 0 : CHK_RET(intraLinks_[remIntraRankId]->TxDataSignal(subStreams_[i - 1]));
483 0 : CHK_RET(intraLinks_[remIntraRankId]->RxDataSignal(subStreams_[i - 1]));
484 0 : }
485 0 : HCCL_INFO("[%s] intra-server allgather success", __func__);
486 0 : return HCCL_SUCCESS;
487 : }
488 :
489 0 : HcclResult AllReduceMultiDeterPipeline::RunAsyncAllgatherPipeline()
490 : {
491 0 : HCCL_INFO("[%s] begin, userRank[%u]", __func__, userRank_);
492 : // 机间 ring algo 逆时针,从后往前
493 0 : u32 prevInterRankId = GetNextServerIdByStep(1);
494 0 : u32 nextInterRankId = GetPreServerIdByStep(1);
495 0 : LINK prevInterLink = serverLinks_[prevInterRankId];
496 0 : LINK nextInterLink = serverLinks_[nextInterRankId];
497 0 : for (u32 step = 0; step < serverSize_; step++) {
498 0 : HCCL_INFO("[%s] allgather pipeline, userRank[%u], step[%u/%u]", __func__, userRank_, step, serverSize_ - 1);
499 0 : CHK_RET(MainRecordSub(0, subStreamNum_));
500 0 : CHK_RET(SubWaitMain(0, subStreamNum_));
501 0 : if (step < serverSize_ - 1) {
502 0 : CHK_RET(RunAllGatherInterServer(step, prevInterLink, nextInterLink));
503 0 : CHK_RET(prevInterLink->PostFinAck(mainStream_));
504 0 : CHK_RET(nextInterLink->WaitFinAck(mainStream_));
505 : // inter的最后一步需要barrier确保数据发完
506 0 : if (step == serverSize_ - STEP_OFFSET_TWO) {
507 0 : CHK_RET(ExecuteBarrier(prevInterLink, nextInterLink, mainStream_));
508 : }
509 : }
510 0 : CHK_RET(RunAllGatherIntraServer(step));
511 0 : CHK_RET(SubRecordMain(0, subStreamNum_));
512 0 : CHK_RET(MainWaitSub(0, subStreamNum_));
513 0 : u32 cclbufferFlag = (serverSize_ + step - 1) % PARITY_BASE;
514 0 : u32 sliceId = ((serverId_ + step) % serverSize_) * intraRankSize_ + intraRankId_;
515 0 : DeviceMem srcMem;
516 0 : if (cclbufferFlag == serverSizeParity_) {
517 0 : CHK_RET(GetLocalOutCclbufferDeviceMem(sliceId, srcMem, false));
518 : } else {
519 0 : CHK_RET(GetLocalInCclbufferDeviceMem(sliceId, srcMem, false));
520 : }
521 0 : DeviceMem dstMem;
522 0 : CHK_RET(GetLocalUserOutDeviceMem(sliceId, dstMem));
523 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_));
524 0 : HCCL_DEBUG(
525 : "[%s] step[%u], cclbufferFlag[%u], sliceId[%u], cclBufferSrcMem: ptr[%p], size[%llu]", __func__, step,
526 : cclbufferFlag, sliceId, srcMem.ptr(), srcMem.size());
527 0 : }
528 0 : HCCL_INFO("[%s] end, userRank[%u]", __func__, userRank_);
529 0 : return HCCL_SUCCESS;
530 0 : }
531 :
532 : // 实现为确定性reduce scatter pipeline + all gather pipeline
533 0 : HcclResult AllReduceMultiDeterPipeline::RunAsync()
534 : {
535 0 : HCCL_INFO(
536 : "[AllReduceMultiDeterPipeline] run begin: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] "
537 : "cclBuffer[%p].",
538 : userRank_, userRankSize_, usrInMemPtr_, usrOutMemPtr_, outCclBuffer_.ptr());
539 0 : CHK_SMART_PTR_NULL(dispatcher_);
540 0 : CHK_RET(RunAsyncReduceScatterPipeline());
541 0 : CHK_RET(RunAsyncAllgatherPipeline());
542 0 : HCCL_INFO("[AllReduceMultiDeterPipeline] AllReduceMultiDeterPipeline success userRank[%u] ", userRank_);
543 0 : return HCCL_SUCCESS;
544 : }
545 :
546 : // 适配新CollExecutor接口
547 0 : HcclResult AllReduceMultiDeterPipeline::Prepare(
548 : HcomCollOpInfo* opInfo, DeviceMem& inBuffer, DeviceMem& outBuffer, const u64 count,
549 : const std::vector<Slice>& slices, const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo,
550 : Stream& mainStream, std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
551 : std::vector<std::shared_ptr<LocalNotify>>& notifySub)
552 : {
553 : // opInfo
554 0 : opInfo_ = opInfo;
555 0 : dataType_ = opInfo_->dataType;
556 0 : unitSize_ = SIZE_TABLE[opInfo_->dataType];
557 0 : memSliceSize_ = opInfo_->count * unitSize_; // 一整块rank的内存大小
558 0 : usrInMemPtr_ = opInfo_->inputAddr;
559 0 : usrOutMemPtr_ = opInfo_->outputAddr;
560 0 : reductionOp_ = opInfo_->reduceOp;
561 :
562 : // stream
563 0 : mainStream_ = mainStream;
564 0 : subStreams_ = subStream;
565 0 : subStreamNum_ = subStreams_.size();
566 0 : CHK_RET(PrepareTopoInfo(level0CommInfo, level1CommInfo));
567 0 : all2allStreamBegin_ = 0;
568 0 : all2allStreamSize_ = intraRankSize_ - 1; // alltoall 从流只需要 intraRankSize_ - 1条
569 0 : reduceMainStreamIdx_ = intraRankSize_ - 1;
570 0 : reduceStreamBegin_ = intraRankSize_ - 1;
571 0 : reduceStreamSize_ = MAX_REDUCE_STREAM_NUM; // reduce 从流只需要MAX_REDUCE_STREAM_NUM条
572 0 : HCCL_INFO(
573 : "[%s] stream: all2allStreamBegin[%u], size[%u], reduceStreamBegin[%u], size[%u], reduceMainStreamIdx[%u]",
574 : __func__, all2allStreamBegin_, all2allStreamSize_, reduceStreamBegin_, reduceStreamSize_, reduceMainStreamIdx_);
575 :
576 : // streamNotify, size: n
577 0 : streamNotifyMain_ = notifyMain;
578 0 : if (streamNotifyMain_.size() < intraRankSize_) {
579 0 : HCCL_ERROR(
580 : "[%s] rank[%u] streamNotifyMain_ size[%u] is smaller than intraRankSize[%u]", __func__, userRank_,
581 : streamNotifyMain_.size(), intraRankSize_);
582 0 : return HCCL_E_INTERNAL;
583 : }
584 0 : streamNotifySub_ = notifySub;
585 0 : if (streamNotifySub_.size() < intraRankSize_) {
586 0 : HCCL_ERROR(
587 : "[%s] rank[%u] streamNotifySub_ size[%u] is smaller than intraRankSize[%u]", __func__, userRank_,
588 : streamNotifySub_.size(), intraRankSize_);
589 0 : return HCCL_E_INTERNAL;
590 : }
591 :
592 : // 此次reduce scatter数据信息
593 0 : inCclBuffer_ = inBuffer;
594 0 : outCclBuffer_ = outBuffer;
595 0 : bufferSize_ = inBuffer.size();
596 0 : slices_ = slices;
597 : // allreduce count为此次处理的数据总数
598 0 : curSize_ = slices_[userRank_].size;
599 0 : count_ = slices_[userRank_].size / unitSize_;
600 0 : lastSize_ = slices_[userRankSize_ - 1].size;
601 0 : isLastRank_ = (userRank_ == userRankSize_ - 1) ? true : false;
602 : // serverSize_是偶数,与正常allreduce pipeline中的allgather pipeline流程一样;若为奇数,则颠倒内存
603 0 : serverSizeParity_ = (serverSize_ % PARITY_BASE == 0) ? 1 : 0;
604 0 : perRankAvgDataSize_ = count * unitSize_ / userRankSize_;
605 0 : if (slices_.size() != userRankSize_) {
606 0 : HCCL_ERROR("[%s] slices size[%llu] not match userRankSize[%u]", __func__, slices_.size(), userRankSize_);
607 0 : return HCCL_E_INTERNAL;
608 : }
609 0 : HCCL_INFO(
610 : "[%s] this time: bufferSize[%u], count[%u], curSize[%u], lastSize[%u], slicesNum[%u] "
611 : "serverSizeParity[%u]",
612 : __func__, bufferSize_, count_, curSize_, lastSize_, slices_.size(), serverSizeParity_);
613 0 : return HCCL_SUCCESS;
614 : }
615 :
616 0 : u64 AllReduceMultiDeterPipeline::GetLocalReduceSerialThresh() { return perRankAvgDataSize_; }
617 :
618 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_MULTI_DETERMINISTIC_PIPELINE, AllReduceMultiDeterPipeline);
619 : } // namespace hccl
|