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 "broadcast_star.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 : // Gather的入口函数
16 0 : BroadcastStar::BroadcastStar(const HcclDispatcher dispatcher) : AlgTemplateBase(dispatcher) {}
17 :
18 0 : BroadcastStar::~BroadcastStar() {}
19 :
20 0 : HcclResult BroadcastStar::Prepare(
21 : DeviceMem& inputMem, DeviceMem& outputMem, DeviceMem& scratchMem, const u64 count, const HcclDataType dataType,
22 : const Stream& stream, const HcclReduceOp reductionOp, const u32 root, const std::vector<Slice>& slices,
23 : const u64 baseOffset, std::vector<u32> nicRankList, u32 userRank)
24 : {
25 0 : userRank_ = userRank;
26 0 : return AlgTemplateBase::Prepare(
27 0 : inputMem, outputMem, scratchMem, count, dataType, stream, reductionOp, root, slices, baseOffset, nicRankList);
28 : }
29 :
30 : HcclResult
31 0 : BroadcastStar::RunAsync(const u32 rank, const u32 rankSize, const std::vector<std::shared_ptr<Transport>>& links)
32 : {
33 : // task下发接口
34 0 : CHK_SMART_PTR_NULL(dispatcher_);
35 : // ==1的处理
36 0 : if (rankSize == 1) {
37 0 : if (inputMem_ != outputMem_) {
38 0 : HcclResult ret = HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
39 0 : CHK_PRT_RET(
40 : ret != HCCL_SUCCESS,
41 : HCCL_ERROR(
42 : "[BroadcastStar][RunAsync]rank[%u] copy input[%p] to output[%p] failed", rank, inputMem_.ptr(),
43 : outputMem_.ptr()),
44 : ret);
45 : }
46 0 : return HCCL_SUCCESS;
47 : }
48 : // links本rank_id与通信域内其它rank的通信连接,rankSize本executor所在通信域的rank个数
49 0 : if (links.size() < rankSize) {
50 0 : HCCL_ERROR(
51 : "[BroadcastStar][RunAsync]rank[%u] linksize[%llu] is less than rankSize[%u]", rank, links.size(), rankSize);
52 0 : return HCCL_E_INTERNAL;
53 : }
54 :
55 0 : Slice sendSlice;
56 0 : sendSlice.offset = 0;
57 0 : sendSlice.size = dataBytes_;
58 0 : Slice recvSlice;
59 0 : recvSlice.offset = 0;
60 0 : recvSlice.size = dataBytes_;
61 0 : if (rank == root_) {
62 : // root 给其他rank发
63 : HcclResult ret;
64 0 : for (u32 dstRank = 0; dstRank < rankSize; dstRank++) {
65 0 : ret = RunSendBroadcast(dstRank, sendSlice, links);
66 0 : CHK_PRT_RET(
67 : ret != HCCL_SUCCESS,
68 : HCCL_ERROR(
69 : "[BroadcastStar][RunAsync] root [%u] send broadcast to"
70 : "other rank[%u] run failed!",
71 : root_, dstRank),
72 : ret);
73 : }
74 : } else {
75 : // 非root 接收来自root的数据
76 0 : HcclResult ret = RunRecvBroadcast(root_, rank, recvSlice, links);
77 0 : CHK_PRT_RET(
78 : ret == HCCL_E_AGAIN, HCCL_WARNING("[BroadcastStar][RunAsync]group has been destroyed. Break!"), ret);
79 0 : CHK_PRT_RET(
80 : ret != HCCL_SUCCESS,
81 : HCCL_ERROR(
82 : "[BroadcastStar][RunAsync] dstrank [%u] recv broadcast from"
83 : "root [%u] run failed!",
84 : rank, root_),
85 : ret);
86 : }
87 0 : HCCL_INFO("BroadBastStar finished: rank[%u]", rank);
88 0 : return HCCL_SUCCESS;
89 : }
90 :
91 0 : HcclResult BroadcastStar::RunRecvBroadcast(
92 : const u32 srcRank, const u32 dstRank, const Slice& slice, const std::vector<LINK>& links)
93 : {
94 : // 非root 接受数据
95 0 : DeviceMem dst;
96 0 : if (slice.size > 0) {
97 0 : if (srcRank >= links.size()) {
98 0 : HCCL_ERROR("[RunRecvBroadcast] root [%u] is out of range, linksize[%llu]", srcRank, links.size());
99 0 : return HCCL_E_INTERNAL;
100 : }
101 0 : dst = outputMem_.range(slice.offset, slice.size);
102 0 : HCCL_DEBUG(
103 : "rank [%u] will recv with output's offset[%llu], size[%llu], dstmem[%p]", dstRank, slice.offset, slice.size,
104 : dst.ptr());
105 :
106 0 : if (links[srcRank]->IsTransportRoce()) {
107 0 : CHK_RET(links[srcRank]->RxEnv(stream_));
108 : } else {
109 0 : CHK_RET(links[srcRank]->TxAck(stream_));
110 : }
111 :
112 0 : HcclResult ret = links[srcRank]->RxAsync(UserMemType::OUTPUT_MEM, slice.offset, dst.ptr(), slice.size, stream_);
113 0 : CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[RunRecvBroadcast]group has been destroyed. Break!"), ret);
114 0 : CHK_PRT_RET(
115 : ret != HCCL_SUCCESS,
116 : HCCL_ERROR(
117 : "[RunRecvBroadcast]root rank[%u] rx async to dstrank[%u] run "
118 : "failed",
119 : srcRank, dstRank),
120 : ret);
121 :
122 0 : if (!links[srcRank]->IsTransportRoce()) {
123 0 : ret = ExecuteBarrier(links[srcRank], stream_); // 多server走rdma可以不用
124 0 : CHK_PRT_RET(
125 : ret != HCCL_SUCCESS,
126 : HCCL_ERROR("[RunRecvBroadcast]dstRank[%u] Broadcast star run tempAlg barrier failed", dstRank), ret);
127 0 : ret = links[srcRank]->RxWaitDone(stream_); // 多server走rdma可以不用
128 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[RunRecvBroadcast]RxWaitDone failed"), ret);
129 : }
130 : }
131 0 : return HCCL_SUCCESS;
132 0 : }
133 :
134 0 : HcclResult BroadcastStar::RunSendBroadcast(const u32 dstRank, const Slice& slice, const std::vector<LINK>& links)
135 : {
136 0 : DeviceMem src;
137 : // root发送数据
138 0 : if (slice.size > 0 && dstRank != root_) {
139 0 : src = inputMem_.range(slice.offset, slice.size);
140 :
141 0 : if (links[dstRank]->IsTransportRoce()) {
142 0 : CHK_RET(links[dstRank]->TxEnv(src.ptr(), slice.size, stream_));
143 : } else {
144 0 : CHK_RET(links[dstRank]->RxAck(stream_));
145 : }
146 :
147 0 : HcclResult ret = links[dstRank]->TxAsync(UserMemType::OUTPUT_MEM, slice.offset, src.ptr(), slice.size, stream_);
148 0 : CHK_PRT_RET(
149 : ret != HCCL_SUCCESS,
150 : HCCL_ERROR("[RunSendBroadcast]rank[%u] tx async with output's offset[%llu] failed", dstRank, slice.offset),
151 : ret);
152 :
153 0 : if (!links[dstRank]->IsTransportRoce()) {
154 0 : ret = ExecuteBarrierSrcRank(links[dstRank], stream_);
155 0 : CHK_PRT_RET(
156 : ret != HCCL_SUCCESS,
157 : HCCL_ERROR("[RunSendBroadcast] srcRank[%u] broadcast star run tempAlg barrier failed", root_), ret);
158 :
159 0 : ret = links[dstRank]->TxWaitDone(stream_);
160 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[RunSendBroadcast]TxWaitDone failed"), ret);
161 : }
162 : }
163 0 : return HCCL_SUCCESS;
164 0 : }
165 :
166 0 : HcclResult BroadcastStar::ExecuteBarrierSrcRank(std::shared_ptr<Transport> link, Stream& stream) const
167 : {
168 0 : CHK_RET(link->RxAck(stream));
169 :
170 0 : CHK_RET(link->TxAck(stream));
171 :
172 0 : CHK_RET(link->RxDataSignal(stream));
173 :
174 0 : CHK_RET(link->TxDataSignal(stream));
175 :
176 0 : return HCCL_SUCCESS;
177 : }
178 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_BROADCAST_STAR, BroadcastStar);
179 : } // namespace hccl
|