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 "aligned_reduce_scatter_double_ring.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 32 : AlignedReduceScatterDoubleRing::AlignedReduceScatterDoubleRing(
16 32 : const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
17 : {
18 32 : }
19 :
20 32 : AlignedReduceScatterDoubleRing::~AlignedReduceScatterDoubleRing()
21 : {
22 32 : }
23 :
24 32 : HcclResult AlignedReduceScatterDoubleRing::Prepare(DeviceMem &inputMem, DeviceMem &outputMem,
25 : DeviceMem &scratchMem, const u64 count, const HcclDataType dataType, const Stream &stream,
26 : const std::vector<std::vector<Slice>> &multRingsSlices, const HcclReduceOp reductionOp, const u32 root,
27 : const u64 baseOffset, const bool disableDMAReduce, const u64 reduceAttrBitMap, const HcomCollOpInfo *opInfo,
28 : const u32 userRank, std::vector<Stream> &subStreams, const std::vector<std::shared_ptr<LocalNotify>> &mainSignals,
29 : const std::vector<std::shared_ptr<LocalNotify>> &subSignals, const std::vector<std::vector<u32>> &ringsOrders,
30 : const std::vector<std::vector<Slice>> &userMemInputSlicesOfDoubleRing)
31 : {
32 32 : reduceAttr_ = reduceAttrBitMap;
33 32 : opInfo_ = opInfo;
34 32 : userRank_ = userRank;
35 32 : subStreams_ = subStreams;
36 32 : mainSignals_ = mainSignals;
37 32 : subSignals_ = subSignals;
38 32 : ringsOrders_ = ringsOrders;
39 32 : userMemInputSlicesOfDoubleRing_ = userMemInputSlicesOfDoubleRing;
40 32 : return AlgTemplateBase::Prepare(inputMem, outputMem, scratchMem, count, dataType, stream, multRingsSlices,
41 32 : reductionOp, root, baseOffset, disableDMAReduce);
42 : }
43 :
44 : // reduce scatter ring direct算法的函数入口
45 0 : HcclResult AlignedReduceScatterDoubleRing::RunAsync(const u32 rank, const u32 rankSize,
46 : const std::vector<LINK> &links)
47 : {
48 : // 基本的检查
49 0 : CHK_RET(CheckParameters(rank, rankSize, links));
50 :
51 : // 判断rank_size == 1的情况,并拷贝
52 0 : if (rankSize == 1) {
53 0 : CHK_RET(OneRankMemcpy());
54 0 : return HCCL_SUCCESS;
55 : }
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 : CHK_RET(LaunchTaskExtend(dispatcher_, stream_, subStreams_));
69 :
70 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished: rank[%u] end", rank);
71 0 : return HCCL_SUCCESS;
72 : }
73 :
74 0 : HcclResult AlignedReduceScatterDoubleRing::CheckParameters(const u32 rank, const u32 rankSize,
75 : const std::vector<LINK> &links)
76 : {
77 0 : CHK_PTR_NULL(opInfo_);
78 0 : CHK_RET(CheckConcurrentDirectParameters(rank, rankSize, links));
79 : // 判断subStreams数量是否正确
80 0 : CHK_PRT_RET(
81 : subStreams_.size() < 1,
82 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] subStreams size[%u] is less than 1", subStreams_.size()),
83 : HCCL_E_PARA);
84 0 : for (auto &s : subStreams_) {
85 0 : CHK_PTR_NULL(s.ptr());
86 : }
87 : // 判断mainSignals数量是否正确
88 0 : CHK_PRT_RET(
89 : mainSignals_.size() < 1,
90 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] mainSignals size[%u] is less than 1", mainSignals_.size()),
91 : HCCL_E_PARA);
92 : // 判断subSignals数量是否正确
93 0 : CHK_PRT_RET(
94 : subSignals_.size() < 1,
95 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] subSignals size[%u] is less than 1", subSignals_.size()),
96 : HCCL_E_PARA);
97 : // 判断ringsOrder数量是否正确
98 0 : for (u32 ringIndex = 0; ringIndex < ringsOrders_.size(); ringIndex++) {
99 0 : CHK_PRT_RET(ringsOrders_[ringIndex].size() != rankSize,
100 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] ringsOrders[%u] size[%u] is not equal to rank size[%u]",
101 : ringIndex, ringsOrders_[ringIndex].size(), rankSize),
102 : HCCL_E_PARA);
103 : }
104 : // 判断userMemInputSlices数量是否正确
105 0 : for (u32 ringIndex = 0; ringIndex < userMemInputSlicesOfDoubleRing_.size(); ringIndex++) {
106 0 : CHK_PRT_RET(userMemInputSlicesOfDoubleRing_[ringIndex].size() % rankSize != 0,
107 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] userMemInputSlicesOfDoubleRing[%u] size[%u] can not divided by size[%u]",
108 : ringIndex, userMemInputSlicesOfDoubleRing_[ringIndex].size(), rankSize),
109 : HCCL_E_PARA);
110 : }
111 0 : u32 mainSliceSize = multRingsSlices_[ALIGNED_MAIN_RING_INDEX].size() / rankSize;
112 0 : u32 subSliceSize = multRingsSlices_[ALIGNED_SUB_RING_INDEX].size() / rankSize;
113 0 : CHK_PRT_RET(mainSliceSize != subSliceSize,
114 : HCCL_ERROR("[AlignedReduceScatterDoubleRing] mainSliceSize[%u] is not equal to subSliceSize[%u].",
115 : mainSliceSize, subSliceSize),
116 : HCCL_E_PARA);
117 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished to CheckParameters");
118 0 : return HCCL_SUCCESS;
119 : }
120 :
121 0 : HcclResult AlignedReduceScatterDoubleRing::OneRankMemcpy()
122 : {
123 0 : CHK_RET(MainRecordSub()); // 主流通知从流开始通信
124 0 : CHK_RET(SubWaitMain()); // 从流等待主流通知
125 0 : for (u32 ringIndex = 0; ringIndex < multRingsSlices_.size(); ringIndex++) {
126 0 : for (u32 sliceIdx = 0; sliceIdx < multRingsSlices_[ringIndex].size(); sliceIdx++) {
127 0 : const Slice &srcSlice = userMemInputSlicesOfDoubleRing_[ringIndex][sliceIdx];
128 0 : const Slice &dstSlice = multRingsSlices_[ringIndex][sliceIdx];
129 0 : DeviceMem src = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcSlice.offset, srcSlice.size);
130 0 : DeviceMem dst;
131 0 : if (opInfo_->outputAddr != nullptr) {
132 : // opInfo_->outputAddr != nullptr指示要将输出发送至user output
133 0 : u64 stepOffset = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][0]].offset;
134 0 : HCCL_DEBUG("Memcpy operation: stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at userMemOut_",
135 : userRank_, stepOffset, dstSlice.size);
136 0 : dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + stepOffset, dstSlice.size);
137 : } else {
138 : // opInfo_->outputAddr == nullptr指示要将输出发送至CCL buffer
139 0 : HCCL_DEBUG("Memcpy operation: stream[main], rank[%u] starts to rcv offset[%llu], size[%llu] at outputMem_",
140 : userRank_, dstSlice.offset, dstSlice.size);
141 0 : dst = outputMem_.range(dstSlice.offset, dstSlice.size);
142 : }
143 0 : if (ringIndex == 1) {
144 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
145 : } else {
146 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, subStreams_[0]));
147 : }
148 0 : }
149 0 : HCCL_DEBUG("[AlignedReduceScatterDoubleRing][OneRankMemcpy] ringIndex[%u] Memcpy success", ringIndex);
150 : }
151 0 : CHK_RET(SubRecordMain()); // 从流通知主流通信完成
152 0 : CHK_RET(MainWaitSub()); // 主流等待从流通知
153 0 : return HCCL_SUCCESS;
154 : }
155 :
156 0 : HcclResult AlignedReduceScatterDoubleRing::InitSenderReducer()
157 : {
158 : // 创建reducer & sender
159 0 : senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
160 0 : CHK_SMART_PTR_NULL(senderInfo_);
161 :
162 0 : reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
163 0 : CHK_SMART_PTR_NULL(reducerInfo_);
164 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished to InitSenderReducer");
165 0 : return HCCL_SUCCESS;
166 : }
167 :
168 0 : HcclResult AlignedReduceScatterDoubleRing::GetInitializedNeighborLinks(const u32 rank, const u32 rankSize,
169 : const std::vector<LINK> &links)
170 : {
171 : // 收集左邻居信息
172 0 : leftLink_ = links[(rank + rankSize - 1) % rankSize];
173 0 : CHK_SMART_PTR_NULL(leftLink_);
174 :
175 : // 收集右邻居信息
176 0 : rightLink_ = links[(rank + 1) % rankSize];
177 0 : CHK_SMART_PTR_NULL(rightLink_);
178 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished to GetInitializedNeighborLinks");
179 0 : return HCCL_SUCCESS;
180 : }
181 :
182 0 : HcclResult AlignedReduceScatterDoubleRing::SetSlices(const u32 rank, const u32 rankSize)
183 : {
184 0 : for (u32 ringIndex = 0; ringIndex < multRingsSlices_.size(); ringIndex++) {
185 0 : if (multRingsSlices_[ringIndex].size() == 0) {
186 0 : multRingsSlices_[ringIndex].resize(rankSize);
187 :
188 : // 生成std::vector<Slice> multRingsSlices_[ringIndex]
189 0 : u64 sliceSize = count_ * SIZE_TABLE[dataType_];
190 :
191 0 : for (u32 i = 0; i < rankSize; i++) {
192 0 : multRingsSlices_[ringIndex][i].size = sliceSize;
193 : // 用于DMA消减过程中,消除src与dst不对位的风险
194 0 : multRingsSlices_[ringIndex][i].offset = RoundUpWithDivisor(i * sliceSize, HCCL_MIN_SLICE_ALIGN);
195 :
196 0 : HCCL_DEBUG("multRingsSlices_[%u], rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu]",
197 : ringIndex, rank, i, multRingsSlices_[ringIndex][i].offset, i,
198 : multRingsSlices_[ringIndex][i].size);
199 : }
200 : }
201 0 : for (u32 i = 0; i < multRingsSlices_[ringIndex].size(); i++) {
202 0 : HCCL_DEBUG(
203 : "[AlignedReduceScatterDoubleRing][SetSlices] multRingsSlices_[%u], rank[%u], "
204 : "slices[%u].offset=[%llu], slices[%u].size=[%llu]",
205 : ringIndex, rank, i, multRingsSlices_[ringIndex][i].offset, i, multRingsSlices_[ringIndex][i].size);
206 : }
207 : // 最后一步搬到userMemOut_的offset, 不同的ring环offset不一样
208 : u64 toUserMemOffset;
209 0 : if (ringIndex == 0) {
210 0 : toUserMemOffset = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][0]].offset;
211 : } else {
212 0 : const auto &prevRingSlice = multRingsSlices_[ringIndex - 1][ringsOrders_[ringIndex - 1][rank]];
213 0 : const auto &slice = multRingsSlices_[ringIndex][ringsOrders_[ringIndex][rank]];
214 0 : toUserMemOffset = slice.offset - prevRingSlice.offset;
215 : }
216 0 : HCCL_DEBUG("[AlignedReduceScatterDoubleRing][SetSlices] rank[%u], ring[%u], toUserMemOffset[%u]", rank,
217 : ringIndex, toUserMemOffset);
218 0 : lastStepOffsets_.emplace_back(toUserMemOffset);
219 : }
220 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished to SetSlices");
221 0 : return HCCL_SUCCESS;
222 : }
223 :
224 0 : HcclResult AlignedReduceScatterDoubleRing::PrepareInitSlices(const u32 rankSize,
225 : u64 ringIndex, u32 discontinuousSliceSize, u32 discontinuousSliceIdx, u32 initSlice0Idx, u32 initSlice1Idx,
226 : DeviceMem &dstInit, DeviceMem &srcInit, DeviceMem &dstSubInit, DeviceMem &srcSubInit)
227 : {
228 : // 第-1步,片内将部分数据从userIn搬到cclIn
229 0 : const Slice &srcInitSlice0 = userMemInputSlicesOfDoubleRing_[ringIndex][initSlice0Idx * discontinuousSliceSize + discontinuousSliceIdx];
230 : srcInit
231 0 : = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice0.offset, srcInitSlice0.size);
232 0 : const Slice &dstInitSlice0 = multRingsSlices_[ringIndex][initSlice0Idx * discontinuousSliceSize + discontinuousSliceIdx];
233 0 : dstInit = inputMem_.range(dstInitSlice0.offset, dstInitSlice0.size);
234 :
235 0 : const Slice &srcInitSlice1 = userMemInputSlicesOfDoubleRing_[ringIndex][initSlice1Idx * discontinuousSliceSize + discontinuousSliceIdx];
236 : srcSubInit
237 0 : = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + srcInitSlice1.offset, srcInitSlice1.size);
238 0 : const Slice &dstInitSlice1 = multRingsSlices_[ringIndex][initSlice1Idx * discontinuousSliceSize + discontinuousSliceIdx];
239 0 : dstSubInit = inputMem_.range(dstInitSlice1.offset, dstInitSlice1.size);
240 : // 第-1步并发
241 0 : if (rankSize == TWO_RANK_SIZE && opInfo_->outputAddr != nullptr) {
242 0 : HCCL_DEBUG(
243 : "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
244 : "userMemInput to offset[%llu], size[%llu] on userMemOut_",
245 : userRank_, srcInitSlice1.offset, srcInitSlice1.size, lastStepOffsets_[ringIndex], dstInitSlice1.size);
246 0 : dstInit = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
247 0 : dstInitSlice1.size);
248 : } else {
249 0 : HCCL_DEBUG(
250 : "Memcpy operation: step[-1] stream[main] src rank[%u] starts to copy(rcv) offset[%llu], size[%llu] on "
251 : "userMemInput to offset[%llu], size[%llu] on CCL",
252 : userRank_, srcInitSlice1.offset, srcInitSlice1.size, dstInitSlice1.offset, dstInitSlice1.size);
253 : }
254 0 : HCCL_DEBUG("Memcpy operation: step[-1] stream[sub] src rank[%u] starts to copy(rcv) offset[%llu], "
255 : " size[%llu] on userMemInput to offset[%llu], size[%llu] on CCL",
256 : userRank_, srcInitSlice0.offset, srcInitSlice0.size, dstInitSlice0.offset, dstInitSlice0.size);
257 0 : return HCCL_SUCCESS;
258 : }
259 :
260 0 : HcclResult AlignedReduceScatterDoubleRing::MemcpyInitSlicesOnMainStreams(
261 : u64 ringIndex, DeviceMem &dstInit, DeviceMem &srcInit)
262 : {
263 0 : if (ringIndex == 1) {
264 0 : CHK_RET(MainWaitSub());
265 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
266 0 : CHK_RET(MainRecordSub());
267 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, stream_));
268 : } else {
269 0 : CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
270 0 : CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
271 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstInit, srcInit, subStreams_[0]));
272 : }
273 0 : return HCCL_SUCCESS;
274 : }
275 :
276 0 : HcclResult AlignedReduceScatterDoubleRing::MemcpyInitSlices(
277 : u64 ringIndex, DeviceMem &dstInit, DeviceMem &srcInit, DeviceMem &dstSubInit, DeviceMem &srcSubInit)
278 : {
279 0 : CHK_RET(MemcpyInitSlicesOnMainStreams(ringIndex, dstInit, srcInit));
280 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB && (!disableDMAReduce_)) {
281 0 : HCCL_DEBUG("[AlignedReduceScatterDoubleRing][MemcpyInitSlices] no graph mode");
282 0 : CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
283 0 : CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
284 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstSubInit, srcSubInit, subStreams_[ringIndex + 1]));
285 : } else {
286 0 : HCCL_DEBUG("[AlignedReduceScatterDoubleRing][MemcpyInitSlices] graph mode");
287 0 : CHK_RET(MemcpyInitSlicesOnMainStreams(ringIndex, dstSubInit, srcSubInit));
288 : }
289 0 : return HCCL_SUCCESS;
290 : }
291 :
292 0 : HcclResult AlignedReduceScatterDoubleRing::RunInitStep(const u32 rank, const u32 rankSize)
293 : {
294 : //主环初始indexes
295 0 : u32 initSlice0Idx = (rankSize - rank - 1 + rankSize) % rankSize;
296 0 : u32 initSlice1Idx = (rankSize - rank - DMA_REDUCE_TWO_OFFSET + rankSize) % rankSize;
297 : // 从环初始indexes
298 0 : u32 subInitSlice0Idx = (rank + rankSize - 1) % rankSize;
299 0 : u32 subInitSlice1Idx = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
300 0 : u32 discontinuousSliceSize = multRingsSlices_[ALIGNED_SUB_RING_INDEX].size() / rankSize;
301 0 : DeviceMem dstInit;
302 0 : DeviceMem srcInit;
303 0 : DeviceMem dstSubInit;
304 0 : DeviceMem srcSubInit;
305 0 : DeviceMem subDstInit;
306 0 : DeviceMem subSrcInit;
307 0 : DeviceMem subDstSubInit;
308 0 : DeviceMem subSrcSubInit;
309 0 : for (u32 discontinuousSliceIdx = 0; discontinuousSliceIdx < discontinuousSliceSize; discontinuousSliceIdx++) {
310 0 : CHK_RET(PrepareInitSlices(rankSize, ALIGNED_SUB_RING_INDEX,
311 : discontinuousSliceSize, discontinuousSliceIdx, subInitSlice0Idx, subInitSlice1Idx,
312 : subDstInit, subSrcInit, subDstSubInit, subSrcSubInit));
313 0 : CHK_RET(PrepareInitSlices(rankSize, ALIGNED_MAIN_RING_INDEX,
314 : discontinuousSliceSize, discontinuousSliceIdx, initSlice0Idx, initSlice1Idx,
315 : dstInit, srcInit, dstSubInit, srcSubInit));
316 0 : HCCL_DEBUG("Memcpy operation: step[-1] starts on ring[%u]", ALIGNED_SUB_RING_INDEX);
317 0 : CHK_RET(MemcpyInitSlices(ALIGNED_SUB_RING_INDEX, subDstInit, subSrcInit, subDstSubInit, subSrcSubInit));
318 0 : HCCL_DEBUG("Memcpy operation: step[-1] starts on ring[%u]", ALIGNED_MAIN_RING_INDEX);
319 0 : CHK_RET(MemcpyInitSlices(ALIGNED_MAIN_RING_INDEX, dstInit, srcInit, dstSubInit, srcSubInit));
320 : }
321 0 : return HCCL_SUCCESS;
322 0 : }
323 :
324 0 : HcclResult AlignedReduceScatterDoubleRing::PrepareRunMainStream(u32 ringIndex, Stream &stream,
325 : LINK &preLink, LINK &nextLink)
326 : {
327 0 : HCCL_DEBUG("AlignedReduceScatterDoubleRing PrepareRunMainStream start");
328 0 : if (ringIndex == 1) {
329 0 : stream = stream_;
330 0 : preLink = rightLink_;
331 0 : nextLink = leftLink_;
332 : } else {
333 0 : stream = subStreams_[0];
334 0 : preLink = leftLink_;
335 0 : nextLink = rightLink_;
336 : }
337 0 : HCCL_DEBUG("AlignedReduceScatterDoubleRing PrepareRunMainStream end");
338 0 : return HCCL_SUCCESS;
339 : }
340 :
341 0 : HcclResult AlignedReduceScatterDoubleRing::PreSync(const u32 ringIndex)
342 : {
343 0 : if (ringIndex == 1) {
344 0 : CHK_RET(MainWaitSub());
345 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
346 0 : CHK_RET(MainRecordSub());
347 : } else {
348 0 : CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
349 0 : CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
350 : }
351 0 : return HCCL_SUCCESS;
352 : }
353 :
354 0 : HcclResult AlignedReduceScatterDoubleRing::PrepareDeviceMems(
355 : const u32 step, const u32 ringIndex, const u32 rankSize,
356 : const u32 txSliceIdx, const u32 rxSliceIdx, const u32 subSliceIdx,
357 : std::vector<SenderMemoryInfo> &txReduceMems, std::vector<ReducerMemoryInfo> &rxReduceMems,
358 : std::vector<DeviceMem> &localSrcMems, std::vector<DeviceMem> &localDstMems)
359 : {
360 0 : u32 sliceSize = multRingsSlices_[ringIndex].size() / rankSize;
361 0 : for (u32 sliceIdx = 0; sliceIdx < sliceSize; sliceIdx++) {
362 0 : const Slice &rxSlice = multRingsSlices_[ringIndex][rxSliceIdx * sliceSize + sliceIdx];
363 0 : const Slice &cclSlice = multRingsSlices_[ringIndex][subSliceIdx * sliceSize + sliceIdx];
364 0 : const Slice &txSlice = multRingsSlices_[ringIndex][txSliceIdx * sliceSize + sliceIdx];
365 0 : const Slice &subSlice = userMemInputSlicesOfDoubleRing_[ringIndex][subSliceIdx * sliceSize + sliceIdx];
366 : // PrepareReduceDeviceMems
367 : // Ack
368 0 : DeviceMem dst;
369 0 : if (step == rankSize - DMA_REDUCE_TWO_OFFSET && opInfo_->outputAddr != nullptr) {
370 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
371 : "at userMemOut_", step, userRank_, lastStepOffsets_[ringIndex], rxSlice.size);
372 0 : dst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
373 0 : rxSlice.size);
374 : } else {
375 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], dst rank[%u] starts to rcv offset[%llu], size[%llu] "
376 : "at inputMem_",
377 : step, userRank_, rxSlice.offset, rxSlice.size);
378 0 : dst = inputMem_.range(rxSlice.offset, rxSlice.size);
379 : }
380 : // 在inline reduce场景, 需要利用scratchMem_暂存
381 0 : DeviceMem srcMemTemp = scratchMem_.range(rxSlice.offset, rxSlice.size);
382 0 : DeviceMem srcMem = inputMem_.range(txSlice.offset, txSlice.size);
383 0 : HCCL_DEBUG("Reduce operation: step[%u] stream[main], receiver starts to rcv offset[%llu], size[%llu]",
384 : step, txSlice.offset, txSlice.size);
385 0 : rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlice.offset, dst, dst, srcMemTemp});
386 0 : txReduceMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSlice.offset, srcMem});
387 :
388 : // PrepareLocalCopyDeviceMems
389 0 : DeviceMem localSrt;
390 0 : DeviceMem localDst;
391 0 : if (step == rankSize - DMA_REDUCE_TWO_OFFSET) {
392 : // do nothing
393 0 : } else if (step == rankSize - DMA_REDUCE_THREE_OFFSET && opInfo_->outputAddr != nullptr) {
394 0 : HCCL_DEBUG("Memcpy operation: step[%u] subStream[%u], src rank[%u] sends offset[%llu], size[%llu], "
395 : "dst rank[%u] starts to rcv offset[%llu], size[%llu], "
396 : "from userMemIn_ to userMemOut_", step, ringIndex + 1, userRank_, subSlice.offset, subSlice.size,
397 : userRank_, lastStepOffsets_[ringIndex], subSlice.size);
398 0 : localSrt = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSlice.offset,
399 0 : subSlice.size);
400 0 : localDst = DeviceMem::create(static_cast<u8 *>(opInfo_->outputAddr) + lastStepOffsets_[ringIndex],
401 0 : subSlice.size);
402 : } else {
403 0 : HCCL_DEBUG("Memcpy operation: step[%u] subStream[%u], src rank[%u] sends offset[%llu], size[%llu], "
404 : "dst rank[%u] starts to rcv offset[%llu], size[%llu], "
405 : "from userMemIn_ to inputMem_", step, ringIndex + 1, userRank_, subSlice.offset, subSlice.size,
406 : userRank_, cclSlice.offset, cclSlice.size);
407 0 : localSrt = DeviceMem::create(static_cast<u8 *>(opInfo_->inputAddr) + subSlice.offset,
408 0 : subSlice.size);
409 0 : localDst = inputMem_.range(cclSlice.offset, cclSlice.size);
410 : }
411 0 : localSrcMems.emplace_back(localSrt);
412 0 : localDstMems.emplace_back(localDst);
413 0 : }
414 0 : return HCCL_SUCCESS;
415 : }
416 :
417 0 : HcclResult AlignedReduceScatterDoubleRing::RxAsyncMemcpy(
418 : const u32 ringIndex, RxMemoryInfo& mem, Stream &stream, const LINK &link)
419 : {
420 : // PreSync
421 0 : CHK_RET(PreSync(ringIndex));
422 0 : CHK_PTR_NULL(mem.dst);
423 0 : void *srcMemPtr = nullptr;
424 0 : CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
425 :
426 0 : DeviceMem srcDevMem(static_cast<s8 *>(srcMemPtr) + mem.srcOffset, mem.len);
427 0 : DeviceMem dstDevMem(static_cast<s8 *>(mem.dst), mem.len);
428 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstDevMem, srcDevMem,
429 : stream, link->GetRemoteRank(), link->GetLinkType()));
430 0 : return HCCL_SUCCESS;
431 0 : }
432 :
433 0 : HcclResult AlignedReduceScatterDoubleRing::ReducerRun(const u32 ringIndex, const HcclDispatcher dispatcher,
434 : const LINK &link,
435 : ReducerMemoryInfo &reduceMem, Stream &stream)
436 : {
437 0 : CHK_PTR_NULL(stream.ptr());
438 0 : bool isSpInlineReduce = link->IsSpInlineReduce();
439 0 : HcclResult ret = HCCL_SUCCESS;
440 0 : if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttr_))) {
441 0 : void *remoteMem = nullptr;
442 0 : CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
443 0 : const u64 dataBytes = reduceMem.remoteRcvTemp.size();
444 0 : CHK_RET(PreSync(ringIndex));
445 0 : CHK_RET(
446 : HcclReduceAsync(dispatcher, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
447 : dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(),
448 : link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
449 :
450 0 : if (reduceMem.localsrc != reduceMem.localdst) {
451 0 : ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
452 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
453 : HCCL_ERROR("[AlignedReduceScatterDoubleRing][Run]memcpy_async localSrc[%p] localDst[%p] failed",
454 : reduceMem.localsrc.ptr(), reduceMem.localdst.ptr()),
455 : ret);
456 : }
457 0 : } else {
458 0 : RxMemoryInfo rxMem = RxMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
459 0 : reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size() };
460 :
461 0 : u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
462 0 : DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
463 0 : RxWithReduceMemoryInfo rxWithReduceMem = RxWithReduceMemoryInfo{ UserMemType::INPUT_MEM, reduceMem.remoteMemOffset,
464 0 : reduceMem.remoteRcvTemp.ptr(), reduceMem.remoteRcvTemp.size(), reduceSrc.ptr(), reduceMem.localdst.ptr(),
465 0 : dataCount };
466 0 : CHK_RET(RxAsyncMemcpy(ringIndex, rxMem, stream, link));
467 0 : RxWithReduceMemoryInfo &rxReduceMem = rxWithReduceMem;
468 0 : if (ringIndex == ALIGNED_SUB_RING_INDEX) {
469 0 : CHK_PRT_RET(stream != subStreams_[0],
470 : HCCL_ERROR("[%s] subStreams_[0] should be used for ringIndex=%d", __func__, ALIGNED_SUB_RING_INDEX), HCCL_E_INTERNAL);
471 0 : CHK_RET(LocalNotify::Post(subStreams_[0], dispatcher_, mainSignals_[0], profilerInput_.stage));
472 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[0], profilerInput_.stage));
473 : }
474 0 : CHK_RET(HcclReduceAsync(dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_,
475 : reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
476 : reduceAttr_));
477 0 : if (ringIndex == ALIGNED_SUB_RING_INDEX) {
478 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[0], profilerInput_.stage));
479 0 : CHK_RET(LocalNotify::Wait(subStreams_[0], dispatcher_, subSignals_[0], profilerInput_.stage));
480 : }
481 0 : }
482 0 : return HCCL_SUCCESS;
483 : }
484 :
485 0 : HcclResult AlignedReduceScatterDoubleRing::LocalMemcpy(const u32 step, const u32 rankSize, const u32 ringIndex,
486 : DeviceMem &localSrcMem, DeviceMem &localDstMem)
487 : {
488 : // 通过校验流数判断是单算子模式还是图模式
489 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB && (!disableDMAReduce_)) {
490 0 : CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
491 0 : CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
492 0 : if (localSrcMem != localDstMem && step != rankSize - DMA_REDUCE_TWO_OFFSET) {
493 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, subStreams_[ringIndex + 1]));
494 : }
495 : } else {
496 : // 图模式
497 0 : CHK_RET(PreSync(ringIndex));
498 0 : if (localSrcMem != localDstMem && step != rankSize - DMA_REDUCE_TWO_OFFSET) {
499 0 : if (ringIndex == 1) {
500 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, stream_));
501 : } else {
502 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMem, localSrcMem, subStreams_[0]));
503 : }
504 : }
505 : }
506 0 : return HCCL_SUCCESS;
507 : }
508 :
509 0 : HcclResult AlignedReduceScatterDoubleRing::RunSubStream(
510 : const u32 step, const u32 rankSize, u32 ringIndex,
511 : std::vector<DeviceMem> &localSrcMems, std::vector<DeviceMem> &localDstMems)
512 : {
513 0 : for (u32 sliceIdx = 0; sliceIdx < localSrcMems.size(); sliceIdx++) {
514 0 : CHK_RET(LocalNotify::Post(subStreams_[ringIndex + 1], dispatcher_, mainSignals_[ringIndex + 1], profilerInput_.stage));
515 0 : CHK_RET(LocalNotify::Wait(subStreams_[ringIndex + 1], dispatcher_, subSignals_[ringIndex + 1], profilerInput_.stage));
516 0 : if (step != rankSize - DMA_REDUCE_TWO_OFFSET) {
517 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, localDstMems[sliceIdx], localSrcMems[sliceIdx], subStreams_[ringIndex + 1]));
518 : }
519 : }
520 0 : return HCCL_SUCCESS;
521 : }
522 :
523 0 : HcclResult AlignedReduceScatterDoubleRing::RunAllStreams(const u32 step, const u32 rankSize,
524 : std::vector<SenderMemoryInfo> &mainTxReduceMems, std::vector<ReducerMemoryInfo> &mainRxReduceMems,
525 : std::vector<SenderMemoryInfo> &subTxReduceMems, std::vector<ReducerMemoryInfo> &subRxReduceMems,
526 : std::vector<DeviceMem> &mainLocalSrcMems, std::vector<DeviceMem> &mainLocalDstMems,
527 : std::vector<DeviceMem> &subLocalSrcMems, std::vector<DeviceMem> &subLocalDstMems)
528 : {
529 : (void)subTxReduceMems;
530 : (void)mainTxReduceMems;
531 0 : Stream mainStream;
532 0 : LINK mainPreLink;
533 0 : LINK mainNextLink;
534 0 : Stream subStream;
535 0 : LINK subPreLink;
536 0 : LINK subNextLink;
537 0 : CHK_RET(PrepareRunMainStream(ALIGNED_MAIN_RING_INDEX, mainStream, mainPreLink, mainNextLink));
538 0 : HCCL_DEBUG("Reduce: step[%u] ring[%u], src rank[%u] starts to send slice to dst rank[%u]",
539 : step, ALIGNED_MAIN_RING_INDEX, mainPreLink->GetRemoteRank(), mainNextLink->GetRemoteRank());
540 0 : CHK_RET(PrepareRunMainStream(ALIGNED_SUB_RING_INDEX, subStream, subPreLink, subNextLink));
541 0 : HCCL_DEBUG("Reduce: step[%u] ring[%u], src rank[%u] starts to send slice to dst rank[%u]",
542 : step, ALIGNED_SUB_RING_INDEX, subPreLink->GetRemoteRank(), subNextLink->GetRemoteRank());
543 :
544 0 : CHK_RET(mainNextLink->TxAck(mainStream));
545 0 : CHK_RET(mainPreLink->RxAck(mainStream));
546 0 : CHK_RET(subNextLink->TxAck(subStream));
547 0 : CHK_RET(subPreLink->RxAck(subStream));
548 :
549 0 : u32 sliceSize = multRingsSlices_[ALIGNED_MAIN_RING_INDEX].size() / rankSize;
550 0 : for (u32 memIdx = 0; memIdx < sliceSize; memIdx++) {
551 0 : CHK_RET(ReducerRun(ALIGNED_MAIN_RING_INDEX, dispatcher_, mainPreLink, mainRxReduceMems[memIdx], mainStream));
552 0 : CHK_RET(ReducerRun(ALIGNED_SUB_RING_INDEX, dispatcher_, subPreLink, subRxReduceMems[memIdx], subStream));
553 0 : CHK_RET(LocalMemcpy(step, rankSize, ALIGNED_MAIN_RING_INDEX, mainLocalSrcMems[memIdx], mainLocalDstMems[memIdx]));
554 0 : CHK_RET(LocalMemcpy(step, rankSize, ALIGNED_SUB_RING_INDEX, subLocalSrcMems[memIdx], subLocalDstMems[memIdx]));
555 : }
556 0 : CHK_RET(mainPreLink->TxDataSignal(mainStream));
557 0 : CHK_RET(mainNextLink->RxDataSignal(mainStream));
558 0 : CHK_RET(subPreLink->TxDataSignal(subStream));
559 0 : CHK_RET(subNextLink->RxDataSignal(subStream));
560 0 : return HCCL_SUCCESS;
561 0 : }
562 :
563 0 : HcclResult AlignedReduceScatterDoubleRing::PreRunStreams(
564 : const u32 step, const u32 rankSize,
565 : const u32 txSliceIdxMain, const u32 rxSliceIdxMain, const u32 subSliceIdxMain,
566 : const u32 txSliceIdxSub, const u32 rxSliceIdxSub, const u32 subSliceIdxSub,
567 : std::vector<SenderMemoryInfo> &txReduceMemsMain,
568 : std::vector<ReducerMemoryInfo> &rxReduceMemsMain,
569 : std::vector<SenderMemoryInfo> &txReduceMemsSub,
570 : std::vector<ReducerMemoryInfo> &rxReduceMemsSub,
571 : std::vector<DeviceMem> &localSrcMemsMain, std::vector<DeviceMem> &localDstMemsMain,
572 : std::vector<DeviceMem> &localSrcMemsSub, std::vector<DeviceMem> &localDstMemsSub)
573 : {
574 0 : CHK_RET(PrepareDeviceMems(step, ALIGNED_MAIN_RING_INDEX, rankSize,
575 : txSliceIdxMain, rxSliceIdxMain, subSliceIdxMain,
576 : txReduceMemsMain, rxReduceMemsMain,
577 : localSrcMemsMain, localDstMemsMain));
578 0 : CHK_RET(PrepareDeviceMems(step, ALIGNED_SUB_RING_INDEX, rankSize,
579 : txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub,
580 : txReduceMemsSub, rxReduceMemsSub,
581 : localSrcMemsSub, localDstMemsSub));
582 0 : return HCCL_SUCCESS;
583 : }
584 :
585 0 : HcclResult AlignedReduceScatterDoubleRing::RunReduceScatter(const u32 rank, const u32 rankSize)
586 : {
587 0 : HCCL_INFO("AlignedReduceScatterDoubleRing starts, the input param rank[%u]", rank);
588 : // 空拷贝用于后续操作附着
589 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
590 : // 主环主流通知从环主流开始通信
591 0 : CHK_RET(MainRecordSub());
592 : // 从环主流等待主环主流通知
593 0 : CHK_RET(SubWaitMain());
594 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) {
595 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
596 0 : CHK_RET(ExecEmptyTasks());
597 0 : CHK_RET(RunInitStep(rank, rankSize));
598 : }
599 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
600 0 : CHK_RET(ExecEmptyTasks());
601 :
602 : // 例如rank[0,1,2,3]中,rank0的rxSliceIdx = 2,txSliceIdx = 3, subSliceIdx = 1
603 : // 从环初始indexes
604 0 : u32 txSliceIdxSub = (rank + rankSize - 1) % rankSize;
605 0 : u32 rxSliceIdxSub = (rank + rankSize - DMA_REDUCE_TWO_OFFSET) % rankSize;
606 0 : u32 subSliceIdxSub = (rank + rankSize - DMA_REDUCE_THREE_OFFSET) % rankSize;
607 0 : HCCL_DEBUG("[RunReduceScatter]txSliceIdxSub is [%u], rxSliceIdxSub is [%u], subSliceIdxSub is [%u]",
608 : txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub);
609 : // 主环初始indexes
610 0 : u32 txSliceIdxMain = (rankSize - rank - 1 + rankSize) % rankSize;
611 0 : u32 rxSliceIdxMain = (rankSize - rank - DMA_REDUCE_TWO_OFFSET + rankSize) % rankSize;
612 0 : u32 subSliceIdxMain = (rankSize - rank - DMA_REDUCE_THREE_OFFSET + rankSize) % rankSize;
613 :
614 0 : for (u32 step = 0; step < rankSize - 1; step++) {
615 : // 并发
616 0 : std::vector<SenderMemoryInfo> txReduceMemsMain;
617 0 : std::vector<ReducerMemoryInfo> rxReduceMemsMain;
618 0 : std::vector<SenderMemoryInfo> txReduceMemsSub;
619 0 : std::vector<ReducerMemoryInfo> rxReduceMemsSub;
620 0 : std::vector<DeviceMem> localDstMemsMain;
621 0 : std::vector<DeviceMem> localSrcMemsMain;
622 0 : std::vector<DeviceMem> localSrcMemsSub;
623 0 : std::vector<DeviceMem> localDstMemsSub;
624 0 : CHK_RET(PreRunStreams(step, rankSize,
625 : txSliceIdxMain, rxSliceIdxMain, subSliceIdxMain,
626 : txSliceIdxSub, rxSliceIdxSub, subSliceIdxSub,
627 : txReduceMemsMain, rxReduceMemsMain, txReduceMemsSub, rxReduceMemsSub,
628 : localSrcMemsMain, localDstMemsMain, localSrcMemsSub, localDstMemsSub));
629 0 : CHK_RET(RunAllStreams(step, rankSize, txReduceMemsMain, rxReduceMemsMain, txReduceMemsSub, rxReduceMemsSub,
630 : localSrcMemsMain, localDstMemsMain, localSrcMemsSub, localDstMemsSub));
631 : // 更新索引
632 0 : txSliceIdxSub = (txSliceIdxSub + rankSize - 1) % rankSize;
633 0 : rxSliceIdxSub = (rxSliceIdxSub + rankSize - 1) % rankSize;
634 0 : subSliceIdxSub = (subSliceIdxSub + rankSize - 1) % rankSize;
635 0 : txSliceIdxMain = (txSliceIdxMain + rankSize - 1) % rankSize;
636 0 : rxSliceIdxMain = (rxSliceIdxMain + rankSize - 1) % rankSize;
637 0 : subSliceIdxMain = (subSliceIdxMain + rankSize - 1) % rankSize;
638 0 : }
639 : // 从环主流通知主环主流通信完成
640 0 : CHK_RET(SubRecordMain());
641 : // 主环主流等待从环主流通知
642 0 : CHK_RET(MainWaitSub());
643 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, stream_, dispatcher_));
644 0 : CHK_RET(ExecEmptyTasks());
645 0 : HCCL_INFO("AlignedReduceScatterDoubleRing finished to RunReduceScatter");
646 0 : return HCCL_SUCCESS;
647 : }
648 :
649 0 : HcclResult AlignedReduceScatterDoubleRing::GetActiveSubstreamNum(u32 &activeSubstreamNum)
650 : {
651 0 : activeSubstreamNum = subStreams_.size();
652 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB &&
653 0 : disableDMAReduce_) {
654 0 : if (subStreams_.size() <= 2) {
655 0 : HCCL_ERROR("[AlignedReduceScatterDoubleRing][GetActiveSubstreamNum]subStreams_.size()[%zu] <= 2",
656 : subStreams_.size());
657 0 : return HCCL_E_PARA;
658 : }
659 0 : activeSubstreamNum = subStreams_.size() - 2;
660 : }
661 0 : return HCCL_SUCCESS;
662 : }
663 :
664 :
665 0 : HcclResult AlignedReduceScatterDoubleRing::ExecEmptyTasks()
666 : {
667 0 : u32 activeSubstreamNum = 0;
668 0 : CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
669 0 : for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
670 0 : CHK_RET(AlgTemplateBase::ExecEmptyTask(inputMem_, outputMem_, subStreams_[signalIndex], dispatcher_));
671 : }
672 0 : return HCCL_SUCCESS;
673 : }
674 :
675 : // 主流通知从流干活
676 0 : HcclResult AlignedReduceScatterDoubleRing::MainRecordSub()
677 : {
678 0 : u32 activeSubstreamNum = 0;
679 0 : CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
680 0 : for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
681 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, subSignals_[signalIndex],
682 : profilerInput_.stage));
683 : }
684 0 : return HCCL_SUCCESS;
685 : }
686 : // 从流等待主流
687 0 : HcclResult AlignedReduceScatterDoubleRing::SubWaitMain()
688 : {
689 0 : u32 activeSubstreamNum = 0;
690 0 : CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
691 0 : for (u32 streamIndex = 0; streamIndex < activeSubstreamNum; streamIndex++) {
692 0 : CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, subSignals_[streamIndex],
693 : profilerInput_.stage));
694 : }
695 0 : return HCCL_SUCCESS;
696 : }
697 : // 主流等待从流
698 0 : HcclResult AlignedReduceScatterDoubleRing::MainWaitSub()
699 : {
700 0 : u32 activeSubstreamNum = 0;
701 0 : CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
702 0 : for (u32 signalIndex = 0; signalIndex < activeSubstreamNum; signalIndex++) {
703 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, mainSignals_[signalIndex], profilerInput_.stage));
704 : }
705 0 : return HCCL_SUCCESS;
706 : }
707 : // 从流告诉主流活干完了
708 0 : HcclResult AlignedReduceScatterDoubleRing::SubRecordMain()
709 : {
710 0 : u32 activeSubstreamNum = 0;
711 0 : CHK_RET(GetActiveSubstreamNum(activeSubstreamNum));
712 0 : for (u32 streamIndex = 0; streamIndex < activeSubstreamNum; streamIndex++) {
713 0 : CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, mainSignals_[streamIndex],
714 : profilerInput_.stage));
715 : }
716 0 : return HCCL_SUCCESS;
717 : }
718 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_DB_RING, AlignedReduceScatterDoubleRing);
719 : } // namespace hccl
|