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 "coll_broadcast_mesh_executor.h"
12 :
13 : namespace hccl {
14 :
15 1 : CollBroadcastMeshExecutor::CollBroadcastMeshExecutor(
16 1 : const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
17 1 : : CollBroadcastExecutor(dispatcher, topoMatcher)
18 1 : {}
19 :
20 0 : HcclResult CollBroadcastMeshExecutor::CalcStreamNum(u32& streamNum)
21 : {
22 0 : u32 totalStreamNum = 0U;
23 0 : if (algType_.algoLevel0 == AlgTypeLevel0::ALG_LEVEL0_4P_MESH) {
24 0 : totalStreamNum = LEVEL0_PLANE_NUM_IN_4PMESH;
25 : } else {
26 0 : if ((workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE)
27 0 : && (topoAttr_.deviceType == DevType::DEV_TYPE_910B) && topoAttr_.isSingleMeshAggregation) {
28 0 : totalStreamNum = topoAttr_.deviceNumPerAggregation;
29 : } else {
30 0 : totalStreamNum = topoAttr_.deviceNumPerAggregation - 1;
31 : }
32 : }
33 0 : streamNum = totalStreamNum > 0 ? totalStreamNum - 1 : 0;
34 :
35 0 : HCCL_INFO("[CollBroadcastMeshExecutor][CalcStreamNum] tag[%s] streamNum_[%u]", tag_.c_str(), streamNum);
36 0 : return HCCL_SUCCESS;
37 : }
38 :
39 0 : HcclResult CollBroadcastMeshExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
40 : {
41 0 : TransportMemType inputType = TransportMemType::RESERVED;
42 0 : TransportMemType outputType = TransportMemType::RESERVED;
43 0 : CHK_RET(CalcTransportMemType(inputType, outputType));
44 0 : CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
45 0 : CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
46 0 : return HCCL_SUCCESS;
47 : }
48 :
49 0 : HcclResult CollBroadcastMeshExecutor::CalcLevel0CommInfo(
50 : TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
51 : {
52 0 : CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_MESH);
53 0 : commParaLevel0.meshSinglePlane = true;
54 0 : CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
55 0 : return HCCL_SUCCESS;
56 0 : }
57 :
58 0 : HcclResult CollBroadcastMeshExecutor::KernelRun(const OpParam& param, ExecMem& execMem)
59 : {
60 0 : HCCL_CONFIG_INFO(HCCL_ALG, "[CollBroadcastMeshExecutor][KernelRun] userRank[%u] starts.", topoAttr_.userRank);
61 0 : u32 perDataSize = SIZE_TABLE[param.DataDes.dataType];
62 :
63 0 : bool isUsedRegister = false;
64 0 : std::unique_ptr<AlgTemplateBase> level0TempAlg1;
65 0 : std::unique_ptr<AlgTemplateBase> level1TempAlg;
66 0 : std::unique_ptr<AlgTemplateBase> level0TempAlg2;
67 :
68 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
69 0 : u32 commIndex = level0CommInfo.localRank;
70 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
71 :
72 0 : level0TempAlg1 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_SCATTER_MESH, dispatcher_);
73 0 : CHK_SMART_PTR_NULL(level0TempAlg1);
74 0 : CHK_RET(level0TempAlg1->Prepare(level0CommInfo.localRank, level0CommInfo.localRankSize));
75 0 : level0TempAlg1->CloseBarrier();
76 :
77 : /* 内层topo:all_reduce */
78 : /* 外层所有rank均参与内层的broadcast计算,所以此处对rank不作限制,但是每个rank需找到自己所在的内层通信域 */
79 0 : std::vector<Slice> slice;
80 0 : CHK_RET(GetRankSliceSize(param.DataDes.dataType, execMem.count, level0CommInfo.localRankSize, slice));
81 :
82 0 : CHK_PRT_RET(
83 : slice.empty(), HCCL_ERROR("[BroadCastOperator][BroadCastMeshExecutor]got slice is empty"), HCCL_E_INTERNAL);
84 :
85 0 : SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
86 :
87 0 : u64 curSize = execMem.count * SIZE_TABLE[param.DataDes.dataType];
88 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
89 0 : HCCL_DEBUG(
90 : "broadcast mesh: curSize[%llu] deviceNumPerAggregation[%u] commLevel0Size[%u]", curSize,
91 : topoAttr_.deviceNumPerAggregation, level0CommInfo.localRankSize);
92 0 : if (curSize / topoAttr_.deviceNumPerAggregation <= NHR_BCAST_SMALL_SIZE) {
93 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
94 0 : TemplateType::TEMPLATE_BROADCAST_NHR_ONESHOT, dispatcher_);
95 : } else {
96 : level1TempAlg
97 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
98 : }
99 0 : HCCL_INFO("broadcast mesh: using nhr algo inter-server.");
100 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
101 0 : isUsedRegister = true;
102 : level1TempAlg
103 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR_V1, dispatcher_);
104 0 : HCCL_INFO("broadcast mesh: using nhr_v1 algo inter-server.");
105 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
106 0 : const u32 level1RankSize = level1CommInfo.localRankSize;
107 0 : HCCL_DEBUG("[CollBroadcastMeshExecutor][KernelRun]level1RankSize is %u", level1RankSize);
108 0 : if (ShouldUseBinaryBroadcastOfNB(
109 0 : curSize / topoAttr_.deviceNumPerAggregation, level1RankSize, topoAttr_.userRankSize,
110 0 : topoAttr_.deviceNumPerAggregation)) {
111 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
112 0 : TemplateType::TEMPLATE_BROADCAST_NB_BINARY, dispatcher_);
113 : } else {
114 : level1TempAlg
115 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
116 : }
117 0 : HCCL_INFO("broadcast mesh: using nonuniform-bruck algo inter-server.");
118 : } else {
119 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
120 0 : TemplateType::TEMPLATE_BROADCAST_RECURSIVE_HD, dispatcher_);
121 0 : HCCL_INFO("broadcast mesh: using Recursive halving-doubling algo inter-server.");
122 : }
123 0 : CHK_SMART_PTR_NULL(level1TempAlg);
124 :
125 : /* 外层topo:all_gather */
126 0 : if (topoAttr_.deviceType == DevType::DEV_TYPE_910B) {
127 0 : level0TempAlg2 = AlgTemplateRegistry::Instance().GetAlgTemplate(
128 0 : TemplateType::TEMPLATE_ALL_GATHER_MESH_ATOMIC, dispatcher_);
129 : } else {
130 : level0TempAlg2
131 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_MESH, dispatcher_);
132 : }
133 0 : CHK_SMART_PTR_NULL(level0TempAlg2);
134 0 : CHK_RET(level0TempAlg2->Prepare(
135 : algResResp_->slaveStreams, algResResp_->notifiesMain, algResResp_->notifiesAux, topoAttr_.userRank, nullptr,
136 : level0CommInfo.localRank, level0CommInfo.localRankSize));
137 :
138 : /* 节点内执行器 stage0 */
139 0 : u32 rootRank = 0;
140 0 : HcclResult ret = GetRankByUserRank(COMM_LEVEL0, COMM_INDEX_0, param.root, rootRank);
141 0 : CHK_PRT_RET(
142 : ret != HCCL_SUCCESS,
143 : HCCL_ERROR("[BroadCastOperator][BroadCastMeshExecutor]invalid root[%u] to get userrank", param.root), ret);
144 :
145 0 : if (ret == HCCL_SUCCESS) {
146 0 : CHK_RET(level0TempAlg1->Prepare(
147 : execMem.inputMem, execMem.outputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, param.stream,
148 : HCCL_REDUCE_RESERVED, rootRank, slice));
149 :
150 0 : u32 rankSize = level0CommInfo.localRankSize;
151 0 : CHK_RET(level0TempAlg1->RegisterProfiler(
152 : (0 << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
153 : + level0CommInfo.localRank,
154 : PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, param.stream));
155 :
156 0 : CHK_RET(RunTemplate(level0TempAlg1, level0CommInfo));
157 : } else {
158 0 : HCCL_ERROR("[BroadCastOperator][BroadCastMeshExecutor]invalid root[%u] to get userrank", param.root);
159 : }
160 0 : HCCL_INFO("[BroadCastOperator][BroadCastMeshExecutor] stage0 run success");
161 0 : u64 hdCount = slice[level0CommInfo.localRank].size / perDataSize;
162 : /* 节点间执行器 stage1 */
163 :
164 0 : u32 subUserrankRoot = topoMatcher_->GetSubRootUserRank(topoAttr_.userRank, param.root);
165 0 : CHK_PRT_RET(
166 : subUserrankRoot == INVALID_VALUE_RANKID,
167 : HCCL_ERROR(
168 : "[BroadCastOperator][BroadCastMeshExecutor]subUserrankRoot[%u] is invalid,userRank[%u],root[%u]",
169 : subUserrankRoot, topoAttr_.userRank, param.root),
170 : HCCL_E_INTERNAL);
171 :
172 0 : u32 subRoot = 0;
173 0 : CHK_RET(GetRankByUserRank(COMM_LEVEL1, commIndex, subUserrankRoot, subRoot));
174 :
175 : // 增加偏移参数
176 0 : if (isUsedRegister) {
177 0 : PrepareData prepareData;
178 0 : prepareData.inputMem = execMem.inputMem;
179 0 : prepareData.outputMem = execMem.outputMem;
180 0 : prepareData.scratchMem = execMem.outputMem;
181 0 : prepareData.count = hdCount;
182 0 : prepareData.dataType = param.DataDes.dataType;
183 0 : prepareData.stream = param.stream;
184 0 : prepareData.reductionOp = HCCL_REDUCE_RESERVED;
185 0 : prepareData.root = subRoot;
186 0 : prepareData.baseOffset = slice[level0CommInfo.localRank].offset;
187 0 : CHK_RET(level1TempAlg->Prepare(prepareData));
188 0 : } else {
189 0 : CHK_RET(level1TempAlg->Prepare(
190 : execMem.inputMem, execMem.outputMem, execMem.outputMem, hdCount, param.DataDes.dataType, param.stream,
191 : HCCL_REDUCE_RESERVED, subRoot, std::vector<Slice>(0), slice[level0CommInfo.localRank].offset));
192 : }
193 :
194 0 : u32 rankSize = level1CommInfo.localRankSize;
195 0 : CHK_RET(level1TempAlg->RegisterProfiler(
196 : (0 << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
197 : + level1CommInfo.localRank,
198 : PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET, param.stream));
199 :
200 0 : CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
201 0 : HCCL_INFO("[BroadCastOperator][BroadCastMeshExecutor] stage1 run success");
202 :
203 : /* 节点内执行器 stage2 */
204 : {
205 0 : if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB) { // offline
206 0 : for (u32 streamIndex = 0; streamIndex < algResResp_->slaveStreams.size(); streamIndex++) {
207 0 : CHK_RET(StreamActiveManager::GetInstance(topoAttr_.deviceLogicId)
208 : .StreamActive(algResResp_->slaveStreams[streamIndex].ptr(), param.stream.ptr()));
209 : }
210 : }
211 0 : CHK_RET(level0TempAlg2->Prepare(
212 : execMem.outputMem, execMem.outputMem, execMem.outputMem, execMem.count, param.DataDes.dataType,
213 : param.stream, HCCL_REDUCE_RESERVED, LEVEL0_BRIDGE_RANK_ID, slice));
214 :
215 0 : u32 rankSize = level0CommInfo.localRankSize;
216 0 : CHK_RET(level0TempAlg2->RegisterProfiler(
217 : (0 << PROF_RINGINDEX_OFFSET_OF_PLANEID) + (rankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID)
218 : + level0CommInfo.localRank,
219 : PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
220 :
221 0 : CHK_RET(RunTemplate(level0TempAlg2, level0CommInfo));
222 : }
223 :
224 0 : HCCL_INFO("[BroadCastOperator][BroadCastMeshExecutor] stage2 run success");
225 0 : return HCCL_SUCCESS;
226 0 : }
227 0 : HcclResult CollBroadcastMeshExecutor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
228 : {
229 0 : CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
230 0 : SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
231 :
232 0 : u32 commIndex = level0CommInfo.localRank; // 找到rank所在的节点间平面
233 0 : CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
234 :
235 0 : level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
236 0 : return HCCL_SUCCESS;
237 0 : }
238 :
239 0 : HcclResult CollBroadcastMeshExecutor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
240 : {
241 0 : if (level1RankSize > 1) {
242 0 : if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
243 : level1TempAlg
244 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR, dispatcher_);
245 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
246 : level1TempAlg
247 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NHR_V1, dispatcher_);
248 0 : HCCL_INFO("broadcast mesh: using nhr_v1 algo inter-server.");
249 0 : } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
250 : level1TempAlg
251 0 : = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_BROADCAST_NB, dispatcher_);
252 0 : HCCL_INFO("broadcast mesh: using nonuniform-bruck algo inter-server.");
253 : } else {
254 0 : level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
255 0 : TemplateType::TEMPLATE_BROADCAST_RECURSIVE_HD, dispatcher_);
256 0 : HCCL_INFO("broadcast mesh: using Recursive halving-doubling algo inter-server.");
257 : }
258 0 : CHK_SMART_PTR_NULL(level1TempAlg);
259 : }
260 0 : return HCCL_SUCCESS;
261 : }
262 : REGISTER_EXEC("BroadCastMeshExecutor", BroadcastMesh, CollBroadcastMeshExecutor);
263 :
264 : } // namespace hccl
|