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 "reduce_scatter_mesh_mix_single_stream.h"
12 : #include "alg_template_register.h"
13 :
14 : namespace hccl {
15 1 : ReduceScatterMeshMixSingleStream::ReduceScatterMeshMixSingleStream(const HcclDispatcher dispatcher)
16 : : AlgTemplateBase(dispatcher),
17 1 : reduceAttr_(0),
18 1 : streamIndex_(0)
19 1 : {}
20 :
21 2 : ReduceScatterMeshMixSingleStream::~ReduceScatterMeshMixSingleStream() {}
22 :
23 1 : HcclResult ReduceScatterMeshMixSingleStream::Prepare(u64 reduceAttrBitMap, u32 streamIndex)
24 : {
25 1 : reduceAttr_ = reduceAttrBitMap;
26 1 : streamIndex_ = streamIndex;
27 1 : return HCCL_SUCCESS;
28 : }
29 :
30 0 : HcclResult ReduceScatterMeshMixSingleStream::RunSourceReducer(
31 : const LINK& link, const std::vector<Slice>& txSlices, const std::vector<Slice>& dstSlices)
32 : {
33 : // 发送inputmem
34 0 : std::vector<SenderMemoryInfo> txMems;
35 0 : for (u64 i = 0; i < txSlices.size(); i++) {
36 0 : DeviceMem srcMem = inputMem_.range(txSlices[i].offset, txSlices[i].size);
37 0 : HCCL_DEBUG(
38 : "[ReduceScatterMeshMixSingleStream][RunSourceReducer] send inputmem range[%llu], size[%llu] "
39 : "tx dstmem offset[%llu]",
40 : txSlices[i].offset, txSlices[i].size, dstSlices[i].offset);
41 0 : txMems.emplace_back(SenderMemoryInfo{baseOffset_ + dstSlices[i].offset, srcMem});
42 0 : }
43 :
44 0 : CHK_RET(senderInfo_->run(link, txMems, stream_));
45 0 : return HCCL_SUCCESS;
46 0 : }
47 :
48 0 : HcclResult ReduceScatterMeshMixSingleStream::RunDestReducer(
49 : const LINK& link, const std::vector<Slice>& rxSlices, const std::vector<Slice>& dstSlices)
50 : {
51 : // 使用scratchmem接收数据,并同inputmem数据做reduce
52 0 : std::vector<ReducerMemoryInfo> rxReduceMems;
53 0 : for (u64 i = 0; i < rxSlices.size(); i++) {
54 0 : DeviceMem dstMem = inputMem_.range(rxSlices[i].offset, rxSlices[i].size);
55 0 : DeviceMem srcMem = scratchMem_.range(dstSlices[i].offset, dstSlices[i].size);
56 0 : HCCL_DEBUG(
57 : "[ReduceScatterMeshMixSingleStream][RunDestReducer] rcv offset[%llu], size[%llu] ,then reduce with "
58 : "offset[%llu] size[%llu] ",
59 : dstSlices[i].offset, dstSlices[i].size, rxSlices[i].offset, rxSlices[i].size);
60 0 : rxReduceMems.emplace_back(ReducerMemoryInfo{baseOffset_ + rxSlices[i].offset, dstMem, dstMem, srcMem});
61 0 : }
62 :
63 0 : CHK_RET(reducerInfo_->run(dispatcher_, link, rxReduceMems, stream_));
64 0 : return HCCL_SUCCESS;
65 0 : }
66 :
67 0 : HcclResult ReduceScatterMeshMixSingleStream::RunReduceScatter(
68 : const u32 rank, const u32 rankSize, const std::vector<LINK>& links, const std::vector<Slice>& inputSlices,
69 : const std::vector<Slice>& scratchSlices)
70 : {
71 0 : u32 interRankSize = inputSlices.size() / rankSize;
72 0 : std::vector<u32> txRankOpOrder;
73 0 : std::vector<u32> rxRankOpOrder;
74 : // 计算默认的每轮接收的源rank和发送的目的rank
75 0 : for (u32 round = 1; round < rankSize; round++) {
76 0 : u32 srcRank = ForwardRank(rank, rankSize, round);
77 0 : u32 dstRank = BackwardRank(rank, rankSize, round);
78 0 : HCCL_INFO("<multiDie>RunReduceScatter:srcRank[%u] dstRank[%u]", srcRank, dstRank);
79 0 : rxRankOpOrder.push_back(srcRank);
80 0 : txRankOpOrder.push_back(dstRank);
81 : }
82 :
83 0 : HcclResult ret = HCCL_SUCCESS;
84 0 : for (u32 round = 1; round < rankSize; round++) {
85 : // 不同的stream依次轮训默认的顺序数组
86 0 : u32 orderIndex = (round + streamIndex_ - 1) % (rankSize - 1);
87 0 : u32 srcRank = rxRankOpOrder[orderIndex];
88 0 : s32 dstRank = txRankOpOrder[orderIndex];
89 0 : CHK_SMART_PTR_NULL(links[srcRank]);
90 0 : HCCL_INFO("rank[%u] will tx_ack to rank[%u]", rank, srcRank);
91 0 : ret = links[srcRank]->TxAck(stream_);
92 0 : CHK_PRT_RET(
93 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ReduceScatter]rank[%u] tx ack to rank[%u] failed", rank, srcRank),
94 : ret);
95 0 : CHK_SMART_PTR_NULL(links[dstRank]);
96 0 : HCCL_INFO("rank[%u] will rx_ack from rank[%d]", rank, dstRank);
97 :
98 0 : ret = links[dstRank]->RxAck(stream_);
99 0 : CHK_PRT_RET(
100 : ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ReduceScatter]rank[%u] rx ack from rank[%d] failed", rank, dstRank),
101 : ret);
102 0 : HCCL_INFO(
103 : "rank:%u round[%u] send to rank:[%d], inputSlices offset[%llu]"
104 : "size[%llu] scratchSlice offset[%llu] size[%llu] ",
105 : rank, round, dstRank, inputSlices[dstRank].offset, inputSlices[dstRank].size, scratchSlices[dstRank].offset,
106 : scratchSlices[dstRank].size);
107 : // 发送数据
108 0 : std::vector<Slice> txSlices;
109 0 : std::vector<Slice> txDstSlices;
110 0 : for (u32 i = dstRank * interRankSize; i < (dstRank + 1) * interRankSize; i++) {
111 0 : txSlices.push_back(inputSlices[i]);
112 0 : txDstSlices.push_back(scratchSlices[i]);
113 : }
114 0 : ret = RunSourceReducer(links[dstRank], txSlices, txDstSlices);
115 0 : CHK_PRT_RET(
116 : ret != HCCL_SUCCESS,
117 : HCCL_ERROR("[Run][ReduceScatter]rank:%u round[%u] reducer src run failed", rank, round), ret);
118 0 : HCCL_INFO(
119 : "rank[%u] round[%u] rx from rank[%u], inSlicesoffset[%llu] size[%llu] "
120 : "scratchSlices offset[%llu] size[%llu]",
121 : rank, round, srcRank, inputSlices[rank].offset, inputSlices[rank].size, scratchSlices[rank].offset,
122 : scratchSlices[rank].size);
123 :
124 0 : std::vector<Slice> rxSlices;
125 0 : std::vector<Slice> rxDstSlices;
126 0 : for (u32 i = rank * interRankSize; i < (rank + 1) * interRankSize; i++) {
127 0 : rxSlices.push_back(inputSlices[i]);
128 0 : rxDstSlices.push_back(scratchSlices[i]);
129 : }
130 0 : ret = RunDestReducer(links[srcRank], rxSlices, rxDstSlices);
131 0 : CHK_PRT_RET(
132 : ret != HCCL_SUCCESS,
133 : HCCL_ERROR("[Run][ReduceScatter]rank[%u] round[%u] reducer dst run failed", rank, round), ret);
134 :
135 0 : ret = links[srcRank]->RxWaitDone(stream_);
136 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ReduceScatter]RxWaitDone failed"), ret);
137 0 : ret = links[dstRank]->TxWaitDone(stream_);
138 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][ReduceScatter]TxWaitDone failed"), ret);
139 0 : }
140 0 : if (barrierSwitchOn_) {
141 0 : for (u32 round = 1; round < rankSize; round++) {
142 0 : u32 orderIndex = (round + streamIndex_ - 1) % (rankSize - 1);
143 0 : u32 srcRank = rxRankOpOrder[orderIndex];
144 0 : s32 dstRank = txRankOpOrder[orderIndex];
145 :
146 0 : ret = ExecuteBarrier(links[srcRank], links[dstRank]);
147 0 : CHK_PRT_RET(
148 : ret != HCCL_SUCCESS,
149 : HCCL_ERROR(
150 : "[Run][ReduceScatter]rank[%u] run ReduceScatter executor barrier "
151 : "failed. srcRank:%u dstRank:%d",
152 : rank, srcRank, dstRank),
153 : ret);
154 : }
155 : }
156 :
157 0 : return HCCL_SUCCESS;
158 0 : }
159 :
160 : // reducescatter的入口函数
161 : HcclResult
162 0 : ReduceScatterMeshMixSingleStream::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK>& links)
163 : {
164 0 : CHK_SMART_PTR_NULL(dispatcher_);
165 0 : CHK_PTR_NULL(stream_.ptr());
166 0 : if (!outputMem_ || !inputMem_) {
167 0 : HCCL_ERROR(
168 : "[ReduceScatterMeshMixSingleStream][RunAsync]rank[%u] run_async inputmem or outputmem is null", rank);
169 0 : return HCCL_E_PTR;
170 : }
171 0 : HCCL_INFO(
172 : "[ReduceScatterMeshMixSingleStream][RunAsync]rank[%u] totalrank[%u] inputMem[%p] outputMem[%p] "
173 : "count[%llu]",
174 : rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
175 :
176 : // 创建reducer & sender
177 0 : senderInfo_.reset(new (std::nothrow) Sender(dataType_, reductionOp_, reduceAttr_));
178 0 : CHK_SMART_PTR_NULL(senderInfo_);
179 :
180 0 : reducerInfo_.reset(new (std::nothrow) Reducer(dataType_, reductionOp_, reduceAttr_));
181 0 : CHK_SMART_PTR_NULL(reducerInfo_);
182 0 : if (rankSize == 1) {
183 0 : if (inputMem_ != outputMem_) {
184 0 : return HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_);
185 : }
186 0 : return HCCL_SUCCESS;
187 : }
188 :
189 0 : if (links.size() < rankSize) {
190 0 : HCCL_ERROR("[ReduceScatterMeshMixSingleStream][RunAsync]rank[%u] linksize error", rank);
191 0 : return HCCL_E_INTERNAL;
192 : }
193 :
194 0 : if (streamIndex_ >= rankSize - 1) {
195 0 : HCCL_ERROR(
196 : "[ReduceScatterMeshMixSingleStream][RunAsync]rank[%u] stream index[%u] is out of range when ranksize[%u]",
197 : rank, streamIndex_, rankSize);
198 0 : return HCCL_E_INTERNAL;
199 : }
200 :
201 0 : u32 unitSize = DataUnitSize(dataType_);
202 0 : if (unitSize == 0) {
203 0 : HCCL_ERROR("[ReduceScatterMeshMixSingleStream][RunAsync]rank[%u] unit data size is zero", rank);
204 0 : return HCCL_E_INTERNAL;
205 : }
206 :
207 0 : std::vector<Slice> scratchSlices(slices_);
208 :
209 : // 运行reduce-scatter, mesh算法
210 0 : CHK_RET(RunReduceScatter(rank, rankSize, links, slices_, scratchSlices));
211 :
212 0 : HCCL_INFO("ReduceScatterMeshMixSingleStream finished: rank[%u]", rank);
213 0 : return HCCL_SUCCESS;
214 0 : }
215 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_REDUCESCATTER_MESH_MIX_SS, ReduceScatterMeshMixSingleStream);
216 : } // namespace hccl
|