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_hd_optim_pub.h"
14 : namespace hccl {
15 0 : AllReduceHDOptim::AllReduceHDOptim(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher)
16 : {
17 0 : }
18 :
19 0 : AllReduceHDOptim::~AllReduceHDOptim()
20 : {
21 0 : }
22 :
23 0 : HcclResult AllReduceHDOptim::Prepare(u64 reduceAttrBitMap, std::vector<Stream> &meshStreams,
24 : std::vector<std::shared_ptr<LocalNotify>> &meshSignal, std::vector<std::shared_ptr<LocalNotify>> &meshSignalAux,
25 : u32 userRank, HcomCollOpInfo *opInfo, bool aicpu)
26 : {
27 0 : reduceAttr_ = reduceAttrBitMap;
28 0 : userRank_ = userRank;
29 0 : meshStreams_ = meshStreams;
30 0 : meshSignal_ = &meshSignal;
31 0 : meshSignalAux_ = &meshSignalAux;
32 0 : opInfo_ = opInfo;
33 0 : aicpu_ = aicpu;
34 0 : return HCCL_SUCCESS;
35 : }
36 :
37 0 : HcclResult AllReduceHDOptim::MainRecordSub(u32 streamNum)
38 : {
39 0 : if(aicpu_) {
40 0 : return HCCL_SUCCESS;
41 : }
42 0 : for (u32 signalIndex = 0; signalIndex < streamNum; signalIndex++) {
43 0 : CHK_RET(LocalNotify::Post(stream_, dispatcher_, (*meshSignalAux_)[signalIndex], profilerInput_.stage));
44 : }
45 0 : return HCCL_SUCCESS;
46 : }
47 :
48 0 : HcclResult AllReduceHDOptim::SubWaitMain(u32 streamNum)
49 : {
50 0 : if(aicpu_) {
51 0 : return HCCL_SUCCESS;
52 : }
53 0 : for (u32 streamIndex = 0; streamIndex < streamNum; streamIndex++) {
54 0 : CHK_RET(LocalNotify::Wait(
55 : meshStreams_[streamIndex], dispatcher_, (*meshSignalAux_)[streamIndex], profilerInput_.stage));
56 : }
57 0 : return HCCL_SUCCESS;
58 : }
59 :
60 0 : HcclResult AllReduceHDOptim::MainWaitSub(u32 streamNum)
61 : {
62 0 : if(aicpu_) {
63 0 : return HCCL_SUCCESS;
64 : }
65 0 : for (u32 signalIndex = 0; signalIndex < streamNum; signalIndex++) {
66 0 : CHK_RET(LocalNotify::Wait(stream_, dispatcher_, (*meshSignal_)[signalIndex], profilerInput_.stage));
67 : }
68 0 : return HCCL_SUCCESS;
69 : }
70 :
71 0 : HcclResult AllReduceHDOptim::SubRecordMain(u32 streamNum)
72 : {
73 0 : if(aicpu_) {
74 0 : return HCCL_SUCCESS;
75 : }
76 0 : for (u32 streamIndex = 0; streamIndex < streamNum; streamIndex++) {
77 0 : CHK_RET(
78 : LocalNotify::Post(meshStreams_[streamIndex], dispatcher_, (*meshSignal_)[streamIndex], profilerInput_.stage));
79 : }
80 0 : return HCCL_SUCCESS;
81 : }
82 :
83 : // allreduce算法的函数入口
84 0 : HcclResult AllReduceHDOptim::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
85 : {
86 0 : HcclResult ret = HCCL_SUCCESS;
87 0 : CHK_SMART_PTR_NULL(dispatcher_);
88 0 : CHK_PTR_NULL(stream_.ptr());
89 0 : HCCL_INFO("AllReduceHDOptim run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
90 : rank, rankSize, outputMem_.ptr(), outputMem_.ptr(), count_);
91 :
92 0 : if (links.size() < rankSize) {
93 0 : HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]",
94 : rank, links.size(), rankSize);
95 0 : return HCCL_E_INTERNAL;
96 : }
97 :
98 0 : if (meshStreams_.size() < base) {
99 0 : HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] meshStreams_[%llu] is less than need[%u]",
100 : rank, meshStreams_.size(), base);
101 0 : return HCCL_E_INTERNAL;
102 : }
103 0 : u32 totalSize = SIZE_TABLE[dataType_] * count_;
104 0 : userMemIn = DeviceMem::create(opInfo_->inputAddr, totalSize);
105 0 : userMemOut = DeviceMem::create(opInfo_->outputAddr, totalSize);
106 0 : emptyMem_ = outputMem_.range(0, 0);
107 0 : nSteps = static_cast<u32>(log2(rankSize));
108 0 : stepPow = static_cast<u32>(pow(base, nSteps));
109 :
110 0 : ret = RunPreCopy(rank, rankSize, links);
111 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
112 : HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunPreCopy step" ,
113 : rank, count_), ret);
114 :
115 0 : if (rank < stepPow) {
116 0 : ret = RunAllReduceHDOptim(rank, rankSize, links);
117 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
118 : HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunAllReduceHDOptim step",
119 : rank, count_), ret);
120 : }
121 :
122 0 : if (stepPow != rankSize) {
123 0 : ret = RunFinalStep(rank, rankSize, links);
124 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
125 : HCCL_ERROR("[AllReduceHDOptim][RunAsync]rank[%u] count[%llu] failed RunFinalStep step",
126 : rank, count_), ret);
127 : }
128 :
129 0 : HCCL_INFO("AllReduceHDOptim finished: rank[%u] ranksize[%u]", rank, rankSize);
130 0 : return HCCL_SUCCESS;
131 : }
132 :
133 0 : HcclResult AllReduceHDOptim::RunPreCopy(u32 rank, u32 rankSize, const std::vector<LINK> &links)
134 : {
135 0 : u32 totalSize = SIZE_TABLE[dataType_] * count_;
136 :
137 0 : DeviceMem src = userMemIn.range(0, totalSize);
138 0 : DeviceMem dst = outputMem_.range(0, totalSize);
139 0 : DeviceMem nextDst = outputMem_.range(totalSize, totalSize);
140 :
141 0 : if (stepPow == rankSize) {
142 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, src, stream_));
143 0 : return HCCL_SUCCESS;
144 : }
145 :
146 0 : u32 neighCur = rank ^ (1 << nSteps);
147 0 : if (neighCur < rankSize) {
148 : // reduce写
149 0 : if (rank >= pow(base, nSteps)) {
150 0 : CHK_PTR_NULL(links[neighCur]);
151 0 : CHK_RET(links[neighCur]->RxAck(stream_));
152 0 : void *remMemPtr = nullptr;
153 0 : CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
154 0 : dst = DeviceMem::create(static_cast<u8 *>(remMemPtr), totalSize);
155 0 : CHK_RET(HcclReduceAsync(
156 : dispatcher_,
157 : static_cast<void *>(src.ptr()),
158 : count_,
159 : dataType_,
160 : reductionOp_,
161 : stream_,
162 : static_cast<void *>(dst.ptr()),
163 : links[neighCur]->GetRemoteRank(),
164 : links[neighCur]->GetLinkType(),
165 : INLINE_REDUCE_BIT));
166 0 : CHK_RET(links[neighCur]->TxDataSignal(stream_));
167 : } else {
168 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
169 0 : CHK_RET(links[neighCur]->TxAck(stream_));
170 0 : CHK_RET(links[neighCur]->RxDataSignal(stream_));
171 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, dst, stream_));
172 : }
173 : } else {
174 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
175 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, nextDst, dst, stream_));
176 : }
177 :
178 0 : return HCCL_SUCCESS;
179 0 : }
180 :
181 0 : HcclResult AllReduceHDOptim::RunBetweenStep(
182 : u32 rank, u32 step, u32 neighBefore, u32 neighNext, u32 rankSize, const std::vector<LINK> &links)
183 : {
184 : (void) rank;
185 0 : u32 totalSize = SIZE_TABLE[dataType_] * count_;
186 :
187 0 : DeviceMem src;
188 0 : DeviceMem dst;
189 :
190 0 : if ((step == 1) && (stepPow == rankSize)) {
191 : // 二次幂整第一步写 串行同步
192 0 : CHK_RET(links[neighBefore]->TxDataSignal(stream_));
193 0 : CHK_RET(links[neighBefore]->RxDataSignal(stream_));
194 0 : CHK_RET(MainRecordSub(base));
195 0 : CHK_RET(SubWaitMain(base));
196 0 : } else {
197 0 : CHK_RET(MainRecordSub(base));
198 0 : CHK_RET(SubWaitMain(base));
199 0 : CHK_RET(links[neighBefore]->TxDataSignal(stream_));
200 0 : CHK_RET(links[neighBefore]->RxDataSignal(stream_));
201 : }
202 :
203 0 : src = outputMem_.range(step * totalSize, totalSize);
204 0 : if ((step == (nSteps - 1)) && (static_cast<u32>(pow(base, nSteps)) == rankSize)) {
205 0 : dst = userMemOut.range(0, totalSize);
206 : } else {
207 0 : dst = outputMem_.range((step + 1) * totalSize, totalSize);
208 : }
209 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, aicpu_?stream_:meshStreams_[0]));
210 :
211 0 : CHK_RET(links[neighNext]->TxAck(aicpu_?stream_:meshStreams_[1]));
212 0 : CHK_RET(links[neighNext]->RxAck(aicpu_?stream_:meshStreams_[1]));
213 :
214 0 : CHK_RET(SubRecordMain(base));
215 0 : CHK_RET(MainWaitSub(base));
216 0 : return HCCL_SUCCESS;
217 0 : }
218 :
219 0 : HcclResult AllReduceHDOptim::RunAllReduceHDOptim(u32 rank, u32 rankSize, const std::vector<LINK> &links)
220 : {
221 0 : u32 unitSize = SIZE_TABLE[dataType_];
222 0 : u32 totalSize = unitSize * count_;
223 :
224 0 : DeviceMem src;
225 0 : DeviceMem dst;
226 :
227 0 : u32 neighCur = rank ^ (1 << 0);
228 0 : CHK_RET(links[neighCur]->TxAck(stream_));
229 0 : CHK_RET(links[neighCur]->RxAck(stream_));
230 :
231 0 : u32 neighNext = 0;
232 0 : void *remMemPtr = nullptr;
233 0 : for (u32 step = 1; step <= nSteps; step++) {
234 0 : if ((step != nSteps) || (stepPow != rankSize)) {
235 0 : dst = outputMem_.range(step * totalSize, totalSize);
236 : } else {
237 0 : dst = userMemOut.range(0, totalSize);
238 : }
239 0 : CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
240 0 : src = DeviceMem::create(static_cast<u8 *>(remMemPtr) + (step - 1) * totalSize, totalSize);
241 0 : if ((step == 1) && (stepPow == rankSize)) {
242 : // 二次幂整第一步写
243 0 : src = userMemIn.range(0, totalSize);
244 0 : dst = DeviceMem::create(static_cast<u8 *>(remMemPtr) + totalSize, totalSize);
245 : }
246 :
247 0 : CHK_RET(HcclReduceAsync(dispatcher_,
248 : static_cast<void *>(src.ptr()),
249 : count_,
250 : dataType_,
251 : reductionOp_,
252 : stream_,
253 : static_cast<void *>(dst.ptr()),
254 : links[neighCur]->GetRemoteRank(),
255 : links[neighCur]->GetLinkType(),
256 : INLINE_REDUCE_BIT));
257 :
258 0 : if (step != nSteps) {
259 0 : neighNext = rank ^ (1 << (step));
260 0 : CHK_RET(RunBetweenStep(rank, step, neighCur, neighNext, rankSize, links));
261 0 : neighCur = neighNext;
262 : }
263 : }
264 :
265 0 : CHK_RET(links[neighCur]->TxDataSignal(stream_));
266 0 : CHK_RET(links[neighCur]->RxDataSignal(stream_));
267 :
268 0 : return HCCL_SUCCESS;
269 0 : }
270 :
271 0 : HcclResult AllReduceHDOptim::RunFinalStep(u32 rank, u32 rankSize, const std::vector<LINK> &links)
272 : {
273 0 : u32 unitSize = SIZE_TABLE[dataType_];
274 0 : u32 totalSize = unitSize * count_;
275 :
276 0 : DeviceMem src = outputMem_.range(nSteps * totalSize, totalSize);
277 0 : DeviceMem dst = userMemOut.range(0, totalSize);
278 :
279 0 : u32 neighCur = rank ^ (1 << nSteps);
280 0 : if (neighCur >= rankSize) {
281 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
282 0 : } else if (rank < pow(base, nSteps)) {
283 0 : CHK_RET(links[neighCur]->TxAck(stream_));
284 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
285 0 : CHK_RET(links[neighCur]->RxDataSignal(stream_));
286 : } else {
287 0 : CHK_RET(links[neighCur]->RxAck(stream_));
288 0 : void *remMemPtr = nullptr;
289 0 : CHK_RET(links[neighCur]->GetRemoteMem(UserMemType::OUTPUT_MEM, &remMemPtr));
290 0 : src = DeviceMem::create(static_cast<u8 *>(remMemPtr) + nSteps * totalSize, totalSize);
291 0 : CHK_RET(HcclD2DMemcpyAsync(
292 : dispatcher_, dst, src, stream_, links[neighCur]->GetRemoteRank(), links[neighCur]->GetLinkType()));
293 0 : CHK_RET(links[neighCur]->TxDataSignal(stream_));
294 : }
295 0 : return HCCL_SUCCESS;
296 0 : }
297 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_REDUCE_HD_OPTIM, AllReduceHDOptim);
298 : } // namespace hccl
|