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 "gather_ring.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 0 : GatherRing::GatherRing(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher), interRank_(0), interRankSize_(0)
16 0 : {}
17 :
18 0 : GatherRing::~GatherRing() {}
19 : // 不带入slice时候,计算gather 输出buffer的偏移
20 0 : void GatherRing::PrepareSlicesData(const u32 unitSize, const u64 totalCount, const u32 rankSize)
21 : {
22 0 : slices_.resize(rankSize);
23 0 : u64 sliceSize = totalCount * unitSize;
24 :
25 0 : for (u32 i = 0; i < rankSize; i++) {
26 0 : slices_[i].offset = i * sliceSize;
27 0 : slices_[i].size = sliceSize;
28 0 : HCCL_DEBUG("rank[%u] default slice[%u]: offset: [%llu] size[%llu]", interRank_, i, i * sliceSize, sliceSize);
29 : }
30 0 : }
31 : // root rank只接收数据,需要接收多次
32 0 : HcclResult GatherRing::RunGatherOnRootRank()
33 : {
34 : HcclResult ret;
35 0 : DeviceMem dst;
36 :
37 0 : u32 round = interRankSize_ - 1;
38 : // root第一轮从前一rank的input接收
39 0 : for (u32 i = 1; i <= round; i++) {
40 0 : u32 rcvSlice = (interRank_ - i + interRankSize_) % interRankSize_;
41 0 : u64 rcvOffset = slices_[rcvSlice].offset;
42 0 : u64 rcvSize = slices_[rcvSlice].size;
43 :
44 : // 给前一节点发送同步,以便前一rank进行下一轮的操作
45 0 : ret = linkLeft_->TxAck(stream_);
46 0 : CHK_PRT_RET(
47 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnRootRank]rootrank[%u] tx ack failed", interRank_), ret);
48 0 : dst = outputMem_.range(rcvOffset, rcvSize);
49 :
50 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rcvOffset + baseOffset_, dst.ptr(), rcvSize, stream_);
51 :
52 0 : CHK_PRT_RET(
53 : ret != HCCL_SUCCESS,
54 : HCCL_ERROR(
55 : "[Run][GatherOnRootRank]root rank[%u] rx sync with dstmem[%p] "
56 : "failed",
57 : interRank_, dst.ptr()),
58 : ret);
59 :
60 0 : HCCL_INFO(
61 : "GatherRing rootrank[%u] round[%u] rx data ouputoffset[%llu] size[%llu]", interRank_, i, rcvOffset,
62 : rcvSize);
63 : }
64 0 : return HCCL_SUCCESS;
65 0 : }
66 : // root的后一节点。只是进行发送
67 0 : HcclResult GatherRing::RunGatherOnRootNextRank()
68 : {
69 0 : HcclResult ret = HCCL_SUCCESS;
70 : // 发送自己的input数据向root rank
71 0 : u64 sendOffset = slices_[interRank_].offset;
72 0 : u64 sendSize = slices_[interRank_].size;
73 0 : ret = linkRight_->RxAck(stream_);
74 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnRootNextRank]rank[%u] rx ack failed", interRank_), ret);
75 :
76 0 : DeviceMem sendMem = outputMem_.range(sendOffset, sendSize);
77 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, sendOffset + baseOffset_, sendMem.ptr(), sendSize, stream_);
78 0 : CHK_PRT_RET(
79 : ret != HCCL_SUCCESS,
80 : HCCL_ERROR(
81 : "[Run][GatherOnRootNextRank]rank[%u] TxAsync offset[%llu] size[%llu] "
82 : "failed",
83 : interRank_, sendOffset, sendSize),
84 : ret);
85 0 : HCCL_DEBUG("GatherRing root next rank[%u] tx async offset[%llu] size[%llu]", interRank_, sendOffset, sendSize);
86 0 : ret = linkRight_->TxWaitDone(stream_);
87 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnRootNextRank]TxWaitDone failed"), ret);
88 :
89 0 : return ret;
90 0 : }
91 : // 非root,非root后一节点rank。先发送自己数据,再接收
92 0 : HcclResult GatherRing::RunGatherOnOtherRank()
93 : {
94 0 : DeviceMem src;
95 0 : DeviceMem dst;
96 0 : HcclResult ret = HCCL_SUCCESS;
97 0 : u64 sendOffset = slices_[interRank_].offset;
98 0 : u64 sendSize = slices_[interRank_].size;
99 :
100 : // 把自己数据发送至下一rank
101 0 : ret = linkLeft_->TxAck(stream_);
102 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] tx ack failed", interRank_), ret);
103 0 : ret = linkRight_->RxAck(stream_);
104 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] rx ack failed", interRank_), ret);
105 0 : DeviceMem sendMem = outputMem_.range(sendOffset, sendSize);
106 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, sendOffset + baseOffset_, sendMem.ptr(), sendSize, stream_);
107 0 : CHK_PRT_RET(
108 : ret != HCCL_SUCCESS,
109 : HCCL_ERROR(
110 : "[Run][GatherOnOtherRank]rank[%u] TxAsync offset[%llu] size[%llu] "
111 : "failed",
112 : interRank_, sendOffset, sendSize),
113 : ret);
114 :
115 0 : u32 round = (interRank_ + interRankSize_ - root_ - 1) % interRankSize_;
116 : // 需要接收的和发送的轮数,包含接收自己的数据
117 0 : for (u32 i = 1; i < round; i++) {
118 0 : u32 rcvRank = (interRank_ + interRankSize_ - i) % interRankSize_; // 收到的数据应当是哪个rank的slice
119 0 : u64 rcvOffset = slices_[rcvRank].offset;
120 0 : u64 rcvSize = slices_[rcvRank].size;
121 : // 使用本端oupt接收,第一轮从对端的input,否则从对端的output
122 0 : dst = outputMem_.range(rcvOffset, rcvSize);
123 : // 从前一个节点收数据
124 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rcvOffset + baseOffset_, dst.ptr(), rcvSize, stream_);
125 0 : CHK_PRT_RET(
126 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] rx async failed", interRank_), ret);
127 :
128 0 : ret = linkLeft_->RxWaitDone(stream_);
129 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]RxWaitDone failed"), ret);
130 0 : ret = linkRight_->TxWaitDone(stream_);
131 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]TxWaitDone failed"), ret);
132 :
133 0 : ret = linkLeft_->TxAck(stream_);
134 0 : CHK_PRT_RET(
135 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] round[%u] tx ack failed", interRank_, i),
136 : ret);
137 :
138 : // 从后一rank接收同步信号
139 0 : ret = linkRight_->RxAck(stream_);
140 0 : CHK_PRT_RET(
141 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] round[%u] rx ack failed", interRank_, i),
142 : ret);
143 : // 向后一rank发送数据
144 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, rcvOffset + baseOffset_, dst.ptr(), rcvSize, stream_);
145 0 : CHK_PRT_RET(
146 : ret != HCCL_SUCCESS,
147 : HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] round[%u] tx async failed", interRank_, i), ret);
148 0 : HCCL_DEBUG("GatherRing rank[%u] round[%u] tx async offset[%llu] size[%llu]", interRank_, i, rcvOffset, rcvSize);
149 : }
150 : // 最后一次接收后,再发送,与for循环内动作相比 缺少leftlink->txack,最后一轮收的一定是root next的数据,root的next
151 : // next节点不走for循环
152 0 : u32 rcvSlice = (root_ + 1) % interRankSize_;
153 0 : u64 rcvOffset = slices_[rcvSlice].offset;
154 0 : u64 rcvSize = slices_[rcvSlice].size;
155 0 : dst = outputMem_.range(rcvOffset, rcvSize);
156 :
157 : // 从前一个节点收数据
158 0 : ret = linkLeft_->RxAsync(UserMemType::OUTPUT_MEM, rcvOffset + baseOffset_, dst.ptr(), rcvSize, stream_);
159 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] rx async failed", interRank_), ret);
160 0 : ret = linkLeft_->RxWaitDone(stream_);
161 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]RxWaitDone failed"), ret);
162 0 : ret = linkRight_->TxWaitDone(stream_);
163 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]TxWaitDone failed"), ret);
164 : // 从后一rank接收同步信号
165 0 : ret = linkRight_->RxAck(stream_);
166 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] rx ack failed", interRank_), ret);
167 : // 向后一rank发送数据
168 0 : ret = linkRight_->TxAsync(UserMemType::OUTPUT_MEM, rcvOffset + baseOffset_, dst.ptr(), rcvSize, stream_);
169 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]rank[%u] tx async failed", interRank_), ret);
170 0 : ret = linkRight_->TxWaitDone(stream_);
171 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][GatherOnOtherRank]TxWaitDone failed"), ret);
172 0 : return HCCL_SUCCESS;
173 0 : }
174 : // gather的入口函数
175 : HcclResult
176 0 : GatherRing::RunAsync(const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
177 : {
178 0 : CHK_SMART_PTR_NULL(dispatcher_);
179 0 : CHK_PTR_NULL(stream_.ptr());
180 0 : if (!outputMem_ || !inputMem_) {
181 0 : HCCL_ERROR("[GatherRing][RunAsync]run_async inputmem or outputmem is null");
182 0 : return HCCL_E_PTR;
183 : }
184 :
185 0 : interRank_ = rank;
186 0 : interRankSize_ = rankSize;
187 :
188 0 : HCCL_INFO(
189 : "GatherRing run: rank[%u] totalrank[%u] count[%llu] input[%p] output[%p]", interRank_, interRankSize_, count_,
190 : inputMem_.ptr(), outputMem_.ptr());
191 :
192 : // ranksize为1时,只有当input!=output 时候进行拷贝
193 0 : if (interRankSize_ == 1) {
194 0 : if (inputMem_ != outputMem_) {
195 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
196 : }
197 0 : return HCCL_SUCCESS;
198 : }
199 :
200 0 : u32 unitSize = DataUnitSize(dataType_);
201 0 : CHK_PRT_RET(
202 : unitSize == 0, HCCL_ERROR("[GatherRing][RunAsync]rank[%u] unit data size is zero", rank), HCCL_E_INTERNAL);
203 :
204 : // 带入vecotr为空,计算每个rank的结果偏移和大小
205 0 : if (slices_.size() == 0) {
206 0 : PrepareSlicesData(unitSize, count_, interRankSize_);
207 : }
208 :
209 0 : u32 ringPrevRank = (rank + rankSize - 1) % rankSize;
210 0 : u32 ringNextRank = (rank + 1) % rankSize;
211 :
212 0 : if (links.size() < rankSize) {
213 0 : HCCL_ERROR("[GatherRing][RunAsync]rank[%u] link size[%llu] is less than rank size", rank, links.size());
214 0 : return HCCL_E_INTERNAL;
215 : }
216 :
217 0 : linkLeft_ = links[ringPrevRank];
218 0 : CHK_SMART_PTR_NULL(linkLeft_);
219 :
220 0 : linkRight_ = links[ringNextRank];
221 0 : CHK_SMART_PTR_NULL(linkRight_);
222 :
223 0 : if (interRank_ == root_) {
224 0 : CHK_RET(RunGatherOnRootRank());
225 0 : } else if (ringPrevRank == root_) {
226 0 : CHK_RET(RunGatherOnRootNextRank());
227 : } else {
228 0 : CHK_RET(RunGatherOnOtherRank());
229 : }
230 :
231 0 : if (barrierSwitchOn_) {
232 : // 执行barrier,保证数据收发完成
233 0 : CHK_RET(ExecuteBarrier(linkLeft_, linkRight_));
234 : }
235 :
236 0 : HCCL_INFO("GatherRing finished: rank:[%u] end", interRank_);
237 :
238 0 : return HCCL_SUCCESS;
239 : }
240 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_GATHER_RING, GatherRing);
241 : } // namespace hccl
|