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 "reducer.h"
12 :
13 : namespace hccl {
14 11 : Reducer::Reducer(const HcclDataType dataType, const HcclReduceOp reductionOp, const u64 reduceAttribute)
15 11 : : dataType_(dataType),
16 11 : reductionOp_(reductionOp),
17 11 : reduceAttribute_(reduceAttribute)
18 : {
19 11 : SetPreSyncFunc([]() {
20 0 : return HCCL_SUCCESS;
21 : });
22 11 : SetPostSyncFunc([]() {
23 0 : return HCCL_SUCCESS;
24 : });
25 11 : }
26 :
27 11 : Reducer::~Reducer() {}
28 :
29 11 : void Reducer::SetPreSyncFunc(std::function<HcclResult()> lambda) { preSync_ = std::move(lambda); }
30 :
31 11 : void Reducer::SetPostSyncFunc(std::function<HcclResult()> lambda) { postSync_ = std::move(lambda); }
32 :
33 9 : HcclResult Reducer::run(
34 : const HcclDispatcher dispatcher, const std::shared_ptr<Transport>& link, const u64 remoteMemOffset,
35 : DeviceMem& localSrc, DeviceMem& localDst, DeviceMem& remoteRcvTemp, Stream& stream, DstMemType resultMem,
36 : const UserMemType srcMemType) const
37 : {
38 9 : CHK_PTR_NULL(localSrc.ptr());
39 9 : CHK_PTR_NULL(localDst.ptr());
40 9 : CHK_PTR_NULL(remoteRcvTemp.ptr());
41 9 : CHK_PTR_NULL(stream.ptr());
42 :
43 9 : HcclResult ret = HCCL_SUCCESS;
44 :
45 9 : u64 dataBytes = remoteRcvTemp.size();
46 9 : HCCL_DEBUG(
47 : "localSrc[%p] localDst[%p] remoteRcvtmep[%p] offset[%llu]", localSrc.ptr(), localDst.ptr(), remoteRcvTemp.ptr(),
48 : remoteMemOffset);
49 :
50 : // server 内 reduce 并且 reduceAttribute_ 也支持,走该分支
51 9 : bool isSpInlineReduce = link->IsSpInlineReduce();
52 9 : if (link->IsSupportTransportWithReduce() && (RDMA_REDUCE_BITMASK & reduceAttribute_)) {
53 : // 数据接收端执行接收动作
54 : // RDMA的RxAsync不需要接收端内存信息
55 0 : CHK_RET(link->RxAsync(UserMemType::INPUT_MEM, remoteMemOffset, localDst.ptr(), localDst.size(), stream));
56 0 : if (link->GetSupportDataReceivedAck()) {
57 0 : CHK_RET(link->DataReceivedAck(stream));
58 : }
59 0 : if (resultMem == DstMemType::RESULT_OUTPUT_MEM) {
60 0 : ret = HcclD2DMemcpyAsync(dispatcher, localDst, localSrc, stream);
61 0 : CHK_PRT_RET(
62 : ret != HCCL_SUCCESS,
63 : HCCL_ERROR(
64 : "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", localSrc.ptr(), localDst.ptr()),
65 : ret);
66 : }
67 9 : } else if (link->IsSupportTransportWithReduce() && link->GetLinkType() == LinkType::LINK_STANDARD_ROCE) {
68 0 : u64 dataCount = localDst.size() / SIZE_TABLE[dataType_];
69 0 : DeviceMem& reduceSrc = (localSrc == localDst) ? remoteRcvTemp : localSrc;
70 0 : CHK_RET(link->RxWithReduce(
71 : srcMemType, remoteMemOffset, remoteRcvTemp.ptr(), dataBytes, reduceSrc.ptr(), localDst.ptr(), dataCount,
72 : dataType_, reductionOp_, stream, reduceAttribute_));
73 9 : } else if (isSpInlineReduce && (INLINE_REDUCE_BITMASK & reduceAttribute_)) {
74 : // runtime 的inline reduce 接口参数为数据的字节长度
75 0 : CHK_RET(link->RxDataSignal(stream));
76 0 : void* remoteMem = nullptr;
77 0 : CHK_RET(link->GetRemoteMem(srcMemType, &remoteMem));
78 0 : CHK_RET(HcclReduceAsync(
79 : dispatcher, static_cast<s8*>(remoteMem) + remoteMemOffset, dataBytes / SIZE_TABLE[dataType_], dataType_,
80 : reductionOp_, stream, localSrc.ptr(), link->GetRemoteRank(), link->GetLinkType(), INLINE_REDUCE_BIT));
81 :
82 0 : if (localSrc != localDst) {
83 0 : ret = HcclD2DMemcpyAsync(dispatcher, localDst, localSrc, stream);
84 0 : CHK_PRT_RET(
85 : ret != HCCL_SUCCESS,
86 : HCCL_ERROR(
87 : "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", localSrc.ptr(), localDst.ptr()),
88 : ret);
89 : }
90 0 : if (link->GetSupportDataReceivedAck()) {
91 0 : CHK_RET(link->TxAck(stream));
92 0 : CHK_RET(link->RxAck(stream));
93 0 : CHK_RET(link->TxDataSignal(stream));
94 0 : CHK_RET(link->RxDataSignal(stream));
95 : }
96 0 : } else {
97 : // 从上一个节点接收数据
98 9 : ret = link->RxAsync(UserMemType::INPUT_MEM, remoteMemOffset, remoteRcvTemp.ptr(), dataBytes, stream);
99 9 : CHK_PRT_RET(
100 : ret != HCCL_SUCCESS,
101 : HCCL_ERROR(
102 : "[Reducer][Run]rx_sync remoteRcvTemp[%p] offset[%llu] size[%llu] "
103 : "failed",
104 : remoteRcvTemp.ptr(), remoteMemOffset, dataBytes),
105 : ret);
106 :
107 9 : if (link->GetSupportDataReceivedAck()) {
108 0 : ret = link->DataReceivedAck(stream);
109 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Reducer][Run]rx_sync data received ack failed"), ret);
110 : }
111 :
112 9 : u64 dataCount = localDst.size() / SIZE_TABLE[dataType_];
113 :
114 : // 根据目的内存执行reduce
115 9 : DeviceMem reduceSrc = (localSrc == localDst) ? remoteRcvTemp : localSrc;
116 9 : ret = HcclReduceAsync(
117 9 : dispatcher, reduceSrc.ptr(), dataCount, dataType_, reductionOp_, stream, localDst.ptr(),
118 9 : INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttribute_);
119 :
120 9 : CHK_PRT_RET(
121 : ret != HCCL_SUCCESS,
122 : HCCL_ERROR(
123 : "[Reducer][Run]reduce_async remoteRcvTemp[%p] localSrc[%p] "
124 : "localDst[%p] failed",
125 : remoteRcvTemp.ptr(), localSrc.ptr(), localDst.ptr()),
126 : ret);
127 9 : }
128 :
129 9 : return ret;
130 : }
131 :
132 : HcclResult
133 0 : Reducer::PrepareRxMems(const std::vector<ReducerMemoryInfo>& reducerMems, std::vector<RxMemoryInfo>& rxMems) const
134 : {
135 0 : rxMems.reserve(reducerMems.size());
136 0 : for (const ReducerMemoryInfo& reduceMem : reducerMems) {
137 0 : rxMems.emplace_back(RxMemoryInfo{
138 0 : UserMemType::INPUT_MEM, reduceMem.remoteMemOffset, reduceMem.remoteRcvTemp.ptr(),
139 0 : reduceMem.remoteRcvTemp.size()});
140 : }
141 0 : return HCCL_SUCCESS;
142 : }
143 :
144 0 : HcclResult Reducer::PrepareRxWithReduceMems(
145 : const std::vector<ReducerMemoryInfo>& reducerMems, std::vector<RxWithReduceMemoryInfo>& rxWithReduceMems) const
146 : {
147 0 : rxWithReduceMems.reserve(reducerMems.size());
148 0 : for (const ReducerMemoryInfo& reduceMem : reducerMems) {
149 0 : u64 dataCount = reduceMem.localdst.size() / SIZE_TABLE[dataType_];
150 0 : DeviceMem reduceSrc = (reduceMem.localsrc == reduceMem.localdst) ? reduceMem.remoteRcvTemp : reduceMem.localsrc;
151 :
152 0 : rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{
153 0 : UserMemType::INPUT_MEM, reduceMem.remoteMemOffset, reduceMem.remoteRcvTemp.ptr(),
154 0 : reduceMem.remoteRcvTemp.size(), reduceSrc.ptr(), reduceMem.localdst.ptr(), dataCount});
155 0 : }
156 0 : return HCCL_SUCCESS;
157 : }
158 :
159 0 : HcclResult Reducer::run(
160 : const HcclDispatcher dispatcher, const std::shared_ptr<Transport>& link,
161 : const std::vector<ReducerMemoryInfo>& reducerMems, Stream& stream, DstMemType resultMem) const
162 : {
163 0 : CHK_PTR_NULL(stream.ptr());
164 :
165 0 : LinkType linkType = link->GetLinkType();
166 0 : bool isSpInlineReduce = link->IsSpInlineReduce();
167 0 : bool isSpRdmaReduce = RDMA_REDUCE_BITMASK & reduceAttribute_;
168 0 : bool isSpTransportWithReduce = link->IsSupportTransportWithReduce();
169 0 : HcclResult ret = HCCL_SUCCESS;
170 :
171 0 : if (isSpTransportWithReduce && isSpRdmaReduce) {
172 : // 数据接收端执行接收动作
173 : // RDMA的RxAsync不需要接收端内存信息
174 0 : std::vector<RxMemoryInfo> rxMems;
175 0 : CHK_RET(PrepareRxMems(reducerMems, rxMems));
176 0 : CHK_RET(link->RxAsync(rxMems, stream));
177 0 : if (link->GetSupportDataReceivedAck()) {
178 0 : CHK_RET(link->DataReceivedAck(stream));
179 : }
180 0 : CHK_RET(preSync_());
181 0 : if (resultMem == DstMemType::RESULT_OUTPUT_MEM) {
182 0 : for (ReducerMemoryInfo reduceMem : reducerMems) {
183 0 : ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
184 0 : CHK_PRT_RET(
185 : ret != HCCL_SUCCESS,
186 : HCCL_ERROR(
187 : "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
188 : reduceMem.localdst.ptr()),
189 : ret);
190 0 : }
191 : }
192 0 : CHK_RET(postSync_());
193 0 : } else if (isSpTransportWithReduce && (linkType == LinkType::LINK_STANDARD_ROCE)) {
194 0 : std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
195 0 : CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
196 0 : CHK_RET(preSync_());
197 0 : CHK_RET(link->RxWithReduce(rxWithReduceMems, dataType_, reductionOp_, stream, reduceAttribute_));
198 0 : CHK_RET(postSync_());
199 0 : } else if (isSpInlineReduce && (INLINE_REDUCE_BITMASK & reduceAttribute_)) {
200 0 : CHK_RET(link->RxDataSignal(stream));
201 0 : void* remoteMem = nullptr;
202 0 : CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
203 0 : CHK_RET(preSync_());
204 0 : for (ReducerMemoryInfo reduceMem : reducerMems) {
205 0 : const u64 dataBytes = reduceMem.remoteRcvTemp.size();
206 0 : CHK_RET(HcclReduceAsync(
207 : dispatcher, static_cast<s8*>(remoteMem) + reduceMem.remoteMemOffset, dataBytes / SIZE_TABLE[dataType_],
208 : dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(), link->GetRemoteRank(), link->GetLinkType(),
209 : INLINE_REDUCE_BIT));
210 :
211 0 : if (reduceMem.localsrc != reduceMem.localdst) {
212 0 : ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
213 0 : CHK_PRT_RET(
214 : ret != HCCL_SUCCESS,
215 : HCCL_ERROR(
216 : "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
217 : reduceMem.localdst.ptr()),
218 : ret);
219 : }
220 0 : HCCL_DEBUG("[Reducer][Run]memcpy_async localSrc is [%p]", reduceMem.localsrc.ptr());
221 0 : }
222 0 : CHK_RET(postSync_());
223 0 : } else {
224 0 : std::vector<RxMemoryInfo> rxMems;
225 0 : CHK_RET(PrepareRxMems(reducerMems, rxMems));
226 0 : std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
227 0 : CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
228 0 : CHK_RET(preSync_());
229 0 : CHK_RET(link->RxAsync(rxMems, stream));
230 0 : CHK_RET(postSync_());
231 0 : if (link->GetSupportDataReceivedAck()) {
232 0 : CHK_RET(link->DataReceivedAck(stream));
233 : }
234 0 : for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
235 0 : CHK_RET(HcclReduceAsync(
236 : dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_, reductionOp_, stream,
237 : rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttribute_));
238 : }
239 0 : }
240 :
241 0 : return HCCL_SUCCESS;
242 : }
243 :
244 0 : HcclResult Reducer::run(
245 : const HcclDispatcher dispatcher, const std::shared_ptr<Transport>& link,
246 : const std::vector<ReducerMemoryInfo>& reducerMems, u32 notifyIdx, Stream& stream, DstMemType resultMem) const
247 : {
248 : (void)resultMem;
249 0 : CHK_PTR_NULL(stream.ptr());
250 0 : CHK_SMART_PTR_NULL(link);
251 :
252 0 : bool isSpInlineReduce = link->IsSpInlineReduce();
253 0 : HcclResult ret = HCCL_SUCCESS;
254 :
255 0 : if (isSpInlineReduce && static_cast<bool>((INLINE_REDUCE_BITMASK & reduceAttribute_))) {
256 0 : CHK_RET(link->Wait(notifyIdx, stream));
257 0 : void* remoteMem = nullptr;
258 0 : CHK_RET(link->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
259 0 : CHK_RET(preSync_());
260 0 : for (ReducerMemoryInfo reduceMem : reducerMems) {
261 0 : const u64 dataBytes = reduceMem.remoteRcvTemp.size();
262 0 : CHK_RET(HcclReduceAsync(
263 : dispatcher, static_cast<s8*>(remoteMem) + reduceMem.remoteMemOffset, dataBytes / SIZE_TABLE[dataType_],
264 : dataType_, reductionOp_, stream, reduceMem.localsrc.ptr(), link->GetRemoteRank(), link->GetLinkType(),
265 : INLINE_REDUCE_BIT));
266 0 : HCCL_DEBUG("[Reducer][Run]memcpy_async localSrc[%p]", reduceMem.localsrc.ptr());
267 0 : if (reduceMem.localsrc != reduceMem.localdst) {
268 0 : ret = HcclD2DMemcpyAsync(dispatcher, reduceMem.localdst, reduceMem.localsrc, stream);
269 0 : CHK_PRT_RET(
270 : ret != HCCL_SUCCESS,
271 : HCCL_ERROR(
272 : "[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
273 : reduceMem.localdst.ptr()),
274 : ret);
275 : }
276 0 : }
277 0 : CHK_RET(postSync_());
278 0 : } else {
279 0 : std::vector<RxMemoryInfo> rxMems;
280 0 : CHK_RET(PrepareRxMems(reducerMems, rxMems));
281 :
282 0 : std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
283 0 : CHK_RET(PrepareRxWithReduceMems(reducerMems, rxWithReduceMems));
284 0 : CHK_RET(preSync_());
285 :
286 0 : CHK_RET(link->Wait(notifyIdx, stream));
287 0 : for (auto& mem : rxMems) {
288 0 : CHK_PTR_NULL(mem.dst);
289 0 : void* srcMemPtr = nullptr;
290 0 : CHK_RET(link->GetRemoteMem(mem.srcMemType, &srcMemPtr));
291 0 : DeviceMem srcDevMem(static_cast<s8*>(srcMemPtr) + mem.srcOffset, mem.len);
292 0 : DeviceMem dstDevMem(static_cast<s8*>(mem.dst), mem.len);
293 0 : CHK_RET(HcclD2DMemcpyAsync(
294 : dispatcher, dstDevMem, srcDevMem, stream, link->GetRemoteRank(), link->GetLinkType()));
295 0 : }
296 0 : CHK_RET(postSync_());
297 0 : if (link->GetSupportDataReceivedAck()) {
298 0 : CHK_RET(link->DataReceivedAck(stream));
299 : }
300 0 : for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
301 0 : CHK_RET(HcclReduceAsync(
302 : dispatcher, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_, reductionOp_, stream,
303 : rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP, reduceAttribute_));
304 : }
305 0 : }
306 0 : return HCCL_SUCCESS;
307 : }
308 : } // namespace hccl
|