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 "alg_template_register.h"
12 : #include "all_reduce_chunk_mesh.h"
13 :
14 : namespace hccl {
15 2 : AllReduceChunkMesh::AllReduceChunkMesh(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
16 :
17 4 : AllReduceChunkMesh::~AllReduceChunkMesh() {}
18 :
19 2 : HcclResult AllReduceChunkMesh::Prepare(
20 : u64 reduceAttrBitMap, std::vector<Stream>& meshStreams, std::vector<std::shared_ptr<LocalNotify>>& meshSignal,
21 : std::vector<std::shared_ptr<LocalNotify>>& meshSignalAux, u32 interRank, u32 interRankSize, u32 userRank,
22 : HcomCollOpInfo* opInfo)
23 : {
24 2 : reduceAttr_ = reduceAttrBitMap;
25 2 : localRank_ = interRank;
26 2 : localRankSize_ = interRankSize;
27 2 : userRank_ = userRank;
28 2 : meshStreams_ = meshStreams;
29 2 : meshSignal_ = &meshSignal;
30 2 : meshSignalAux_ = &meshSignalAux;
31 2 : opInfo_ = opInfo;
32 2 : return HCCL_SUCCESS;
33 : }
34 0 : HcclResult AllReduceChunkMesh::MainRecordSub()
35 : {
36 0 : for (u32 signalIndex = 0; signalIndex < meshSignalAux_->size(); signalIndex++) {
37 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
38 : }
39 0 : return HCCL_SUCCESS;
40 : }
41 :
42 0 : HcclResult AllReduceChunkMesh::SubWaitMain()
43 : {
44 0 : for (u32 streamIndex = 0; streamIndex < meshSignalAux_->size(); streamIndex++) {
45 0 : CHK_RET(LocalNotify::Wait(
46 : meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex], profilerInput_.stage));
47 : }
48 0 : return HCCL_SUCCESS;
49 : }
50 :
51 0 : HcclResult AllReduceChunkMesh::MainWaitSub()
52 : {
53 0 : for (u32 signalIndex = 0; signalIndex < meshSignal_->size(); signalIndex++) {
54 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
55 : }
56 0 : return HCCL_SUCCESS;
57 : }
58 :
59 0 : HcclResult AllReduceChunkMesh::SubRecordMain()
60 : {
61 0 : for (u32 streamIndex = 0; streamIndex < meshSignal_->size(); streamIndex++) {
62 0 : CHK_RET(LocalNotify::Post(
63 : meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex], profilerInput_.stage));
64 : }
65 0 : return HCCL_SUCCESS;
66 : }
67 :
68 : // 将数据均分,最小单位是128
69 : HcclResult
70 0 : AllReduceChunkMesh::PrepareSlice(u64 dataCount, u32 unitSize, u32 sliceNum, std::vector<Slice>& dataSlice) const
71 : {
72 0 : u64 totalSize = dataCount * unitSize;
73 0 : Slice temp;
74 0 : dataSlice.clear();
75 0 : dataSlice.reserve(sliceNum);
76 0 : if (sliceNum == 0) {
77 0 : HCCL_ERROR("[Prepare][SliceData]data slice prepare, sliceNum is 0");
78 0 : return HCCL_E_PARA;
79 : }
80 0 : u64 sizePerSlice = (totalSize + sliceNum - 1) / sliceNum; /* 1是为了向上取整 */
81 0 : sizePerSlice = RoundUpWithDivisor(sizePerSlice, HCCL_MIN_SLICE_ALIGN);
82 0 : u64 residueSize = totalSize;
83 0 : u32 i = 0;
84 0 : while (residueSize > 0) {
85 0 : u64 sliceSize = sizePerSlice < residueSize ? sizePerSlice : residueSize;
86 0 : temp.size = sliceSize;
87 0 : temp.offset = totalSize - residueSize;
88 0 : i++;
89 0 : if (sliceSize <= 0) {
90 0 : HCCL_ERROR("[Prepare][SliceData]data_slice_prepare sliceSize[%llu]", sliceSize);
91 0 : return HCCL_E_PARA;
92 : }
93 0 : residueSize -= sliceSize;
94 0 : dataSlice.push_back(temp);
95 : }
96 0 : while (i < sliceNum) {
97 0 : temp.size = 0;
98 0 : temp.offset = totalSize;
99 0 : i++;
100 0 : dataSlice.push_back(temp);
101 : }
102 0 : return HCCL_SUCCESS;
103 : }
104 :
105 0 : HcclResult AllReduceChunkMesh::PrepareAllreduceSliceData()
106 : {
107 0 : u32 unitSize = SIZE_TABLE[dataType_];
108 0 : HcclResult ret = HCCL_SUCCESS;
109 0 : CHK_RET(PrepareSlice(count_, unitSize, localRankSize_, slices_));
110 0 : for (u32 rank = 0; rank < localRankSize_; rank++) {
111 0 : std::vector<Slice> dataSegsSlice;
112 0 : ret = PrepareSlice(slices_[rank].size / unitSize, unitSize, localRankSize_ - 1, dataSegsSlice);
113 0 : sliceMap[rank] = dataSegsSlice;
114 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceChunkMesh][PrepareSlice]rank[%u] failed", rank), ret);
115 0 : }
116 0 : return HCCL_SUCCESS;
117 : }
118 :
119 0 : HcclResult AllReduceChunkMesh::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
120 : {
121 0 : HcclResult ret = HCCL_SUCCESS;
122 0 : CHK_SMART_PTR_NULL(dispatcher_);
123 0 : CHK_PTR_NULL(stream_.ptr());
124 0 : HCCL_INFO(
125 : "AllReduceChunkMesh run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
126 : inputMem_.ptr(), outputMem_.ptr(), count_);
127 :
128 0 : if (links.size() < rankSize) {
129 0 : HCCL_ERROR(
130 : "[AllReduceChunkMesh][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank, links.size(),
131 : rankSize);
132 0 : return HCCL_E_INTERNAL;
133 : }
134 :
135 : // 如果ranksize为1, inline reduce和普通跨片reduce操作一致,从input->output
136 0 : if (rankSize == 1) {
137 0 : if (opInfo_->inputAddr != opInfo_->outputAddr) {
138 0 : DeviceMem userMemIn = DeviceMem::create(opInfo_->inputAddr, count_ * DataUnitSize(dataType_));
139 0 : DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * DataUnitSize(dataType_));
140 0 : ret = HcclD2DMemcpyAsync(dispatcher_, userMemOut, userMemIn, stream_);
141 0 : CHK_PRT_RET(
142 : ret != HCCL_SUCCESS, HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] memcpy async failed", rank), ret);
143 0 : }
144 0 : return ret;
145 : }
146 :
147 0 : ret = PrepareAllreduceSliceData();
148 0 : CHK_PRT_RET(
149 : ret != HCCL_SUCCESS,
150 : HCCL_ERROR("[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in PrepareSliceData step", rank, count_), ret);
151 :
152 0 : ret = RunReduceScatter(rank, rankSize, links);
153 0 : CHK_PRT_RET(
154 : ret != HCCL_SUCCESS,
155 : HCCL_ERROR(
156 : "[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in reducescater "
157 : "step",
158 : rank, count_),
159 : ret);
160 :
161 0 : ret = RunAllGather(rank, rankSize, links);
162 0 : CHK_PRT_RET(
163 : ret != HCCL_SUCCESS,
164 : HCCL_ERROR(
165 : "[AllReduceRing][RunAsync]rank[%u] count[%llu] failed in AllGather "
166 : "step",
167 : rank, count_),
168 : ret);
169 :
170 0 : HCCL_INFO("AllReduceChunkMesh finished: rank[%u] ranksize[%u]", rank, rankSize);
171 0 : return HCCL_SUCCESS;
172 : }
173 :
174 0 : HcclResult AllReduceChunkMesh::RunReduceScatter(u32 rank, u32 rankSize, const std::vector<LINK>& links)
175 : {
176 0 : HCCL_INFO(
177 : "ReduceScatterMeshAtomicOpbase run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank,
178 : rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
179 :
180 : // 数据准备
181 0 : u32 unitSize = DataUnitSize(dataType_);
182 :
183 0 : DeviceMem commMemOut = outputMem_;
184 :
185 0 : DeviceMem src;
186 0 : DeviceMem dst;
187 :
188 0 : src = DeviceMem::create(static_cast<char*>(opInfo_->inputAddr), count_ * unitSize);
189 :
190 0 : if (commMemOut.ptr() == opInfo_->outputAddr) {
191 : // 图模式
192 0 : src = src.range(slices_[rank].offset, slices_[rank].size);
193 0 : dst = commMemOut.range(slices_[rank].offset, slices_[rank].size);
194 : } else {
195 : // 单算子
196 0 : src = src.range(0, count_ * unitSize);
197 0 : dst = commMemOut.range(0, count_ * unitSize);
198 : }
199 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
200 :
201 0 : DeviceMem emptySrc = commMemOut.range(0, 0);
202 0 : DeviceMem emptyDst = commMemOut.range(0, 0);
203 :
204 : // 主从流之前加空拷贝 防止成环
205 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
206 :
207 0 : CHK_RET(MainRecordSub());
208 0 : CHK_RET(SubWaitMain());
209 :
210 0 : for (u32 round = 1; round < rankSize; round++) {
211 0 : u32 dstRank = (round + rank) % rankSize;
212 0 : Stream& subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
213 0 : CHK_RET(links[dstRank]->TxAck(subStream));
214 0 : CHK_RET(links[dstRank]->RxAck(subStream));
215 : }
216 :
217 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
218 :
219 0 : for (u32 round = 1; round < rankSize; round++) {
220 : // 主从流同步
221 :
222 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
223 :
224 0 : CHK_RET(SubRecordMain());
225 0 : CHK_RET(MainWaitSub());
226 :
227 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
228 :
229 0 : CHK_RET(MainRecordSub());
230 0 : CHK_RET(SubWaitMain());
231 :
232 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
233 :
234 : // 跨片reduceinline写
235 0 : for (u32 peer = 1; peer < rankSize; peer++) {
236 0 : u32 gap = (peer + round) > rankSize ? (peer + round - 1) % (rankSize - 1) : (peer + round - 1);
237 0 : u32 dstRank = (gap + rank) % rankSize;
238 0 : Stream& subStream = (peer == localRankSize_ - 1) ? stream_ : meshStreams_[peer - 1];
239 0 : u32 dstSlice = peer - 1;
240 0 : void* remMemPtr = nullptr;
241 :
242 0 : if (commMemOut.ptr() == opInfo_->outputAddr) {
243 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::INPUT_MEM, &remMemPtr));
244 : } else {
245 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
246 : }
247 :
248 0 : src = DeviceMem::create(
249 0 : static_cast<char*>(remMemPtr) + slices_[rank].offset + sliceMap[rank][dstSlice].offset,
250 0 : sliceMap[rank][dstSlice].size);
251 0 : dst = commMemOut.range(
252 0 : slices_[rank].offset + sliceMap[rank][dstSlice].offset, sliceMap[rank][dstSlice].size);
253 0 : CHK_RET(HcclReduceAsync(
254 : dispatcher_, static_cast<void*>(src.ptr()), sliceMap[rank][dstSlice].size / unitSize, dataType_,
255 : reductionOp_, subStream, static_cast<void*>(dst.ptr()), links[dstRank]->GetRemoteRank(),
256 : links[dstRank]->GetLinkType(), INLINE_REDUCE_BIT));
257 : }
258 : }
259 :
260 0 : for (u32 round = 1; round < rankSize; round++) {
261 0 : u32 gap = (round - 1) == 0 ? (rankSize - 1) : (round - 1);
262 0 : u32 dstRank = (rank + gap) % rankSize;
263 0 : Stream& subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
264 0 : CHK_RET(links[dstRank]->TxDataSignal(subStream));
265 0 : CHK_RET(links[dstRank]->RxDataSignal(subStream));
266 : }
267 :
268 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
269 :
270 0 : CHK_RET(SubRecordMain());
271 0 : CHK_RET(MainWaitSub());
272 :
273 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
274 0 : return HCCL_SUCCESS;
275 0 : }
276 :
277 0 : HcclResult AllReduceChunkMesh::RunAllGather(u32 rank, u32 rankSize, const std::vector<LINK>& links)
278 : {
279 0 : HCCL_INFO(
280 : "AllGatherMesh run: rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] count[%llu]", rank, rankSize,
281 : inputMem_.ptr(), outputMem_.ptr(), count_);
282 0 : u32 unitSize = DataUnitSize(dataType_);
283 :
284 0 : DeviceMem userMemOut = DeviceMem::create(opInfo_->outputAddr, count_ * unitSize);
285 0 : DeviceMem commMemOut = DeviceMem::create(outputMem_.ptr(), outputMem_.size());
286 :
287 0 : DeviceMem emptySrc = userMemOut.range(0, 0);
288 0 : DeviceMem emptyDst = commMemOut.range(0, 0);
289 :
290 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
291 :
292 0 : CHK_RET(MainRecordSub());
293 0 : CHK_RET(SubWaitMain());
294 :
295 0 : for (u32 round = 1; round < rankSize; round++) {
296 0 : u32 dstRank = BackwardRank(rank, rankSize, round);
297 0 : Stream& subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
298 0 : CHK_RET(links[dstRank]->TxAck(subStream));
299 0 : CHK_RET(links[dstRank]->RxAck(subStream));
300 : }
301 :
302 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
303 :
304 0 : CHK_RET(SubRecordMain());
305 0 : CHK_RET(MainWaitSub());
306 :
307 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
308 :
309 0 : CHK_RET(MainRecordSub());
310 0 : CHK_RET(SubWaitMain());
311 :
312 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
313 :
314 0 : DeviceMem src;
315 0 : DeviceMem dst;
316 0 : if (opInfo_->outputAddr != outputMem_.ptr()) {
317 0 : dst = userMemOut.range(slices_[rank].offset, slices_[rank].size);
318 0 : src = commMemOut.range(slices_[rank].offset, slices_[rank].size);
319 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, meshStreams_[meshStreams_.size() - 1]));
320 : }
321 :
322 0 : for (u32 round = 1; round < rankSize; round++) {
323 0 : u32 dstRank = BackwardRank(rank, rankSize, round);
324 0 : Stream& subStream = (round == localRankSize_ - 1) ? stream_ : meshStreams_[round - 1];
325 0 : void* remMemPtr = nullptr;
326 0 : CHK_RET(links[dstRank]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
327 0 : src = DeviceMem::create(static_cast<char*>(remMemPtr) + slices_[dstRank].offset, slices_[dstRank].size);
328 0 : dst = userMemOut.range(slices_[dstRank].offset, slices_[dstRank].size);
329 0 : CHK_RET(HcclD2DMemcpyAsync(
330 : dispatcher_, dst, src, subStream, links[dstRank]->GetRemoteRank(), links[dstRank]->GetLinkType()));
331 :
332 0 : CHK_RET(links[dstRank]->TxDataSignal(subStream));
333 0 : CHK_RET(links[dstRank]->RxDataSignal(subStream));
334 0 : HCCL_DEBUG("[AllReduceChunkMesh]round %u success");
335 : }
336 :
337 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
338 :
339 0 : CHK_RET(SubRecordMain());
340 0 : CHK_RET(MainWaitSub());
341 :
342 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, emptyDst, emptySrc, stream_));
343 :
344 0 : HCCL_INFO("[AllGatherMesh] finished: rank[%u]", rank);
345 0 : return HCCL_SUCCESS;
346 0 : }
347 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_CHUNK_MESH, AllReduceChunkMesh);
348 : } // namespace hccl
|