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 "reduce_scatter_nhr.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 0 : ReduceScatterNHR::ReduceScatterNHR(const HcclDispatcher dispatcher)
17 0 : :NHRBase(dispatcher)
18 : {
19 0 : }
20 :
21 0 : ReduceScatterNHR::~ReduceScatterNHR()
22 : {
23 0 : }
24 :
25 0 : HcclResult ReduceScatterNHR::Prepare(u64 reduceAttrBitMap, bool needMerge)
26 : {
27 0 : reduceAttr_ = reduceAttrBitMap;
28 0 : isNeedMerge = needMerge;
29 0 : return HCCL_SUCCESS;
30 : }
31 :
32 0 : HcclResult ReduceScatterNHR::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
33 : {
34 : // 基本的检查
35 0 : CHK_RET(SimpleCheck(rank, rankSize, links));
36 0 : HCCL_INFO("[ReduceScatterNHR][RunAsync] rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
37 : rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
38 :
39 0 : if (isNeedMerge == true) {
40 : // 获取tree映射,存储到类对象的成员变量中
41 0 : GetSliceMap(rankSize);
42 : }
43 :
44 : // 判断rank_size == 1
45 0 : if (rankSize == 1) {
46 0 : if (inputMem_ != outputMem_) {
47 0 : return HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
48 : }
49 0 : return HCCL_SUCCESS;
50 : }
51 :
52 0 : u32 unitSize = DataUnitSize(dataType_);
53 0 : CHK_PRT_RET(unitSize == 0, HCCL_ERROR("[ReduceScatterNHR][RunAsync] rank[%u] unit data size is zero", rank),
54 : HCCL_E_INTERNAL);
55 :
56 0 : std::vector<Slice> outputSlices(slices_);
57 :
58 : // 处理和检查Slices
59 0 : if (slices_.size() == 0) {
60 0 : slices_.resize(rankSize);
61 0 : outputSlices.resize(rankSize);
62 :
63 : // 生成std::vector<Slice> slices_
64 0 : u64 sliceSize = count_ * unitSize;
65 :
66 0 : for (u32 i = 0; i < rankSize; i++) {
67 0 : slices_[i].size = sliceSize;
68 0 : slices_[i].offset = (i * sliceSize);
69 :
70 0 : outputSlices[i].size = sliceSize;
71 0 : outputSlices[i].offset = (inputMem_.size() > outputMem_.size()) ? 0 : (i * sliceSize);
72 0 : HCCL_DEBUG("[ReduceScatterNHR][RunAsync] rank[%u], slices[%u].offset=[%llu] slices[%u].size=[%llu] "
73 : "outputSlices[%u].offset=[%llu], outputSlices[%u].size=[%llu] count_[%llu] unitSize[%llu]",
74 : rank, i, slices_[i].offset, i, slices_[i].size, i, outputSlices[i].offset, i, outputSlices[i].size,
75 : count_, unitSize);
76 : }
77 : }
78 :
79 0 : CHK_RET(CheckSlices(slices_, rankSize));
80 :
81 : // 创建reducer & sender
82 0 : senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
83 0 : CHK_SMART_PTR_NULL(senderInfo_);
84 :
85 0 : reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
86 0 : CHK_SMART_PTR_NULL(reducerInfo_);
87 :
88 0 : if (sliceMap_.size() != rankSize) {
89 0 : GetRankMapping(rankSize, true); // 没有初始化过,说明不是由allreduce或者bcast调入,需要保序
90 : }
91 :
92 : // 运行reduce-scatter, NHR 算法
93 0 : CHK_RET(RunReduceScatterNHR(rank, rankSize, links, slices_, outputSlices));
94 :
95 0 : HCCL_INFO("[ReduceScatterNHR][RunAsync] ReduceScatterNHR finished: rank[%u] end", rank);
96 0 : return HCCL_SUCCESS;
97 0 : }
98 :
99 0 : void ReduceScatterNHR::GetSliceMap(const u32 rankSize)
100 : {
101 0 : std::vector<u32> tree;
102 0 : for (u32 i = 0; i < rankSize; i++) {
103 0 : tree.push_back(i);
104 : }
105 :
106 : // 其他的再进行计算
107 0 : std::vector<u32> tmp(rankSize);
108 0 : u32 nSteps = 0;
109 0 : for (u32 tmp = rankSize - 1; tmp != 0; tmp >>= 1, nSteps++) {
110 : }
111 :
112 0 : u32 len = rankSize;
113 :
114 0 : for (u32 step = 0; step < nSteps; step++) {
115 0 : u32 nSlices = (rankSize - 1 + (1 << step)) / (1 << (step + 1));
116 0 : if (nSlices <= 1) {
117 0 : break;
118 : }
119 :
120 0 : bool endFlag = false;
121 :
122 0 : for (u32 part = 0; part * len < rankSize; part++) {
123 0 : u32 start = part * len;
124 0 : u32 end = std::min(start + len, rankSize);
125 0 : Reorder(start, end, len, tree, tmp);
126 :
127 0 : if (((end - start) & 1) == 1) {
128 0 : endFlag = true;
129 : }
130 : }
131 :
132 0 : for (u32 i = 0; i < rankSize; i++) {
133 0 : tree[i] = tmp[i];
134 : }
135 :
136 0 : if (endFlag) {
137 0 : break;
138 : }
139 :
140 0 : len >>= 1;
141 : }
142 :
143 : // 因为取的是tree中rank的idx,所以直接返回反向的映射
144 0 : sliceMap_.resize(rankSize);
145 0 : for (u32 i = 0; i < rankSize; i++) {
146 0 : sliceMap_[tree[i]] = i;
147 : }
148 :
149 0 : return;
150 0 : }
151 :
152 0 : void ReduceScatterNHR::Reorder(u32 start, u32 end, u32 len, std::vector<u32> &tree, std::vector<u32> &tmp)
153 : {
154 0 : const u32 idxTwo = 2;
155 :
156 0 : for (u32 i = start; i < end; i++) {
157 0 : u32 offset = i - start;
158 0 : if ((offset & 1) == 0) {
159 0 : tmp[start + offset / idxTwo] = tree[i];
160 : } else {
161 0 : tmp[start + (offset + len) / idxTwo] = tree[i];
162 : }
163 : }
164 0 : }
165 :
166 0 : HcclResult ReduceScatterNHR::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
167 : {
168 : // 判断stream, dispatcher是否为空
169 0 : CHK_SMART_PTR_NULL(dispatcher_);
170 0 : CHK_PTR_NULL(stream_.ptr());
171 :
172 : // 检查memory
173 0 : CHK_PRT_RET(!outputMem_ || !inputMem_,
174 : HCCL_ERROR("[ReduceScatterNHR][RunAsync] rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
175 :
176 : // 判断links数量是否正确
177 0 : CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[ReduceScatterNHR][RunAsync] rank[%u] link size[%llu] is "
178 : "less than rank size[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
179 0 : return HCCL_SUCCESS;
180 : }
181 :
182 0 : HcclResult ReduceScatterNHR::CheckSlices(const std::vector<Slice> &checkSlices, const u32 rankSize)
183 : {
184 0 : CHK_PRT_RET(checkSlices.size() % rankSize != 0,
185 : HCCL_ERROR("[ReduceScatterNHR][RunAsync] slices.size[%u] should be divided by rankSize[%u]",
186 : checkSlices.size(), rankSize), HCCL_E_INTERNAL);
187 0 : return HCCL_SUCCESS;
188 : }
189 :
190 0 : HcclResult ReduceScatterNHR::InlineReducer(const LINK &linkLeft, const std::vector<ReducerMemoryInfo> &rxReduceMems)
191 : {
192 0 : HcclResult ret = HCCL_SUCCESS;
193 0 : void *remoteMem = nullptr;
194 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &remoteMem));
195 0 : for (ReducerMemoryInfo reduceMem : rxReduceMems) {
196 0 : const u64 dataBytes = reduceMem.remoteRcvTemp.size();
197 0 : CHK_RET(
198 : HcclReduceAsync(dispatcher_, static_cast<s8 *>(remoteMem) + reduceMem.remoteMemOffset,
199 : dataBytes / SIZE_TABLE[dataType_], dataType_, reductionOp_, stream_, reduceMem.localsrc.ptr(),
200 : linkLeft->GetRemoteRank(), linkLeft->GetLinkType(), INLINE_REDUCE_BIT));
201 :
202 0 : if (reduceMem.localsrc != reduceMem.localdst) {
203 0 : ret = HcclD2DMemcpyAsync(dispatcher_, reduceMem.localdst, reduceMem.localsrc, stream_);
204 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
205 : HCCL_ERROR("[Reducer][Run]memcpy_async localSrc[%p] localDst[%p] failed", reduceMem.localsrc.ptr(),
206 : reduceMem.localdst.ptr()), ret);
207 : }
208 0 : }
209 0 : return HCCL_SUCCESS;
210 : }
211 :
212 0 : HcclResult ReduceScatterNHR::InlineReduceRx(const LINK &linkLeft, std::vector<Slice> &rxSlices,
213 : std::vector<Slice> &rxSlicestemp)
214 : {
215 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
216 0 : for (u64 i = 0; i < rxSlices.size(); i++) {
217 0 : DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
218 0 : DeviceMem srcMemTemp = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
219 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
220 : "offset[%llu] size[%llu] ",
221 : rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
222 0 : rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlices[i].offset, dstMem, dstMem, srcMemTemp});
223 0 : }
224 0 : CHK_RET(InlineReducer(linkLeft, rxReduceMems));
225 0 : return HCCL_SUCCESS;
226 0 : }
227 :
228 0 : HcclResult ReduceScatterNHR::InlineReduceRxLastStep(const LINK &linkLeft, InterServerAlgoStep &stepInfo,
229 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
230 : {
231 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
232 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
233 0 : u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
234 0 : DeviceMem dstMem = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
235 0 : DeviceMem srcMem = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size);
236 0 : DeviceMem tmpMem = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
237 0 : HCCL_DEBUG("[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
238 : "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]",
239 : rxSliceIdx, inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
240 :
241 0 : rxReduceMems.emplace_back(
242 0 : ReducerMemoryInfo { baseOffset_ + inputSlices[rxSliceIdx].offset, srcMem, dstMem, tmpMem });
243 0 : }
244 0 : CHK_RET(InlineReducer(linkLeft, rxReduceMems));
245 0 : return HCCL_SUCCESS;
246 0 : }
247 :
248 0 : HcclResult ReduceScatterNHR::TbeReduceRx(const LINK &linkLeft, std::vector<Slice> &rxSlices,
249 : std::vector<Slice> &rxSlicestemp)
250 : {
251 0 : void *srcMemPtr = nullptr;
252 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &srcMemPtr));
253 0 : std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
254 0 : for (u64 i = 0; i < rxSlices.size(); i++) {
255 0 : DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
256 0 : DeviceMem srcMem(static_cast<s8 *>(srcMemPtr) + baseOffset_ + rxSlices[i].offset, rxSlices[i].size);
257 0 : DeviceMem dstMemScratch = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
258 0 : u64 dataCount = dstMem.size() / SIZE_TABLE[dataType_];
259 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
260 : "offset[%llu] size[%llu] ",
261 : rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
262 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMemScratch, srcMem, stream_, linkLeft->GetRemoteRank(), // left的inputMem拷到本端的scratchMem
263 : linkLeft->GetLinkType()));
264 0 : rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{ UserMemType::INPUT_MEM, baseOffset_ + rxSlices[i].offset,
265 0 : dstMemScratch.ptr(), dstMemScratch.size(), dstMemScratch.ptr(), dstMem.ptr(), dataCount });
266 0 : }
267 0 : for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
268 0 : CHK_RET(HcclReduceAsync(dispatcher_, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_, // 本端scratchMem localReduce到 本端inputMem
269 : reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
270 : reduceAttr_));
271 : }
272 0 : return HCCL_SUCCESS;
273 0 : }
274 :
275 0 : HcclResult ReduceScatterNHR::TbeReduceRxLastStep(const LINK &linkLeft, InterServerAlgoStep &stepInfo,
276 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
277 : {
278 0 : void *srcMemPtr = nullptr;
279 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::INPUT_MEM, &srcMemPtr));
280 0 : std::vector<RxWithReduceMemoryInfo> rxWithReduceMems;
281 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
282 0 : u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
283 0 : DeviceMem srcMemRemote(static_cast<s8 *>(srcMemPtr) + baseOffset_ + inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size); // 对端inputMem
284 0 : DeviceMem dstMem = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size); // 本端outputMem
285 0 : DeviceMem srcMem = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size); // 本端inputMem
286 0 : DeviceMem tmpMem = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size); // 本端scratchMem
287 0 : u64 dataCount = dstMem.size() / SIZE_TABLE[dataType_];
288 0 : HCCL_DEBUG("[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
289 : "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]", rxSliceIdx,
290 : inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
291 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, tmpMem, srcMemRemote, stream_, linkLeft->GetRemoteRank(), // left的inputMem拷到本端的scratchMem
292 : linkLeft->GetLinkType()));
293 0 : DeviceMem reduceSrc = (srcMem == dstMem) ? tmpMem : srcMem;
294 0 : rxWithReduceMems.emplace_back(RxWithReduceMemoryInfo{ UserMemType::INPUT_MEM, baseOffset_ + inputSlices[rxSliceIdx].offset,
295 0 : tmpMem.ptr(), tmpMem.size(), reduceSrc.ptr(), dstMem.ptr(), dataCount });
296 0 : }
297 0 : for (RxWithReduceMemoryInfo rxReduceMem : rxWithReduceMems) {
298 0 : CHK_RET(HcclReduceAsync(dispatcher_, rxReduceMem.reduceSrc, rxReduceMem.reduceDataCount, dataType_, // 本端inputMem localReduce到 本端outputMem(之前拷到本端scratch的数据呢?)
299 : reductionOp_, stream_, rxReduceMem.reduceDst, INVALID_VALUE_RANKID, LinkType::LINK_ONCHIP,
300 : reduceAttr_));
301 : }
302 0 : return HCCL_SUCCESS;
303 0 : }
304 :
305 0 : HcclResult ReduceScatterNHR::RunDestReducerLastStep(const LINK &linkLeft, InterServerAlgoStep &stepInfo,
306 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
307 : {
308 0 : HcclResult ret = HCCL_SUCCESS;
309 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
310 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) { // rst算法的reduce scatter最后一步是一个slice,暂不用合并
311 0 : u32 rxSliceIdx = stepInfo.rxSliceIdxs[i];
312 0 : DeviceMem dstMem = outputMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
313 0 : DeviceMem srcMem = inputMem_.range(inputSlices[rxSliceIdx].offset, inputSlices[rxSliceIdx].size);
314 0 : DeviceMem tmpMem = scratchMem_.range(outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
315 0 : HCCL_DEBUG("[ReduceScatterNHR][RunReduceScatterNHR] final reduce rxSliceIdx[%u] will reduce with "
316 : "inputMem_ offset[%llu] to ouput_mem_ offset[%llu] size[%llu]", rxSliceIdx,
317 : inputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].offset, outputSlices[rxSliceIdx].size);
318 :
319 0 : rxReduceMems.emplace_back(
320 0 : ReducerMemoryInfo { baseOffset_ + inputSlices[rxSliceIdx].offset, srcMem, dstMem, tmpMem });
321 0 : }
322 :
323 0 : ret = reducerInfo_->run(dispatcher_, linkLeft, rxReduceMems, stream_);
324 0 : return ret;
325 0 : }
326 :
327 0 : HcclResult ReduceScatterNHR::GetRxSlices(std::vector<Slice> &rxSlices, std::vector<Slice> &rxSlicestemp,
328 : InterServerAlgoStep &stepInfo, const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
329 : {
330 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
331 0 : rxSlices.push_back(inputSlices[stepInfo.rxSliceIdxs[i]]);
332 0 : rxSlicestemp.push_back(outputSlices[stepInfo.rxSliceIdxs[i]]);
333 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] i[%u] rxSliceIndex[%u] rx offset[%llu] size[%llu]",
334 : i, stepInfo.rxSliceIdxs[i], outputSlices[stepInfo.rxSliceIdxs[i]].offset,
335 : outputSlices[stepInfo.rxSliceIdxs[i]].size);
336 : }
337 :
338 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] rxslices size [%u], rxslices temp size [%u]",
339 : rxSlices.size(), rxSlicestemp.size());
340 :
341 : // 合并连续slices
342 0 : MergeSlices(rxSlices);
343 0 : MergeSlices(rxSlicestemp);
344 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] merged rxslices size [%u], merged rxslices temp size [%u]",
345 : rxSlices.size(), rxSlicestemp.size());
346 0 : return HCCL_SUCCESS;
347 : }
348 :
349 0 : HcclResult ReduceScatterNHR::SdmaReducer(const u32 nSteps, const LINK &linkLeft, InterServerAlgoStep &stepInfo,
350 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
351 : {
352 0 : HcclResult ret = HCCL_SUCCESS;
353 0 : std::vector<Slice> rxSlices;
354 0 : std::vector<Slice> rxSlicestemp;
355 0 : if ((INLINE_REDUCE_BITMASK & reduceAttr_) == 1) { // InlineReduce
356 0 : if (stepInfo.step == (nSteps - 1)) {
357 0 : ret = InlineReduceRxLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
358 : } else {
359 0 : CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
360 0 : ret = InlineReduceRx(linkLeft, rxSlices, rxSlicestemp);
361 : }
362 : } else { // TbeReduce
363 0 : if (stepInfo.step == (nSteps - 1)) {
364 0 : ret = TbeReduceRxLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
365 : } else {
366 0 : CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
367 0 : ret = TbeReduceRx(linkLeft, rxSlices, rxSlicestemp);
368 : }
369 : }
370 0 : return ret;
371 0 : }
372 :
373 0 : HcclResult ReduceScatterNHR::RunReduceScatterNHR(const u32 rank, const u32 rankSize,
374 : const std::vector<LINK> &links,
375 : const std::vector<Slice> &inputSlices,
376 : const std::vector<Slice> &outputSlices)
377 : {
378 0 : bool bRetSize = (inputSlices.size() < rankSize);
379 0 : CHK_PRT_RET(bRetSize, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] inputslice size[%llu] is less "
380 : "than rank size[%u]", rank, outputSlices.size(), rankSize), HCCL_E_INTERNAL);
381 :
382 0 : bRetSize = (outputSlices.size() < rankSize);
383 0 : CHK_PRT_RET(bRetSize, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] outputslice size[%llu] is less "
384 : "than rank size[%u]", rank, outputSlices.size(), rankSize), HCCL_E_INTERNAL);
385 :
386 0 : HcclResult ret = HCCL_SUCCESS;
387 :
388 : // 计算通信步数
389 0 : u32 nSteps = GetStepNumInterServer(rankSize);
390 :
391 : // 逐步编排任务
392 0 : for (u32 step = 0; step < nSteps; step++) {
393 0 : InterServerAlgoStep stepInfo;
394 0 : GetStepInfo(step, nSteps, rank, rankSize, stepInfo);
395 :
396 : // 链的关系没有变化,区别的是发送的slice编号,因为重排tree不影响每棵树节点间的连接关系
397 0 : LINK linkLeft = links[stepInfo.fromRank];
398 0 : CHK_SMART_PTR_NULL(linkLeft);
399 :
400 0 : LINK linkRight = links[stepInfo.toRank];
401 0 : CHK_SMART_PTR_NULL(linkRight);
402 :
403 : // 当前每个数据块发送一次ACK、reduce一次、同步一次
404 0 : HCCL_DEBUG("[ReduceScatterNHR][RunReduceScatterNHR] rank[%u] rankSize[%u] from[%u] to[%u] step[%u] nSteps[%u] "
405 : "nSlices[%u]", rank, rankSize, stepInfo.fromRank, stepInfo.toRank, step, nSteps, stepInfo.nSlices);
406 :
407 0 : if (linkLeft->IsSpInlineReduce() && linkRight->IsSpInlineReduce()) { // SDMA
408 0 : CHK_RET(linkRight->TxAck(stream_));
409 0 : CHK_RET(linkLeft->RxAck(stream_));
410 0 : CHK_RET(SdmaReducer(nSteps, linkLeft, stepInfo, inputSlices, outputSlices));
411 0 : CHK_RET(linkLeft->TxDataSignal(stream_)); // 告知left我读完了
412 0 : CHK_RET(linkRight->RxDataSignal(stream_)); // 等right读完
413 : } else { // RDMA
414 0 : CHK_RET(linkLeft->TxAck(stream_));
415 0 : CHK_RET(linkRight->RxAck(stream_));
416 : // tx
417 0 : ret = RunSourceSender(linkRight, stepInfo, inputSlices, outputSlices);
418 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] Tx failed"), ret);
419 :
420 : // rx
421 0 : if (step == (nSteps - 1)) {
422 0 : ret = RunDestReducerLastStep(linkLeft, stepInfo, inputSlices, outputSlices);
423 : } else {
424 0 : ret = RunDestReducer(linkLeft, stepInfo, inputSlices, outputSlices);
425 : }
426 :
427 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] Rx failed"), ret);
428 0 : ret = linkLeft->PostFinAck(stream_);
429 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] PostFinAck failed"), ret);
430 :
431 0 : ret = linkRight->WaitFinAck(stream_);
432 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHR][RunReduceScatterNHR] WaitFinAck failed"), ret);
433 :
434 0 : if (barrierSwitchOn_) {
435 0 : CHK_RET(ExecuteBarrier(linkLeft, linkRight));
436 : }
437 : }
438 0 : }
439 0 : return HCCL_SUCCESS;
440 : }
441 :
442 0 : HcclResult ReduceScatterNHR::RunSourceSender(const LINK &link, InterServerAlgoStep &stepInfo,
443 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
444 : {
445 0 : std::vector<Slice> txSlices;
446 0 : std::vector<Slice> txSlicestemp;
447 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
448 0 : txSlices.push_back(inputSlices[stepInfo.txSliceIdxs[i]]);
449 0 : txSlicestemp.push_back(outputSlices[stepInfo.txSliceIdxs[i]]);
450 0 : HCCL_DEBUG("[ReduceScatterNHR][RunSourceSender] i[%u] txSliceIndex[%u] tx data offset[%llu] size[%llu]",
451 : i, stepInfo.txSliceIdxs[i], outputSlices[stepInfo.txSliceIdxs[i]].offset,
452 : outputSlices[stepInfo.txSliceIdxs[i]].size);
453 : }
454 0 : HCCL_DEBUG("[ReduceScatterNHR][RunSourceSender] txSlices size [%u], txSlices temp size [%u]",
455 : txSlices.size(), txSlicestemp.size());
456 :
457 : // 合并连续slices
458 0 : MergeSlices(txSlices);
459 0 : MergeSlices(txSlicestemp);
460 0 : HCCL_DEBUG("[ReduceScatterNHR][RunSourceSender] merged txSlices size [%u], merged txSlices temp size [%u]",
461 : txSlices.size(), txSlicestemp.size());
462 :
463 0 : std::vector<SenderMemoryInfo> txMems;
464 0 : for (u64 i = 0; i < txSlices.size(); i++) {
465 0 : DeviceMem srcMem = inputMem_.range(txSlices[i].offset, txSlices[i].size);
466 0 : HCCL_DEBUG("[ReduceScatterNHR][RunSourceSender] send inputmem range[%llu], size[%llu] tx dstmem offset[%llu]",
467 : txSlices[i].offset, txSlices[i].size, txSlicestemp[i].offset);
468 0 : txMems.emplace_back(SenderMemoryInfo{baseOffset_ + txSlicestemp[i].offset, srcMem});
469 0 : }
470 :
471 0 : CHK_RET(senderInfo_->run(link, txMems, stream_));
472 0 : return HCCL_SUCCESS;
473 0 : }
474 :
475 0 : HcclResult ReduceScatterNHR::RunDestReducer(const LINK &link, InterServerAlgoStep &stepInfo,
476 : const std::vector<Slice> &inputSlices, const std::vector<Slice> &outputSlices)
477 : {
478 0 : std::vector<Slice> rxSlices;
479 0 : std::vector<Slice> rxSlicestemp;
480 0 : CHK_RET(GetRxSlices(rxSlices, rxSlicestemp, stepInfo, inputSlices, outputSlices));
481 :
482 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
483 0 : for (u64 i = 0; i < rxSlices.size(); i++) {
484 0 : DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
485 0 : DeviceMem srcMemTemp = scratchMem_.range(rxSlicestemp[i].offset, rxSlicestemp[i].size);
486 0 : HCCL_DEBUG("[ReduceScatterNHR][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
487 : "offset[%llu] size[%llu] ",
488 : rxSlicestemp[i].offset, rxSlicestemp[i].size, rxSlices[i].offset, rxSlices[i].size);
489 0 : rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlices[i].offset, dstMem, dstMem, srcMemTemp});
490 0 : }
491 :
492 0 : CHK_RET(reducerInfo_->run(dispatcher_, link, rxReduceMems, stream_));
493 0 : return HCCL_SUCCESS;
494 0 : }
495 :
496 : // NHR每步的算法描述原理函数
497 0 : HcclResult ReduceScatterNHR::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep &stepInfo)
498 : {
499 : (void)nSteps;
500 0 : stepInfo.txSliceIdxs.clear();
501 0 : stepInfo.rxSliceIdxs.clear();
502 0 : u32 sliceSize = slices_.size() / rankSize;
503 0 : stepInfo.step = step;
504 0 : stepInfo.myRank = rank;
505 :
506 : // 计算通信对象
507 0 : u32 deltaRank = 1 << step;
508 0 : u32 sendTo = (rank + rankSize - deltaRank) % rankSize;
509 0 : u32 recvFrom = (rank + deltaRank) % rankSize;
510 :
511 : // 数据份数和数据编号增量
512 0 : u32 nSlices = (rankSize - 1 + (1 << step)) / (1 << (step + 1));
513 0 : u32 deltaSliceIndex = 1 << (step + 1);
514 0 : u32 txSliceIdx = sendTo; // 第一片rank
515 0 : u32 rxSliceIdx = rank;
516 :
517 0 : for (u32 i = 0; i < nSlices; i++) {
518 0 : for (u32 j = 0; j < sliceSize; j++) {
519 0 : u32 targetTxSliceIdx = sliceMap_[txSliceIdx];
520 0 : stepInfo.txSliceIdxs.push_back(targetTxSliceIdx * sliceSize + j);
521 :
522 0 : u32 targetRxSliceIdx = sliceMap_[rxSliceIdx];
523 0 : stepInfo.rxSliceIdxs.push_back(targetRxSliceIdx * sliceSize + j);
524 :
525 0 : HCCL_DEBUG("[ReduceScatterNHR][GetStepInfo] i[%u] txSliceIdx[%u]->targetTxSliceIdx[%u] rxSliceIdx[%u]->"
526 : "targetRxSliceIdx[%u]", i, txSliceIdx, targetTxSliceIdx, rxSliceIdx, targetRxSliceIdx);
527 : }
528 0 : txSliceIdx = (txSliceIdx + rankSize - deltaSliceIndex) % rankSize;
529 0 : rxSliceIdx = (rxSliceIdx + rankSize - deltaSliceIndex) % rankSize;
530 : }
531 :
532 0 : stepInfo.nSlices = nSlices * sliceSize;
533 0 : stepInfo.toRank = sendTo;
534 0 : stepInfo.fromRank = recvFrom;
535 0 : return HCCL_SUCCESS;
536 : }
537 :
538 0 : HcclResult ReduceScatterNHR::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
539 : const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
540 : {
541 0 : if (rankSize == 1) {
542 0 : return HCCL_SUCCESS;
543 : }
544 0 : if (links.size() < rankSize) {
545 0 : return HCCL_SUCCESS;
546 : }
547 0 : u32 nSteps = 0;
548 0 : for(u32 temp = rankSize - 1; temp != 0; temp >>= 1, ++nSteps){}
549 :
550 0 : for (u32 step = 0; step < nSteps; step++) {
551 0 : u32 deltaRank = 1 << step;
552 0 : u32 sendTo = (rank + rankSize - deltaRank) % rankSize;;
553 0 : LINK linkRight = links[sendTo];
554 0 : CHK_SMART_PTR_NULL(linkRight);
555 :
556 0 : NslbDpAdjInfo adjInfoStep = {0};
557 0 : adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
558 0 : adjInfoStep.phaseId = step + 1;
559 0 : adjInfoStep.rev = 0;
560 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
561 0 : }
562 0 : nslbAdjInfo.dstRankNum = nSteps;
563 0 : return HCCL_SUCCESS;
564 : }
565 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_NHR, ReduceScatterNHR);
566 : } // ~~ namespace hccl
|