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