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 "all_gather_nhr.h"
12 : #include <cmath>
13 : #include "alg_template_register.h"
14 :
15 : namespace hccl {
16 0 : AllGatherNHR::AllGatherNHR(const HcclDispatcher dispatcher) : NHRBase(dispatcher)
17 : {
18 0 : }
19 :
20 0 : AllGatherNHR::~AllGatherNHR()
21 : {
22 0 : }
23 :
24 0 : HcclResult AllGatherNHR::Prepare(bool needMerge)
25 : {
26 0 : isNeedMerge = needMerge;
27 0 : return HCCL_SUCCESS;
28 : }
29 :
30 : // 服务器间allgather的入口函数
31 0 : HcclResult AllGatherNHR::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
32 : {
33 : // 从AllReduce或者Broadcast调用AllGather需要merge
34 0 : if (isNeedMerge) {
35 : // 获取tree映射,存储到类对象的成员变量中
36 0 : GetRankMapping(rankSize);
37 : }
38 0 : CHK_SMART_PTR_NULL(dispatcher_);
39 0 : CHK_PTR_NULL(stream_.ptr());
40 0 : HCCL_INFO("[AllGatherNHR][RunAsync] rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
41 : rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
42 :
43 0 : if (rankSize == 1) {
44 0 : if (inputMem_ != outputMem_) {
45 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, outputMem_, inputMem_, stream_));
46 : }
47 0 : return HCCL_SUCCESS;
48 : }
49 :
50 0 : CHK_PRT_RET(links.size() < rankSize,
51 : HCCL_ERROR("[AllGatherNHR][RunAsync] rank[%u] linkSize is less than rankSize", rank), HCCL_E_INTERNAL);
52 :
53 0 : u32 unitSize = DataUnitSize(dataType_);
54 0 : CHK_PRT_RET(unitSize == 0, HCCL_ERROR("[AllGatherNHR][RunAsync]rank[%u] unitSize is zero", rank), HCCL_E_INTERNAL);
55 :
56 0 : std::vector<Slice> inputSlices(slices_);
57 0 : if (slices_.size() == 0) {
58 0 : slices_.resize(rankSize);
59 0 : inputSlices.resize(rankSize);
60 :
61 0 : u64 sliceSize = count_ * unitSize;
62 0 : for (u32 i = 0; i < rankSize; i++) {
63 0 : slices_[i].size = sliceSize;
64 0 : slices_[i].offset = sliceSize * i;
65 0 : inputSlices[i].size = sliceSize;
66 0 : inputSlices[i].offset = (inputMem_.size() < outputMem_.size()) ? 0 : (sliceSize * i);
67 0 : HCCL_DEBUG("[AllGatherNHR][RunAsync] rank[%u], slices[%u].offset=%llu, slices[%u].size=%llu",
68 : rank, i, slices_[i].offset, i, slices_[i].size);
69 : }
70 : }
71 :
72 0 : if (sliceMap_.size() != rankSize) {
73 0 : GetRankMapping(rankSize, true); // 没有初始化过,说明不是由allreduce或者bcast调入,需要保序
74 : }
75 :
76 : // 双buffer下, 先将input拷贝到output的合适位置
77 0 : if (inputMem_ != outputMem_) {
78 0 : u32 targetIdx = sliceMap_[rank];
79 0 : DeviceMem dst = outputMem_.range(slices_[targetIdx].offset, slices_[targetIdx].size);
80 0 : DeviceMem src = inputMem_.range(inputSlices[targetIdx].offset, inputSlices[targetIdx].size);
81 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
82 0 : }
83 :
84 : // 运行all-gather, ring算法
85 0 : CHK_RET(RunAllGather(rank, rankSize, slices_, links));
86 :
87 0 : HCCL_INFO("[AllGatherNHR][RunAsync] AllGatherNHR finished: rank[%u] end", rank);
88 0 : return HCCL_SUCCESS;
89 0 : }
90 :
91 0 : HcclResult AllGatherNHR::SdmaRx(const LINK &linkLeft, const LINK &linkRight, std::vector<Slice> &rxSlices)
92 : {
93 0 : CHK_RET(linkRight->TxAck(stream_)); // 告知right可以从我这读了
94 0 : CHK_RET(linkLeft->RxAck(stream_)); // 等left告知可以从他那读了
95 0 : void *srcMemPtr = nullptr;
96 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
97 0 : for (const Slice &rxSlice : rxSlices) {
98 0 : DeviceMem dstMem = outputMem_.range(rxSlice.offset, rxSlice.size);
99 0 : DeviceMem srcMem(static_cast<s8 *>(srcMemPtr) + baseOffset_ + rxSlice.offset, rxSlice.size);
100 0 : HCCL_DEBUG("[AllGatherNHR] rx dstMem[%p] range[%llu], size[%llu] ", dstMem.ptr(),
101 : rxSlice.offset, rxSlice.size);
102 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream_, linkLeft->GetRemoteRank(), // Memecpy
103 : linkLeft->GetLinkType()));
104 0 : }
105 0 : CHK_RET(linkLeft->TxDataSignal(stream_)); // 告知left我读完了
106 0 : CHK_RET(linkRight->RxDataSignal(stream_)); // 等right读完
107 :
108 0 : return HCCL_SUCCESS;
109 : }
110 :
111 0 : HcclResult AllGatherNHR::RdmaTxRx(const LINK &linkLeft, const LINK &linkRight, InterServerAlgoStep &stepInfo,
112 : std::vector<Slice> &txSlices, std::vector<Slice> &rxSlices)
113 : {
114 0 : HcclResult ret = HCCL_SUCCESS;
115 0 : CHK_RET(linkLeft->TxAck(stream_));
116 0 : CHK_RET(linkRight->RxAck(stream_));
117 0 : ret = Tx(linkRight, txSlices);
118 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherNHR][RunAllGather] rank[%u] round[%u] "
119 : "tx %u slices failed", stepInfo.myRank, stepInfo.step, stepInfo.nSlices), ret);
120 0 : ret = Rx(linkLeft, rxSlices);
121 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherNHR][RunAllGather] rank[%u] round[%u] "
122 : "rx %u slices failed", stepInfo.myRank, stepInfo.step, stepInfo.nSlices), ret);
123 :
124 0 : ret = linkLeft->PostFinAck(stream_);
125 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherNHR][RunAllGather] PostFinAck failed"), ret);
126 0 : ret = linkRight->WaitFinAck(stream_);
127 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherNHR][RunAllGather] WaitFinAck failed"), ret);
128 0 : return HCCL_SUCCESS;
129 : }
130 :
131 0 : HcclResult AllGatherNHR::RunAllGather(u32 rank, u32 rankSize, const std::vector<Slice> &outputSlices,
132 : const std::vector<LINK> &links)
133 : {
134 0 : CHK_PRT_RET(outputSlices.size() < rankSize,
135 : HCCL_ERROR("[AllGatherNHR][RunAllGather] rank[%u] OutputSlice Size is less than rank size", rank),
136 : HCCL_E_INTERNAL);
137 0 : HcclResult ret = HCCL_SUCCESS;
138 :
139 : // 计算通信步数
140 0 : u32 nSteps = GetStepNumInterServer(rankSize);
141 :
142 : // 逐步编排任务
143 0 : for (u32 step = 0; step < nSteps; step++) {
144 0 : InterServerAlgoStep stepInfo;
145 0 : GetStepInfo(step, nSteps, rank, rankSize, stepInfo);
146 :
147 0 : LINK linkLeft = links[stepInfo.fromRank];
148 0 : CHK_SMART_PTR_NULL(linkLeft);
149 :
150 0 : LINK linkRight = links[stepInfo.toRank];
151 0 : CHK_SMART_PTR_NULL(linkRight);
152 :
153 0 : std::vector<Slice> txSlices;
154 0 : std::vector<Slice> rxSlices;
155 :
156 0 : HCCL_DEBUG("[AllGatherNHR][RunAllGather] rank[%u] rankSize[%u] recvFrom[%u] sendTo[%u] step[%u] nSteps[%u] "
157 : "nSlices[%u]", rank, rankSize, stepInfo.fromRank, stepInfo.toRank, step, nSteps, stepInfo.nSlices);
158 :
159 0 : for (u32 i = 0; i < stepInfo.nSlices; i++) {
160 0 : txSlices.push_back(outputSlices[stepInfo.txSliceIdxs[i]]);
161 0 : rxSlices.push_back(outputSlices[stepInfo.rxSliceIdxs[i]]);
162 :
163 0 : HCCL_DEBUG("[AllGatherNHR][RunAllGather] i[%u] rxSliceIndex[%u] txSliceIndex[%u] rx data offset[%llu] "
164 : "size[%llu]", i, stepInfo.rxSliceIdxs[i], stepInfo.txSliceIdxs[i],
165 : outputSlices[stepInfo.rxSliceIdxs[i]].offset, outputSlices[stepInfo.rxSliceIdxs[i]].size);
166 : }
167 :
168 : // 合并连续slices
169 0 : MergeSlices(rxSlices);
170 0 : MergeSlices(txSlices);
171 :
172 0 : if (linkLeft->IsSpInlineReduce() && linkRight->IsSpInlineReduce()) {
173 0 : ret = SdmaRx(linkLeft, linkRight, rxSlices);
174 : } else {
175 0 : ret = RdmaTxRx(linkLeft, linkRight, stepInfo, txSlices, rxSlices);
176 : }
177 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[AllGatherNHR][RunAllGather] RunAllGather failed"), ret);
178 0 : }
179 0 : return HCCL_SUCCESS;
180 : }
181 :
182 0 : HcclResult AllGatherNHR::Tx(const LINK &link, std::vector<Slice> &txSlices)
183 : {
184 0 : std::vector<TxMemoryInfo> txMems;
185 0 : for (const Slice &txSlice : txSlices) {
186 0 : DeviceMem srcMem = outputMem_.range(txSlice.offset, txSlice.size);
187 0 : HCCL_DEBUG("[AllGatherNHR][Tx] tx srcMem[%p] offset[%llu] size[%llu]", srcMem.ptr(), txSlice.offset,
188 : txSlice.size);
189 0 : txMems.emplace_back(TxMemoryInfo{UserMemType::OUTPUT_MEM, txSlice.offset + baseOffset_,
190 0 : srcMem.ptr(), txSlice.size});
191 0 : }
192 :
193 0 : CHK_RET(link->TxAsync(txMems, stream_));
194 0 : return HCCL_SUCCESS;
195 0 : }
196 :
197 0 : HcclResult AllGatherNHR::Rx(const LINK &link, std::vector<Slice> &rxSlices)
198 : {
199 0 : std::vector<RxMemoryInfo> rxMems;
200 0 : for (const Slice &rxSlice : rxSlices) {
201 0 : DeviceMem dstMem = outputMem_.range(rxSlice.offset, rxSlice.size);
202 0 : HCCL_DEBUG("[AllGatherNHR][Rx] rx dstMem[%p] range[%llu], size[%llu] ", dstMem.ptr(),
203 : rxSlice.offset, rxSlice.size);
204 0 : rxMems.emplace_back(RxMemoryInfo{UserMemType::OUTPUT_MEM, rxSlice.offset + baseOffset_,
205 0 : dstMem.ptr(), rxSlice.size});
206 0 : }
207 :
208 0 : CHK_RET(link->RxAsync(rxMems, stream_));
209 0 : return HCCL_SUCCESS;
210 0 : }
211 :
212 : // NHR每步的算法描述原理函数
213 0 : HcclResult AllGatherNHR::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep &stepInfo)
214 : {
215 0 : stepInfo.txSliceIdxs.clear();
216 0 : stepInfo.rxSliceIdxs.clear();
217 0 : u32 sliceSize = slices_.size() / rankSize;
218 0 : stepInfo.step = step;
219 0 : stepInfo.myRank = rank;
220 :
221 : // 计算通信对象
222 0 : u32 deltaRank = 1 << (nSteps - 1 - step);
223 0 : u32 recvFrom = (rank + rankSize - deltaRank) % rankSize;
224 0 : u32 sendTo = (rank + deltaRank) % rankSize;
225 :
226 : // 数据份数和数据编号增量
227 0 : u32 nSlices = (rankSize - 1 + (1 << (nSteps - 1 - step))) / (1 << (nSteps - step));
228 0 : u32 deltaSliceIndex = 1 << (nSteps - step);
229 0 : u32 txSliceIdx = rank;
230 0 : u32 rxSliceIdx = (rank - (1 << (nSteps - 1 - step)) + rankSize) % rankSize;
231 :
232 0 : stepInfo.nSlices = nSlices * sliceSize;
233 0 : stepInfo.toRank = sendTo;
234 0 : stepInfo.fromRank = recvFrom;
235 :
236 0 : for (u32 i = 0; i < nSlices; i++) {
237 0 : for (u32 j = 0; j < sliceSize; j++) {
238 0 : u32 targetTxSliceIdx = sliceMap_[txSliceIdx];
239 0 : stepInfo.txSliceIdxs.push_back(targetTxSliceIdx * sliceSize + j);
240 :
241 0 : u32 targetRxSliceIdx = sliceMap_[rxSliceIdx];
242 0 : stepInfo.rxSliceIdxs.push_back(targetRxSliceIdx * sliceSize + j);
243 :
244 0 : HCCL_DEBUG("[AllGatherNHR][GetStepInfo] i[%u] txSliceIdx[%u]->targetTxSliceIdx[%u] rxSliceIdx[%u]->"
245 : "targetRxSliceIdx[%u]", i, txSliceIdx, targetTxSliceIdx, rxSliceIdx, targetRxSliceIdx);
246 : }
247 0 : txSliceIdx = (txSliceIdx + rankSize - deltaSliceIndex) % rankSize;
248 0 : rxSliceIdx = (rxSliceIdx + rankSize - deltaSliceIndex) % rankSize;
249 : }
250 0 : return HCCL_SUCCESS;
251 : }
252 :
253 0 : HcclResult AllGatherNHR::GetNslbAdjInfo(const u32 rank, const u32 rankSize,
254 : const std::vector<LINK> &links, AdjInfo& nslbAdjInfo)
255 : {
256 0 : if (rankSize == 1) {
257 0 : return HCCL_SUCCESS;
258 : }
259 0 : if (links.size() < rankSize) {
260 0 : return HCCL_SUCCESS;
261 : }
262 0 : u32 nSteps = 0;
263 0 : for(u32 temp = rankSize - 1; temp != 0; temp >>= 1, ++nSteps){}
264 :
265 0 : for (u32 step = 0; step < nSteps; step++) {
266 0 : u32 deltaRank = 1 << (nSteps - 1 - step);
267 0 : u32 sendTo =(rank + deltaRank) % rankSize;
268 0 : LINK linkRight = links[sendTo];
269 0 : CHK_SMART_PTR_NULL(linkRight);
270 :
271 0 : NslbDpAdjInfo adjInfoStep = {0};
272 0 : adjInfoStep.dstLocalRankId = linkRight->GetRemoteRank();
273 0 : adjInfoStep.phaseId = step + 1;
274 0 : adjInfoStep.rev = 0;
275 0 : nslbAdjInfo.nsAdjInfo.push_back(adjInfoStep);
276 0 : }
277 0 : nslbAdjInfo.dstRankNum = nSteps;
278 0 : return HCCL_SUCCESS;
279 : }
280 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_ALL_GATHER_NHR, AllGatherNHR);
281 : } // namespace hccl
|