1# Copyright 2025-2026 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"""parse config for cost model"""
16import inspect
17import re
18import weakref
19from copy import deepcopy
20from enum import Enum
21from pprint import pformat
22
23from hyper_parallel.auto_parallel.sapp_nd.nd.common.generate_partitions import PartitionGenerator
24from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.logger import logger
25
26
27class AttentionType(Enum):
28 """Attention type enumeration."""
29 MHA = "mha"
30 GQA = "gqa"
31 MLA = "mla"
32
33
34def detect_attention_type(ccfg: "CostModelConfig") -> AttentionType:
35 """Detect attention type from cost model config.
36
37 Detection rules:
38 1. If kv_lora_rank > 0: MLA
39 2. If n_kv < a: GQA
40 3. Otherwise: MHA
41
42 Args:
43 ccfg: Cost model config.
44
45 Returns:
46 AttentionType enum.
47
48 Example:
49 >>> ccfg.kv_lora_rank = 512
50 >>> ccfg.a = 64
51 >>> ccfg.n_kv = 64
52 >>> detect_attention_type(ccfg)
53 <AttentionType.MLA: 'mla'>
54 """
55 if ccfg.kv_lora_rank > 0:
56 return AttentionType.MLA
57 if ccfg.n_kv < ccfg.a:
58 return AttentionType.GQA
59 return AttentionType.MHA
60
61
62def compute_kv_dim(ccfg) -> float:
63 """Return effective KV dimension per TP rank based on attention type.
64
65 When TP is active, KV heads are split across TP ranks, so each
66 rank holds only 1/t of the total KV dimension. MLA is an
67 exception: the compressed latent vector is not split by TP,
68 so kv_lora_rank stays unchanged.
69
70 Args:
71 ccfg: Cost model config with attributes a, n_kv, dh, h, t, kv_lora_rank.
72
73 Returns:
74 Effective KV dimension per TP rank (float).
75 """
76 attention_type = detect_attention_type(ccfg)
77 t = max(1, ccfg.t)
78 if attention_type == AttentionType.MLA:
79 return float(ccfg.kv_lora_rank)
80 if attention_type == AttentionType.GQA:
81 n_kv = min(ccfg.n_kv if ccfg.n_kv > 0 else ccfg.a, ccfg.a)
82 return n_kv * ccfg.dh / t
83 return ccfg.h / t
84
85
86# class CostModelConfig(Config) :
87class CostModelConfig(PartitionGenerator):
88 """cost model variables class"""
89
90 def __init__(
91 self,
92 input_config=None,
93 hook_cls=None,
94 framework=None,
95 source_code=None,
96 ):
97 super().__init__(input_config, hook_cls, framework, source_code)
98 logger.debug(
99 "parser = %s for %s", str(self.parser), str(self.model_name)
100 )
101
102 def __str__(self):
103 return "CostModelConfig attributes:\n" + pformat(
104 {
105 k: v
106 for k, v in vars(self).items()
107 if isinstance(v, (int, float, str, bool))
108 }
109 )
110
111 def __getattr__(self, attr):
112 call_source = inspect.currentframe().f_back.f_code.co_name
113 if attr not in self.__dict__:
114 logger.warning(
115 "[%s] Attribute %s does not exist. "
116 "Value '0' will be assigned.",
117 call_source,
118 attr,
119 )
120 return 0
121 return self.__dict__[attr]
122
123 def __copy__(self):
124 res = object.__new__(type(self))
125 res.__dict__.update(self.__dict__)
126 return res
127
128 def __deepcopy__(self, memo):
129 res = object.__new__(type(self))
130 for k, v in self.__dict__.items():
131 setattr(res, k, deepcopy(v, memo))
132 return res
133
134 def __getstate__(self) -> dict:
135 """Return instance state for multiprocessing serialization."""
136 return self.__dict__.copy()
137
138 def __setstate__(self, state: dict) -> None:
139 """Restore instance state after multiprocessing deserialization."""
140 self.__dict__.update(state)
141
142 def fp_bytes(self, precision):
143 """Return bytes size for datatype"""
144 if precision and isinstance(precision, str):
145 res = re.match(r"[^0-9]*([0-9]+)[^0-9]*", precision)
146 if res:
147 return int(res.group(1)) // 8
148 logger.warning("No bytes detected from FP Precision: %s", precision)
149 return 0
150
151 def print_stages_i(self, stage_id, stage):
152 """for print_stages"""
153 stage_layers = []
154 for chunk in stage:
155 chunk_lay_occ = []
156 if chunk:
157 layer, count = chunk[0], 1
158 for lay_id in range(1, len(chunk)):
159 if chunk[lay_id] == layer:
160 count += 1
161 else:
162 chunk_lay_occ += [f"{count}{layer.name[0]}"]
163 layer, count = chunk[lay_id], 1
164 chunk_lay_occ += [f"{count}{layer.name[0]}"]
165 stage_layers += [chunk_lay_occ]
166 logger.info("stage _%s : %s", stage_id, stage_layers)
167
168 def print_stages(self, stages, spec_stage_id=-1):
169 """Call after generate_partitions"""
170 if spec_stage_id == -1:
171 for stage_id, stage in enumerate(stages):
172 self.print_stages_i(stage_id, stage)
173 elif 0 <= spec_stage_id < len(stages):
174 self.print_stages_i(spec_stage_id, stages[spec_stage_id])
175 else:
176 logger.warning("Incorrect spec_stage_id")
177
178 def count_layers(self, stages):
179 """Count non-embedding and non-output layers in generated stages."""
180 return sum(sum(len(layer) for layer in chunk) for chunk in stages) - 2
181
182 def print_parallelism(self):
183 """strategy pretty printer"""
184 if not self.multimodal:
185 logger.info("%s Parallelism used :", self.model_name)
186 logger.info(
187 "DP %s, TP %s, PP %s, EP %s, CP %s, VPP %s",
188 self.d,
189 self.t,
190 self.p,
191 self.ep,
192 self.cp,
193 self.vp,
194 )
195 logger.info(
196 "d_exp %s, t_exp %s, os_max_shard %s, etp %s",
197 self.d_exp,
198 self.t_exp,
199 self.os_max_shard,
200 self.etp,
201 )
202 logger.info(
203 "shard_grad_exp %s, shard_grad_non_exp %s",
204 self.shard_grad_exp,
205 self.shard_grad_non_exp,
206 )
207 logger.info(
208 "shard_p_os_exp %s, shard_p_os_non_exp %s",
209 self.shard_p_os_exp,
210 self.shard_p_os_non_exp,
211 )
212 logger.info(
213 "shard_embed %s, shard_output_activ %s, shard_rec_input %s",
214 self.shard_embed,
215 self.shard_output_activ,
216 self.shard_recompute_input,
217 )
218 else:
219 for m in self.mm_ccfgs:
220 self.mm_ccfgs[m].print_parallelism()
221
222 def strategy_num_devices(self):
223 """total num devices"""
224 return self.d * self.t * self.cp * self.p
225
226 def is_consistent_pp_config(self):
227 """check if pp/offset/recomputation consistency"""
228
229 def is_valid_cfg(cfg):
230 if cfg is None or isinstance(cfg, (int, bool)):
231 return True
232 if not isinstance(cfg, list) or not cfg:
233 return False
234 if isinstance(cfg[0], int):
235 return len(cfg) == self.p
236 if isinstance(cfg[0], list):
237 return len(cfg) == self.vp and all(
238 isinstance(c, list) and len(c) == self.p for c in cfg
239 )
240 return False
241
242 return (
243 is_valid_cfg(self.offset)
244 and is_valid_cfg(self.full_rec)
245 and is_valid_cfg(self.sel_rec)
246 )
247
248 @staticmethod
249 def __maybe_set_int(target, attr, value):
250 """Set an integer strategy attribute when an override is supplied."""
251 if isinstance(value, int):
252 setattr(target, attr, value)
253
254 def __strategy_target(self, model_name):
255 """Get the config object targeted by a strategy update."""
256 if not self.multimodal:
257 return self
258 if model_name in self.mm_ccfgs:
259 return self.mm_ccfgs[model_name]
260 raise TypeError(
261 f"{self.model_name}: model_name is required (multimodal)"
262 )
263
264 def set_strategy(self, **kwargs):
265 """overwrite parallelism"""
266 model_name = kwargs.get("model_name", None)
267 dp = kwargs.get("dp", None)
268 tp = kwargs.get("mp", None)
269 cp = kwargs.get("cp", None)
270 ep = kwargs.get("ep", None)
271 op = kwargs.get("op", None)
272 etp = kwargs.get("etp", None)
273 pp = kwargs.get("pp", None)
274 vpp = kwargs.get("vpp", None)
275 off = kwargs.get("offset", None)
276 fr = kwargs.get("full_rec", None)
277 sr = kwargs.get("sel_rec", None)
278 m = kwargs.get("mb", None)
279 b = kwargs.get("mbs", None)
280 target_ccfg = self.__strategy_target(model_name)
281
282 for attr, value in (
283 ("d", dp),
284 ("t", tp),
285 ("ep", ep),
286 ("etp", etp),
287 ("cp", cp),
288 ("vp", vpp),
289 ("p", pp),
290 ("m", m),
291 ("b", b),
292 ):
293 self.__maybe_set_int(target_ccfg, attr, value)
294 target_ccfg.sp = target_ccfg.t
295 if op is not None and isinstance(op, int):
296 target_ccfg.os_max_shard = op
297 # Sync has_op with os_max_shard: op<=1 means no optimizer sharding
298 target_ccfg.has_op = op > 1
299 target_ccfg.gbs = target_ccfg.b * target_ccfg.d * target_ccfg.m
300 logger.debug(
301 "in ccfg: DP = %d, TP = %d, EP = %d, CP = %d, "
302 "PP = %d, MB = %d, MBS = %d, VPP = %d",
303 target_ccfg.d,
304 target_ccfg.t,
305 target_ccfg.ep,
306 target_ccfg.cp,
307 target_ccfg.p,
308 target_ccfg.m,
309 target_ccfg.b,
310 target_ccfg.vp,
311 )
312 if hasattr(target_ccfg.parser, "config_shard_emb"):
313 target_ccfg.parser.config_shard_emb()
314 if hasattr(target_ccfg.parser, "config_shard_recompute"):
315 target_ccfg.parser.config_shard_recompute()
316 target_ccfg.parser.config_dp_tp_exp(target_ccfg)
317 target_ccfg.parser.config_optimizer_shard(target_ccfg)
318 target_ccfg.parser.config_comm_flag(target_ccfg)
319 if fr is not None:
320 target_ccfg.full_rec = fr
321 if sr is not None:
322 target_ccfg.sel_rec = sr
323 if isinstance(off, (int, list)):
324 target_ccfg.offset = off
325 if not target_ccfg.is_consistent_pp_config():
326 raise AttributeError(
327 f"{target_ccfg.model_name}: "
328 "Inconsistent pipeline parallel variables "
329 f"pp {target_ccfg.p} vpp {target_ccfg.vp} "
330 f"offset {target_ccfg.offset} "
331 f"full_rec {target_ccfg.full_rec} "
332 f"sel_rec {target_ccfg.sel_rec}"
333 )
334 self.__maybe_set_int(target_ccfg, "cp", cp)
335
336 def get_strategy(self):
337 """return parallelism/recompute strategies"""
338
339 def strategy(mm):
340 return {
341 "dp": mm.d,
342 "tp": mm.t,
343 "pp": mm.p,
344 "ep": mm.ep,
345 "cp": mm.cp,
346 "vpp": mm.vp,
347 "op": mm.os_max_shard,
348 "gbs": mm.b * mm.m * mm.d,
349 "sched": mm.pp_sched,
350 "offset": mm.offset,
351 "full_rec": mm.full_rec,
352 "sel_rec": mm.sel_rec,
353 }
354
355 # logger.output("get_strat ccfg")
356 if self.multimodal:
357 return {mm.model_name: strategy(mm) for mm in self.mm_ccfgs.values()}
358 return strategy(self)
359
360 def layer_custom_config_callback(self, fun):
361 """
362 Use input fun as callback for layer_custom_config
363 Only for overwriting cost model variables
364 """
365 config_ref = weakref.ref(self)
366 for idx, f in enumerate(self.layer_custom_config):
367
368 def wrap(e, hook=f[1]):
369 hook(e)
370 if isinstance(e, CostModelConfig):
371 config = config_ref()
372 if config is None:
373 raise ReferenceError(
374 "CostModelConfig has already been released"
375 )
376 fun(config)
377 else:
378 e.set_ccfg(fun)
379
380 wrap.__name__ = f"{f[1].__name__}_{fun.__name__}"
381 self.layer_custom_config[idx] = (f[0], wrap)