Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / nd / global_config.py: 91%
153 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-04 05:18 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-04 05:18 +0800
1# Copyright 2024 Huawei Technologies Co., Ltd
2#
3# Licensed under the Apache License, Version 2.0 (the "License");
4# you may not use this file except in compliance with the License.
5# You may obtain a copy of the License at
6#
7# http://www.apache.org/licenses/LICENSE-2.0
8#
9# Unless required by applicable law or agreed to in writing, software
10# distributed under the License is distributed on an "AS IS" BASIS,
11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12# See the License for the specific language governing permissions and
13# limitations under the License.
14# ============================================================================
15"""One configuration interface for parallelization"""
17import copy
18from math import gcd
20from hyper_parallel.auto_parallel.sapp_nd.nd.common.arch_hooks import CWrap, check_and_apply_custom_hook
21from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger
22import hyper_parallel.auto_parallel.sapp_nd.nd.dimensions as Dim
23import hyper_parallel.auto_parallel.sapp_nd.nd.common.hardware as Hard
24import hyper_parallel.auto_parallel.sapp_nd.nd.balancing_adapter as BA
27class GlobalConfig:
28 """Union of cost model & parallel config"""
30 def __init__(self, config, dimensions=None, mppb=False):
32 self.wrap = CWrap(config)
33 self.ccfg = self.wrap.ccfg
35 if dimensions is not None:
36 logger.debug("dimensions = %s", str(dimensions))
37 self.dimensions = dimensions
38 else:
39 logger.debug("dimensions = %s", str(Dim.ALL_DIMS))
40 self.dimensions = Dim.ALL_DIMS.copy()
41 logger.debug("self.dimensions = %s", str(self.dimensions))
42 logger.debug("layer_num_for_offset = %d", self.layer_num_for_offset())
43 logger.debug("total layer num = %d", self.total_layer_num())
44 self.balancing = BA.BalancingAdapter(
45 self.layer_num_for_offset(),
46 copy.deepcopy(self.ccfg.offset),
47 copy.deepcopy(self.ccfg.full_rec),
48 mppb,
49 )
51 def dim_val(self, dim, parallel_config):
52 """Get the value of a parallel dimension"""
53 if parallel_config.has_dim(dim):
54 return parallel_config.val(dim)
55 return dim.from_config(self.ccfg)
57 def global_batch_size(self, parallel_config):
58 """Compute global batch size from hyperparameters"""
59 dp = self.dim_val(Dim.DP, parallel_config)
60 pp = self.dim_val(Dim.PP, parallel_config)
61 mb = self.dim_val(Dim.MBN, parallel_config)
62 bs = self.dim_val(Dim.MBS, parallel_config)
63 if pp > 1:
64 logger.info("GBS = %dDP * %dMB * %dBS", dp, mb, bs)
65 return dp * mb * bs
66 logger.info("GBS = %dDP * %dBS", dp, bs)
67 return dp * bs
69 def layer_num_for_offset(self):
70 """Compute layer number including MTP when necessary for offset"""
71 layer_num = self.ccfg.n_lay
72 if self.ccfg.emb_out_in_offset:
73 layer_num += 2
74 if self.ccfg.is_mtp_in_offset:
75 layer_num += self.ccfg.n_mtp
76 return layer_num
78 def total_layer_num(self):
79 """Compute total layer number, always including MTP"""
80 layer_num = self.ccfg.n_lay + self.ccfg.n_mtp
81 return layer_num
83 def adapt_config_balancing(self, new_pp, new_vpp):
84 """Adapt the layer-to-stage assignment to different PP"""
85 logger.debug("new_pp=%d, new_vpp=%d", new_pp, new_vpp)
87 new_recompute_config = self.balancing.treat_recompute(new_pp, new_vpp)
88 logger.debug("adapted recompute config: %s", str(new_recompute_config))
89 new_offset = self.balancing.treat_offset(new_pp, new_vpp)
90 logger.debug("adapted offset: %s", str(new_offset))
91 ok = self.balancing.offset_checker(new_pp, new_vpp, new_offset)
92 if not ok:
93 logger.error("Offset {%s} NOT VALID", str(new_offset))
94 return new_offset, new_recompute_config
96 def adapt_config(self, pp, vpp):
97 """Adapt configuration to different parallel config"""
98 return self.adapt_config_balancing(pp, vpp)
100 def write(self, folder, parallel_config):
101 """Dump config into a yaml file"""
102 if folder:
103 file_name = parallel_config.unique_name()
104 self.ccfg.config.dump(file_name, folder)
106 def moe_valid(self, parallel_config):
107 """Check whether the model is MoE"""
108 expert_num = self.ccfg.n_exp
109 if expert_num > 1:
110 ep = self.dim_val(Dim.EP, parallel_config)
111 dp = self.dim_val(Dim.DP, parallel_config)
112 mp = self.dim_val(Dim.TP, parallel_config)
113 logger.debug(
114 "moe valid ? EP %d <= E %d & EP %d <= DP %d * MP %d",
115 ep,
116 expert_num,
117 ep,
118 dp,
119 mp,
120 )
121 return ep <= min(expert_num, dp * mp)
122 return True
124 def ep_constraints_valid(self, parallel_config):
125 """Check EP-specific divisibility constraints (C1, C2).
127 Runs only for MoE models (n_exp > 1). C1 ensures experts can be
128 evenly partitioned across EP ranks; C2 ensures the expert FFN hidden
129 dim can be evenly sharded by the expert TP degree. Both checks use
130 architecture constants from ``self.ccfg`` and the candidate values
131 from ``parallel_config``.
133 C3 (device count) is intentionally skipped here because the search
134 loop borrows EP from the dp*tp budget, so dp*tp*pp*cp already
135 equals total_devices by construction.
137 Args:
138 parallel_config: candidate ``Dim.Dimensions``.
140 Returns:
141 bool: True if all applicable EP constraints pass (or the model
142 is dense), False otherwise.
143 """
144 if self.ccfg.n_exp <= 1:
145 return True
146 # pylint: disable=C0415
147 from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.validators.ep_constraints import EpConstraints
148 ep = self.dim_val(Dim.EP, parallel_config)
149 r1 = EpConstraints.check_ep_divisibility(self.ccfg.n_exp, ep)
150 if not r1:
151 logger.warning("EP constraint C1 failed: %s", r1.message)
152 return False
153 tp = self.dim_val(Dim.TP, parallel_config)
154 etp = max(getattr(self.ccfg, "etp", 0), 0)
155 t_exp = max(etp, 1) if etp > 1 else max(tp, 1)
156 hff_exp = max(getattr(self.ccfg, "hff_exp", 0), 0)
157 r2 = EpConstraints.check_expert_hidden_divisibility(hff_exp, t_exp)
158 if not r2:
159 logger.warning("EP constraint C2 failed: %s", r2.message)
160 return False
161 return True
163 def make_parallel_config_args(self, **kwargs):
164 """Create a parallel config from parallel values"""
165 logger.debug("dimensions considered: %s", str(self.dimensions))
167 dims = []
168 # dims.append((Dim.DP, dp))
169 for dim in self.dimensions:
170 dims.append((dim, kwargs.get(dim.lname())))
172 has_mbn_not_in = Dim.MBN not in self.dimensions
173 has_pp_in = Dim.PP in self.dimensions
174 has_dp_or_mbs_in = Dim.DP in self.dimensions or Dim.MBS in self.dimensions
175 if has_mbn_not_in and has_pp_in and has_dp_or_mbs_in:
176 dims.append((Dim.MBN, kwargs.get(Dim.MBN.lname())))
177 self.dimensions.append(Dim.MBN)
178 return Dim.Dimensions(dims, all_dims=self.dimensions)
180 def make_parallel_config(self, dtpc_p, mbsn, evos_p):
181 """Create a parallel config from parallel values"""
182 logger.debug("dimensions considered: %s", str(self.dimensions))
183 (dp, mp, pp, cp) = dtpc_p
184 (mbs, mbn) = mbsn
185 (ep, vpp, op, sp) = evos_p
186 return self.make_parallel_config_args(
187 dp=dp,
188 mp=mp,
189 pp=pp,
190 cp=cp,
191 mbs=mbs,
192 mb=mbn,
193 ep=ep,
194 vpp=vpp,
195 op=op,
196 sp=sp,
197 )
199 def set_parallel_config(self, parallel_config):
200 """Set a given parallel configuration in the config"""
201 kwargs = {}
202 ok = True
203 new_pp = self.dim_val(Dim.PP, parallel_config)
204 new_vp = self.dim_val(Dim.VPP, parallel_config)
205 new_offset, new_recompute = self.adapt_config(new_pp, new_vp)
206 kwargs["offset"] = new_offset
207 kwargs["full_rec"] = new_recompute
208 # kwargs["sel_rec"] = sel_rec
209 for dim, value in parallel_config.dims_val.items():
210 kwargs[dim.name.lower()] = value
212 self.ccfg.set_strategy(**kwargs)
213 if not self.ccfg.multimodal:
214 if not self.ccfg.hooks_dict:
215 logger.info(
216 "'hook_cls' not specified,"
217 "search in predefined arch_hooks"
218 )
219 check_and_apply_custom_hook(self.ccfg)
220 else:
221 logger.info("Apply hooks")
222 hook = list(self.ccfg.hooks_dict.values())[0]
223 hook(self.wrap)
225 return ok
227 def space(self, dim, divide, reverse=False):
228 """Generate the space for a given dimension"""
229 if dim in self.dimensions:
230 if dim.get_bound() is not None:
231 logger.debug(
232 "Space of bounded dim %s is %s",
233 str(dim),
234 str(
235 Hard.all_divisors(
236 divide, reverse=reverse, max_bound=dim.get_bound()
237 )
238 ),
239 )
240 return Hard.all_divisors(
241 divide, reverse=reverse, max_bound=dim.get_bound()
242 )
243 logger.debug(
244 "Space of dim %s is %s",
245 str(dim),
246 str(Hard.all_divisors(divide, reverse=reverse)),
247 )
248 return Hard.all_divisors(divide, reverse=reverse)
249 logger.debug(
250 "Space of original dim %s is [%s]",
251 str(dim),
252 str(dim.from_config(self.ccfg)),
253 )
254 return [dim.from_config(self.ccfg)]
256 def range_space(self, dim, bound):
257 """Generate the space for a given dimension"""
258 if dim in self.dimensions:
259 return range(1, bound + 1)
260 return [dim.from_config(self.ccfg)]
262 def bool_space(self, dim):
263 """Generate the space for a given boolean dimension"""
264 if dim in self.dimensions:
265 return [False, True]
266 return [dim.from_config(self.ccfg)]
268 def max_op(self, dp, tp, ep):
269 """Compute bound for dimension OP"""
270 if (
271 isinstance(self.ccfg.optimizer, str)
272 and "muon" not in self.ccfg.optimizer.lower()
273 ):
274 return dp
275 if self.ccfg.n_exp and self.ccfg.n_exp > 1:
276 exp_gcd = gcd(dp * tp // max(tp, ep), self.ccfg.n_exp)
277 else:
278 exp_gcd = dp
280 dc_kv_valid = self.ccfg.dc_kv and self.ccfg.dc_kv > 1
281 dhr_valid = self.ccfg.dhr and self.ccfg.dhr > 1
282 if dc_kv_valid and dhr_valid:
283 att_gcd = gcd(self.ccfg.h, self.ccfg.dc_kv + self.ccfg.dhr)
284 else:
285 att_gcd = self.ccfg.h
286 return gcd(exp_gcd, att_gcd)