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_unified_march.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 : static const u32 NEIGHBORS_NUM_TWO = 2; // 2: 邻居数量
16 : static const u32 NEIGHBORS_NUM_ONE = 1; // 1: 邻居数量
17 : static const u32 DIVISOR_NUM_TWO = 2;
18 :
19 1 : ReduceScatterUnifiedMarch::ReduceScatterUnifiedMarch(const HcclDispatcher dispatcher)
20 1 : : AlgTemplateBase(dispatcher)
21 1 : {}
22 :
23 2 : ReduceScatterUnifiedMarch::~ReduceScatterUnifiedMarch() {}
24 :
25 1 : HcclResult ReduceScatterUnifiedMarch::Prepare(Stream &mainStream, SubCommInfo &level0CommInfo,
26 : DeviceMem &userInput, DeviceMem &userOutput, DeviceMem &usrInMem,
27 : DeviceMem &scratchMem, u64 totalCount, std::vector<Stream> &subStreams,
28 : const std::vector<std::shared_ptr<LocalNotify>> &meshSignalMainToSub,
29 : const std::vector<std::shared_ptr<LocalNotify>> &meshSignalSubToMain,
30 : const HcclDataType dataType, const HcclReduceOp reductionOp,
31 : const std::vector<std::vector<Slice>> &multRingsUserMemSlice, u64 reduceAttrBitMap)
32 : {
33 1 : reduceAttr_ = reduceAttrBitMap;
34 1 : mainStream_ = mainStream;
35 1 : intraRank_ = level0CommInfo.localRank;
36 1 : intraRankSize_ = level0CommInfo.localRankSize;
37 1 : CHK_PRT_RET(intraRankSize_ == 0 || (intraRankSize_ % DIVISOR_NUM_TWO != 0),
38 : HCCL_ERROR("[ReduceScatterUnifiedMarch][Prepare]intraRankSize_ is zero or not divisible by 2"),
39 : HCCL_E_PARA);
40 1 : links_ = level0CommInfo.links;
41 :
42 1 : userInput_ = userInput;
43 1 : userOutput_ = userOutput;
44 1 : usrInMem_ = usrInMem;
45 1 : scratchMem_ = scratchMem;
46 1 : HCCL_INFO("userInput_[%p] size[%llu], userOutput_[%p] size[%llu], usrInMem_[%p] size[%llu], scratchMem_[%p] size[%llu]",
47 : userInput_.ptr(), userInput_.size(), userOutput_.ptr(), userOutput_.size(), usrInMem_.ptr(), usrInMem_.size(),
48 : scratchMem_.ptr(), scratchMem_.size());
49 :
50 1 : subStreams_ = subStreams;
51 1 : meshSignalMainToSub_ = meshSignalMainToSub;
52 1 : meshSignalSubToMain_ = meshSignalSubToMain;
53 1 : CHK_PRT_RET(subStreams_.size() < NEIGHBORS_NUM_TWO || meshSignalMainToSub_.size() < NEIGHBORS_NUM_TWO ||
54 : meshSignalSubToMain_.size() < NEIGHBORS_NUM_TWO,
55 : HCCL_ERROR("[AllGatherUnifiedMarch] subStreams_ size[%u] or meshSignalMainToSub_ size[%u] or "\
56 : "meshSignalSubToMain_ size[%u] is less than 2",
57 : subStreams_.size(), meshSignalMainToSub_.size(), meshSignalSubToMain_.size()), HCCL_E_PARA);
58 :
59 1 : totalCount_ = totalCount;
60 1 : dataType_ = dataType;
61 1 : reductionOp_ = reductionOp;
62 1 : blockDataByte_ = totalCount_ * SIZE_TABLE[dataType_];
63 1 : multRingsUserMemSlice_ = multRingsUserMemSlice;
64 1 : CHK_PRT_RET(multRingsUserMemSlice_[0].size() % intraRankSize_ != 0,
65 : HCCL_ERROR("[ReduceScatterUnifiedMarch] multRingsUserMemSlice_[0] size[%u] can not be divided by rank size[%u]",
66 : multRingsUserMemSlice_[0].size(), intraRankSize_), HCCL_E_PARA);
67 :
68 1 : return HCCL_SUCCESS;
69 : }
70 :
71 0 : std::string ReduceScatterUnifiedMarch::GetStreamIndexString()
72 : {
73 0 : std::string res = "";
74 0 : for (u32 streamIndex = 0; streamIndex < subStreams_.size(); streamIndex++) {
75 0 : res += std::to_string(streamIndex) + ", ";
76 : }
77 0 : return res;
78 0 : }
79 :
80 : // 主流通知所有从流
81 0 : HcclResult ReduceScatterUnifiedMarch::NotifySubStreamStart(u32 streamSize)
82 : {
83 0 : CHK_PRT_RET(streamSize > subStreams_.size() || streamSize > meshSignalSubToMain_.size(),
84 : HCCL_ERROR("[ReduceScatterUnifiedMarch][NotifySubStreamStart] streamSize[%u] is out of range"\
85 : "subStreams_ size[%zu] or meshSignalSubToMain_ size[%zu]", streamSize, subStreams_.size(),
86 : meshSignalSubToMain_.size()), HCCL_E_PARA);
87 0 : for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
88 0 : CHK_RET(LocalNotify::Post(mainStream_, dispatcher_, meshSignalSubToMain_[streamIndex], INVALID_VALUE_STAGE));
89 0 : CHK_RET(LocalNotify::Wait(subStreams_[streamIndex], dispatcher_, meshSignalSubToMain_[streamIndex],
90 : INVALID_VALUE_STAGE));
91 : }
92 0 : HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifySubStreamStart] intraRank [%u] main stream notify substream [%s]",
93 : intraRank_, GetStreamIndexString().c_str());
94 0 : return HCCL_SUCCESS;
95 : }
96 :
97 0 : HcclResult ReduceScatterUnifiedMarch::WaitSubStreamFinish(u32 streamSize)
98 : {
99 0 : CHK_PRT_RET(streamSize > subStreams_.size() || streamSize > meshSignalMainToSub_.size(),
100 : HCCL_ERROR("[ReduceScatterUnifiedMarch][WaitSubStreamFinish] streamSize[%u] is out of range"\
101 : "subStreams_ size[%zu] or meshSignalMainToSub_ size[%zu]",
102 : streamSize, subStreams_.size(), meshSignalMainToSub_.size()), HCCL_E_PARA);
103 0 : for (u32 streamIndex = 0; streamIndex < streamSize; streamIndex++) {
104 0 : CHK_RET(LocalNotify::Post(subStreams_[streamIndex], dispatcher_, meshSignalMainToSub_[streamIndex],
105 : INVALID_VALUE_STAGE));
106 0 : CHK_RET(LocalNotify::Wait(mainStream_, dispatcher_, meshSignalMainToSub_[streamIndex],
107 : INVALID_VALUE_STAGE));
108 : }
109 0 : HCCL_DEBUG("[ReduceScatterUnifiedMarch][WaitSubStreamFinish] intraRank [%u] main stream wait substream [%s]",
110 : intraRank_, GetStreamIndexString().c_str());
111 0 : return HCCL_SUCCESS;
112 : }
113 :
114 0 : HcclResult ReduceScatterUnifiedMarch::NotifyNeighborsStart(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
115 : {
116 : // 图模式保持使用Post/Wait接口
117 0 : if (GetWorkflowMode() != HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
118 : // notify是否越界由平台侧保证
119 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
120 0 : if (neighborRankId == 0) {
121 0 : CHK_RET(nextIntralLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
122 0 : CHK_RET(prevIntraLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait
123 0 : } else if (neighborRankId == 1) {
124 0 : CHK_RET(prevIntraLink->Post(notifyIdx_, subStreams_[neighborRankId])); // AckRecord
125 0 : CHK_RET(nextIntralLink->Wait(notifyIdx_, subStreams_[neighborRankId])); // AckWait
126 : }
127 : }
128 0 : HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done",
129 : intraRank_, neighbors);
130 0 : return HCCL_SUCCESS;
131 : }
132 :
133 : // 一条流负责一个环
134 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
135 : // 交替使用Ack和DataSignal两种notify
136 0 : const u32 NOTIFY_IDX_TWO = 2;
137 0 : if (neighborRankId == 0) {
138 0 : if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
139 0 : CHK_RET(nextIntralLink->TxAck(subStreams_[neighborRankId])); // AckRecord
140 0 : CHK_RET(prevIntraLink->RxAck(subStreams_[neighborRankId])); // AckWait
141 : } else {
142 0 : CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
143 0 : CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId])); // DataWait
144 : }
145 0 : } else if (neighborRankId == 1) {
146 0 : if (notifyIdx_ % NOTIFY_IDX_TWO == 0) {
147 0 : CHK_RET(prevIntraLink->TxAck(subStreams_[neighborRankId])); // AckRecord
148 0 : CHK_RET(nextIntralLink->RxAck(subStreams_[neighborRankId])); // AckWait
149 : } else {
150 0 : CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
151 0 : CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId])); // DataWait
152 : }
153 : }
154 : }
155 0 : HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsStart] intraRank[%u] switch on [%u]neigbhbors done",
156 : intraRank_, neighbors);
157 0 : return HCCL_SUCCESS;
158 : }
159 :
160 0 : HcclResult ReduceScatterUnifiedMarch::NotifyNeighborsEnd(LINK& prevIntraLink, LINK& nextIntralLink, u32 neighbors)
161 : {
162 0 : for (u32 neighborRankId = 0; neighborRankId < neighbors; neighborRankId++) {
163 0 : if (neighborRankId == 0) {
164 0 : CHK_RET(prevIntraLink->TxDataSignal(subStreams_[neighborRankId])); // DataRecord
165 0 : CHK_RET(nextIntralLink->RxDataSignal(subStreams_[neighborRankId]));
166 0 : } else if (neighborRankId == 1) {
167 0 : CHK_RET(nextIntralLink->TxDataSignal(subStreams_[neighborRankId]));
168 0 : CHK_RET(prevIntraLink->RxDataSignal(subStreams_[neighborRankId]));
169 : }
170 : }
171 0 : HCCL_DEBUG("[ReduceScatterUnifiedMarch][NotifyNeighborsEnd] intraRank[%u] notifys [%u]neigbhbors reduce done",
172 : intraRank_, neighbors);
173 0 : return HCCL_SUCCESS;
174 : }
175 :
176 0 : HcclResult ReduceScatterUnifiedMarch::DoSerialReduce(void* remDMAMemPtr, void* dstAddr, u64 memSize,
177 : u64 dataCount, Stream &tmpStream, LINK& tmpLink, u64 remoteOffsetByte)
178 : {
179 0 : for (u32 sliceIdx = 0; sliceIdx < (multRingsUserMemSlice_[0].size() / intraRankSize_); sliceIdx++) {
180 0 : DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(remDMAMemPtr) + remoteOffsetByte +
181 0 : multRingsUserMemSlice_[0][sliceIdx].offset, memSize);
182 : DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(dstAddr) +
183 0 : multRingsUserMemSlice_[0][sliceIdx].offset, memSize);
184 :
185 0 : if ((INLINE_REDUCE_BITMASK & reduceAttr_) == 1) { // inlineReduce
186 0 : struct hccl::Transport::Buffer remoteBuf;
187 0 : remoteBuf.addr = srcMem.ptr();
188 0 : remoteBuf.size = srcMem.size();
189 0 : struct hccl::Transport::Buffer localBuf;
190 0 : localBuf.addr = dstMem.ptr();
191 0 : localBuf.size = dstMem.size();
192 0 : HCCL_DEBUG("intralRank[%u] slice[%u] offset[%llu] do inlinereduce with remoteBuf[addr[%p], size[%llu]] and "\
193 : "localBuf[addr[%p], size[%llu]]", intraRank_, sliceIdx, multRingsUserMemSlice_[0][sliceIdx].offset,
194 : remoteBuf.addr, remoteBuf.size, localBuf.addr, localBuf.size);
195 0 : CHK_RET(tmpLink->ReadReduceSync(localBuf, remoteBuf, dataType_, reductionOp_, tmpStream));
196 : } else { // TBE_reduce
197 : // left的inputMem拷到本端的scratchMem
198 0 : DeviceMem tempMem = scratchMem_.range(remoteOffsetByte, srcMem.size());
199 0 : struct hccl::Transport::Buffer remoteBuf;
200 0 : remoteBuf.addr = srcMem.ptr();
201 0 : remoteBuf.size = srcMem.size();
202 0 : struct hccl::Transport::Buffer localBuf;
203 0 : localBuf.addr = tempMem.ptr();
204 0 : localBuf.size = tempMem.size();
205 0 : HCCL_DEBUG("intralRank[%u] slice[%u] offset[%llu] do SDMA read with remoteBuf[addr[%p], size[%llu]] and "\
206 : "localBuf[addr[%p], size[%llu]]", intraRank_, sliceIdx, multRingsUserMemSlice_[0][sliceIdx].offset,
207 : remoteBuf.addr, remoteBuf.size, localBuf.addr, localBuf.size);
208 0 : CHK_RET(tmpLink->ReadSync(localBuf, remoteBuf, tmpStream));
209 0 : CHK_RET(HcclReduceAsync(dispatcher_, tempMem.ptr(), dataCount, dataType_, reductionOp_, tmpStream,
210 : dstMem.ptr(), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttr_));
211 0 : }
212 0 : }
213 0 : return HCCL_SUCCESS;
214 : }
215 :
216 0 : HcclResult ReduceScatterUnifiedMarch::RunSingleSliceRead(u32 ringPrevRank, u32 ringNextRank, u32 step, u32 totalStep)
217 : {
218 0 : LINK prevIntraLink = links_[ringPrevRank];
219 0 : CHK_SMART_PTR_NULL(prevIntraLink);
220 0 : LINK nextIntralLink = links_[ringNextRank];
221 0 : CHK_SMART_PTR_NULL(nextIntralLink);
222 0 : u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
223 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, neighbors));
224 :
225 : //拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
226 0 : CHK_RET(WaitSubStreamFinish(neighbors));
227 0 : CHK_RET(NotifySubStreamStart(neighbors));
228 :
229 : // 从前向rank读取数据
230 0 : void* preRemDMAMemPtr = nullptr;
231 0 : CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::INPUT_MEM, &preRemDMAMemPtr));
232 0 : u32 preDataIndex = (intraRank_ + intraRankSize_ - step - totalStep) % intraRankSize_;
233 0 : u64 preOffsetByte = preDataIndex * blockDataByte_;
234 0 : void* preDstAddr = static_cast<u8 *>(userInput_.ptr()) + preOffsetByte;
235 :
236 0 : CHK_RET(DoSerialReduce(preRemDMAMemPtr, preDstAddr, blockDataByte_, totalCount_,
237 : subStreams_[0], prevIntraLink, preOffsetByte));
238 0 : HCCL_INFO("[ReduceScatterUnifiedMarch][RunSingleSliceRead] intralRank [%u] reduce with ringPrevRank [%u] done",
239 : intraRank_, ringPrevRank);
240 :
241 : // 从后向rank读取数据
242 0 : if (neighbors > NEIGHBORS_NUM_ONE) {
243 0 : void* nextRemDMAMemPtr = nullptr;
244 0 : CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::INPUT_MEM, &nextRemDMAMemPtr));
245 0 : u32 nextDataIndex = (intraRank_ + totalStep + step) % intraRankSize_;
246 0 : u64 nextOffsetByte = nextDataIndex * blockDataByte_;
247 0 : void* nextDstAddr = static_cast<u8 *>(userInput_.ptr()) + nextOffsetByte;
248 :
249 0 : CHK_RET(DoSerialReduce(nextRemDMAMemPtr, nextDstAddr, blockDataByte_, totalCount_,
250 : subStreams_[1], nextIntralLink, nextOffsetByte));
251 0 : HCCL_INFO("[ReduceScatterUnifiedMarch][RunSingleSliceRead] intralRank [%u]"\
252 : "reduce with ringNextRank [%u] done", intraRank_, ringNextRank);
253 : }
254 :
255 : /* 2卡 场景,在最后一步的notifyDone */
256 0 : if (step == 0) {
257 0 : CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, neighbors));
258 : }
259 0 : notifyIdx_++;
260 :
261 0 : return HCCL_SUCCESS;
262 0 : }
263 :
264 0 : HcclResult ReduceScatterUnifiedMarch::RunHalfSliceRead(u32 ringPrevRank, u32 ringNextRank, u32 step, u32 totalStep)
265 : {
266 0 : LINK prevIntraLink = links_[ringPrevRank];
267 0 : CHK_SMART_PTR_NULL(prevIntraLink);
268 0 : LINK nextIntralLink = links_[ringNextRank];
269 0 : CHK_SMART_PTR_NULL(nextIntralLink);
270 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
271 :
272 : //拉齐 从流record主流、主流record从流 保证从流同时开始做SDMA
273 0 : CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
274 0 : CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
275 :
276 : // 从前向rank读取数据
277 0 : void* preRemDMAMemPtr = nullptr;
278 0 : CHK_RET(prevIntraLink->GetRemoteMem(UserMemType::INPUT_MEM, &preRemDMAMemPtr));
279 0 : u32 temIdx = (step == 0) ? totalStep : 0;
280 0 : u32 preDataIndex = (intraRank_ + intraRankSize_ - temIdx) % intraRankSize_;
281 : // 考虑总数据量不能被整除的情况
282 0 : u32 partOneCount = (step != totalStep) ? (totalCount_ / DIVISOR_NUM_TWO) : (totalCount_ - totalCount_ / DIVISOR_NUM_TWO);
283 0 : u64 partOneSize = partOneCount * SIZE_TABLE[dataType_];
284 0 : u64 preOffsetByte = (step != totalStep) ? (preDataIndex * blockDataByte_) :
285 0 : (preDataIndex * blockDataByte_ + totalCount_ / DIVISOR_NUM_TWO * SIZE_TABLE[dataType_]);
286 0 : void* preDstAddr = static_cast<u8 *>(userInput_.ptr()) + preOffsetByte;
287 :
288 0 : CHK_RET(DoSerialReduce(preRemDMAMemPtr, preDstAddr, partOneSize, partOneCount,
289 : subStreams_[0], prevIntraLink, preOffsetByte));
290 0 : HCCL_INFO("[ReduceScatterUnifiedMarch][RunHalfSliceRead] intralRank [%u] reduce with ringPrevRank [%u] done",
291 : intraRank_, ringPrevRank);
292 :
293 : // 从后向rank读取数据
294 0 : void* nextRemDMAMemPtr = nullptr;
295 0 : CHK_RET(nextIntralLink->GetRemoteMem(UserMemType::INPUT_MEM, &nextRemDMAMemPtr));
296 0 : temIdx = (step == 0) ? totalStep : 0;
297 0 : u32 nextDataIndex = (intraRank_ + temIdx) % intraRankSize_;
298 0 : u32 partTwoCount = totalCount_ - partOneCount;
299 0 : u64 partTwoSize = partTwoCount * SIZE_TABLE[dataType_];
300 0 : u64 nextOffsetByte = (step != totalStep) ? (nextDataIndex * blockDataByte_ + partOneSize) :
301 0 : nextDataIndex * blockDataByte_;
302 0 : void* nextDstAddr = static_cast<u8 *>(userInput_.ptr()) + nextOffsetByte;
303 :
304 0 : CHK_RET(DoSerialReduce(nextRemDMAMemPtr, nextDstAddr, partTwoSize, partTwoCount,
305 : subStreams_[1], nextIntralLink, nextOffsetByte));
306 0 : HCCL_INFO("[ReduceScatterUnifiedMarch][RunHalfSliceRead] intralRank [%u] reduce with ringNextRank [%u] done",
307 : intraRank_, ringNextRank);
308 :
309 : /* 4卡及以上的 场景,在最后一步的notifyDone */
310 : // 单算子使用Ack/Datasignal接口,必须保证两者交替使用
311 0 : if (step == totalStep) {
312 0 : const u32 NOTIFY_IDX_TWO = 2;
313 0 : if (notifyIdx_ % NOTIFY_IDX_TWO != 0 && GetWorkflowMode() == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
314 0 : notifyIdx_++;
315 0 : CHK_RET(WaitSubStreamFinish(NEIGHBORS_NUM_TWO));
316 0 : CHK_RET(NotifySubStreamStart(NEIGHBORS_NUM_TWO));
317 0 : CHK_RET(NotifyNeighborsStart(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
318 : }
319 0 : CHK_RET(NotifyNeighborsEnd(prevIntraLink, nextIntralLink, NEIGHBORS_NUM_TWO));
320 : }
321 0 : notifyIdx_++;
322 :
323 0 : return HCCL_SUCCESS;
324 0 : }
325 :
326 0 : HcclResult ReduceScatterUnifiedMarch::RunAsync()
327 : {
328 0 : HcclOpMetaInfoDef opMeta = HcclOpMetaInfo::GetOneForReduceScatter();
329 0 : CHK_RET(InitTask(dispatcher_, mainStream_, opMeta.isEnableCache, opMeta.GetCacheKey()));
330 :
331 : // 获取link的收、发
332 0 : u32 ringPrevRank = (intraRank_ + intraRankSize_ - 1) % intraRankSize_;
333 0 : u32 ringNextRank = (intraRank_ + 1) % intraRankSize_;
334 :
335 0 : u32 neighbors = (ringPrevRank == ringNextRank) ? NEIGHBORS_NUM_ONE : NEIGHBORS_NUM_TWO;
336 0 : CHK_RET(NotifySubStreamStart(neighbors));
337 :
338 : // 计算所需的总步骤
339 0 : u32 totalStep = intraRankSize_ / DIVISOR_NUM_TWO + 1;
340 0 : u32 step = 0;
341 0 : if(totalStep == DIVISOR_NUM_TWO) {
342 0 : CHK_RET(RunSingleSliceRead(ringPrevRank, ringNextRank, step, totalStep));
343 : } else {
344 : // 进行第1步收发
345 0 : CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, step, totalStep));
346 :
347 : // 进行第k步收发
348 0 : step++;
349 0 : for (; step < totalStep - DIVISOR_NUM_TWO; step++) {
350 0 : CHK_RET(RunSingleSliceRead(ringPrevRank, ringNextRank, step, totalStep));
351 : }
352 :
353 : // 进行第totalStep - 1步收发
354 0 : CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, step, totalStep));
355 :
356 : // 进行第totalStep步收发
357 0 : CHK_RET(RunHalfSliceRead(ringPrevRank, ringNextRank, totalStep, totalStep));
358 : }
359 :
360 0 : CHK_RET(WaitSubStreamFinish(neighbors));
361 0 : CHK_RET(LaunchTaskExtend(dispatcher_, mainStream_, subStreams_));
362 :
363 0 : HCCL_INFO("[ReduceScatterUnifiedMarch][RunAsync] finished.");
364 0 : return HCCL_SUCCESS;
365 : }
366 :
367 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_UNIFIED_MARCH, ReduceScatterUnifiedMarch);
368 : } // namespace hccl
|