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