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