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