Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / nd / common / cost_model_preprocess.py: 94%

175 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-21 04:29 +0800

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)