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