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_nhr_oneshot.h"
12 : #include <cmath>
13 : #include "alg_template_register.h"
14 :
15 : namespace hccl {
16 3 : BroadcastNHROneshot::BroadcastNHROneshot(const HcclDispatcher dispatcher)
17 3 : : NHRBase(dispatcher), localBaseOffset_(0), isForAllReduce_(false)
18 : {
19 3 : }
20 :
21 3 : BroadcastNHROneshot::~BroadcastNHROneshot()
22 : {
23 3 : }
24 :
25 3 : HcclResult BroadcastNHROneshot::RunAsync(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
26 : {
27 : // 基本的检查
28 3 : CHK_RET(SimpleCheck(rank, rankSize, links));
29 3 : HCCL_INFO("[BroadcastNHROneshot][RunAsync] rank[%u] ranksize[%u] inputMem[%p] outputMem[%p] count[%llu]",
30 : rank, rankSize, inputMem_.ptr(), outputMem_.ptr(), count_);
31 :
32 3 : u32 unitSize = DataUnitSize(dataType_);
33 3 : CHK_PRT_RET(unitSize == 0, HCCL_ERROR("[BroadcastNHROneshot][RunAsync] unitSize is zero"), HCCL_E_INTERNAL);
34 :
35 3 : if (!isForAllReduce_) {
36 0 : localBaseOffset_ = baseOffset_; // broadcast的本地偏移量和baseOffset_一致
37 : }
38 :
39 : // 双buffer下, 先将input拷贝到output的合适位置
40 3 : if (inputMem_ != outputMem_ && rank == root_) {
41 3 : u64 totalSize = count_ * SIZE_TABLE[dataType_];
42 3 : DeviceMem src = inputMem_.range(localBaseOffset_, totalSize);
43 3 : DeviceMem dst = outputMem_.range(localBaseOffset_, totalSize);
44 3 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dst, src, stream_));
45 3 : }
46 :
47 : // 如果ranksize为1, 从input->output就结束
48 3 : if (rankSize == 1) {
49 0 : return HCCL_SUCCESS;
50 : }
51 :
52 : // 运行bcast, rst算法
53 3 : CHK_RET(RunBroadcastNHROneshot(rank, rankSize, links));
54 :
55 3 : HCCL_INFO("[BroadcastNHROneshot][RunAsync] finished: rank[%u] ranksize[%u]", rank, rankSize);
56 3 : return HCCL_SUCCESS;
57 : }
58 :
59 3 : HcclResult BroadcastNHROneshot::RunAsyncForAllReduce(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
60 : {
61 3 : isForAllReduce_ = true;
62 3 : return RunAsync(rank, rankSize, links);
63 : }
64 :
65 3 : HcclResult BroadcastNHROneshot::SimpleCheck(const u32 rank, const u32 rankSize, const std::vector<LINK> &links)
66 : {
67 : // 判断stream, dispatcher是否为空
68 3 : CHK_SMART_PTR_NULL(dispatcher_);
69 3 : CHK_PTR_NULL(stream_.ptr());
70 :
71 : // 检查memory
72 3 : CHK_PRT_RET(!outputMem_ || !inputMem_,
73 : HCCL_ERROR("[BroadcastNHROneshot][SimpleCheck] rank[%u] inputmem or outputmem is null", rank), HCCL_E_PTR);
74 :
75 : // 判断links数量是否正确
76 3 : CHK_PRT_RET(links.size() < rankSize, HCCL_ERROR("[BroadcastNHROneshot][SimpleCheck] rank[%u] link size[%llu] is "
77 : "less than rank size[%u]", rank, links.size(), rankSize), HCCL_E_INTERNAL);
78 3 : return HCCL_SUCCESS;
79 : }
80 :
81 0 : HcclResult BroadcastNHROneshot::SdmaRx(LINK &linkLeft, LINK &linkRight, InterServerAlgoStep &stepInfo,
82 : const std::vector<LINK> &links)
83 : {
84 0 : u64 totalSize = count_ * SIZE_TABLE[dataType_];
85 0 : DeviceMem srcMem = outputMem_.range(localBaseOffset_, totalSize);
86 :
87 0 : if (linkRight != nullptr) {
88 0 : CHK_RET(linkRight->TxAck(stream_));
89 : }
90 0 : if (linkLeft != nullptr) {
91 0 : CHK_RET(linkLeft->RxAck(stream_));
92 0 : void *srcMemPtr = nullptr;
93 0 : CHK_RET(linkLeft->GetRemoteMem(UserMemType::OUTPUT_MEM, &srcMemPtr));
94 0 : DeviceMem srcMemLeft(static_cast<s8 *>(srcMemPtr) + baseOffset_, totalSize);
95 0 : CHK_RET(HcclD2DMemcpyAsync(dispatcher_, srcMem, srcMemLeft, stream_, linkLeft->GetRemoteRank(), // Memecpy
96 : linkLeft->GetLinkType()));
97 0 : CHK_RET(linkLeft->TxDataSignal(stream_));
98 0 : }
99 0 : if (linkRight != nullptr) {
100 0 : CHK_RET(linkRight->RxDataSignal(stream_));
101 : }
102 0 : return HCCL_SUCCESS;
103 0 : }
104 :
105 9 : HcclResult BroadcastNHROneshot::RdmaTxRx(LINK &linkLeft, LINK &linkRight, InterServerAlgoStep &stepInfo,
106 : const std::vector<LINK> &links)
107 : {
108 9 : u64 totalSize = count_ * SIZE_TABLE[dataType_];
109 9 : DeviceMem srcMem = outputMem_.range(localBaseOffset_, totalSize);
110 :
111 9 : if (linkLeft != nullptr) {
112 0 : CHK_RET(linkLeft->TxAck(stream_));
113 : }
114 :
115 9 : if (linkRight != nullptr) {
116 9 : CHK_RET(linkRight->RxAck(stream_));
117 9 : CHK_RET(linkRight->TxAsync(UserMemType::OUTPUT_MEM, baseOffset_, srcMem.ptr(), srcMem.size(), stream_));
118 9 : CHK_RET(linkRight->WaitFinAck(stream_));
119 : }
120 :
121 9 : if (linkLeft != nullptr) {
122 0 : CHK_RET(linkLeft->RxAsync(UserMemType::OUTPUT_MEM, baseOffset_, srcMem.ptr(), srcMem.size(), stream_));
123 0 : CHK_RET(linkLeft->PostFinAck(stream_));
124 : }
125 9 : return HCCL_SUCCESS;
126 9 : }
127 :
128 3 : HcclResult BroadcastNHROneshot::RunBroadcastNHROneshot(u32 rank, u32 rankSize, const std::vector<LINK> &links)
129 : {
130 : // 计算通信步数
131 3 : u32 nSteps = GetStepNumInterServer(rankSize);
132 3 : HCCL_DEBUG("[BroadcastNHROneshot][RunBroadcastNHROneshot] rank[%u] rankSize[%u] nSteps[%u]",
133 : rank, rankSize, nSteps);
134 :
135 : // 逐步编排任务
136 12 : for (u32 step = 0; step < nSteps; step++) {
137 9 : InterServerAlgoStep stepInfo;
138 9 : GetStepInfo(step, nSteps, rank, rankSize, stepInfo);
139 :
140 9 : HCCL_DEBUG("[BroadcastNHROneshot][RunBroadcastNHROneshot] recvFrom[%u] sendTo[%u] step[%u]",
141 : stepInfo.fromRank, stepInfo.toRank, step);
142 :
143 9 : LINK linkLeft;
144 9 : LINK linkRight;
145 9 : if (stepInfo.txSliceIdxs.size() > 0) {
146 9 : linkRight = links[stepInfo.toRank];
147 9 : CHK_SMART_PTR_NULL(linkRight);
148 : }
149 9 : if (stepInfo.rxSliceIdxs.size() > 0) {
150 0 : linkLeft = links[stepInfo.fromRank];
151 0 : CHK_SMART_PTR_NULL(linkLeft);
152 : }
153 :
154 18 : if ((linkRight != nullptr && linkRight->IsSpInlineReduce()) ||
155 9 : (linkLeft != nullptr && linkLeft->IsSpInlineReduce())) {
156 0 : CHK_RET(SdmaRx(linkLeft, linkRight, stepInfo, links));
157 : } else {
158 9 : CHK_RET(RdmaTxRx(linkLeft, linkRight, stepInfo, links));
159 : }
160 9 : }
161 3 : return HCCL_SUCCESS;
162 : }
163 :
164 : // NHR每步的算法描述原理函数
165 9 : HcclResult BroadcastNHROneshot::GetStepInfo(u32 step, u32 nSteps, u32 rank, u32 rankSize, InterServerAlgoStep &stepInfo)
166 : {
167 9 : stepInfo.txSliceIdxs.clear();
168 9 : stepInfo.rxSliceIdxs.clear();
169 9 : stepInfo.nSlices = 1;
170 9 : stepInfo.toRank = rankSize;
171 9 : stepInfo.fromRank = rankSize;
172 9 : stepInfo.step = step;
173 9 : stepInfo.myRank = rank;
174 :
175 9 : u32 nRanks = (rankSize - 1 + (1 << (nSteps - 1 - step))) / (1 << (nSteps - step)); // 本步需要进行收/发的rank数
176 :
177 9 : u32 deltaRoot = (rank + rankSize - root_) % rankSize;
178 :
179 9 : u32 deltaRankPair = 1 << (nSteps - 1 - step);
180 9 : u32 deltaRankGroup = 1 << (nSteps - step);
181 :
182 9 : if (deltaRoot / deltaRankGroup < nRanks) {
183 9 : if (deltaRoot % deltaRankGroup == 0) {
184 9 : stepInfo.toRank = (rank + deltaRankPair) % rankSize;
185 9 : stepInfo.txSliceIdxs.push_back(0);
186 : }
187 :
188 9 : if ((deltaRoot + deltaRankPair) % deltaRankGroup == 0) {
189 0 : stepInfo.fromRank = (rank + rankSize - deltaRankPair) % rankSize;
190 0 : stepInfo.rxSliceIdxs.push_back(0);
191 : }
192 : }
193 9 : return HCCL_SUCCESS;
194 : }
195 :
196 : REGISTER_TEMPLATE(TemplateType::TEMPLATE_BROADCAST_NHR_ONESHOT, BroadcastNHROneshot);
197 : } // namespace hccl
|