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 <cmath>
12 : #include "alg_template_register.h"
13 : #include "all_reduce_local_reduce.h"
14 :
15 : namespace hccl {
16 2 : AllReduceLocalReduce::AllReduceLocalReduce(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
17 :
18 4 : AllReduceLocalReduce::~AllReduceLocalReduce() {}
19 :
20 2 : HcclResult AllReduceLocalReduce::Prepare(
21 : u64 reduceAttrBitMap, std::vector<Stream>& meshStreams, std::vector<std::shared_ptr<LocalNotify>>& meshSignal,
22 : std::vector<std::shared_ptr<LocalNotify>>& meshSignalAux, u32 interRank, u32 interRankSize, u32 userRank,
23 : HcomCollOpInfo* opInfo)
24 : {
25 2 : reduceAttr_ = reduceAttrBitMap;
26 2 : localRank_ = interRank;
27 2 : localRankSize_ = interRankSize;
28 2 : userRank_ = userRank;
29 2 : meshStreams_ = meshStreams;
30 2 : meshSignal_ = &meshSignal;
31 2 : meshSignalAux_ = &meshSignalAux;
32 2 : opInfo_ = opInfo;
33 2 : return HCCL_SUCCESS;
34 : }
35 0 : HcclResult AllReduceLocalReduce::MainRecordSub()
36 : {
37 0 : for (u32 signalIndex = 0; signalIndex < meshSignalAux_->size(); signalIndex++) {
38 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
39 : }
40 0 : return HCCL_SUCCESS;
41 : }
42 :
43 0 : HcclResult AllReduceLocalReduce::SubWaitMain()
44 : {
45 0 : for (u32 streamIndex = 0; streamIndex < meshSignalAux_->size(); streamIndex++) {
46 0 : CHK_RET(LocalNotify::Wait(
47 : meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex], profilerInput_.stage));
48 : }
49 0 : return HCCL_SUCCESS;
50 : }
51 :
52 0 : HcclResult AllReduceLocalReduce::MainWaitSub()
53 : {
54 0 : for (u32 signalIndex = 0; signalIndex < meshSignal_->size(); signalIndex++) {
55 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
56 : }
57 0 : return HCCL_SUCCESS;
58 : }
59 :
60 0 : HcclResult AllReduceLocalReduce::SubRecordMain()
61 : {
62 0 : for (u32 streamIndex = 0; streamIndex < meshSignal_->size(); streamIndex++) {
63 0 : CHK_RET(LocalNotify::Post(
64 : meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex], profilerInput_.stage));
65 : }
66 0 : return HCCL_SUCCESS;
67 : }
68 :
69 : // 将数据均分,最小单位是128
70 1 : HcclResult AllReduceLocalReduce::PrepareSlice(
71 : u64 dataCount, u32 unitSize, u32 sliceNum, std::vector<Slice>& dataSlice, std::vector<Slice>& startSlice)
72 : {
73 1 : Slice temp;
74 1 : Slice startTemp;
75 1 : u64 totalSize = dataCount * unitSize;
76 1 : dataSlice.clear();
77 1 : dataSlice.reserve(sliceNum);
78 1 : if (sliceNum == 0) {
79 0 : HCCL_ERROR("[Prepare][SliceData]data slice prepare, sliceNum is 0");
80 0 : return HCCL_E_PARA;
81 : }
82 1 : u64 sizePerSliceOri = (totalSize + sliceNum - 1) / sliceNum; /* 1是为了向上取整 */
83 1 : u64 sizeLimit = 0;
84 1 : if (outputMem_.ptr() == opInfo_->outputAddr) {
85 1 : sizeLimit = outputMem_.size();
86 : } else {
87 0 : sizeLimit = totalSize;
88 : }
89 :
90 1 : u64 sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, HCCL_MIN_SLICE_ALIGN_910B); // 512B对齐
91 1 : if (sizePerSlice * (localRankSize_ - 1) > sizeLimit) {
92 1 : sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, HCCL_MIN_SLICE_ALIGN_ONCHIP);
93 : }
94 1 : if (sizePerSlice * (localRankSize_ - 1) > sizeLimit) {
95 1 : sizePerSlice = RoundUpWithDivisor(sizePerSliceOri, unitSize);
96 : }
97 1 : u64 residueSize = totalSize;
98 1 : u32 i = 0;
99 7 : while (residueSize > 0) {
100 6 : u64 sliceSize = sizePerSlice < residueSize ? sizePerSlice : residueSize;
101 6 : temp.size = sliceSize;
102 6 : temp.offset = totalSize - residueSize;
103 6 : i++;
104 6 : if (sliceSize <= 0) {
105 0 : HCCL_ERROR("[Prepare][SliceData]data_slices_prepare sliceSize[%llu]", sliceSize);
106 0 : return HCCL_E_PARA;
107 : }
108 6 : residueSize -= sliceSize;
109 6 : dataSlice.push_back(temp);
110 6 : if (i != sliceNum) {
111 6 : startTemp.size = sizePerSlice;
112 6 : startTemp.offset = 0;
113 6 : startSlice.push_back(startTemp);
114 : } else {
115 0 : startTemp.size = sizePerSlice;
116 0 : startTemp.offset = sizePerSlice;
117 0 : startSlice.push_back(startTemp);
118 : }
119 : }
120 3 : while (i < sliceNum) {
121 2 : temp.size = 0;
122 2 : temp.offset = totalSize;
123 2 : i++;
124 2 : dataSlice.push_back(temp);
125 2 : startTemp.size = 0;
126 2 : startTemp.offset = 0;
127 2 : startSlice.push_back(startTemp);
128 : }
129 1 : return HCCL_SUCCESS;
130 : }
131 :
132 0 : HcclResult AllReduceLocalReduce::PrepareAllreduceSliceData()
133 : {
134 0 : return PrepareSlice(count_, DataUnitSize(dataType_), localRankSize_, slices_, startOffset);
135 : }
136 :
137 : // ringallreduce算法的函数入口
138 0 : HcclResult AllReduceLocalReduce::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
139 : {
140 0 : HcclResult ret = HCCL_SUCCESS;
141 0 : CHK_SMART_PTR_NULL(dispatcher_);
142 0 : CHK_PTR_NULL(stream_.ptr());
143 0 : HCCL_INFO(
144 : "AllReduceLocalReduce run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
145 : inputMem_.ptr(), outputMem_.ptr(), count_);
146 :
147 0 : if (links.size() < rankSize) {
148 0 : HCCL_ERROR(
149 : "[AllReduceLocalReduce][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank, links.size(),
150 : rankSize);
151 0 : return HCCL_E_INTERNAL;
152 : }
153 :
154 : // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
155 0 : if (rankSize == 1) {
156 0 : if (opInfo_->inputAddr != opInfo_->outputAddr) {
157 0 : DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * DataUnitSize(dataType_));
158 0 : DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * DataUnitSize(dataType_));
159 0 : ret = HcclD2DMemcpyAsync(dispatcher_, userMemOut, userMemIn, stream_);
160 0 : CHK_PRT_RET(
161 : ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] memcpy async failed", rank),
162 : ret);
163 0 : }
164 0 : return ret;
165 : }
166 :
167 0 : ret = PrepareAllreduceSliceData();
168 0 : CHK_PRT_RET(
169 : ret != HCCL_SUCCESS,
170 : HCCL_ERROR(
171 : "[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in PrepareSliceData step", rank, count_),
172 : ret);
173 :
174 0 : ret = RunReduceScatter(rank, rankSize, links);
175 0 : CHK_PRT_RET(
176 : ret != HCCL_SUCCESS,
177 : HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in ReduceScatter step", rank, count_),
178 : ret);
179 :
180 0 : ret = RunAllGather(rank, rankSize, links);
181 0 : CHK_PRT_RET(
182 : ret != HCCL_SUCCESS,
183 : HCCL_ERROR("[AllReduceLocalReduce][RunAsync]rank[%u] count[%llu] failed in AllGather step", rank, count_), ret);
184 :
185 0 : HCCL_INFO("AllReduceLocalReduce finished: rank[%u] ranksize[%u]", rank, rankSize);
186 0 : return HCCL_SUCCESS;
187 : }
188 :
189 0 : HcclResult AllReduceLocalReduce::RunReduceScatter(u32 rank, u32 rankSize, const std::vector<LINK>& links)
190 : {
191 0 : HCCL_INFO(
192 : "ReduceScatterMeshLocalReduce run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank,
193 : rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
194 :
195 : // 数据准备
196 0 : u32 unitSize = DataUnitSize(dataType_);
197 0 : DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * unitSize);
198 0 : DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
199 :
200 0 : DeviceMem src;
201 0 : DeviceMem dst;
202 :
203 0 : src = DeviceMem::create(static_cast<char*>(opInfo_->inputAddr) + slices_[rank].offset, slices_[rank].size);
204 0 : dst = commMemOut.range(slices_[rank].offset, slices_[rank].size);
205 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
206 :
207 0 : DeviceMem emptySrc = userMemIn.range(0, 0);
208 0 : DeviceMem emptyDst = commMemOut.range(0, 0);
209 :
210 0 : CHK_RET(MainRecordSub());
211 0 : CHK_RET(SubWaitMain());
212 :
213 0 : for (u32 round = 1; round < rankSize; round++) {
214 0 : u32 dstRank = (round + rank) % rankSize;
215 0 : Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
216 :
217 0 : CHK_RET(links[dstRank]->TxAck(subStream));
218 0 : CHK_RET(links[dstRank]->RxAck(subStream));
219 : }
220 0 : HCCL_DEBUG("[ReduceScatterMeshLocalReduce] D2DMemcpy start");
221 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
222 :
223 0 : CHK_RET(SubRecordMain());
224 0 : CHK_RET(MainWaitSub());
225 :
226 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
227 :
228 0 : CHK_RET(MainRecordSub());
229 0 : CHK_RET(SubWaitMain());
230 :
231 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
232 :
233 0 : for (u32 round = 1; round < rankSize; round++) {
234 0 : Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
235 0 : void* remMemPtr = nullptr;
236 0 : u32 dstRank = (rank + round) % rankSize;
237 0 : u32 dstSlice = (dstRank + round) % (rankSize - 1);
238 0 : if (dstRank == (rankSize - 1)) {
239 0 : dstSlice = (dstSlice + rankSize - 1 - 1) % (rankSize - 1);
240 : }
241 0 : if (round == (rankSize - 1)) {
242 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
243 :
244 0 : dstSlice = (dstRank == (rankSize - 1)) ? (dstRank - 1) : dstRank;
245 0 : dst = DeviceMem::create(
246 0 : static_cast<char*>(remMemPtr) + startOffset[dstRank].offset + dstSlice * startOffset[dstRank].size,
247 0 : slices_[dstRank].size);
248 0 : src = userMemIn.range(slices_[dstRank].offset, slices_[dstRank].size);
249 :
250 0 : HCCL_INFO(
251 : "AllReducelocalreduce reduce dst offset1 %llu offset2 %llu size %llu, rank %u, dstrank %u",
252 : startOffset[dstRank].offset, dstSlice * startOffset[dstRank].size, slices_[dstRank].size, rank,
253 : dstRank);
254 :
255 0 : HCCL_INFO(
256 : "AllReducelocalreduce reduce src offset %llu size %llu, rank %u, dstrank %u", slices_[dstRank].offset,
257 : slices_[dstRank].size, rank, dstRank);
258 :
259 0 : CHK_RET(HcclReduceAsync(
260 : dispatcher_, static_cast<void*>(src.ptr()), slices_[dstRank].size / unitSize, dataType_, reductionOp_,
261 : subStream, static_cast<void*>(dst.ptr()), links[dstRank]->GetRemoteRank(),
262 : links[dstRank]->GetLinkType(), INLINE_REDUCE_BIT));
263 : } else {
264 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
265 :
266 0 : dst = DeviceMem::create(
267 0 : static_cast<char*>(remMemPtr) + startOffset[dstRank].offset + dstSlice * startOffset[dstRank].size,
268 0 : slices_[dstRank].size);
269 0 : src = userMemIn.range(slices_[dstRank].offset, slices_[dstRank].size);
270 :
271 0 : HCCL_INFO(
272 : "AllReducelocalreduce memcpy dst offset1 %llu offset2 %llu size %llu, rank %u, dstrank %u",
273 : startOffset[dstRank].offset, dstSlice * startOffset[dstRank].size, slices_[dstRank].size, rank,
274 : dstRank);
275 :
276 0 : HCCL_INFO(
277 : "AllReducelocalreduce memcpy src offset %llu size %llu, rank %u, dstrank %u", slices_[dstRank].offset,
278 : slices_[dstRank].size, rank, dstRank);
279 :
280 0 : CHK_RET(HcclD2DMemcpyAsync(
281 : dispatcher_, dst, src, subStream, links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
282 : }
283 :
284 0 : CHK_RET(links[dstRank]->TxDataSignal(subStream));
285 0 : CHK_RET(links[dstRank]->RxDataSignal(subStream));
286 : }
287 :
288 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
289 :
290 0 : CHK_RET(SubRecordMain());
291 0 : CHK_RET(MainWaitSub());
292 :
293 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
294 :
295 0 : HcclResult ret = HCCL_SUCCESS;
296 0 : ret = RunLocalReduce(rank, rankSize);
297 :
298 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceLocalReduce]rank[%u] ReduceScatter failed", rank), ret);
299 0 : return HCCL_SUCCESS;
300 0 : }
301 :
302 0 : HcclResult AllReduceLocalReduce::RunLocalReduce(u32 rank, u32 rankSize)
303 : {
304 0 : DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
305 0 : u32 power = static_cast<u32>(log2(rankSize - 1));
306 0 : u32 rankPower = static_cast<u32>(pow(2, power));
307 0 : u32 unitSize = SIZE_TABLE[dataType_];
308 0 : u64 align = startOffset[rank].size;
309 0 : u64 totalSize = slices_[rank].size;
310 0 : DeviceMem src;
311 0 : DeviceMem dst;
312 0 : for (u32 i = 0u; i < rankSize - rankPower - 1; ++i) {
313 0 : u64 size = totalSize;
314 0 : if (rank < rankPower) {
315 0 : src = commMemOut.range(startOffset[rank].offset + (rankPower + i) * align, size);
316 0 : dst = commMemOut.range(startOffset[rank].offset + i * align, size);
317 0 : HCCL_INFO(
318 : "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
319 : startOffset[rank].offset + (rankPower + i) * align, startOffset[rank].offset + i * align, size);
320 : } else {
321 0 : dst = commMemOut.range(startOffset[rank].offset + (rankPower + i) * align, size);
322 0 : src = commMemOut.range(startOffset[rank].offset + i * align, size);
323 0 : HCCL_INFO(
324 : "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
325 : startOffset[rank].offset + i * align, startOffset[rank].offset + (rankPower + i) * align, size);
326 : }
327 0 : CHK_RET(HcclReduceAsync(
328 : dispatcher_, static_cast<void*>(src.ptr()), size / unitSize, dataType_, reductionOp_, stream_,
329 : static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
330 : }
331 0 : u32 center = rank < rankPower ? rank : (rank - rankPower + 1);
332 0 : center = std::min(center, rankPower - 1);
333 0 : u64 offset = rank < rankPower ? 0 : ((rankSize - rankPower - 1) * align);
334 0 : offset += startOffset[rank].offset;
335 0 : for (u32 round = 0; round < power; round++) {
336 0 : u32 slices_num = static_cast<u32>(rankPower / pow(2, round + 1));
337 0 : u64 size = totalSize;
338 0 : if (center < slices_num) {
339 0 : for (auto i = 0u; i < slices_num; ++i) {
340 0 : src = commMemOut.range(offset + (slices_num + i) * align, size);
341 0 : dst = commMemOut.range(offset + i * align, size);
342 0 : HCCL_INFO(
343 : "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank,
344 : offset + (slices_num + i) * align, offset + i * align, size);
345 0 : CHK_RET(HcclReduceAsync(
346 : dispatcher_, static_cast<void*>(src.ptr()), src.size() / unitSize, dataType_, reductionOp_, stream_,
347 : static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
348 : }
349 : } else {
350 0 : for (auto i = 0u; i < slices_num; ++i) {
351 0 : dst = commMemOut.range(offset + (slices_num + i) * align, size);
352 0 : src = commMemOut.range(offset + i * align, size);
353 0 : HCCL_INFO(
354 : "[RunLocalReduce]LocalReduce rank[%u] src[%llu], dst[%llu] size[%llu]", rank, offset + i * align,
355 : offset + (slices_num + i) * align, size);
356 0 : CHK_RET(HcclReduceAsync(
357 : dispatcher_, static_cast<void*>(src.ptr()), src.size() / unitSize, dataType_, reductionOp_, stream_,
358 : static_cast<void*>(dst.ptr()), INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, INLINE_REDUCE_BIT));
359 : }
360 0 : offset = offset + slices_num * align;
361 0 : center -= slices_num;
362 : }
363 : }
364 0 : return HCCL_SUCCESS;
365 0 : }
366 :
367 0 : HcclResult AllReduceLocalReduce::RunAllGather(u32 rank, u32 rankSize, const std::vector<LINK>& links)
368 : {
369 0 : HCCL_INFO(
370 : "AllGatherMesh run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
371 : inputMem_.ptr(), outputMem_.ptr(), count_);
372 0 : u32 unitSize = DataUnitSize(dataType_);
373 0 : DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * unitSize);
374 0 : DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
375 :
376 0 : DeviceMem src;
377 0 : DeviceMem dst;
378 :
379 0 : DeviceMem emptySrc = commMemOut.range(0, 0);
380 0 : DeviceMem emptyDst = userMemOut.range(0, 0);
381 :
382 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
383 :
384 0 : CHK_RET(MainRecordSub());
385 0 : CHK_RET(SubWaitMain());
386 :
387 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
388 :
389 0 : for (u32 round = 1; round < rankSize; round++) {
390 0 : u32 dstRank = BackwardRank(rank, rankSize, round);
391 0 : Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
392 0 : CHK_RET(links[dstRank]->TxAck(subStream));
393 0 : CHK_RET(links[dstRank]->RxAck(subStream));
394 : }
395 :
396 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
397 :
398 0 : CHK_RET(SubRecordMain());
399 0 : CHK_RET(MainWaitSub());
400 :
401 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
402 :
403 0 : CHK_RET(MainRecordSub());
404 0 : CHK_RET(SubWaitMain());
405 :
406 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
407 :
408 0 : if (userMemOut.ptr() != commMemOut.ptr()) {
409 0 : src = commMemOut.range(slices_[rank].offset, slices_[rank].size);
410 0 : dst = userMemOut.range(slices_[rank].offset, slices_[rank].size);
411 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, meshStreams_[meshStreams_.size() - 1]));
412 : }
413 :
414 0 : for (u32 round = 1; round < rankSize; round++) {
415 0 : u32 dstRank = BackwardRank(rank, rankSize, round);
416 0 : Stream& subStream = (round == rankSize - 1) ? stream_ : meshStreams_[round - 1];
417 0 : void* remMemPtr = nullptr;
418 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
419 :
420 0 : src = DeviceMem::create(static_cast<char*>(remMemPtr) + dstRank * slices_[0].size, slices_[dstRank].size);
421 0 : dst = userMemOut.range(slices_[dstRank].offset, slices_[dstRank].size);
422 :
423 0 : CHK_RET(HcclD2DMemcpyAsync(
424 : dispatcher_, dst, src, subStream, links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
425 0 : CHK_RET(links[dstRank]->TxDataSignal(subStream));
426 0 : CHK_RET(links[dstRank]->RxDataSignal(subStream));
427 : }
428 :
429 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
430 :
431 0 : CHK_RET(SubRecordMain());
432 0 : CHK_RET(MainWaitSub());
433 :
434 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
435 :
436 0 : HCCL_INFO("AllGatherMesh finished: rank[%u]", rank);
437 0 : return HCCL_SUCCESS;
438 0 : }
439 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_LOCAL_REDUCE, AllReduceLocalReduce);
440 : } // namespace hccl
|