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