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