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_ring_concurrent_direct.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : ReduceScatterRingConcurrentDirect::ReduceScatterRingConcurrentDirect(
16 0 : const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
17 : {
18 0 : }
19 :
20 0 : ReduceScatterRingConcurrentDirect::~ReduceScatterRingConcurrentDirect()
21 : {
22 0 : }
23 :
24 0 : HcclResult ReduceScatterRingConcurrentDirect::Prepare(const u64 reduceAttrBitMap, const HcomCollOpInfo *opInfo,
25 : const u32 userRank, std::vector<Stream> &subStreams,
26 : const std::vector<std::shared_ptr<LocalNotify>> &mainSignals,
27 : const std::vector<std::shared_ptr<LocalNotify>> &subSignals,
28 : const std::vector<u32> &ringsOrder,
29 : const std::vector<Slice> &userMemInputSlices, bool isSdma)
30 : {
31 0 : reduceAttr_ = reduceAttrBitMap;
32 0 : opInfo_ = opInfo;
33 0 : userRank_ = userRank;
34 0 : subStreams_ = subStreams;
35 0 : mainSignals_ = mainSignals;
36 0 : subSignals_ = subSignals;
37 0 : ringsOrder_ = ringsOrder;
38 0 : userSlices_ = userMemInputSlices;
39 0 : isSdma_ = isSdma;
40 0 : return HCCL_SUCCESS;
41 : }
42 :
43 : // reduce scatter ring direct算法的函数入口
44 0 : HcclResult ReduceScatterRingConcurrentDirect::RunAsync(const u32 rank, const u32 rankSize,
45 : const std::vector<LINK> &links)
46 : {
47 : // 基本的检查
48 0 : CHK_RET(CheckParameters(rank, rankSize, links));
49 :
50 : // 判断rank_size == 1的情况,并拷贝
51 0 : if (rankSize == 1) {
52 0 : CHK_RET(OneRankMemcpy());
53 0 : return HCCL_SUCCESS;
54 : }
55 0 : HCCL_DEBUG("ReduceScatterRingConcurrentDirect starts: rank[%u]", rank);
56 : // 收集本地mem信息
57 0 : CHK_RET(InitSenderReducer());
58 :
59 : // 收集邻居信息
60 0 : CHK_RET(GetInitializedNeighborLinks(rank, rankSize, links));
61 :
62 : // 填充slice_
63 0 : CHK_RET(SetSlices(rank, rankSize));
64 :
65 : // 运行reduce-scatter, ring算法
66 0 : CHK_RET(RunReduceScatter(rank, rankSize));
67 :
68 0 : if (barrierSwitchOn_) {
69 : // 执行barrier,保证数据收发完成
70 0 : CHK_RET(ExecuteBarrier(leftLink_, rightLink_));
71 : }
72 :
73 0 : CHK_RET(LaunchTaskExtend(dispatcher_, stream_, subStreams_));
74 :
75 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect finished: rank[%u]", rank);
76 0 : return HCCL_SUCCESS;
77 : }
78 :
79 0 : HcclResult ReduceScatterRingConcurrentDirect::CheckParameters(const u32 rank, const u32 rankSize,
80 : const std::vector<LINK> &links)
81 : {
82 0 : CHK_PTR_NULL(opInfo_);
83 0 : CHK_RET(CheckConcurrentDirectParameters(rank, rankSize, links));
84 : // 判断subStreams数量是否正确
85 0 : CHK_PRT_RET(
86 : subStreams_.size() < 1,
87 : HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subStreams size[%u] is less than 1", subStreams_.size()),
88 : HCCL_E_PARA);
89 0 : for (auto &s : subStreams_) {
90 0 : CHK_PTR_NULL(s.ptr());
91 : }
92 : // 判断mainSignals数量是否正确
93 0 : CHK_PRT_RET(
94 : mainSignals_.size() < 1,
95 : HCCL_ERROR("[ReduceScatterRingConcurrentDirect] mainSignals size[%u] is less than 1", mainSignals_.size()),
96 : HCCL_E_PARA);
97 : // 判断subSignals数量是否正确
98 0 : CHK_PRT_RET(
99 : subSignals_.size() < 1,
100 : HCCL_ERROR("[ReduceScatterRingConcurrentDirect] subSignals size[%u] is less than 1", subSignals_.size()),
101 : HCCL_E_PARA);
102 : // 判断ringsOrder数量是否正确
103 0 : CHK_PRT_RET(ringsOrder_.size() != rankSize,
104 : HCCL_ERROR("[ReduceScatterRingConcurrentDirect] ringsOrder size[%u] is not equal to rank size[%u]",
105 : ringsOrder_.size(), rankSize),
106 : HCCL_E_PARA);
107 : // 判断userMemInputSlices数量是否正确
108 0 : CHK_PRT_RET(userSlices_.size() % rankSize != 0,
109 : HCCL_ERROR("[ReduceScatterRingConcurrentDirect] userMemInputSlices size[%u] can not divided by size[%u]",
110 : userSlices_.size(), rankSize),
111 : HCCL_E_PARA);
112 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect CheckParameters success");
113 0 : return HCCL_SUCCESS;
114 : }
115 :
116 0 : HcclResult ReduceScatterRingConcurrentDirect::OneRankMemcpy()
117 : {
118 0 : for (u32 sliceIdx = 0; sliceIdx < slices_.size(); sliceIdx++) {
119 0 : const Slice &srcSlice = userSlices_[sliceIdx];
120 0 : const Slice &dstSlice = slices_[sliceIdx];
121 0 : DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcSlice.offset, srcSlice.size);
122 0 : DeviceMem dst;
123 0 : if (opInfo_->outputAddr != nullptr) {
124 : // opInfo_->outputAddr != nullptr指示要将输出发送至user output
125 0 : u64 stepOffset = slices_[ringsOrder_[0]].offset;
126 0 : HCCL_DEBUG("[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at userMemOut_",
127 : userRank_, stepOffset, dstSlice.size);
128 0 : dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + stepOffset, dstSlice.size);
129 : } else {
130 : // opInfo_->outputAddr == nullptr指示要将输出发送至CCL buffer
131 0 : HCCL_DEBUG("[Memcpy operation] stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at outputMem_",
132 : userRank_, dstSlice.offset, dstSlice.size);
133 0 : dst = outputMem_.range(dstSlice.offset, dstSlice.size);
134 : }
135 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
136 0 : }
137 0 : return HCCL_SUCCESS;
138 : }
139 :
140 0 : HcclResult ReduceScatterRingConcurrentDirect::InitSenderReducer()
141 : {
142 : // 创建reducer & sender
143 0 : senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
144 0 : CHK_SMART_PTR_NULL(senderInfo_);
145 :
146 0 : reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
147 0 : CHK_SMART_PTR_NULL(reducerInfo_);
148 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect finished to InitSenderReducer");
149 0 : return HCCL_SUCCESS;
150 : }
151 :
152 0 : HcclResult ReduceScatterRingConcurrentDirect::GetInitializedNeighborLinks(const u32 rank, const u32 rankSize,
153 : const std::vector<LINK> &links)
154 : {
155 : // 收集左邻居信息
156 0 : leftLink_ = links[(rank + rankSize - 1) % rankSize];
157 0 : CHK_SMART_PTR_NULL(leftLink_);
158 :
159 : // 收集右邻居信息
160 0 : rightLink_ = links[(rank + 1) % rankSize];
161 0 : CHK_SMART_PTR_NULL(rightLink_);
162 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect finished to GetInitializedNeighborLinks");
163 0 : return HCCL_SUCCESS;
164 : }
165 :
166 0 : HcclResult ReduceScatterRingConcurrentDirect::SetSlices(const u32 rank, const u32 rankSize)
167 : {
168 0 : if (slices_.size() == 0) {
169 0 : slices_.resize(rankSize);
170 :
171 : // 生成std::vector<Slice> slices_
172 0 : u64 sliceSize = count_ * SIZE_TABLE[dataType_];
173 : ;
174 :
175 0 : for (u32 i = 0; i < rankSize; i++) {
176 0 : slices_[i].size = sliceSize;
177 : // 用于DMA消减过程中,消除src与dst不对位的风险
178 0 : slices_[i].offset = RoundUpWithDivisor(i * sliceSize, HCCL_MIN_SLICE_ALIGN);
179 :
180 0 : HCCL_DEBUG("rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]", rank, i, slices_[i].offset, i,
181 : slices_[i].size);
182 : }
183 : }
184 0 : if (UNLIKELY(HcclCheckLogLevel(DLOG_DEBUG))) {
185 0 : for (u32 i = 0; i < slices_.size(); i++) {
186 0 : HCCL_DEBUG(
187 : "[ReduceScatterRingConcurrentDirect][SetSlices]rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]",
188 : rank, i, slices_[i].offset, i, slices_[i].size);
189 : }
190 : }
191 : // 最后一步搬到userMemOut_的offset, 不同的ring环offset不一样
192 0 : lastStepOffset_ = slices_[ringsOrder_[0]].offset;
193 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect finished to SetSlices");
194 0 : return HCCL_SUCCESS;
195 : }
196 :
197 0 : HcclResult ReduceScatterRingConcurrentDirect::RunInitStep(const u32 rank, const u32 rankSize)
198 : {
199 : // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3
200 0 : u32 initSlice0Idx = 0;
201 0 : u32 initSlice1Idx = 0;
202 0 : initSlice0Idx = (rank + rankSize - 1) % rankSize;
203 0 : initSlice1Idx = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
204 0 : u32 sliceSize = slices_.size() / rankSize;
205 :
206 0 : for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
207 : // 第-1步,片内将部分数据从userIn搬到cclIn
208 0 : const Slice &srcInitSlice0 = userSlices_[initSlice0Idx * sliceSize + sliceIdx];
209 : DeviceMem srcSubInit
210 0 : = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice0.offset, srcInitSlice0.size);
211 0 : const Slice &dstInitSlice0 = slices_[initSlice0Idx * sliceSize + sliceIdx];
212 0 : DeviceMem dstSubInit = inputMem_.range(dstInitSlice0.offset, dstInitSlice0.size);
213 0 : const Slice &srcInitSlice1 = userSlices_[initSlice1Idx * sliceSize + sliceIdx];
214 : DeviceMem srcInit
215 0 : = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice1.offset, srcInitSlice1.size);
216 0 : const Slice &dstInitSlice1 = slices_[initSlice1Idx * sliceSize + sliceIdx];
217 0 : DeviceMem dstInit = inputMem_.range(dstInitSlice1.offset, dstInitSlice1.size);
218 : // 第-1步并发
219 0 : CHK_RET(MainRecordSub()); // 主流通知从流开始通信
220 0 : CHK_RET(SubWaitMain()); // 从流等待主流通知
221 0 : if (rankSize == TWO_RANK_SIZE && opInfo_->outputAddr != nullptr) {
222 0 : HCCL_DEBUG(
223 : "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
224 : "userMemInput to offset[%llu], size[%llu] on userMemOut_",
225 : userRank_, srcInitSlice1.offset, srcInitSlice1.size, lastStepOffset_, dstInitSlice1.size);
226 0 : dstInit = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
227 0 : dstInitSlice1.size);
228 : } else {
229 0 : HCCL_DEBUG(
230 : "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
231 : "userMemInput to offset[%llu], size[%llu] on CCL",
232 : userRank_, srcInitSlice1.offset, srcInitSlice1.size, dstInitSlice1.offset, dstInitSlice1.size);
233 : }
234 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, stream_));
235 0 : HCCL_DEBUG("Memcpy operation: step[-1] stream[sub] src rank[%u] starts to copy(rcv) offset[%llu], "
236 : " size[%llu] on userMemInput to offset[%llu], size[%llu] on CCL",
237 : userRank_, srcInitSlice0.offset, srcInitSlice0.size, dstInitSlice0.offset, dstInitSlice0.size);
238 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstSubInit, srcSubInit, subStreams_[0]));
239 0 : CHK_RET(SubRecordMain()); // 从流通知主流通信完成
240 0 : CHK_RET(MainWaitSub()); // 主流等待从流通知
241 0 : }
242 0 : return HCCL_SUCCESS;
243 : }
244 :
245 0 : HcclResult ReduceScatterRingConcurrentDirect::PreSync()
246 : {
247 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[0], profilerInput_.stage));
248 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
249 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[0], profilerInput_.stage));
250 0 : return HCCL_SUCCESS;
251 : }
252 :
253 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerSpInlineSlice(const HcclDispatcher dispatcher,
254 : const LINK &link, void *remoteMem, ReducerMemoryInfo reduceMem, Stream &stream)
255 : {
256 0 : const u64 dataBytes = reduceMem.remoteRcvTemp.size();
257 0 : CHK_RET(
258 : HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
259 : dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
260 : link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
261 :
262 0 : if (reduceMem.localsrc != reduceMem.localdst) {
263 0 : HcclResult ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
264 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
265 : HCCL_ERROR("[ReducerRun]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
266 : reduceMem.localdst.ptr()),
267 : ret);
268 : }
269 0 : return HCCL_SUCCESS;
270 : }
271 :
272 : // 仅rdma场景调用:主流一次性PreSync后批量下发inline reduce任务
273 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunSpInlineReduce(
274 : const HcclDispatcher dispatcher, const LINK &link,
275 : const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
276 : {
277 0 : CHK_RET(link->RxDataSignal(stream));
278 0 : void *remoteMem = nullptr;
279 0 : CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
280 0 : CHK_RET(PreSync());
281 0 : for (ReducerMemoryInfo reduceMem : reducerMems) {
282 0 : CHK_RET(ReducerSpInlineSlice(dispatcher, link, remoteMem, reduceMem, stream));
283 0 : }
284 0 : return HCCL_SUCCESS;
285 : }
286 :
287 : // sdma场景主流单个slice的远端读任务
288 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerSdmaRemoteReadSlice(
289 : const LINK &link, const ReducerMemoryInfo &reduceMem, Stream &stream)
290 : {
291 0 : RxMemoryInfo mem{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
292 0 : reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() };
293 0 : CHK_PTR_NULL(mem.dst);
294 0 : void *srcMemPtr = nullptr;
295 0 : CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
296 :
297 0 : DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + mem.srcOffset, mem.len);
298 0 : DeviceMem dstDevMem(static_cast<s8 *>(mem.dst), mem.len);
299 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstDevMem, srcDevMem,
300 : stream, link->GetRemoteRank(), link->GetLinkType()));
301 0 : return HCCL_SUCCESS;
302 0 : }
303 :
304 : // 数据接收确认 + 主流批量本地reduce任务
305 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerLocalReduceSuffix(
306 : const HcclDispatcher dispatcher, const LINK &link,
307 : const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
308 : {
309 0 : if (link->GetSupportDataReceivedAck()) {
310 0 : CHK_RET(link->DataReceivedAck(stream));
311 : }
312 0 : for (ReducerMemoryInfo reduceMem : reducerMems) {
313 0 : u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
314 0 : DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
315 0 : CHK_RET(HcclReduceAsync(dispatcher, reduceSrc.ptr(), dataCount, dataType_,
316 : reductionOp_, stream, reduceMem.localdst.ptr(), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
317 : reduceAttr_));
318 0 : }
319 0 : return HCCL_SUCCESS;
320 : }
321 :
322 : // 仅rdma场景调用:主流一次性PreSync、RxAsync后批量下发本地reduce任务
323 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRunNoSpInlineReduce(
324 : const HcclDispatcher dispatcher, const LINK &link,
325 : const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
326 : {
327 0 : std::vector<RxMemoryInfo> rxMems;
328 0 : for (const ReducerMemoryInfo &reduceMem : reducerMems) {
329 0 : rxMems.emplace_back(RxMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
330 0 : reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() });
331 : }
332 0 : CHK_RET(PreSync());
333 0 : CHK_RET(link->RxAsync(rxMems, stream));
334 0 : CHK_RET(ReducerLocalReduceSuffix(dispatcher, link, reducerMems, stream));
335 0 : return HCCL_SUCCESS;
336 0 : }
337 :
338 0 : HcclResult ReduceScatterRingConcurrentDirect::ReducerRun(const HcclDispatcher dispatcher, const LINK &link,
339 : const std::vector<ReducerMemoryInfo> &reducerMems, Stream &stream)
340 : {
341 0 : CHK_PTR_NULL(stream.ptr());
342 0 : bool isSpInlineReduce = link->IsSpInlineReduce();
343 0 : if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_))) {
344 0 : CHK_RET(ReducerRunSpInlineReduce(dispatcher, link, reducerMems, stream));
345 0 : } else {
346 0 : CHK_RET(ReducerRunNoSpInlineReduce(dispatcher, link, reducerMems, stream));
347 : }
348 0 : return HCCL_SUCCESS;
349 : }
350 :
351 0 : HcclResult ReduceScatterRingConcurrentDirect::RunMainStreamTx(const u32 step,
352 : const std::vector<Slice> &txSliceVector, const std::vector<Slice> &rxSliceVector,
353 : const u32 rank, const u32 rankSize, std::vector<ReducerMemoryInfo> &rxReduceMems)
354 : {
355 : (void) rank;
356 0 : CHK_RET(leftLink_->TxAck(stream_));
357 0 : CHK_RET(rightLink_->RxAck(stream_));
358 0 : u32 sliceSize = slices_.size() / rankSize;
359 :
360 : // 通信,如果是最后一步,则做消减拷贝
361 0 : std::vector<SenderMemoryInfo> txMems;
362 0 : DeviceMem dst;
363 0 : for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
364 : // Ack
365 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], src rank[%u] starts to send offset[%llu] size[%llu] "
366 : "from leftMem_",
367 : step, leftLink_->GetRemoteRank(), rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
368 0 : if (isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
369 0 : HCCL_DEBUG("[RunReduceScatter] sdma DMAReduce step");
370 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
371 : "at userMemOut_",
372 : step, userRank_, lastStepOffset_, rxSliceVector[sliceIdx].size);
373 0 : dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
374 0 : rxSliceVector[sliceIdx].size);
375 : } else {
376 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
377 : "at inputMem_",
378 : step, userRank_, rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
379 0 : dst = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
380 0 : if (!isSdma_ && step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
381 0 : HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce step");
382 0 : finalSrc_ = inputMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
383 0 : finalDst_ = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
384 0 : rxSliceVector[sliceIdx].size);
385 : }
386 : }
387 : // 在inline reduce场景, 需要利用scratchMem_暂存
388 0 : DeviceMem srcMemTemp = scratchMem_.range(rxSliceVector[sliceIdx].offset, rxSliceVector[sliceIdx].size);
389 0 : DeviceMem srcMem = inputMem_.range(txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
390 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], senderInfo_ rank[%u] starts to rcv offset[%llu], "
391 : " size[%llu]",
392 : step, rightLink_->GetRemoteRank(), txSliceVector[sliceIdx].offset, txSliceVector[sliceIdx].size);
393 0 : rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSliceVector[sliceIdx].offset,
394 : dst, dst, srcMemTemp});
395 0 : txMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSliceVector[sliceIdx].offset, srcMem});
396 0 : }
397 0 : CHK_RET(senderInfo_->run(rightLink_, txMems, stream_));
398 0 : return HCCL_SUCCESS;
399 0 : }
400 :
401 : // 仅rdma场景调用:主流Tx前缀 + ReducerRun整段下发
402 0 : HcclResult ReduceScatterRingConcurrentDirect::RunMainStream(const u32 step, std::vector<Slice> txSliceVector,
403 : std::vector<Slice> rxSliceVector, const u32 rank, const u32 rankSize)
404 : {
405 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
406 0 : CHK_RET(RunMainStreamTx(step, txSliceVector, rxSliceVector, rank, rankSize, rxReduceMems));
407 0 : CHK_RET(ReducerRun(dispatcher_, leftLink_, rxReduceMems, stream_));
408 0 : return HCCL_SUCCESS;
409 0 : }
410 :
411 : // 从流提前下发Post(mainSignals),使主流Wait(mainSignals)可在队列未积压时及时通过
412 0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStreamPrePost()
413 : {
414 0 : CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
415 0 : return HCCL_SUCCESS;
416 : }
417 :
418 : // 仅rdma场景调用:从流Wait(subSignals)后批量下发本步拷贝任务
419 0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStream(const u32 step, std::vector<Slice> subSliceVector,
420 : std::vector<Slice> cclSliceVector, const u32 rank, const u32 rankSize)
421 : {
422 : (void) rank;
423 0 : CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
424 0 : for (u32 sliceIdx = 0; sliceIdx < subSliceVector.size(); sliceIdx++) {
425 0 : CHK_RET(RunSubStreamSlice(step, sliceIdx, subSliceVector, cclSliceVector, rankSize));
426 : }
427 0 : return HCCL_SUCCESS;
428 : }
429 :
430 : // 从流单个slice的拷贝任务
431 0 : HcclResult ReduceScatterRingConcurrentDirect::RunSubStreamSlice(const u32 step, const u32 sliceIdx,
432 : const std::vector<Slice> &subSliceVector, const std::vector<Slice> &cclSliceVector, const u32 rankSize)
433 : {
434 0 : HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], src rank[%u] starts to send offset[%llu], size[%llu] "
435 : "from userMemIn_", step, userRank_, subSliceVector[sliceIdx].offset, subSliceVector[sliceIdx].size);
436 0 : DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSliceVector[sliceIdx].offset,
437 0 : subSliceVector[sliceIdx].size);
438 0 : DeviceMem dst;
439 0 : if (step == rankSize - DMA_REDUCE_TWO_OFFSET) {
440 : // do nothing
441 0 : } else if (isSdma_ && step == rankSize - DMA_REDUCE_THREE_OFFSET && opInfo_->outputAddr != nullptr) {
442 0 : HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
443 : "to userMemOut_", step, userRank_, lastStepOffset_, subSliceVector[sliceIdx].size);
444 0 : dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffset_,
445 0 : subSliceVector[sliceIdx].size);
446 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
447 0 : } else {
448 0 : HCCL_DEBUG("Memcpy operation: step[%u] stream[sub], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
449 : "to inputMem_",
450 : step, userRank_, cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
451 0 : dst = inputMem_.range(cclSliceVector[sliceIdx].offset, cclSliceVector[sliceIdx].size);
452 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
453 : }
454 0 : return HCCL_SUCCESS;
455 0 : }
456 :
457 : // 仅sdma场景调用:RxDataSignal后,主从流逐slice交替下发,最后补本地reduce后缀
458 0 : HcclResult ReduceScatterRingConcurrentDirect::RunSdmaStepConcurrent(const u32 step,
459 : const std::vector<ReducerMemoryInfo> &rxReduceMems, const std::vector<Slice> &subSliceVector,
460 : const std::vector<Slice> &cclSliceVector, const u32 rank, const u32 rankSize)
461 : {
462 : (void) rank;
463 0 : CHK_RET(leftLink_->RxDataSignal(stream_));
464 0 : bool isSpInlineReduce = leftLink_->IsSpInlineReduce() && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_));
465 0 : void *remoteMem = nullptr;
466 0 : if (isSpInlineReduce) {
467 0 : CHK_RET(leftLink_->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
468 : }
469 : // 每个slice按 从流Post(mainSignals) -> 主流Wait/Empty/Post -> 从流Wait(subSignals) -> 从流拷贝 ->
470 : // 主流reduce/远端读 的顺序交替下发
471 0 : for (u32 sliceIdx = 0; sliceIdx < subSliceVector.size(); sliceIdx++) {
472 0 : CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
473 0 : CHK_RET(PreSync());
474 0 : CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
475 0 : CHK_RET(RunSubStreamSlice(step, sliceIdx, subSliceVector, cclSliceVector, rankSize));
476 0 : if (isSpInlineReduce) {
477 0 : CHK_RET(ReducerSpInlineSlice(dispatcher_, leftLink_, remoteMem, rxReduceMems[sliceIdx], stream_));
478 : } else {
479 0 : CHK_RET(ReducerSdmaRemoteReadSlice(leftLink_, rxReduceMems[sliceIdx], stream_));
480 : }
481 : }
482 0 : if (!isSpInlineReduce) {
483 0 : CHK_RET(ReducerLocalReduceSuffix(dispatcher_, leftLink_, rxReduceMems, stream_));
484 : }
485 0 : return HCCL_SUCCESS;
486 : }
487 :
488 0 : HcclResult ReduceScatterRingConcurrentDirect::RunReduceScatter(const u32 rank, const u32 rankSize)
489 : {
490 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect starts, the input param rank[%u]", rank);
491 : // 空拷贝用于后续操作附着
492 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
493 :
494 0 : CHK_RET(RunInitStep(rank, rankSize));
495 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
496 0 : CHK_RET(MainRecordSub()); // 主流通知从流开始通信
497 0 : CHK_RET(SubWaitMain()); // 从流等待主流通知
498 : // 空拷贝用于主从流任务并发
499 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
500 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, subStreams_[0], dispatcher_));
501 0 : u32 sliceSize = slices_.size() / rankSize;
502 :
503 : // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3, subSliceIdx = 1
504 0 : u32 txSliceIdx = (rank + rankSize - 1) % rankSize;
505 0 : u32 rxSliceIdx = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
506 0 : u32 subSliceIdx = (rank + rankSize - DMA_REDUCE_THREE_OFFSET) % rankSize;
507 0 : for (u32 step = 0; step < rankSize - 1; step++) {
508 0 : std::vector<Slice> rxSliceVector;
509 0 : std::vector<Slice> cclSliceVector;
510 0 : std::vector<Slice> txSliceVector;
511 0 : std::vector<Slice> subSliceVector;
512 0 : for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
513 0 : rxSliceVector.push_back(slices_[rxSliceIdx * sliceSize + sliceIdx]);
514 0 : cclSliceVector.push_back(slices_[subSliceIdx * sliceSize + sliceIdx]);
515 0 : txSliceVector.push_back(slices_[txSliceIdx * sliceSize + sliceIdx]);
516 0 : subSliceVector.push_back(userSlices_[subSliceIdx * sliceSize + sliceIdx]);
517 : }
518 :
519 : // dispatcher_aicpu 单条流的任务队列存在上限,队列满后host会阻塞下发,因此主流与从流的任务必须
520 : // 交替下发:从流Wait(subSignals)依赖主流Post(subSignals),主流Wait(mainSignals)依赖从流
521 : // Post(mainSignals)。若先集中下发某一条流的全部任务,队列被占满后host阻塞,而队列中等待的信号
522 : // 又需要另一条流尚未下发的任务来产生,两条流互相死等。以下保证每个Wait与其配对的Post在小窗口
523 : // 内先后完成下发,且每条流上的任务序列保持不变。
524 0 : if (!isSdma_) {
525 : // 从流先Post(mainSignals),主流整段下发完成后,从流再Wait并下发本步拷贝任务
526 0 : CHK_RET(RunSubStreamPrePost());
527 : // 主流
528 0 : CHK_RET(RunMainStream(step, txSliceVector, rxSliceVector, rank, rankSize));
529 : // 从流
530 0 : CHK_RET(RunSubStream(step, subSliceVector, cclSliceVector, rank, rankSize));
531 : } else {
532 : // sdma场景:主流Tx前缀下发后,主从流逐slice交替下发
533 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
534 0 : CHK_RET(RunMainStreamTx(step, txSliceVector, rxSliceVector, rank, rankSize, rxReduceMems));
535 0 : CHK_RET(RunSdmaStepConcurrent(step, rxReduceMems, subSliceVector, cclSliceVector, rank, rankSize));
536 0 : }
537 :
538 : // 更新索引
539 0 : subSliceIdx = (subSliceIdx + rankSize - 1) % rankSize;
540 0 : txSliceIdx = (txSliceIdx + rankSize - 1) % rankSize;
541 0 : rxSliceIdx = (rxSliceIdx + rankSize - 1) % rankSize;
542 0 : }
543 0 : CHK_RET(SubRecordMain()); // 从流通知主流通信完成
544 0 : CHK_RET(MainWaitSub()); // 主流等待从流通知
545 0 : if (!isSdma_ && opInfo_->outputAddr != nullptr) {
546 0 : HCCL_DEBUG("[RunReduceScatter] rdma DMAReduce last step");
547 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, finalDst_, finalSrc_, stream_));
548 : }
549 0 : HCCL_INFO("ReduceScatterRingConcurrentDirect finished to RunReduceScatter");
550 0 : return HCCL_SUCCESS;
551 : }
552 :
553 : // 主流通知从流干活
554 0 : HcclResult ReduceScatterRingConcurrentDirect::MainRecordSub()
555 : {
556 0 : for (u32 signalIndex = 0; signalIndex < subSignals_.size(); signalIndex++) {
557 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[signalIndex],
558 : profilerInput_.stage));
559 : }
560 0 : return HCCL_SUCCESS;
561 : }
562 : // 从流等待主流
563 0 : HcclResult ReduceScatterRingConcurrentDirect::SubWaitMain()
564 : {
565 0 : for (u32 streamIndex = 0; streamIndex < subSignals_.size(); streamIndex++) {
566 0 : CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, subSignals_[streamIndex],
567 : profilerInput_.stage));
568 : }
569 0 : return HCCL_SUCCESS;
570 : }
571 : // 主流等待从流
572 0 : HcclResult ReduceScatterRingConcurrentDirect::MainWaitSub()
573 : {
574 0 : for (u32 signalIndex = 0; signalIndex < mainSignals_.size(); signalIndex++) {
575 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[signalIndex], profilerInput_.stage));
576 : }
577 0 : return HCCL_SUCCESS;
578 : }
579 : // 从流告诉主流活干完了
580 0 : HcclResult ReduceScatterRingConcurrentDirect::SubRecordMain()
581 : {
582 0 : for (u32 streamIndex = 0; streamIndex < mainSignals_.size(); streamIndex++) {
583 0 : CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, mainSignals_[streamIndex],
584 : profilerInput_.stage));
585 : }
586 0 : return HCCL_SUCCESS;
587 : }
588 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_RING_DIRECT, ReduceScatterRingConcurrentDirect);
589 : } // namespace hccl
|