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 "reduce_scatter_multi_deter_pipeline.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : ReduceScatterMultiDeterPipeline::ReduceScatterMultiDeterPipeline(const HcclDispatcher dispatcher)
16 0 : : MultiDeterPipeline(dispatcher)
17 0 : {}
18 :
19 0 : ReduceScatterMultiDeterPipeline::~ReduceScatterMultiDeterPipeline() {}
20 :
21 0 : HcclResult ReduceScatterMultiDeterPipeline::GetRemoteCclbufferDeviceMem(
22 : u32 inputSliceIndex, LINK link, u32 outputSliceIndex, DeviceMem& remoteMem)
23 : {
24 0 : u64 inputSliceOffset = memSliceSize_ * inputSliceIndex + offset_;
25 0 : u64 eachOffset = eachRankCclbufferSize_;
26 0 : u64 outputSliceOffset = eachOffset * outputSliceIndex;
27 0 : u64 outputInSliceOffset = (HCCL_MIN_SLICE_ALIGN_910B + (inputSliceOffset % HCCL_MIN_SLICE_ALIGN_910B)
28 : - (outputSliceOffset % HCCL_MIN_SLICE_ALIGN_910B))
29 0 : % HCCL_MIN_SLICE_ALIGN_910B;
30 0 : void* remoteMemPtr = nullptr;
31 0 : CHK_RET(link->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr)); // 图模式不一定是input,统一output
32 0 : u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
33 0 : u8* intraSrcAddr = beginAddrU8 + outputSliceOffset + outputInSliceOffset;
34 0 : remoteMem = DeviceMem::create(intraSrcAddr, curSize_);
35 0 : if (remoteMem.ptr() == nullptr) {
36 0 : HCCL_ERROR(
37 : "[%s] outputSliceOffset + outputInSliceOffset + curSize_ = [%llu] > cclBufferSize[%llu]", __func__,
38 : outputSliceOffset + outputInSliceOffset + curSize_, cclBuffer_.size());
39 0 : return HCCL_E_MEMORY;
40 : }
41 0 : HCCL_DEBUG(
42 : "[%s] rank[%u], beginAddr[%p], outputSliceOffset[%llu](outputSliceIndex * eachOffset[%llu]), "
43 : "outputInSliceOffset[%llu], curSize[%llu], totalBufferSize[%llu]",
44 : __func__, outputSliceIndex, remoteMem.ptr(), outputSliceOffset, eachOffset, outputInSliceOffset, curSize_,
45 : cclBuffer_.size());
46 0 : return HCCL_SUCCESS;
47 : }
48 :
49 : // RDMA 发送时顶格收,所以不需要128K对齐,故sliceOffset为0
50 : // SDMA 发送后 取本地 sliceOffset = slices_[rankIdInAllRanks].offset偏移处存放地址
51 : HcclResult
52 0 : ReduceScatterMultiDeterPipeline::GetLocalCclbufferDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem, u64 sliceOffset)
53 : {
54 0 : u64 eachOffset = eachRankCclbufferSize_; // 当前轮有效数据大小 + HCCL_MIN_SLICE_ALIGN_910B作为偏移
55 0 : u64 rdmaOffset = eachOffset * rankIdInAllRanks;
56 0 : u64 sdmaOffset = sliceOffset;
57 0 : u64 offset = sliceOffset == 0 ? rdmaOffset : sdmaOffset;
58 0 : localMem = cclBuffer_.range(offset, curSize_);
59 0 : if (localMem.ptr() == nullptr) {
60 0 : HCCL_ERROR(
61 : "[%s] get localMem failed, rdmaOffset + curSize_"
62 : "= [%llu] or sdmaOffset + curSize = [%llu] > cclBufferSize[%llu]",
63 : __func__, rdmaOffset + curSize_, sdmaOffset + curSize_, cclBuffer_.size());
64 0 : return HCCL_E_MEMORY;
65 : }
66 0 : HCCL_DEBUG(
67 : "[%s] rank[%u], beginAddr[%p], offset[%llu](rdmaOffset[%u] or sdmaOffset[%llu]), "
68 : "curSize[%llu], totalBufferSize[%llu]",
69 : __func__, rankIdInAllRanks, localMem.ptr(), sliceOffset, rdmaOffset, sdmaOffset, curSize_, cclBuffer_.size());
70 0 : return HCCL_SUCCESS;
71 : }
72 :
73 0 : HcclResult ReduceScatterMultiDeterPipeline::GetLocalUserInDeviceMem(u32 rankIdInAllRanks, DeviceMem& localMem)
74 : {
75 0 : u8* beginAddrU8 = static_cast<u8*>(usrInMemPtr_);
76 0 : u64 eachOffset = memSliceSize_;
77 0 : u8* intraSrcAddr = beginAddrU8 + (rankIdInAllRanks * eachOffset); // 不用 + offset_,因为usrInMem_已经加过了
78 0 : localMem = DeviceMem::create(intraSrcAddr, curSize_);
79 0 : if (localMem.ptr() == nullptr) {
80 0 : HCCL_ERROR(
81 : "[%s] get localMem failed, rankIdInAllRanks * eachOffset + curSize_ = [%u] is too big", __func__,
82 : rankIdInAllRanks * eachOffset + curSize_, cclBuffer_.size());
83 0 : return HCCL_E_MEMORY;
84 : }
85 0 : HCCL_DEBUG(
86 : "[%s] ranks[%u], offset[%llu](rankIdInAllRanks * eachOffset[%llu])"
87 : "intraSrcAddr[%p], memSliceSize[%llu], usrInMemPtr[%p]",
88 : __func__, rankIdInAllRanks, rankIdInAllRanks * eachOffset, eachOffset, intraSrcAddr, memSliceSize_,
89 : usrInMemPtr_);
90 0 : return HCCL_SUCCESS;
91 : }
92 :
93 0 : HcclResult ReduceScatterMultiDeterPipeline::RunLocalCopy()
94 : {
95 : // 每张卡将自己的input的第user rank块数据搬到output,例如0A 1B 2C
96 0 : DeviceMem userIn;
97 0 : CHK_RET(GetLocalUserInDeviceMem(userRank_, userIn));
98 0 : DeviceMem userOut = DeviceMem::create(usrOutMemPtr_, curSize_);
99 : // 使用主流搬迁卡内数据
100 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, userOut, userIn, mainStream_));
101 0 : HCCL_DEBUG(
102 : "[%s] intra-card copy data from [%u] to usrOutMem[%p] size[%llu]", __func__, userRank_, usrOutMemPtr_,
103 : curSize_);
104 0 : return HCCL_SUCCESS;
105 0 : }
106 :
107 : // 机内alltoall full mesh收集数据, #step表示pairwise的第step步
108 0 : HcclResult ReduceScatterMultiDeterPipeline::RunIntraAlltoallPreSync(u32 step)
109 : {
110 0 : HCCL_DEBUG("[%s] intra-server alltoall begin, step[%u]", __func__, step);
111 0 : HCCL_DEBUG(
112 : "[%s] intra-server SDMA send begin, [serverId, intraRankId] = [%u, %u]", __func__, serverId_, intraRankId_);
113 0 : for (u32 i = 0; i < intraRankSize_ - 1; ++i) {
114 0 : HCCL_DEBUG(
115 : "[%s] intra-server SDMA send begin, userRank[%u] step[%u] pro[%u/%u]", __func__, userRank_, i, i + 1,
116 : intraRankSize_ - 1);
117 : // 从机内rankId为recvIntraRankId收集数据,也发给机内rankId为sendIntraRankId数据
118 0 : u32 recvIntraRankId = GetPreIntraRankIdByStep(i + 1);
119 0 : u32 sendIntraRankId = GetNextIntraRankIdByStep(i + 1);
120 0 : LINK recvIntraLink = intraLinks_[recvIntraRankId];
121 0 : LINK sendIntraLink = intraLinks_[sendIntraRankId];
122 0 : CHK_RET(sendIntraLink->TxAck(subStreams_[i]));
123 0 : CHK_RET(sendIntraLink->RxAck(subStreams_[i]));
124 0 : }
125 : // 增加主从流同步,目的是让SDMA同时进行
126 0 : CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
127 0 : CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
128 0 : CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
129 0 : CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
130 0 : return HCCL_SUCCESS;
131 : }
132 :
133 0 : HcclResult ReduceScatterMultiDeterPipeline::BatchPostNotifyForStreams(
134 : const std::vector<std::vector<std::pair<u32, u32>>>& streamTasks, bool isStartPhase, bool useMainStream)
135 : {
136 0 : if (useMainStream) {
137 0 : HCCL_DEBUG("[%s] use mainStrem, skip notify wait", __func__);
138 0 : return HCCL_SUCCESS;
139 : }
140 0 : for (u32 s = 0; s < MAX_REDUCE_STREAM_NUM; s++) {
141 0 : if (streamTasks[s].empty())
142 0 : continue; // 无任务的流跳过
143 0 : u32 streamIdx = reduceStreamBegin_ + s;
144 0 : if (reduceMainStreamIdx_ == streamIdx) {
145 0 : continue;
146 : }
147 0 : if (isStartPhase) {
148 : // 启动阶段:主流→子流 通知(Post主流,Wait子流)
149 0 : CHK_RET(LocalNotify::Post(
150 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
151 0 : CHK_RET(LocalNotify::Wait(
152 : subStreams_[streamIdx], dispatcher_, streamNotifySub_[streamIdx], profilerInput_.stage));
153 0 : HCCL_DEBUG("[%s] stream[%u] start phase notify done", __func__, streamIdx);
154 : } else {
155 : // 同步阶段:子流→主流 通知(Post子流,Wait主流)
156 0 : CHK_RET(LocalNotify::Post(
157 : subStreams_[streamIdx], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
158 0 : CHK_RET(LocalNotify::Wait(
159 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[streamIdx], profilerInput_.stage));
160 0 : HCCL_DEBUG("[%s] stream[%u] sync phase notify done", __func__, streamIdx);
161 : }
162 : }
163 0 : return HCCL_SUCCESS;
164 : }
165 :
166 : // 机内localreduce首先按序收集所有内存块,接着二分归并reduce,最多使用4条流并行
167 0 : HcclResult ReduceScatterMultiDeterPipeline::RunIntraLocalReduce(u32 step)
168 : {
169 0 : HCCL_DEBUG("[%s] inter-server local reduce begin, step[%u]", __func__, step);
170 0 : u32 recvServerId = GetPreServerIdByStep(step); // 从上一个收
171 0 : u32 sendServerId = GetNextServerIdByStep(step); // 发给发下一个
172 0 : std::vector<DeviceMem> reduceMem;
173 0 : std::vector<bool> isReduceBlock;
174 0 : u32 retIndex = 0;
175 0 : isReduceBlock.resize(intraRankSize_);
176 : // 机内,最后rank的规约结果放在倒数第2块,其他放在倒数第1块
177 0 : if (intraRankId_ == intraRankSize_ - 1) {
178 0 : retIndex = intraRankSize_ - SECOND_TO_LAST;
179 : } else {
180 0 : retIndex = intraRankSize_ - 1;
181 : }
182 : // 机内第0步,规约到usrOut,即rank所在机间的intraRankId_处
183 0 : if (step == 0) {
184 0 : retIndex = intraRankId_;
185 : }
186 0 : HCCL_DEBUG(
187 : "[%s] intra-server local reduce, retIndex[%u], intraRankSize[%u] intraRankId[%u]", __func__, retIndex,
188 : intraRankSize_, intraRankId_);
189 0 : reduceMem.resize(intraRankSize_);
190 0 : u32 userInIdx = GetRankIdx(sendServerId, intraRankId_);
191 : // 最后一块留给allreduce
192 0 : u32 idx = 0;
193 0 : const u32 serverId = (step == 0) ? serverId_ : recvServerId;
194 0 : for (u32 i = 0; i < intraRankSize_; ++i) {
195 0 : if (i == intraRankId_) {
196 0 : if (step == 0) {
197 : // step=0:归到usrOut
198 0 : isReduceBlock[i] = true;
199 0 : DeviceMem usrOutInraMem = DeviceMem::create(usrOutMemPtr_, curSize_);
200 0 : reduceMem[i] = std::move(usrOutInraMem);
201 0 : HCCL_DEBUG("[%s] inter-server local reduce, NO.%u reduceMem stores userOut", __func__, i);
202 0 : } else {
203 : // step≠0:填充usrIn
204 0 : isReduceBlock[i] = false;
205 0 : DeviceMem usrInIntraMem;
206 0 : CHK_RET(GetLocalUserInDeviceMem(userInIdx, usrInIntraMem));
207 0 : reduceMem[i] = std::move(usrInIntraMem);
208 0 : HCCL_DEBUG(
209 : "[%s] inter-server local reduce, NO.%u reduceMem stores userInIdx[%u],", __func__, i, userInIdx);
210 0 : }
211 0 : continue;
212 0 : }
213 : // 其他情况:统一处理CCLBuffer
214 0 : isReduceBlock[i] = true;
215 0 : const u32 outCCLbufferIdx = GetRankIdx(serverId, idx);
216 0 : DeviceMem cclbufferIntraMem;
217 0 : CHK_RET(GetLocalCclbufferDeviceMem(outCCLbufferIdx, cclbufferIntraMem, slices_[outCCLbufferIdx].offset));
218 0 : reduceMem[i] = std::move(cclbufferIntraMem);
219 0 : HCCL_DEBUG(
220 : "[%s] inter-server local reduce, NO.%u reduceMem stores outCCLbufferIdx[%u]", __func__, i, outCCLbufferIdx);
221 0 : idx++;
222 0 : }
223 0 : CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, false));
224 0 : HCCL_INFO("[%s] intra-server step[%u] run local reduce success", __func__, step);
225 0 : return HCCL_SUCCESS;
226 0 : }
227 :
228 0 : HcclResult ReduceScatterMultiDeterPipeline::RunInterSend(u32 step)
229 : {
230 0 : HCCL_DEBUG("[%s] inter-server RDMA write begin, step[%u]", __func__, step);
231 : // 使用主流进行rdma
232 0 : u32 recvServerId = GetPreServerIdByStep(step); // 从上一个收
233 0 : u32 sendServerId = GetNextServerIdByStep(step); // 发给下一个
234 0 : LINK recvInterLink = serverLinks_[recvServerId];
235 0 : LINK sendInterLink = serverLinks_[sendServerId];
236 : // 跨机 输出的位置是buffer的第(Rn+Sn-n)%Sn组分块的第1块
237 : // 跨机 输入是input的第(R1+1)%S1组数据(往后1)的倒数第2块
238 0 : u32 reduceRetIndex = intraRankSize_ - SECOND_TO_LAST;
239 0 : u32 recvCCLbufferIdx = GetRankIdx(recvServerId, 0); // 本端rank0:收的位置
240 0 : u32 sendCCLbufferIdx = GetRankIdx(recvServerId, reduceRetIndex); // 本端rank0:发的位置
241 : // 跨机写对端
242 0 : u32 sendRankId = GetRankIdx(sendServerId, intraRankId_); // 发送到sendRankId
243 0 : u32 recvRankId = GetRankIdx(recvServerId, intraRankId_); // 从recvRankId接收
244 0 : u32 remoteRecvCCLbufferIdx = GetRankIdx(serverId_, 0); // 接收端:收的位置
245 : // 2机又从serverId_收也从serverId_发
246 0 : u32 remoteSendCCLbufferIdx = serverSize_ == MIN_SERVER_NUM ?
247 0 : GetRankIdx(serverId_, reduceRetIndex) :
248 0 : GetRankIdx(sendServerId, reduceRetIndex); // 发送端:发的位置
249 0 : DeviceMem recvMem;
250 0 : DeviceMem sendMem;
251 0 : CHK_RET(GetLocalCclbufferDeviceMem(recvCCLbufferIdx, recvMem, eachRankCclbufferSize_ * recvCCLbufferIdx));
252 0 : CHK_RET(GetLocalCclbufferDeviceMem(sendCCLbufferIdx, sendMem, slices_[sendCCLbufferIdx].offset));
253 0 : HCCL_DEBUG(
254 : "[%s] inter-server RDMA write begin, cclbufferIdx: [%u] send to [%u], [%u] recv from [%u]", __func__,
255 : sendCCLbufferIdx, remoteRecvCCLbufferIdx, recvCCLbufferIdx, remoteSendCCLbufferIdx);
256 0 : HCCL_DEBUG(
257 : "[%s] inter-server RDMA write begin, rankId: [%u] send to [%u], [%u] recv from [%u]", __func__, userRank_,
258 : sendRankId, userRank_, recvRankId);
259 : // A + X 单机16卡为SDMA读语义
260 0 : if (recvInterLink->IsSpInlineReduce() && sendInterLink->IsSpInlineReduce()) {
261 0 : CHK_RET(sendInterLink->TxAck(mainStream_));
262 0 : CHK_RET(recvInterLink->RxAck(mainStream_));
263 0 : DeviceMem dstMem = std::move(recvMem);
264 0 : DeviceMem srcMem;
265 0 : void* remoteMemPtr = nullptr;
266 0 : CHK_RET(recvInterLink->GetRemoteMem(UserMemType::OUTPUT_MEM, &remoteMemPtr));
267 0 : u8* beginAddrU8 = static_cast<u8*>(remoteMemPtr);
268 0 : u8* intraSrcAddr = beginAddrU8 + slices_[remoteSendCCLbufferIdx].offset;
269 0 : srcMem = DeviceMem::create(intraSrcAddr, curSize_);
270 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, mainStream_, recvRankId, recvInterLink->GetLinkType()));
271 0 : CHK_RET(sendInterLink->TxDataSignal(mainStream_));
272 0 : CHK_RET(recvInterLink->RxDataSignal(mainStream_));
273 0 : } else {
274 0 : CHK_RET(recvInterLink->TxAck(mainStream_));
275 0 : CHK_RET(sendInterLink->RxAck(mainStream_));
276 :
277 0 : CHK_RET(sendInterLink->TxAsync(
278 : UserMemType::OUTPUT_MEM, remoteRecvCCLbufferIdx * eachRankCclbufferSize_, sendMem.ptr(), curSize_,
279 : mainStream_));
280 0 : CHK_RET(recvInterLink->RxAsync(
281 : UserMemType::OUTPUT_MEM, slices_[remoteSendCCLbufferIdx].offset, recvMem.ptr(), curSize_, mainStream_));
282 :
283 0 : CHK_RET(recvInterLink->PostFinAck(mainStream_));
284 0 : CHK_RET(sendInterLink->WaitFinAck(mainStream_));
285 : }
286 0 : HCCL_INFO("[%s] inter-server step[%u] run RDMA send success", __func__, step);
287 0 : return HCCL_SUCCESS;
288 0 : }
289 :
290 0 : HcclResult ReduceScatterMultiDeterPipeline::RunFinalReduce()
291 : {
292 : // 主从流同步
293 0 : HCCL_DEBUG("[%s] intra-server final reduce begin", __func__);
294 0 : std::vector<DeviceMem> reduceMem;
295 0 : std::vector<bool> isReduceBlock;
296 0 : u32 retIndex = serverId_;
297 0 : DeviceMem usrOutInraMem = DeviceMem::create(usrOutMemPtr_, curSize_);
298 0 : isReduceBlock.resize(serverSize_);
299 0 : reduceMem.resize(serverSize_);
300 :
301 0 : HCCL_DEBUG("[%s] intra-server retIndex[%u], interRankSize[%u]", __func__, retIndex, serverSize_);
302 : // 收集每个机子的数据进行最后的redeuce
303 0 : for (u32 i = 0; i < serverSize_; ++i) {
304 0 : if (i == serverId_) {
305 0 : isReduceBlock[i] = true;
306 0 : reduceMem[i] = std::move(usrOutInraMem);
307 0 : HCCL_DEBUG("[%s] inter-server final local reduce, NO.%u reduceMem stores userOut", __func__, i);
308 0 : continue;
309 : }
310 : // 所有local reduce数据为第serverId大块的第0块
311 0 : isReduceBlock[i] = true;
312 0 : u32 cclbufferIdx = GetRankIdx(i, 0);
313 0 : DeviceMem cclbufferIntraMem;
314 0 : CHK_RET(GetLocalCclbufferDeviceMem(cclbufferIdx, cclbufferIntraMem, 0));
315 0 : reduceMem[i] = std::move(cclbufferIntraMem);
316 0 : HCCL_DEBUG(
317 : "[%s] inter-server final local reduce, NO.%u reduceMem stores cclbufferIdx[%u]", __func__, i, cclbufferIdx);
318 0 : }
319 0 : CHK_RET(LocalReduce(reduceMem, isReduceBlock, retIndex, true));
320 0 : HCCL_INFO("[%s] intra-server run final local reduce success", __func__);
321 0 : return HCCL_SUCCESS;
322 0 : }
323 :
324 0 : HcclResult ReduceScatterMultiDeterPipeline::AlltoallSync(u32 step, bool isStartPhase)
325 : {
326 0 : if (isStartPhase) {
327 0 : CHK_RET(MainRecordSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
328 0 : CHK_RET(SubWaitMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
329 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
330 : } else {
331 0 : CHK_RET(SubRecordMain(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
332 0 : CHK_RET(MainWaitSub(all2allStreamBegin_, all2allStreamBegin_ + all2allStreamSize_));
333 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
334 : }
335 0 : return HCCL_SUCCESS;
336 : }
337 :
338 0 : HcclResult ReduceScatterMultiDeterPipeline::LocalReduceSync(u32 step, bool isStartPhase)
339 : {
340 0 : if (isStartPhase) {
341 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, streamNotifySub_[reduceMainStreamIdx_], -1));
342 0 : CHK_RET(LocalNotify::Wait(
343 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifySub_[reduceMainStreamIdx_],
344 : INVALID_VALUE_STAGE));
345 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] begin sync", __func__, userRank_, step, allSteps_);
346 : } else {
347 0 : CHK_RET(LocalNotify::Post(
348 : subStreams_[reduceMainStreamIdx_], dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], -1));
349 0 : CHK_RET(
350 : LocalNotify::Wait(mainStream_, dispatcher_, streamNotifyMain_[reduceMainStreamIdx_], INVALID_VALUE_STAGE));
351 0 : HCCL_DEBUG("[%s] userRank[%u], step[%u/%u] end sync", __func__, userRank_, step, allSteps_);
352 : }
353 0 : return HCCL_SUCCESS;
354 : }
355 :
356 : // 每个server内首先要进行alltoall full mesh收集数据,再进行机内local reduce,最后发送给指定server
357 0 : HcclResult ReduceScatterMultiDeterPipeline::RunAsync()
358 : {
359 0 : CHK_RET(RunAsyncReduceScatterPipeline());
360 0 : return HCCL_SUCCESS;
361 : }
362 :
363 : // 适配新CollExecutor接口
364 0 : HcclResult ReduceScatterMultiDeterPipeline::Prepare(
365 : HcomCollOpInfo* opInfo, DeviceMem& cclBuffer, const u64 count, const u64 offset, const std::vector<Slice>& slices,
366 : const SubCommInfo& level0CommInfo, const SubCommInfo& level1CommInfo, Stream& mainStream,
367 : std::vector<Stream>& subStream, std::vector<std::shared_ptr<LocalNotify>>& notifyMain,
368 : std::vector<std::shared_ptr<LocalNotify>>& notifySub)
369 : {
370 : // stream
371 0 : subStreams_ = subStream;
372 0 : mainStream_ = mainStream;
373 0 : subStreamNum_ = subStreams_.size();
374 0 : CHK_RET(PrepareTopoInfo(level0CommInfo, level1CommInfo));
375 0 : all2allStreamBegin_ = 0;
376 0 : all2allStreamSize_ = intraRankSize_ - 1; // alltoall 从流只需要 intraRankSize_ - 1条
377 0 : reduceStreamBegin_ = intraRankSize_ - 1;
378 0 : reduceMainStreamIdx_ = intraRankSize_ - 1;
379 0 : reduceStreamSize_ = MAX_REDUCE_STREAM_NUM; // reduce 从流只需要MAX_REDUCE_STREAM_NUM条
380 0 : HCCL_INFO(
381 : "[%s] stream: all2allStreamBegin[%u], size[%u], reduceStreamBegin[%u], reduceMainStreamIdx[%u], size[%u]",
382 : __func__, all2allStreamBegin_, all2allStreamSize_, reduceStreamBegin_, reduceMainStreamIdx_, reduceStreamSize_);
383 :
384 : // opInfo
385 0 : opInfo_ = opInfo;
386 0 : reductionOp_ = opInfo_->reduceOp;
387 0 : usrInMemPtr_ = opInfo_->inputAddr;
388 0 : usrOutMemPtr_ = opInfo_->outputAddr;
389 0 : dataType_ = opInfo_->dataType;
390 0 : unitSize_ = SIZE_TABLE[opInfo_->dataType];
391 0 : memSliceSize_ = opInfo_->count * unitSize_; // 一整块rank的内存大小
392 :
393 : // streamNotify, size: n
394 0 : streamNotifySub_ = notifySub;
395 0 : if (streamNotifySub_.size() < intraRankSize_) {
396 0 : HCCL_ERROR(
397 : "[%s] rank[%u] streamNotifySub_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
398 : streamNotifySub_.size(), intraRankSize_);
399 0 : return HCCL_E_INTERNAL;
400 : }
401 0 : streamNotifyMain_ = notifyMain;
402 0 : if (streamNotifyMain_.size() < intraRankSize_) {
403 0 : HCCL_ERROR(
404 : "[%s] rank[%u] streamNotifyMain_ size [%u] error, is smaller than intraRankSize[%u]", __func__, userRank_,
405 : streamNotifyMain_.size(), intraRankSize_);
406 0 : return HCCL_E_INTERNAL;
407 : }
408 0 : HCCL_INFO(
409 : "[%s] notify: streamNum[%u], streamNotifySubNum[%u], streamNotifyMainNum[%u]", __func__, subStreams_.size(),
410 : streamNotifySub_.size(), streamNotifyMain_.size());
411 :
412 : // 此次reduce scatter数据信息
413 0 : cclBuffer_ = cclBuffer;
414 0 : count_ = count;
415 0 : curSize_ = count_ * unitSize_;
416 0 : bufferSize_ = cclBuffer.size();
417 0 : offset_ = offset;
418 0 : slices_ = slices;
419 0 : eachRankCclbufferSize_ = curSize_ + HCCL_MIN_SLICE_ALIGN_910B;
420 0 : if (slices_.size() != userRankSize_) {
421 0 : HCCL_ERROR("[%s] slices size[%llu] not match userRankSize[%u]", __func__, slices_.size(), userRankSize_);
422 0 : return HCCL_E_INTERNAL;
423 : }
424 0 : HCCL_INFO(
425 : "[%s] this time: bufferSize[%u], count[%u], curSize[%u], offset[%u], slicesNum[%u]", __func__, bufferSize_,
426 : count_, curSize_, offset_, slices_.size());
427 0 : return HCCL_SUCCESS;
428 : }
429 :
430 0 : u64 ReduceScatterMultiDeterPipeline::GetLocalReduceSerialThresh() { return curSize_; }
431 :
432 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_MULTI_DETERMINISTIC_PIPELINE, ReduceScatterMultiDeterPipeline);
433 : } // namespace hccl
|