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