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