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_v1.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 :
16 0 : ReduceScatterNHRV1::ReduceScatterNHRV1(const HcclDispatcher dispatcher)
17 0 : : NHRV1Base(dispatcher)
18 : {
19 0 : }
20 :
21 0 : ReduceScatterNHRV1::~ReduceScatterNHRV1()
22 : {
23 0 : }
24 :
25 0 : HcclResult ReduceScatterNHRV1::Prepare(u64 reduceAttrBitMap, HcomCollOpInfo *opInfo)
26 : {
27 : (void)opInfo;
28 0 : reduceAttr_ = reduceAttrBitMap;
29 0 : return HCCL_SUCCESS;
30 : }
31 :
32 0 : HcclResult ReduceScatterNHRV1::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("ReduceScatterNHRV1 run: rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
37 : rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
38 :
39 : // 判断rank_size == 1
40 0 : if (rankSize == 1) {
41 0 : if (inputMem_ != outputMem_) {
42 0 : return HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
43 : }
44 0 : return HCCL_SUCCESS;
45 : }
46 :
47 : // 处理和检查Slices
48 0 : if (slices_.size() == 0) {
49 0 : CHK_RET(SetDefaultSlices(rank, rankSize));
50 : }
51 0 : CHK_RET(CheckSlices(rankSize));
52 :
53 : // 获取通信关系
54 0 : RingInfo info = GetRingInfo(rankSize);
55 :
56 : // 垂直方向做Ring
57 0 : CHK_RET(RunReduceScatterOnVertical(rank, links, info));
58 :
59 : // 水平方向做Ring
60 0 : CHK_RET(RunReduceScatterOnHorizontal(rank, links, info));
61 :
62 : // 额外的搬运(从(x,sqrt-1)搬运到(x,sqrt))
63 : /* 一个可能的优化点:
64 : 以8节点为例: 0 1 2
65 : 3 4 5
66 : 6 7
67 : 当前的做法是:{0,1}、{3,4}、{6、7}做水平Ring,0/3/6/7拿到各自的那份结果,1拿到1和2的结果,4拿到4和5的结果,
68 : 最后1把2的结果发给2,4把5的结果发给5
69 : 一种可能的更优做法是:以ReduceOp=Sum为例,首先把2和5的数据都置为0,
70 : 然后直接{0,1,2}、{3,4,5}、{6,7}做水平Ring,避免不等分Ring和额外的拷贝步骤,但需要调用TBE-asign
71 : */
72 0 : CHK_RET(RunLastCopyStep(rank, links, info));
73 :
74 : // 搬运数据到OutputMem
75 0 : CHK_RET(RunCopyDataToOutputMem(rank));
76 :
77 0 : HCCL_INFO("ReduceScatterNHRV1 finished: rank[%u] end", rank);
78 0 : return HCCL_SUCCESS;
79 0 : }
80 :
81 0 : HcclResult ReduceScatterNHRV1::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
82 : {
83 : // 判断stream, dispatcher是否为空
84 0 : CHK_SMART_PTR_NULL(dispatcher_);
85 0 : CHK_PTR_NULL(stream_.ptr());
86 :
87 : // 检查memory
88 0 : CHK_PRT_RET(!outputMem_ || !inputMem_,
89 : HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
90 :
91 : // 判断links数量是否正确
92 0 : CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] link size[%llu] is less than "\
93 : "rank size[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
94 0 : return HCCL_SUCCESS;
95 : }
96 :
97 0 : HcclResult ReduceScatterNHRV1::SetDefaultSlices(const u32 rank, const u32 rankSize)
98 : {
99 0 : u32 unitSize = DataUnitSize(dataType_);
100 0 : CHK_PRT_RET(unitSize == 0,
101 : HCCL_ERROR("[ReduceScatterNHRV1]rank[%u] unit data size is zero", rank), HCCL_E_INTERNAL);
102 :
103 0 : slices_.resize(rankSize);
104 0 : u64 sliceSize = count_ * unitSize;
105 0 : for (u32 i = 0; i < rankSize; i++) {
106 0 : slices_[i].size = sliceSize;
107 0 : slices_[i].offset = (i * sliceSize);
108 0 : HCCL_DEBUG("rank[%u], slices[%u].offset=[%llu], slices[%u].size=[%llu] ",
109 : rank, i, slices_[i].offset, i, slices_[i].size);
110 : }
111 0 : return HCCL_SUCCESS;
112 : }
113 :
114 0 : HcclResult ReduceScatterNHRV1::CheckSlices(const u32 rankSize)
115 : {
116 0 : CHK_PRT_RET(slices_.size() != rankSize,
117 : HCCL_ERROR("[ReduceScatterNHRV1]slices.size[%u] should be equal to rankSize[%u]",
118 : slices_.size(), rankSize), HCCL_E_INTERNAL);
119 :
120 0 : for (u32 idx = 1; idx < slices_.size(); idx++) {
121 0 : CHK_PRT_RET(slices_[idx-1].offset + slices_[idx-1].size != slices_[idx].offset,
122 : HCCL_ERROR("[ReduceScatterNHRV1]only support continuous slices, but get "\
123 : "slices[%u].offset[%u], slices[%u].size[%u], slices[%u].offset[%u]",
124 : idx-1, slices_[idx-1].offset, idx-1, slices_[idx-1].size, idx, slices_[idx].offset), HCCL_E_INTERNAL);
125 : }
126 0 : return HCCL_SUCCESS;
127 : }
128 :
129 0 : HcclResult ReduceScatterNHRV1::RunReduceScatterBrokenRing(const u32 rank, const std::vector<LINK> &links,
130 : const std::vector<Slice> &slices)
131 : {
132 0 : std::unique_ptr<AlgTemplateBase> tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
133 0 : TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
134 0 : CHK_SMART_PTR_NULL(tempAlg);
135 0 : CHK_RET(tempAlg->Prepare(reduceAttr_));
136 :
137 0 : if (!barrierSwitchOn_) {
138 0 : tempAlg->CloseBarrier();
139 : }
140 :
141 0 : CHK_RET(tempAlg->Prepare(inputMem_, inputMem_, scratchMem_, count_, dataType_,
142 : stream_, reductionOp_, root_, slices));
143 :
144 0 : CHK_RET(tempAlg->RegisterProfiler(
145 : profilerInput_.planeID, profilerInput_.stage, profilerInput_.step, stream_));
146 :
147 0 : return tempAlg->RunAsync(rank, links.size(), links);
148 0 : }
149 :
150 0 : HcclResult ReduceScatterNHRV1::RunReduceScatterOnVertical(const u32 rank, const std::vector<LINK> &links,
151 : const RingInfo &info)
152 : {
153 0 : u32 hIndex = info.GetHIndex(rank); // 查找自己位于第几列
154 :
155 : // 构造新的links和slices
156 0 : std::vector<LINK> subLinks;
157 0 : std::vector<Slice> subSlices;
158 0 : u32 sliceIndexOffset = 0; // slice数量的累计偏移
159 0 : u32 hIndexForRing = (hIndex < info.GetRowSize()) ? hIndex : info.GetVIndex(rank); // 属于第几个垂直Ring
160 0 : u32 vSizeForRing = info.GetVSizeByHIndex(hIndexForRing); // 所属垂直Ring的大小
161 0 : for (u32 vIdx = 0; vIdx < vSizeForRing; vIdx++) { // 处理垂直方向上的Rank
162 : // 增加link
163 0 : u32 rankInRing = info.GetRank(vIdx, hIndexForRing);
164 0 : CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Vertical] rank[%u] out of range, "\
165 : "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
166 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] links[%u]=%u", rank, vIdx, rankInRing);
167 0 : subLinks.push_back(links[rankInRing]);
168 :
169 : // 寻找要合并的slice
170 0 : u32 nSlices = info.GetHSizeByVIndex(vIdx);
171 0 : Slice& headSlice = slices_[sliceIndexOffset];
172 0 : Slice& tailSlice = slices_[sliceIndexOffset + nSlices - 1];
173 :
174 : // 增加slice
175 0 : Slice slice;
176 0 : slice.offset = headSlice.offset;
177 0 : slice.size = tailSlice.offset + tailSlice.size - headSlice.offset;
178 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
179 : rank, vIdx, slice.offset, vIdx, slice.size);
180 0 : subSlices.push_back(slice);
181 :
182 : // 更新偏移
183 0 : sliceIndexOffset += nSlices;
184 : }
185 :
186 : // -- 可能还涉及跳跃的一个链接,比如8节点
187 : // ---- 0 1 2
188 : // ---- 3 4 5
189 : // ---- 6 7
190 : // -- 两个垂直Ring分别是{0,3,6,2}和{1,4,7,5},而不是{0,3,6}和{1,4,7}
191 0 : if (info.GetHSizeByVIndex(hIndexForRing) > info.GetRowSize()) {
192 : // 添加link
193 0 : u32 rankInRing = info.GetRank(hIndexForRing, info.GetRowSize());
194 0 : CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Vertical] rank[%u] out of range, "\
195 : "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
196 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] links[%u]=%u", rank, subLinks.size(), rankInRing);
197 0 : subLinks.push_back(links[rankInRing]);
198 :
199 : // 添加slice
200 0 : Slice slice;
201 0 : slice.offset = 0;
202 0 : slice.size = 0;
203 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
204 : rank, subLinks.size(), slice.offset, subLinks.size(), slice.size);
205 0 : subSlices.push_back(slice);
206 : }
207 :
208 : // 长度不足2,直接跳过
209 0 : if (subLinks.size() < 2) {
210 0 : return HCCL_SUCCESS;
211 : }
212 :
213 : // 计算在垂直Ring中的rank号
214 0 : u32 subRank = (hIndex == hIndexForRing) ? info.GetVIndex(rank) : vSizeForRing;
215 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Vertical] rank[%u] subRank=%u", rank, subRank);
216 :
217 : // 执行Broken Ring ReduceScatter
218 0 : return RunReduceScatterBrokenRing(subRank, subLinks, subSlices);
219 0 : }
220 :
221 0 : HcclResult ReduceScatterNHRV1::RunReduceScatterOnHorizontal(const u32 rank, const std::vector<LINK> &links,
222 : const RingInfo &info)
223 : {
224 0 : u32 hIndex = info.GetHIndex(rank);
225 0 : if (hIndex >= info.GetRowSize()) {
226 0 : return HCCL_SUCCESS;
227 : }
228 :
229 : // 构造新的links和slices
230 0 : u32 vIndex = info.GetVIndex(rank);
231 0 : std::vector<LINK> subLinks;
232 0 : std::vector<Slice> subSlices;
233 0 : for (u32 hIdx = 0; hIdx < info.GetRowSize(); hIdx++) {
234 : // 增加link
235 0 : u32 rankInRing = info.GetRank(vIndex, hIdx);
236 0 : CHK_PRT_RET(rankInRing >= links.size(), HCCL_ERROR("[ReduceScatterNHRV1][Horizontal] rank[%u] out of range, "\
237 : "rankInRing=%u, links.size=%u", rank, rankInRing, links.size()), HCCL_E_INTERNAL);
238 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] links[%u]=%u", rank, hIdx, rankInRing);
239 0 : subLinks.push_back(links[rankInRing]);
240 :
241 : // 增肌slice(CheckSlices()已经约束slices_里的Slice都是连续的)
242 : // -- 比如8节点在做完水平Ring后
243 : // ---- 0 1 2
244 : // ---- 3 4 5
245 : // ---- 6 7
246 : // -- 0/3/6/7节点只拿到自己那份ReduceScatter结果,而1拿到1和2的两份数据、4拿到4和5的两份数据
247 0 : u64 sliceSize = slices_[rankInRing].size;
248 0 : if (hIdx == info.GetRowSize() - 1 && info.GetHSizeByVIndex(vIndex) > info.GetRowSize()) {
249 0 : sliceSize += slices_[rankInRing+1].size;
250 : }
251 :
252 0 : Slice slice;
253 0 : slice.offset = slices_[rankInRing].offset;
254 0 : slice.size = sliceSize;
255 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] subSlices[%u].offset=%llu, subSlices[%u].size=%llu",
256 : rank, hIdx, slice.offset, hIdx, slice.size);
257 0 : subSlices.push_back(slice);
258 : }
259 :
260 : // 长度不足2,直接跳过
261 0 : if (subLinks.size() < 2) {
262 0 : return HCCL_SUCCESS;
263 : }
264 :
265 : // 计算在水平Ring中的rank号
266 0 : u32 subRank = hIndex;
267 0 : HCCL_DEBUG("[ReduceScatterNHRV1][Horizontal] rank[%u] subRank=%u", rank, subRank);
268 :
269 : // 执行Broken Ring ReduceScatter
270 0 : return RunReduceScatterBrokenRing(subRank, subLinks, subSlices);
271 0 : }
272 :
273 0 : HcclResult ReduceScatterNHRV1::RunLastCopyStep(const u32 rank, const std::vector<LINK> &links, const RingInfo &info)
274 : {
275 : HcclResult ret;
276 :
277 0 : u32 hIndex = info.GetHIndex(rank); // 查找自己位于第几列
278 0 : u32 vIndex = info.GetVIndex(rank); // 查找自己位于第几行
279 0 : if (hIndex >= info.GetRowSize() - 1 && info.GetHSizeByVIndex(vIndex) > info.GetRowSize()) {
280 0 : u32 peerRank = (hIndex == info.GetRowSize() - 1) ? rank + 1 : rank - 1;
281 :
282 : // 检查指针
283 0 : CHK_SMART_PTR_NULL(links[peerRank]);
284 :
285 : // TxAck
286 0 : ret = links[peerRank]->TxAck(stream_);
287 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
288 : HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]rank[%u] tx ack from peerank[%u] failed",
289 : rank, peerRank), ret);
290 :
291 : // RxAck
292 0 : ret = links[peerRank]->RxAck(stream_);
293 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
294 : HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]rank[%u] rx ack from peerank[%u] failed",
295 : rank, peerRank), ret);
296 :
297 0 : if (hIndex == info.GetRowSize() - 1) { // 发数据
298 0 : Slice& txSlice = slices_[peerRank];
299 0 : DeviceMem srcMem = inputMem_.range(txSlice.offset, txSlice.size);
300 0 : HCCL_DEBUG("tx srcMem[%p] range[%llu] size[%llu] ", srcMem.ptr(), txSlice.offset, txSlice.size);
301 0 : CHK_RET(ExecuteTxSync(links[peerRank], UserMemType::INPUT_MEM, txSlice.offset + baseOffset_,
302 : srcMem.ptr(), srcMem.size(), stream_));
303 :
304 0 : ret = links[peerRank]->TxWaitDone(stream_);
305 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]TxWaitDone failed"), ret);
306 0 : } else { // 收数据
307 0 : Slice& rxSlice = slices_[rank];
308 0 : DeviceMem dstMem = inputMem_.range(rxSlice.offset, rxSlice.size);
309 0 : HCCL_DEBUG("rx dstMem[%p] range[%llu], size[%llu] ", dstMem.ptr(), rxSlice.offset, rxSlice.size);
310 0 : CHK_RET(ExecuteRxSync(links[peerRank], UserMemType::INPUT_MEM, rxSlice.offset + baseOffset_,
311 : dstMem.ptr(), dstMem.size(), stream_));
312 :
313 0 : ret = links[peerRank]->RxWaitDone(stream_);
314 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ReduceScatterNHRV1][RunLastCopyStep]RxWaitDone failed"), ret);
315 0 : }
316 :
317 : // 如果不Barrier,SDMA结果在数据量超过CCL Buffer之后结果不正确
318 0 : CHK_RET(ExecuteBarrier(links[peerRank], stream_));
319 : }
320 0 : return HCCL_SUCCESS;
321 : }
322 :
323 0 : HcclResult ReduceScatterNHRV1::RunCopyDataToOutputMem(const u32 rank)
324 : {
325 0 : if (inputMem_ != outputMem_) {
326 0 : Slice& srcSlice = slices_[rank];
327 0 : DeviceMem dst = outputMem_.range(0, srcSlice.size);
328 0 : DeviceMem src = inputMem_.range(srcSlice.offset, srcSlice.size);
329 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
330 0 : }
331 0 : return HCCL_SUCCESS;
332 : }
333 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, ReduceScatterNHRV1);
334 : } // ~~ namespace hccl
|