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_ring.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : ScatterRing::ScatterRing(const HcclDispatcher dispatcher)
16 0 : : AlgTemplateBase(dispatcher), interRank_(0), interRankSize_(0)
17 : {
18 0 : }
19 :
20 0 : ScatterRing::~ScatterRing()
21 : {
22 0 : }
23 :
24 0 : HcclResult ScatterRing::RunScatterOnRootRank()
25 : {
26 0 : DeviceMem src;
27 0 : DeviceMem dst;
28 : // rank存放scatter 结果的偏移
29 0 : u64 scatterOffset = slices_[interRank_].offset;
30 0 : u64 scatterResult = slices_[interRank_].size;
31 :
32 0 : HcclResult ret = HCCL_SUCCESS;
33 : // 需要判断input不等于outputmem,scatter 输入只有一个input时不用拷贝
34 0 : if (inputMem_ != outputMem_) {
35 0 : src = inputMem_.range(scatterOffset, scatterResult);
36 0 : dst = outputMem_.range(scatterOffset, scatterResult);
37 :
38 0 : HCCL_DEBUG("rootrank[%u] copy input[%p] to output[%p] scatter_offset[%llu] copysize[%llu]", \
39 : interRank_, src.ptr(), dst.ptr(), scatterOffset, scatterResult);
40 0 : ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
41 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
42 : HCCL_ERROR("[Run][ScatterOnRootRank]root rank[%u] memcpy async from input[%p] "\
43 : "failed to output[%p]", interRank_, inputMem_.ptr(), outputMem_.ptr()), ret);
44 : }
45 :
46 : // 数据向下一个rank发送,依次发送后继所有rank的数据
47 0 : for (u32 i = 1; i < interRankSize_; i++) {
48 0 : u32 preRank = (interRank_ - i + interRankSize_) % interRankSize_;
49 0 : scatterOffset = slices_[preRank].offset;
50 0 : scatterResult = slices_[preRank].size;
51 :
52 0 : src = inputMem_.range(scatterOffset, scatterResult);
53 : // 等待后一节点同步信号,进行下一轮操作
54 0 : CHK_RET(linkRight_->RxAck(stream_));
55 :
56 : // 向root rank的后一rank发送
57 0 : HCCL_DEBUG(" root rank[%u] sendto dstrank[%u] from srcmem offset[%llu] size[%llu]", \
58 : interRank_, preRank, scatterOffset, scatterResult);
59 0 : CHK_RET(linkRight_->TxAsync(UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, src.ptr(),
60 : scatterResult, stream_));
61 :
62 0 : HCCL_DEBUG("root rank[%u] will rx_ack", interRank_);
63 0 : ret = linkRight_->TxWaitDone(stream_);
64 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnRootRank]TxWaitDone failed"), ret);
65 : }
66 0 : return HCCL_SUCCESS;
67 0 : }
68 0 : HcclResult ScatterRing::RunScatterOnEndRank()
69 : {
70 0 : DeviceMem src;
71 0 : DeviceMem dst;
72 0 : u64 scatterOffset = slices_[interRank_].offset;
73 0 : u64 scatterResult = slices_[interRank_].size;
74 : // 给前一节点发送同步,以便前一rank进行下一轮的操作
75 0 : CHK_RET(linkLeft_->TxAck(stream_));
76 :
77 0 : dst = outputMem_.range(scatterOffset, scatterResult);
78 0 : HCCL_DEBUG("last rank[%u] rx data ouputoffset[%llu] size[%llu]", \
79 : interRank_, scatterOffset, scatterResult);
80 0 : HcclResult ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(),
81 0 : scatterResult, stream_);
82 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnEndRank]last rank[%u] rx sync failed", \
83 : interRank_), ret);
84 0 : ret = linkLeft_->RxWaitDone(stream_);
85 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnRootRank]RxWaitDone failed"), ret);
86 0 : return HCCL_SUCCESS;
87 0 : }
88 0 : HcclResult ScatterRing::RunScatterOnMidRank()
89 : {
90 0 : DeviceMem src;
91 0 : DeviceMem dst;
92 0 : DeviceMem dstLast;
93 : // 与root的rank号之差 + 接收的轮数 = rank_size, 每个rank 接收的次数为 root_+ranksize-rank%interRankSize_
94 0 : u32 round = (root_ + interRankSize_ - interRank_) % interRankSize_;
95 0 : HCCL_DEBUG("rank:[%u] will receive %u rounds data", interRank_, round);
96 :
97 0 : UserMemType memType = (interRank_ == ((root_ + 1) % interRankSize_)) ?
98 : UserMemType::INPUT_MEM : UserMemType::OUTPUT_MEM;
99 :
100 0 : HcclResult ret = HCCL_SUCCESS;
101 : // 需要接收的和发送的轮数,包含接收自己的数据
102 0 : for (u32 i = 1; i <= round; i++) {
103 0 : u32 dataRank = (interRank_ + round - i) % interRankSize_; // 收到的数据应当是哪个rank的
104 0 : u64 scatterOffset = slices_[dataRank].offset;
105 0 : u64 scatterResult = slices_[dataRank].size;
106 :
107 0 : u32 lastDataRank = (interRank_ + round - i + 1) % interRankSize_; // 加1得到发送的数据应当是哪个rank的
108 0 : u64 scatterLastOffset = slices_[lastDataRank].offset;
109 0 : u64 scatterLastResult = slices_[lastDataRank].size;
110 :
111 0 : dst = outputMem_.range(scatterOffset, scatterResult);
112 0 : dstLast = outputMem_.range(scatterLastOffset, scatterLastResult);
113 :
114 0 : if (i != 1) {
115 : // 给前一节点发送同步,以便前一rank进行下一轮的操作
116 0 : ret = linkLeft_->TxAck(stream_);
117 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx ack failed",
118 : interRank_, i), ret);
119 : // 从后一rank接收同步信号
120 0 : ret = linkRight_->RxAck(stream_);
121 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u]round[%u] rx ack failed",
122 : interRank_, i), ret);
123 : // 向后一rank发送数据
124 0 : HCCL_DEBUG("rank[%u] round[%u] tx async offset[%llu] size[%llu]", interRank_, \
125 : i, scatterLastOffset, scatterLastResult);
126 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, scatterLastOffset + baseOffset_, dstLast.ptr(),
127 0 : scatterLastResult, stream_);
128 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx async failed",
129 : interRank_, i), ret);
130 : } else { // 最后一轮接收数据,拷贝到自己的outputmem
131 : // 给前一节点发送同步,以便前一rank进行下一轮的操作
132 0 : ret = linkLeft_->TxAck(stream_);
133 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]rank[%u] round[%u] tx ack failed",
134 : interRank_, i), ret);
135 : }
136 0 : HCCL_DEBUG("rank[%u] round[%u] rcv with rank[%u]'s offset[%llu] size[%llu]", \
137 : interRank_, i, dataRank, scatterOffset, scatterResult);
138 0 : CHK_RET(linkLeft_->RxAsync(memType, scatterOffset + baseOffset_, dst.ptr(), scatterResult, stream_));
139 0 : ret = linkRight_->TxWaitDone(stream_);
140 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]TxWaitDone failed"), ret);
141 0 : ret = linkLeft_->RxWaitDone(stream_);
142 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ScatterOnMidRank]RxWaitDone failed"), ret);
143 : }
144 0 : return HCCL_SUCCESS;
145 0 : }
146 :
147 0 : void ScatterRing::PrepareSlicesData(const u32 unitSize, const u64 totalCount, const u32 rankSize) const
148 : {
149 0 : slices_.resize(rankSize);
150 0 : u64 sliceSize = (totalCount / rankSize) * unitSize;
151 :
152 0 : for (u32 i = 0; i < rankSize; i++) {
153 0 : slices_[i].offset = i * sliceSize;
154 0 : slices_[i].size = sliceSize;
155 0 : HCCL_DEBUG("rank[%u] default slice[%u]: offset: [%llu] size[%llu]", interRank_, i, i * sliceSize, sliceSize);
156 : }
157 0 : }
158 :
159 : // scatter的入口函数
160 0 : HcclResult ScatterRing::RunAsync(const u32 rank, const u32 rankSize,
161 : const std::vector<std::shared_ptr<Transport> > &links)
162 : {
163 0 : CHK_SMART_PTR_NULL(dispatcher_);
164 0 : CHK_PTR_NULL(stream_.ptr());
165 0 : if (!outputMem_ || !inputMem_) {
166 0 : HCCL_ERROR("[ScatterRing][RunAsync]run_async inputmem or outputmem is null");
167 0 : return HCCL_E_PTR;
168 : }
169 :
170 0 : interRank_ = rank;
171 0 : interRankSize_ = rankSize;
172 :
173 0 : HCCL_INFO("ScatterRing run: rank[%u] totalrank[%u] count[%llu] input[%p] output[%p]",
174 : interRank_, interRankSize_, count_, inputMem_.ptr(), outputMem_.ptr());
175 :
176 : // ranksize为1时,只有当input!=output 时候进行拷贝
177 0 : if (interRankSize_ == 1) {
178 0 : if (inputMem_ != outputMem_) {
179 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
180 : }
181 0 : return HCCL_SUCCESS;
182 : }
183 :
184 0 : u32 unitSize = DataUnitSize(dataType_);
185 0 : CHK_PRT_RET(unitSize == 0, HCCL_ERROR("[ScatterRing][RunAsync]rank[%u] unit data size is zero", rank),
186 : HCCL_E_INTERNAL);
187 :
188 : // 带入vecotr为空,计算每个rank的结果偏移和大小
189 0 : if (slices_.size() == 0) {
190 0 : PrepareSlicesData(unitSize, count_, interRankSize_);
191 : }
192 :
193 : // 获取link的收、发缓存, 计算chunk_size
194 0 : u32 ringPrevRank = (rank + rankSize - 1) % rankSize;
195 0 : u32 ringNextRank = (rank + 1) % rankSize;
196 :
197 0 : if (links.size() < rankSize) {
198 0 : HCCL_ERROR("[ScatterRing][RunAsync]rank[%u] link size[%llu] is less than rank size", rank, links.size());
199 0 : return HCCL_E_INTERNAL;
200 : }
201 :
202 0 : linkLeft_ = links[ringPrevRank];
203 0 : CHK_SMART_PTR_NULL(linkLeft_);
204 :
205 0 : linkRight_ = links[ringNextRank];
206 0 : CHK_SMART_PTR_NULL(linkRight_);
207 :
208 0 : CHK_RET(ScatterSlicesPrep(rankSize, nicRankList_.size()));
209 :
210 : // 单环场景下 nicRankList_ 长度默认为 8。
211 : // 多环场景下 nicRankList_ 长度为网口数量。此时若 rankSize != nicRankList_ 则为网口裁剪场景
212 0 : if (rankSize != HCCL_NIC_MAX_NUM || nicRankList_.size() == HCCL_NIC_MAX_NUM) {
213 : // 非网口裁剪场景:
214 : // root rank向其他rank发送数据,
215 0 : if (interRank_ == root_) {
216 0 : CHK_RET(RunScatterOnRootRank());
217 0 : } else if (ringNextRank == root_) { // 最后一个节点只负责接收数据,拷贝至outputmem
218 0 : CHK_RET(RunScatterOnEndRank());
219 : } else {
220 0 : CHK_RET(RunScatterOnMidRank());
221 : }
222 : } else {
223 : // 网口裁剪场景:当前仅在 910A 8P_RING (4环),且网口不满配情况下使用
224 0 : CHK_RET(RunScatterChunk(rank, rankSize, slices_));
225 : }
226 :
227 0 : if (barrierSwitchOn_) {
228 : // 执行barrier,保证数据收发完成
229 0 : CHK_RET(ExecuteBarrier(linkLeft_, linkRight_));
230 : }
231 0 : HCCL_INFO("ScatterRing finished: rank:[%u] end", interRank_);
232 :
233 0 : return HCCL_SUCCESS;
234 : }
235 :
236 0 : HcclResult ScatterRing::RunScatterChunk(const u32 rank, const u32 rankSize, const std::vector<Slice> &outputSlices)
237 : {
238 : HcclResult ret;
239 0 : DeviceMem dst;
240 0 : u32 sendSliceLen = rankSliceLists_[rank].size();
241 0 : u32 chunkSize = HCCL_NIC_MAX_NUM / nicRankList_.size();
242 0 : if (sendSliceLen >= chunkSize) {
243 0 : CHK_RET(HeadScatterChunk(rank, rankSize, outputSlices));
244 0 : for (u32 midRankIdx = 1; midRankIdx < sendSliceLen - 1; midRankIdx++) {
245 0 : ret = MidScatterChunk(rank, rankSize, midRankIdx, outputSlices);
246 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
247 : HCCL_ERROR("[Run][ScatterChunk]rank[%u] run mid[%u] ReduceScatter chunk failed",
248 : rank, midRankIdx), HCCL_E_INTERNAL);
249 : }
250 0 : CHK_RET(TailScatterChunk(rank, rankSize, sendSliceLen - 1, outputSlices));
251 0 : } else if (rankSliceLists_[(rank + rankSize - 1) % rankSize].size() != 0) {
252 0 : for (u32 sliceIdx = 0; sliceIdx < chunkSize; sliceIdx++) {
253 0 : std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), rank);
254 0 : u32 nicIdx = distance(nicRankList_.begin(), iterNic);
255 0 : u32 chunkStart = nicIdx * chunkSize;
256 0 : u32 rxSliceIndex = chunkStart + sliceIdx;
257 0 : u64 rxScatterOffset = slices_[rxSliceIndex].offset;
258 0 : u64 rxScatterResult = slices_[rxSliceIndex].size;
259 0 : dst = outputMem_.range(rxScatterOffset, rxScatterResult);
260 0 : CHK_RET(linkLeft_->TxAck(stream_));
261 :
262 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(),
263 0 : rxScatterResult, stream_);
264 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
265 : HCCL_ERROR("[Run][ScatterChunk]rank[%u] Left Link rx outputSlices[%u] Failed",
266 : rank, rxSliceIndex), ret);
267 : }
268 : }
269 0 : return HCCL_SUCCESS;
270 0 : }
271 :
272 0 : HcclResult ScatterRing::HeadScatterChunk(u32 rank, u32 rankSize, const std::vector<Slice> &outputSlices)
273 : {
274 : HcclResult ret;
275 0 : DeviceMem dst;
276 0 : u32 rxSliceIndex = rankSliceLists_[rank][0];
277 0 : u32 txSliceIndex = rxSliceIndex;
278 0 : u64 scatterOffset = slices_[rxSliceIndex].offset;
279 0 : u64 scatterResult = slices_[rxSliceIndex].size;
280 0 : dst = outputMem_.range(scatterOffset, scatterResult);
281 0 : std::vector<u32> preRankSlices(rankSliceLists_[(rank - 1 + rankSize) % rankSize]);
282 0 : std::vector<u32>::iterator iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rxSliceIndex);
283 0 : if (iterSlice != preRankSlices.end()) {
284 0 : CHK_RET(linkLeft_->TxAck(stream_));
285 :
286 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(),
287 0 : scatterResult, stream_);
288 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][HeadScatterChunk]rank[%u] Left Link rx "\
289 : "outputSlices[%u] Failed", rank, rxSliceIndex), ret);
290 : }
291 0 : iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rankSliceLists_[rank][1]);
292 0 : if (iterSlice != preRankSlices.end()) {
293 0 : CHK_RET(MidScatterChunk(rank, rankSize, 0, outputSlices));
294 : } else {
295 0 : CHK_RET(linkRight_->RxAck(stream_));
296 :
297 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, scatterOffset + baseOffset_, dst.ptr(),
298 0 : scatterResult, stream_);
299 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][HeadScatterChunk]rank[%u] Right Link tx "\
300 : "outputSlices[%u] Failed", rank, txSliceIndex), ret);
301 : }
302 0 : return HCCL_SUCCESS;
303 0 : }
304 :
305 0 : HcclResult ScatterRing::MidScatterChunk(u32 rank, u32 rankSize, u32 sliceIdx, const std::vector<Slice> &outputSlices)
306 : {
307 : (void)outputSlices;
308 : HcclResult ret;
309 0 : DeviceMem dst;
310 0 : u32 rxSliceIndex = rankSliceLists_[rank][sliceIdx + 1];
311 0 : u32 txSliceIndex = rankSliceLists_[rank][sliceIdx];
312 0 : u64 rxScatterOffset = slices_[rxSliceIndex].offset;
313 0 : u64 rxScatterResult = slices_[rxSliceIndex].size;
314 0 : u64 txScatterOffset = slices_[txSliceIndex].offset;
315 0 : u64 txScatterResult = slices_[txSliceIndex].size;
316 :
317 0 : std::vector<u32> preRankSlices(rankSliceLists_[(rank - 1 + rankSize) % rankSize]);
318 0 : std::vector<u32>::iterator iterSlice = std::find(preRankSlices.begin(), preRankSlices.end(), rxSliceIndex);
319 0 : if (iterSlice != preRankSlices.end()) {
320 0 : CHK_RET(linkLeft_->TxAck(stream_));
321 :
322 0 : dst = outputMem_.range(txScatterOffset, txScatterResult);
323 0 : CHK_RET(linkRight_->RxAck(stream_));
324 :
325 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(),
326 0 : txScatterResult, stream_);
327 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][MidScatterChunk]rank[%u] Right Link tx "\
328 : "outputSlices[%u] Failed", rank, txSliceIndex), ret);
329 0 : dst = outputMem_.range(rxScatterOffset, rxScatterResult);
330 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(),
331 0 : rxScatterResult, stream_);
332 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][MidScatterChunk]rank[%u] Left Link rx "\
333 : "outputSlices[%u] Failed", rank, rxSliceIndex), ret);
334 : } else {
335 0 : dst = outputMem_.range(txScatterOffset, txScatterResult);
336 0 : CHK_RET(linkRight_->RxAck(stream_));
337 :
338 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(),
339 0 : txScatterResult, stream_);
340 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][MidScatterChunk]rank[%u] Right Link tx "\
341 : "outputSlices[%u] Failed", rank, txSliceIndex), ret);
342 : }
343 0 : return HCCL_SUCCESS;
344 0 : }
345 :
346 0 : HcclResult ScatterRing::TailScatterChunk(u32 rank, u32 rankSize, u32 sliceIdx, const std::vector<Slice> &outputSlices)
347 : {
348 : (void)rankSize;
349 : (void)outputSlices;
350 : HcclResult ret;
351 0 : DeviceMem dst;
352 0 : u32 chunkSize = HCCL_NIC_MAX_NUM / nicRankList_.size();
353 0 : u32 txSliceIndex = rankSliceLists_[rank][sliceIdx];
354 0 : u64 txScatterOffset = slices_[txSliceIndex].offset;
355 0 : u64 txScatterResult = slices_[txSliceIndex].size;
356 0 : std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), rank);
357 0 : if (iterNic != nicRankList_.end() && rank != root_) {
358 0 : u32 nicIdx = distance(nicRankList_.begin(), iterNic);
359 0 : u32 chunkStart = nicIdx * chunkSize;
360 0 : u32 rxSliceIndex = chunkStart;
361 0 : u64 rxScatterOffset = slices_[rxSliceIndex].offset;
362 0 : u64 rxScatterResult = slices_[rxSliceIndex].size;
363 0 : CHK_RET(linkLeft_->TxAck(stream_));
364 :
365 0 : dst = outputMem_.range(txScatterOffset, txScatterResult);
366 0 : CHK_RET(linkRight_->RxAck(stream_));
367 :
368 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(),
369 0 : txScatterResult, stream_);
370 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][TailScatterChunk]rank[%u] Right Link tx "\
371 : "outputSlices[%u] Failed", rank, txSliceIndex), ret);
372 0 : dst = outputMem_.range(rxScatterOffset, rxScatterResult);
373 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(),
374 0 : rxScatterResult, stream_);
375 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][TailScatterChunk]rank[%u] Left Link rx "\
376 : "outputSlices[%u] Failed", rank, rxSliceIndex), ret);
377 :
378 0 : for (u32 sliceIdx = 1; sliceIdx < chunkSize; sliceIdx++) {
379 0 : rxSliceIndex = chunkStart + sliceIdx;
380 0 : u64 rxScatterOffset = slices_[rxSliceIndex].offset;
381 0 : u64 rxScatterResult = slices_[rxSliceIndex].size;
382 0 : dst = outputMem_.range(rxScatterOffset, rxScatterResult);
383 0 : CHK_RET(linkLeft_->TxAck(stream_));
384 :
385 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rxScatterOffset + baseOffset_, dst.ptr(),
386 0 : rxScatterResult, stream_);
387 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][TailScatterChunk]rank[%u] Left Link rx "\
388 : "outputSlices[%u] Failed", rank, rxSliceIndex), ret);
389 : }
390 : } else {
391 0 : dst = outputMem_.range(txScatterOffset, txScatterResult);
392 0 : CHK_RET(linkRight_->RxAck(stream_));
393 :
394 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, txScatterOffset + baseOffset_, dst.ptr(),
395 0 : txScatterResult, stream_);
396 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[ScatterRing][TailScatterChunk]rank[%u] Right Link tx "\
397 : "outputSlices[%u] Failed", rank, txSliceIndex), ret);
398 : }
399 0 : return HCCL_SUCCESS;
400 0 : }
401 :
402 0 : HcclResult ScatterRing::ScatterSlicesPrep(u32 rankSize, u32 nicSize)
403 : {
404 0 : u32 chunkSize = HCCL_NIC_MAX_NUM / nicSize;
405 0 : for (u32 rankIdx = 0; rankIdx < rankSize; rankIdx++) {
406 0 : std::vector<u32> sliceList; // 单个rank上的发送slice编号
407 0 : for (u32 nicDis = 1; nicDis <= rankSize; nicDis++) { // 递减从root遍历至当前rank的位置
408 0 : u32 nicIdx = (root_ + rankSize - nicDis) % rankSize;
409 0 : if (rankIdx == nicIdx) {
410 0 : break;
411 : }
412 0 : std::vector<u32>::iterator iterNic = std::find(nicRankList_.begin(), nicRankList_.end(), nicIdx);
413 0 : if (iterNic != nicRankList_.end()) { // 当前rank为网口所在位置,将网口对应的chunksize份silce放入sliceList
414 0 : u32 nicListIdx = distance(nicRankList_.begin(), iterNic);
415 0 : for (u32 chunkIdx = 0; chunkIdx < chunkSize; chunkIdx++) {
416 0 : sliceList.push_back(chunkSize * nicListIdx + chunkIdx);
417 : }
418 : }
419 : }
420 0 : rankSliceLists_.push_back(sliceList);
421 0 : }
422 0 : return HCCL_SUCCESS;
423 : }
424 0 : HcclResult ScatterRing::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
425 : const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
426 : {
427 0 : if (rankSize == 1) {
428 0 : return HCCL_E_NOT_SUPPORT;
429 : }
430 0 : u32 ringNextRank = (rank + 1) % rankSize;
431 0 : LINK nslbNext = links[ringNextRank];
432 :
433 0 : NslbDpAdjInfo adjInfoStep = {0};
434 0 : nslbAdjInfo.dstRankNum = 1;
435 0 : adjInfoStep.dstLocalRankId = nslbNext->GetRemoteRank();
436 0 : adjInfoStep.phaseId = 1;
437 0 : adjInfoStep.rev = 0;
438 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
439 :
440 0 : return HCCL_SUCCESS;
441 0 : }
442 :
443 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_SCATTER_RING, ScatterRing);
444 : } // namespace hccl
|