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