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 "scatter_nhr.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : ScatterNHR::ScatterNHR(const HcclDispatcher dispatcher) : NHRBase(dispatcher), interRank_(0), interRankSize_(0) {}
16 :
17 0 : ScatterNHR::~ScatterNHR() {}
18 :
19 0 : HcclResult ScatterNHR::Prepare(bool needMerge)
20 : {
21 0 : isNeedMerge = needMerge;
22 0 : return HCCL_SUCCESS;
23 : }
24 :
25 : // scatter的入口函数
26 : HcclResult
27 0 : ScatterNHR::RunAsync(const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
28 : {
29 : // 从Broadcast调用Scatter需要merge
30 0 : if (isNeedMerge) {
31 : // 获取tree映射,存储到类对象的成员变量中
32 0 : GetRankMapping(rankSize);
33 : }
34 0 : CHK_SMART_PTR_NULL(dispatcher_);
35 0 : CHK_PTR_NULL(stream_.ptr());
36 0 : if (!outputMem_ || !inputMem_) {
37 0 : HCCL_ERROR("[ScatterNHR][RunAsync] run_async inputmem or outputmem is null");
38 0 : return HCCL_E_PTR;
39 : }
40 :
41 0 : interRank_ = rank;
42 0 : interRankSize_ = rankSize;
43 :
44 : // ranksize为1时,只有当input!=output 时候进行拷贝
45 0 : if (interRankSize_ == 1) {
46 0 : if (inputMem_ != outputMem_) {
47 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
48 : }
49 0 : return HCCL_SUCCESS;
50 : }
51 :
52 0 : u32 unitSize = DataUnitSize(dataType_);
53 0 : CHK_PRT_RET(
54 : unitSize == 0, HCCL_ERROR("[ScatterNHR][RunAsync] rank[%u] unit data size is zero", rank), HCCL_E_INTERNAL);
55 :
56 : // 带入vecotr为空,计算每个rank的结果偏移和大小
57 0 : if (slices_.size() == 0) {
58 0 : PrepareSlicesData(unitSize, count_, interRankSize_);
59 : }
60 :
61 0 : CHK_PRT_RET(
62 : links.size() < rankSize,
63 : HCCL_ERROR("[ScatterNHR][RunAsync] rank[%u] link size[%llu] is less than rank size", rank, links.size()),
64 : HCCL_E_INTERNAL);
65 :
66 0 : if (sliceMap_.size() != rankSize) {
67 0 : GetRankMapping(rankSize, true); // 没有初始化过,说明不是由allreduce或者bcast调入,需要保序
68 : }
69 :
70 0 : DeviceMem src;
71 :
72 0 : HcclResult ret = HCCL_SUCCESS;
73 : // 需要判断input不等于outputmem,scatter 输入只有一个input时不用拷贝
74 0 : if (inputMem_ != outputMem_) {
75 0 : u32 targetIdx = sliceMap_[interRank_];
76 :
77 0 : src = inputMem_.range(slices_[targetIdx].offset, slices_[targetIdx].size);
78 0 : ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, src, stream_);
79 0 : CHK_PRT_RET(
80 : ret != HCCL_SUCCESS,
81 : HCCL_ERROR(
82 : "[ScatterNHR][RunAsync] root rank[%u] memcpy async from input[%p] "
83 : "failed to output[%p]",
84 : interRank_, inputMem_.ptr(), outputMem_.ptr()),
85 : ret);
86 : }
87 :
88 : // 运行scatter, NHR 算法
89 0 : CHK_RET(RunScatterNHR(links));
90 0 : return HCCL_SUCCESS;
91 0 : }
92 :
93 0 : HcclResult ScatterNHR::SdmaRx(
94 : LINK& linkLeft, LINK& linkRight, InterServerAlgoStep& stepInfo, [[maybe_unused]] const std::vector<LINK>& links)
95 : {
96 0 : if (linkRight != nullptr) {
97 0 : CHK_RET(linkRight->TxAck(stream_));
98 : }
99 0 : if (linkLeft != nullptr) {
100 0 : CHK_RET(linkLeft->RxAck(stream_));
101 0 : std::vector<Slice> rxSlices;
102 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
103 0 : rxSlices.push_back(slices_[stepInfo.rxSliceIdxs[i]]);
104 : }
105 0 : MergeSlices(rxSlices);
106 0 : void* srcMemPtr = nullptr;
107 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
108 0 : for (const Slice& rxSlice : rxSlices) {
109 0 : DeviceMem dstMem = outputMem_.range(rxSlice.offset, rxSlice.size);
110 0 : DeviceMem srcMem(static_cast<s8*>(srcMemPtr) + baseOffset_ + rxSlice.offset, rxSlice.size);
111 0 : HCCL_DEBUG(
112 : "[ScatterNHR] rx dstMem[%p] range[%llu], size[%llu] ", dstMem.ptr(), rxSlice.offset, rxSlice.size);
113 0 : CHK_RET(HcclD2DMemcpyAsync(
114 : dispatcher_, dstMem, srcMem, stream_, linkLeft->GetRemoteRank(), // Memecpy
115 : linkLeft->GetLinkType()));
116 0 : }
117 0 : CHK_RET(linkLeft->TxDataSignal(stream_)); // 告知left读完了
118 0 : }
119 0 : if (linkRight != nullptr) {
120 0 : CHK_RET(linkRight->RxDataSignal(stream_)); // 等right读完
121 : }
122 0 : return HCCL_SUCCESS;
123 : }
124 :
125 0 : HcclResult ScatterNHR::RdmaTxRx(
126 : LINK& linkLeft, LINK& linkRight, InterServerAlgoStep& stepInfo, [[maybe_unused]] const std::vector<LINK>& links)
127 : {
128 0 : HcclResult ret = HCCL_SUCCESS;
129 :
130 0 : if (linkLeft != nullptr) {
131 0 : CHK_RET(linkLeft->TxAck(stream_));
132 : }
133 :
134 0 : if (linkRight != nullptr) {
135 0 : CHK_RET(linkRight->RxAck(stream_));
136 :
137 0 : std::vector<Slice> txSlices;
138 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
139 0 : txSlices.push_back(slices_[stepInfo.txSliceIdxs[i]]);
140 : }
141 0 : ret = Tx(linkRight, txSlices);
142 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterNHR][RunScatterNHR] Tx failed"), ret);
143 :
144 0 : CHK_RET(linkRight->TxWaitDone(stream_));
145 :
146 0 : ret = linkRight->WaitFinAck(stream_);
147 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterNHR][RunScatterNHR] WaitFinAck failed"), ret);
148 0 : }
149 :
150 0 : if (linkLeft != nullptr) {
151 0 : std::vector<Slice> rxSlices;
152 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
153 0 : rxSlices.push_back(slices_[stepInfo.rxSliceIdxs[i]]);
154 : }
155 0 : ret = Rx(linkLeft, rxSlices);
156 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterNHR][RunScatterNHR] Rx failed"), ret);
157 :
158 0 : CHK_RET(linkLeft->RxWaitDone(stream_));
159 :
160 0 : ret = linkLeft->PostFinAck(stream_);
161 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterNHR][RunScatterNHR] PostFinAck failed"), ret);
162 0 : }
163 0 : return HCCL_SUCCESS;
164 : }
165 :
166 0 : HcclResult ScatterNHR::RunScatterNHR(const std::vector<std::shared_ptr<Transport>>& links)
167 : {
168 : // 计算通信步数
169 0 : u32 nSteps = GetStepNumInterServer(interRankSize_);
170 :
171 : // 逐步编排任务
172 0 : for (u32 step = 0; step < nSteps; step++) {
173 0 : InterServerAlgoStep stepInfo;
174 0 : GetStepInfo(step, nSteps, interRank_, interRankSize_, stepInfo);
175 :
176 0 : HCCL_DEBUG(
177 : "[ScatterNHR][RunScatterNHR] rank[%u] recvFrom[%u] sendTo[%u] step[%u]", interRank_, stepInfo.fromRank,
178 : stepInfo.toRank, step);
179 :
180 0 : LINK linkLeft;
181 0 : LINK linkRight;
182 0 : if (stepInfo.txSliceIdxs.size() > 0) {
183 0 : linkRight = links[stepInfo.toRank];
184 0 : CHK_SMART_PTR_NULL(linkRight);
185 : }
186 0 : if (stepInfo.rxSliceIdxs.size() > 0) {
187 0 : linkLeft = links[stepInfo.fromRank];
188 0 : CHK_SMART_PTR_NULL(linkLeft);
189 : }
190 :
191 0 : if ((linkRight != nullptr && linkRight->IsSpInlineReduce())
192 0 : || (linkLeft != nullptr && linkLeft->IsSpInlineReduce())) {
193 0 : CHK_RET(SdmaRx(linkLeft, linkRight, stepInfo, links));
194 : } else {
195 0 : CHK_RET(RdmaTxRx(linkLeft, linkRight, stepInfo, links));
196 : }
197 0 : }
198 0 : return HCCL_SUCCESS;
199 : }
200 :
201 0 : void ScatterNHR::PrepareSlicesData(const u32 unitSize, const u64 totalCount, const u32 rankSize) const
202 : {
203 0 : slices_.resize(rankSize);
204 0 : u64 sliceSize = (totalCount / rankSize) * unitSize;
205 :
206 0 : for (u32 i = 0; i < rankSize; i++) {
207 0 : slices_[i].offset = i * sliceSize;
208 0 : slices_[i].size = sliceSize;
209 : }
210 0 : return;
211 : }
212 :
213 0 : HcclResult ScatterNHR::Tx(const LINK& link, std::vector<Slice>& txSlices)
214 : {
215 0 : std::vector<TxMemoryInfo> txMems;
216 :
217 0 : HCCL_DEBUG("[ScatterNHR][Tx] txSlices size [%u]", txSlices.size());
218 : // 合并连续slices
219 0 : MergeSlices(txSlices);
220 0 : HCCL_DEBUG("[ScatterNHR][Tx] merged txSlices size [%u]", txSlices.size());
221 :
222 0 : for (const Slice& txSlice : txSlices) {
223 0 : DeviceMem srcMem = outputMem_.range(txSlice.offset, txSlice.size);
224 0 : HCCL_DEBUG("[ScatterNHR][Tx] tx srcMem[%p] range[%llu] size[%llu]", srcMem.ptr(), txSlice.offset, txSlice.size);
225 0 : txMems.emplace_back(
226 0 : TxMemoryInfo{UserMemType::OUTPUT_MEM, txSlice.offset + baseOffset_, srcMem.ptr(), txSlice.size});
227 0 : }
228 :
229 0 : CHK_RET(link->TxAsync(txMems, stream_));
230 0 : return HCCL_SUCCESS;
231 0 : }
232 :
233 0 : HcclResult ScatterNHR::Rx(const LINK& link, std::vector<Slice>& rxSlices)
234 : {
235 0 : std::vector<RxMemoryInfo> rxMems;
236 :
237 0 : HCCL_DEBUG("[ScatterNHR][Rx] rxslices size [%u]", rxSlices.size());
238 : // 合并连续slices
239 0 : MergeSlices(rxSlices);
240 0 : HCCL_DEBUG("[ScatterNHR][Rx] merged rxslices size [%u]", rxSlices.size());
241 :
242 0 : for (const Slice& rxSlice : rxSlices) {
243 0 : DeviceMem dstMem = outputMem_.range(rxSlice.offset, rxSlice.size);
244 0 : HCCL_DEBUG("[ScatterNHR][Rx] rx dstMem[%p] range[%llu] size[%llu]", dstMem.ptr(), rxSlice.offset, rxSlice.size);
245 0 : rxMems.emplace_back(
246 0 : RxMemoryInfo{UserMemType::OUTPUT_MEM, rxSlice.offset + baseOffset_, dstMem.ptr(), rxSlice.size});
247 0 : }
248 :
249 0 : CHK_RET(link->RxAsync(rxMems, stream_));
250 0 : return HCCL_SUCCESS;
251 0 : }
252 :
253 : // NHR每步的算法描述原理函数
254 0 : HcclResult ScatterNHR::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep& stepInfo)
255 : {
256 0 : stepInfo.txSliceIdxs.clear();
257 0 : stepInfo.rxSliceIdxs.clear();
258 0 : stepInfo.nSlices = 0;
259 0 : stepInfo.toRank = rankSize;
260 0 : stepInfo.fromRank = rankSize;
261 0 : stepInfo.step = step;
262 0 : stepInfo.myRank = rank;
263 :
264 0 : u32 deltaRoot = (root_ + rankSize - rank) % rankSize;
265 0 : u32 deltaRankPair = 1 << step;
266 :
267 : // 数据份数和数据编号增量
268 0 : u32 nSlices = (rankSize - 1 + (1 << step)) / (1 << (step + 1));
269 0 : u32 deltaSliceIndex = 1 << (step + 1);
270 :
271 : // 判断是否是2的幂
272 0 : u32 nRanks = 0; // 本步需要进行收/发的rank数
273 0 : bool isPerfect = (rankSize & (rankSize - 1)) == 0;
274 0 : if (!isPerfect && step == nSteps - 1) {
275 0 : nRanks = rankSize - deltaRankPair;
276 : } else {
277 0 : nRanks = deltaRankPair;
278 : }
279 :
280 0 : if (deltaRoot < nRanks) { // 需要发
281 0 : u32 sendTo = (rank + rankSize - deltaRankPair) % rankSize;
282 0 : u32 txSliceIdx = sendTo;
283 0 : for (u32 i = 0; i < nSlices; i++) {
284 0 : u32 targetTxSliceIdx = sliceMap_[txSliceIdx];
285 0 : stepInfo.txSliceIdxs.push_back(targetTxSliceIdx);
286 0 : txSliceIdx = (txSliceIdx + rankSize - deltaSliceIndex) % rankSize;
287 : }
288 :
289 0 : stepInfo.toRank = sendTo;
290 0 : stepInfo.nSlices = nSlices;
291 0 : } else if (deltaRoot >= deltaRankPair && deltaRoot < nRanks + deltaRankPair) { // 需要收
292 0 : u32 recvFrom = (rank + deltaRankPair) % rankSize;
293 0 : u32 rxSliceIdx = rank;
294 0 : for (u32 i = 0; i < nSlices; i++) {
295 0 : u32 targetRxSliceIdx = sliceMap_[rxSliceIdx];
296 0 : stepInfo.rxSliceIdxs.push_back(targetRxSliceIdx);
297 0 : rxSliceIdx = (rxSliceIdx + rankSize - deltaSliceIndex) % rankSize;
298 : }
299 :
300 0 : stepInfo.fromRank = recvFrom;
301 0 : stepInfo.nSlices = nSlices;
302 : }
303 0 : return HCCL_SUCCESS;
304 : }
305 :
306 : HcclResult
307 0 : ScatterNHR::GetNslbAdjInfo(const u32 rank, const u32 rankSize, const std::vector<LINK>& links, AdjInfo& nslbAdjInfo)
308 : {
309 0 : if (rankSize == 1) {
310 0 : return HCCL_SUCCESS;
311 : }
312 0 : if (links.size() < rankSize) {
313 0 : return HCCL_SUCCESS;
314 : }
315 0 : u32 nSteps = 0;
316 0 : for (u32 temp = rankSize - 1; temp != 0; temp >>= 1, ++nSteps) {
317 : }
318 :
319 0 : u32 deltaRoot = (rankSize - rank) % rankSize;
320 0 : for (u32 step = 0; step < nSteps; step++) {
321 0 : u32 deltaRankPair = 1 << step;
322 0 : u32 nRanks = 0;
323 0 : bool isPerfect = (rankSize & (rankSize - 1)) == 0;
324 0 : if (!isPerfect && step == nSteps - 1) {
325 0 : nRanks = rankSize - deltaRankPair;
326 : } else {
327 0 : nRanks = deltaRankPair;
328 : }
329 0 : if (deltaRoot >= nRanks) {
330 0 : continue;
331 : }
332 :
333 0 : u32 sendTo = (rank + rankSize - deltaRankPair) % rankSize;
334 0 : LINK linkRight = links[sendTo];
335 0 : CHK_SMART_PTR_NULL(linkRight);
336 :
337 0 : NslbDpAdjInfo adjInfoStep = {};
338 0 : adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
339 0 : adjInfoStep.phaseId = step + 1;
340 0 : adjInfoStep.rev = 0;
341 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
342 0 : }
343 0 : nslbAdjInfo.dstRankNum = nslbAdjInfo.nsAdjInfo.size();
344 0 : return HCCL_SUCCESS;
345 : }
346 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_SCATTER_NHR, ScatterNHR);
347 : } // namespace hccl
|