Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / memory_estimation / _backbone.py: 59%

412 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-04 05:18 +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"""backbone memory estimation module""" 

16from __future__ import annotations 

17from typing import TYPE_CHECKING 

18import os 

19import ast 

20import math 

21import logging 

22import inspect 

23import pprint 

24import importlib 

25import matplotlib.pyplot as plt 

26from PIL import Image 

27from hyper_parallel.auto_parallel.sapp_nd.nd.common.config import Config 

28from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger as nd_logger 

29from hyper_parallel.auto_parallel.sapp_nd.nd.common.cost_model_preprocess import CostModelConfig 

30from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.logger import logger 

31from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._context import Context, MemType 

32from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.utils import EvalUtils 

33from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.body import EvalBody 

34from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.evaluators.comm import EvalLayerComm 

35from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._bwd_overhead import _BackwardOverhead 

36from hyper_parallel.auto_parallel.sapp_nd.memory_estimation._ppb import _PPB 

37from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.hook_base import MemEvalHook 

38 

39if TYPE_CHECKING: 

40 from typing import Any, Dict, Tuple 

41 

42current_dir = os.path.dirname(os.path.abspath(__file__)) 

43EVAL_YML = os.path.join(current_dir, "configs_eval/default.yaml") 

44 

45 

46class _Backbone: 

47 """backbone class""" 

48 

49 def __init__(self, config: Any, **kwargs): 

50 self._child_cls = self 

51 self.mb = EvalUtils.mb 

52 self.eval_cfg = Config(kwargs.get("eval_yml", EVAL_YML)) 

53 self._ctx = kwargs.get("ctx", Context()) 

54 self._ccfg = kwargs.get("ccfg", None) 

55 self.framework = kwargs.get("framework", None) 

56 self.source_code = kwargs.get("source_code", None) 

57 self.hook_cls, self.config_path = None, None 

58 if not self._ccfg: 

59 self.hook_cls = kwargs.get("hook_cls", self.eval_cfg.hook_class) 

60 if isinstance(self.hook_cls, str): 

61 self._load_eval_yaml_hook_cls(self.hook_cls) 

62 if self.hook_cls and not isinstance(self.hook_cls, MemEvalHook): 

63 raise AttributeError( 

64 f"'{self.hook_cls}' is not a MemEvalHook instance" 

65 ) 

66 if config: 

67 if kwargs.get("log_level", 1) == 0: 

68 # logger.setLevel(logging.CRITICAL) 

69 nd_logger.setLevel(logging.CRITICAL) 

70 self.update_config(config) 

71 else: 

72 raise AttributeError("missing config") 

73 self.evaluator_instances = None 

74 self.ppb = None 

75 self._overhead_obj = _BackwardOverhead( 

76 self, self._ccfg, self._ctx, self._inner_dynamic_mem 

77 ) 

78 self._ppb_obj = _PPB(self.eval_cfg, self._inner_dynamic_mem) 

79 

80 @property 

81 def ccfg(self) -> CostModelConfig: 

82 """read-only""" 

83 return self._ccfg 

84 

85 @property 

86 def ctx(self) -> Context: 

87 """read-only""" 

88 return self._ctx 

89 

90 def _load_eval_yaml_hook_cls(self, hook_cls): 

91 """hook_class in eval yaml""" 

92 target_mod_path = None 

93 try: 

94 # search in folder 'hooks' 

95 hooks_dir = os.path.join(current_dir, "hooks") 

96 for f in os.listdir(hooks_dir): 

97 if f.endswith(".py"): 

98 mod_path = f"hyper_parallel.auto_parallel.sapp_nd.memory_estimation.hooks.{f.split('.')[0]}" 

99 spec = importlib.util.find_spec(mod_path) 

100 if spec is None or spec.origin is None: 

101 continue 

102 with open(spec.origin, "r", encoding="utf-8") as mf: 

103 source = mf.read() 

104 tree = ast.parse(source) 

105 mod_cls = None 

106 for node in ast.walk(tree): 

107 if isinstance(node, ast.ClassDef) and node.name == hook_cls: 

108 mod_cls = node 

109 break 

110 if mod_cls: 

111 target_mod_path = mod_path 

112 break 

113 if target_mod_path: 

114 module = importlib.import_module(target_mod_path) 

115 self.hook_cls = getattr(module, hook_cls)() 

116 except (ModuleNotFoundError, ImportError) as e: 

117 print(e) 

118 

119 # Peak memory estimation 

120 

121 def update_config(self, new_config: Any) -> None: 

122 """processing input config""" 

123 if self._ccfg is not None: 

124 self._ccfg.update_config(new_config, self.hook_cls, self.framework, self.source_code) 

125 else: 

126 self._ccfg = CostModelConfig(new_config, self.hook_cls, self.framework, self.source_code) 

127 if isinstance(new_config, str): 

128 if not self.config_path: 

129 logger.info( 

130 "%s Process config file: %s", 

131 "=" * 30, 

132 new_config.split("/")[-1], 

133 ) 

134 self.config_path = new_config 

135 self.evaluator_instances = None 

136 self.ppb = None 

137 

138 def _inner_static_mem(self) -> float: 

139 """static memory evaluation for backbone estimation""" 

140 if self._ctx.current_node in self._ctx.node_eval: 

141 p = self._ctx.eval.stat.p(self._ccfg, self._ctx) 

142 ost = self._ctx.eval.stat.os(self._ccfg, self._ctx) 

143 grad = self._ctx.eval.stat.grad(self._ccfg, self._ctx) 

144 res = p + ost + grad 

145 self._ctx.save2log("_param", res) 

146 # Log routed/shared expert param breakdown for MoE layers 

147 if self._ccfg.n_exp > 1 and self.is_regular_layer(self._ctx.current_node): 

148 result = self._ctx.eval.num_p(self._ccfg, self._ctx) 

149 if isinstance(result, tuple) and len(result) == 3: 

150 _, routed_exp_p, shared_exp_p = result 

151 routed_exp_mem = ( 

152 routed_exp_p / self._ccfg.ep 

153 * self._ccfg.bytes_p / self._ccfg.shard_p_os_exp 

154 ) 

155 shared_exp_mem = ( 

156 shared_exp_p 

157 * self._ccfg.bytes_p / self._ccfg.shard_p_os_exp_partial 

158 ) 

159 self._ctx.save2log("routed_exp_param", routed_exp_mem) 

160 self._ctx.save2log("shared_exp_param", shared_exp_mem) 

161 return p + ost + grad 

162 return 0 

163 

164 def _inner_dynamic_mem( 

165 self, ppb=False, default_micro_factor=None 

166 ) -> tuple[float, float]: 

167 """dynamic memory evaluation for backbone estimation""" 

168 if self._ctx.current_node in self._ctx.node_eval: 

169 if ppb: 

170 micro_factor = 1 

171 elif default_micro_factor: 

172 micro_factor = default_micro_factor 

173 else: 

174 sched = self._ccfg.pp_sched 

175 micro_factor = self._ctx.pp_micro_eval[sched]( 

176 self._ccfg, self._ctx 

177 ) 

178 self._ctx.micro_factor = max(1, micro_factor) 

179 activation = self._ctx.eval.dyn.activation(self._ccfg, self._ctx) 

180 self._ctx.save2log("_activ", activation) 

181 comm = self._inner_comm_mem(micro_factor) 

182 return activation, comm 

183 return 0, 0 

184 

185 def _inner_comm_mem(self, micro_factor) -> float: 

186 """dynamic memory evaluation for backbone estimation""" 

187 if self._ctx.current_node in self._ctx.node_eval: 

188 comm_eval_field = self._ctx.eval.dyn.comm 

189 comm_cat = { 

190 "dp": MemType.AG_COMM, 

191 "tp": MemType.AG_COMM, 

192 "cp": MemType.AG_COMM, 

193 "ep": MemType.A2A_COMM, 

194 } 

195 comm_mem = {} 

196 for k, fun in vars(comm_eval_field).items(): 

197 if fun is None: 

198 continue 

199 sig_param = inspect.signature(fun).parameters.values() 

200 if ( 

201 # any( 

202 # p.kind == inspect.Parameter.VAR_KEYWORD 

203 # for p in sig_param 

204 # ) 

205 # and len(sig_param) > 2 

206 len(sig_param) > 2 

207 ): 

208 comm = fun(self._ccfg, self._ctx, micro_factor) 

209 else: 

210 comm = fun(self._ccfg, self._ctx) 

211 comm_mem[k] = comm 

212 res = EvalUtils.eval_expr_insight( 

213 expr=self._ctx.comm_expr, 

214 ctx=self._ctx, 

215 mem_val=comm_mem, 

216 mem_cat=comm_cat, 

217 ) 

218 self._ctx.save2log("_comm", res) 

219 return res 

220 return 0 

221 

222 def __subplot_stages(self, stage_i, stat_mem_i, dyn_mem_i, i) -> None: 

223 """Plot static and dynamic memory for one pipeline stage.""" 

224 _, ax = plt.subplots(figsize=(5, 8)) 

225 bottoms = {"stat": 0, "dyn": 0} 

226 color_stat = plt.get_cmap("cividis") 

227 color_dyn = plt.get_cmap("plasma") 

228 total_lay = sum(len(chunk) for chunk in stage_i) 

229 color_denominator = max(1, total_lay - 1) 

230 for chunk_id, chunk in enumerate(stage_i): 

231 for lay_id, lay_type in enumerate(chunk): 

232 real_id = self._ctx.real_lay_ids[chunk_id][i][lay_id] 

233 name = f"Lay_{real_id}_{lay_type.name[0]}" 

234 idx = chunk_id * len(dyn_mem_i) + lay_id 

235 stat = self.mb(stat_mem_i[chunk_id][lay_id]) 

236 dyn = self.mb(dyn_mem_i[chunk_id][lay_id]) 

237 ax.bar( 

238 0.1, 

239 [dyn], 

240 bottom=bottoms["dyn"], 

241 label=name, 

242 color=color_dyn(idx / color_denominator), 

243 width=0.15, 

244 linewidth=0.5, 

245 edgecolor="black", 

246 ) 

247 ax.text( 

248 0.2, 

249 bottoms["dyn"] + dyn / 2, 

250 f"DYN_{name}", 

251 va="center", 

252 fontsize=5, 

253 ) 

254 ax.bar( 

255 0.5, 

256 [stat], 

257 bottom=bottoms["stat"], 

258 label=name, 

259 color=color_stat(idx / color_denominator), 

260 width=0.15, 

261 linewidth=0.5, 

262 edgecolor="black", 

263 ) 

264 ax.text( 

265 0.6, 

266 bottoms["stat"] + stat / 2, 

267 f"STAT_{name}", 

268 va="center", 

269 fontsize=5, 

270 ) 

271 bottoms["dyn"] += dyn 

272 bottoms["stat"] += stat 

273 ax.axhline( 

274 self.get_max_device_memory(), 

275 color="red", 

276 linewidth=1, 

277 ls="dotted", 

278 ) 

279 ax.axhline( 

280 bottoms["dyn"] + bottoms["stat"], 

281 color="blue", 

282 linewidth=1, 

283 ls="dashed", 

284 ) 

285 ax.text( 

286 0, 

287 self.get_max_device_memory(), 

288 "Device Memory", 

289 color="red", 

290 va="bottom", 

291 fontsize=8, 

292 ) 

293 ax.text( 

294 0.62, 

295 bottoms["dyn"] + bottoms["stat"], 

296 "Prediction total", 

297 color="blue", 

298 va="bottom", 

299 fontsize=8, 

300 ) 

301 ax.set_xticks([]) 

302 ax.set_ylabel("Size (MB)") 

303 ax.set_xlim([0, 0.8]) 

304 ax.set_xlabel(f"Stage_{i}") 

305 

306 def __plot_stages(self, stages, stat_mems, dyn_mems) -> None: 

307 """plot bars for estimations""" 

308 

309 logger.info("Plotting predictions in plots/") 

310 if not os.path.exists("plots"): 

311 os.makedirs("plots") 

312 imgs = [] 

313 for stage_id in range(self._ccfg.p): 

314 self.__subplot_stages( 

315 stages[stage_id], 

316 stat_mems[stage_id], 

317 dyn_mems[stage_id], 

318 stage_id, 

319 ) 

320 img = f"plots/MemPlot_stage_{stage_id}.png" 

321 imgs += [(stage_id, img)] 

322 logger.info("save plot: %s", img) 

323 plt.savefig(img, dpi=300, bbox_inches="tight") 

324 plt.clf() 

325 plt.close() 

326 # Concatenate every stage plots 

327 if imgs: 

328 imgs = sorted(imgs) 

329 canvas = [Image.open(i) for _, i in imgs] 

330 stage_canva = Image.new( 

331 "RGB", 

332 ( 

333 canvas[0].size[0] 

334 * 2 ** math.ceil(math.log2(len(canvas)) / 2), 

335 canvas[0].size[1] 

336 * 2 ** math.floor(math.log2(len(canvas)) / 2), 

337 ), 

338 ) 

339 offset_x, offset_y = 0, 0 

340 for i in canvas: 

341 stage_canva.paste(i, (offset_x, offset_y)) 

342 offset_x += i.size[0] 

343 if offset_x >= stage_canva.size[0]: 

344 offset_x = 0 

345 offset_y += i.size[1] 

346 stage_canva.save("plots/MemPlot_all_stages.png") 

347 logger.info("save stage plot: plots/MemPlot_all_stages.png") 

348 

349 def __update_stage_logs(self, stage_logs: list, stage_id: int) -> None: 

350 """update stage's temporary buffers from ctx's buffers""" 

351 if not stage_logs[stage_id].node_compute_log: 

352 stage_logs[stage_id].node_compute_log = {} 

353 stage_logs[stage_id].node_compute_log.update( 

354 self._ctx.node_compute_log 

355 ) 

356 if not stage_logs[stage_id].accu_mem_type: 

357 stage_logs[stage_id].accu_mem_type = { 

358 mt: 0 for mt in list(MemType) 

359 } 

360 for mem_type in list(MemType): 

361 val = self._ctx.accu_mem_type[mem_type] 

362 stage_logs[stage_id].accu_mem_type[mem_type] += val 

363 

364 def __preprocess_layer_custom_config_list(self, stages: list) -> list: 

365 """flatten layer_custom_config for backbone estimation""" 

366 flatten = sum( 

367 [[f[1]] * f[0] for f in self._ccfg.layer_custom_config], [] 

368 ) 

369 total_n_lay = self._ccfg.n_lay + self._ccfg.n_mtp 

370 total_n_lay_stages = self._ccfg.count_layers(stages) 

371 if not self._ccfg.multimodal and total_n_lay != total_n_lay_stages: 

372 raise(AttributeError( 

373 f"Mismatch of num_layers between parsed value ({total_n_lay})" 

374 f" and generated partitions ({total_n_lay_stages})" 

375 f" => offset may be incorrect ({self._ccfg.offset})" 

376 )) 

377 if self._ccfg.n_lay > 0 and len(flatten) != total_n_lay: 

378 raise AttributeError( 

379 f"layer_custom_config occurrences ({len(flatten)})" 

380 f" != num_layers ({total_n_lay})" 

381 ) 

382 if self._ccfg.pp_sched == "zero_bubble_v": 

383 n_layer_first_chunk = ( 

384 sum(len(s[0]) for s in stages) - 1 

385 ) # Except embedding layer 

386 flatten = ( 

387 flatten[:n_layer_first_chunk] 

388 + flatten[n_layer_first_chunk:][::-1] 

389 ) 

390 return flatten 

391 

392 def _estimate_backbone(self, *args) -> Tuple[list, Dict]: 

393 """Evaluator's main function for estimation""" 

394 stages = args[0] 

395 spec_stage_id = args[3] 

396 

397 if spec_stage_id >= self._ccfg.p or spec_stage_id < 0: 

398 spec_stage_id = -1 

399 # Process partition generation 

400 if not stages: 

401 stages = self._ccfg.generate_partitions_vpp() 

402 # multimodal=self._ccfg.multimodal 

403 # ) 

404 if not self._ccfg.multimodal: 

405 return self.__estimate_stages_backbone( 

406 stages, args[1], args[2], spec_stage_id, args[4] 

407 ) 

408 res = [] 

409 original_ccfg = self._ccfg 

410 common_lc = [] 

411 self.evaluator_instances = [] 

412 # Build common layer_custom_config + Build temporary evaluators 

413 for m in self._ccfg.mm_order: 

414 self._ccfg.mm_ccfgs[m].config = original_ccfg.config 

415 if not self._child_cls: 

416 raise AttributeError("expected non null _child_cls") 

417 tmp_evaluator = type(self._child_cls)( 

418 None, 

419 ccfg=self._ccfg.mm_ccfgs[m], 

420 trace_fun=getattr(self, "toggle_func_trace", False), 

421 ) 

422 tmp_evaluator.import_eval_yaml() 

423 num_layer = tmp_evaluator.get_num_layers() 

424 # assert not isinstance(num_layer, tuple) 

425 strategy = tmp_evaluator.get_strategy() 

426 full_rec = strategy["full_rec"] 

427 offset = strategy["offset"] 

428 self._ccfg.hooks_dict[m](tmp_evaluator) 

429 strategy = tmp_evaluator.get_strategy() 

430 if ( 

431 tmp_evaluator.get_num_layers() != num_layer 

432 or strategy["full_rec"] != full_rec 

433 or strategy["offset"] != offset 

434 ): 

435 # Reverify num layers, recomp, offset after hook 

436 stages[m] = ( 

437 tmp_evaluator.ccfg.generate_partitions_vpp_unimodal() 

438 ) 

439 tmp_evaluator.set_layer_custom(None) 

440 common_lc += self._ccfg.mm_ccfgs[m].layer_custom_config 

441 self.evaluator_instances += [tmp_evaluator] 

442 if args[1]: 

443 logger.info("Submodule %s", self._ccfg.mm_ccfgs[m].model_name) 

444 tmp_evaluator.print_ctx() 

445 self.print_stages(stages[m]) 

446 self.set_layer_custom(common_lc) 

447 if args[1]: 

448 logger.info( 

449 "Combined layer_custom_config for %s\n%s", 

450 self._ccfg.model_name, 

451 pprint.pformat(self._ccfg.layer_custom_config, compact=True), 

452 ) 

453 logger.info( 

454 "Sub evaluator instances for %s\n%s", 

455 self._ccfg.model_name, 

456 pprint.pformat(self.evaluator_instances, compact=True), 

457 ) 

458 

459 res = self.__estimate_stages_backbone( 

460 self._ccfg.combine_partition_multimodal(stages), 

461 args[1], 

462 args[2], 

463 spec_stage_id, 

464 args[4], 

465 ) 

466 self._ccfg = original_ccfg 

467 return res 

468 

469 def __estimate_stages_backbone(self, *args) -> Tuple[list, Dict]: 

470 """Evaluator's main function for stage estimation""" 

471 stages = args[0] 

472 verbose = args[1] 

473 compute_ppb = args[2] 

474 spec_stage_id = args[3] 

475 

476 if verbose: 

477 logger.info("Partition of layers :") 

478 self._ccfg.print_stages(stages, spec_stage_id) 

479 insights = [] 

480 # Compute peak memory 

481 flatten = self.__preprocess_layer_custom_config_list(stages) 

482 if verbose: 

483 logger.info( 

484 "Flatten layer_custom_config\n%s", 

485 pprint.pformat( 

486 [f if not f else f.__name__ for f in flatten], compact=True 

487 ), 

488 ) 

489 

490 stage_misc = { 

491 "stat": [[[0 for _ in c] for c in s] for s in stages], 

492 "dyn": [[[0 for _ in c] for c in s] for s in stages], 

493 "logs": [Config({}) for _ in range(self._ccfg.p)], 

494 } 

495 # tmp_ppb_lay_desc = [] # PPB purpose 

496 ppb_lay_desc = [] 

497 record_lay_types = {} 

498 self.__chunk_stage_lay_loops( 

499 flatten, 

500 stages, 

501 record_lay_types, 

502 stage_misc, 

503 verbose, 

504 compute_ppb, 

505 ppb_lay_desc, # tmp_ppb_lay_desc, 

506 ) 

507 self.__postprocess_stages( 

508 stages, 

509 record_lay_types, 

510 stage_misc, 

511 verbose, 

512 spec_stage_id, 

513 insights, 

514 ) 

515 # PPB Input 

516 ppb_input = None 

517 if compute_ppb == 1: 

518 self._ppb_obj.ppb_combine_bodies(ppb_lay_desc) 

519 ppb_input = {"layers_description": ppb_lay_desc} 

520 elif compute_ppb == 2: 

521 self._ppb_obj.ppb_combine_bodies_new(ppb_lay_desc) 

522 ppb_input = {"layers_description_new": ppb_lay_desc} 

523 if args[4]: # Plot 

524 self.__plot_stages(stages, stage_misc["stat"], stage_misc["dyn"]) 

525 return insights, ppb_input 

526 

527 def __update_evaluator(self, node, verbose): 

528 if self.evaluator_instances and node == self._ctx.head_node: 

529 tmp_eval = self.evaluator_instances.pop(0) 

530 self._ccfg = tmp_eval.ccfg 

531 self._ctx.copy_tmp_buff(tmp_eval.ctx) 

532 self._ctx = tmp_eval.ctx 

533 if verbose: 

534 logger.info( 

535 "Update ccfg and ctx, module : %s", 

536 self._ccfg.model_name, 

537 ) 

538 

539 def __update_next_layer_custom_function(self, *args): 

540 """Apply the next layer custom hook before evaluating a layer.""" 

541 flatten, verbose = args[0], args[1] 

542 record_lay_types = args[2] 

543 stage_id, chunk_id, lay_id = args[3], args[4], args[5] 

544 node = args[6] 

545 if self.is_regular_layer(node) and flatten: 

546 hook = flatten.pop(0) 

547 if hook: 

548 if verbose: 

549 logger.info("Apply hook %s", hook.__name__) 

550 record_lay_types[(stage_id, chunk_id, lay_id)] = ( 

551 self._ccfg, 

552 self._ctx, 

553 hook, 

554 ) 

555 hook(self) 

556 else: 

557 record_lay_types[(stage_id, chunk_id, lay_id)] = ( 

558 self._ccfg, 

559 self._ctx, 

560 lambda _: None, 

561 ) 

562 else: 

563 record_lay_types[(stage_id, chunk_id, lay_id)] = ( 

564 self._ccfg, 

565 self._ctx, 

566 lambda _: None, 

567 ) 

568 if verbose: 

569 logger.info( 

570 "stage_id=%s, chunk_id=%s, lay_id=%s, node=%s", 

571 stage_id, 

572 chunk_id, 

573 lay_id, 

574 node, 

575 ) 

576 self._ccfg.print_parallelism() 

577 

578 def __chunk_stage_lay_loops(self, *args): 

579 """Evaluate every layer in every stage and chunk.""" 

580 flatten, stages, record_lay_types = args[0], args[1], args[2] 

581 sm = args[3] 

582 verbose, compute_ppb, ppb_lay_desc = args[4], args[5], args[6] 

583 self._ctx.real_lay_ids = [] 

584 count = 0 

585 for chunk_id in range(self._ccfg.vp): 

586 self._ctx.real_lay_ids += [[]] 

587 for stage_id in range(self._ccfg.p): 

588 self._ctx.real_lay_ids[chunk_id] += [[]] 

589 self._ctx.init_tmp_buff() 

590 for lay_id in range(len(stages[stage_id][chunk_id])): 

591 node = stages[stage_id][chunk_id][lay_id] 

592 if self.is_regular_layer(node): 

593 self._ctx.real_lay_ids[chunk_id][stage_id] += [count] 

594 count += 1 

595 else: 

596 self._ctx.real_lay_ids[chunk_id][stage_id] += [""] 

597 # Update evaluator (multimodal) 

598 self.__update_evaluator(node, verbose) 

599 # Update next layer custom function 

600 self.__update_next_layer_custom_function( 

601 flatten, 

602 verbose, 

603 record_lay_types, 

604 stage_id, 

605 chunk_id, 

606 lay_id, 

607 node, 

608 ) 

609 self._ctx.current_stage_id = stage_id 

610 self._ctx.current_chunk_id = chunk_id 

611 self._ctx.current_lay_id = lay_id 

612 self._ctx.current_node = node 

613 static_mem = self._inner_static_mem() 

614 sm["stat"][stage_id][chunk_id][lay_id] = static_mem 

615 sm["dyn"][stage_id][chunk_id][lay_id] = sum( 

616 self._inner_dynamic_mem() 

617 ) 

618 if verbose: 

619 logger.info("pp micro factor for dynamic: %s",self._ctx.micro_factor) 

620 # PPB Purpose 

621 if compute_ppb == 1: 

622 desc = self._ppb_obj.lay_ppb( 

623 self._ccfg, 

624 self._ctx, 

625 sm["stat"][stage_id][chunk_id][lay_id], 

626 ) 

627 self._ppb_obj.add_to_ppb_list(ppb_lay_desc, desc) 

628 elif compute_ppb == 2: 

629 desc = self._ppb_obj.lay_ppb_new( 

630 self._ccfg, 

631 self._ctx, 

632 sm["stat"][stage_id][chunk_id][lay_id], 

633 ) 

634 self._ppb_obj.add_to_ppb_list(ppb_lay_desc, desc) 

635 self.__update_stage_logs(sm["logs"], stage_id) 

636 

637 def __postprocess_stages(self, *args): 

638 """Build memory insights from raw stage evaluation buffers.""" 

639 stages, record_lay_types = args[0], args[1] 

640 sm = args[2] 

641 verbose, spec_stage_id = args[3], args[4] 

642 insights = args[5] 

643 for stage_id in range(self._ccfg.p): 

644 ins = {} # Mem Insights purpose 

645 ins["Static"] = sum( 

646 sum(mem for mem in c) for c in sm["stat"][stage_id] 

647 ) 

648 ins["Dynamic"] = sum( 

649 sum(mem for mem in c) for c in sm["dyn"][stage_id] 

650 ) 

651 self._ctx.init_tmp_buff() 

652 if not self._ccfg.freeze: 

653 ins["Dynamic"] += self._overhead_obj.estimate( 

654 stages, stage_id, record_lay_types 

655 ) 

656 self.__update_stage_logs(sm["logs"], stage_id) 

657 safety_buffer = 1024 * 1024 * 1024 # 1 GB 

658 if ins["Dynamic"] > 0: 

659 ins["Dynamic"] += safety_buffer 

660 stage_accu = sm["logs"][stage_id].accu_mem_type 

661 ins["ModelParameters"] = self.mb(stage_accu[MemType.MODEL_PARAM]) 

662 ins["OptimizerStates"] = self.mb(stage_accu[MemType.OPTIM_STATE]) 

663 ins["AccumulGradients"] = self.mb(stage_accu[MemType.ACCU_GRAD]) 

664 ins["Attn"] = self.mb(stage_accu[MemType.ATTN_ACTIV]) 

665 ins["FFn"] = self.mb(stage_accu[MemType.FFN_ACTIV]) 

666 ins["Norm"] = self.mb(stage_accu[MemType.NORM_ACTIV]) 

667 ins["AllGather Comm"] = self.mb(stage_accu[MemType.AG_COMM]) 

668 ins["All2All Comm"] = self.mb(stage_accu[MemType.A2A_COMM]) 

669 

670 if self._ccfg.cp > 1: 

671 cp_memory = EvalBody.act_cp_layer(self._ccfg, self._ctx) 

672 cp_comm_buffer = EvalLayerComm.cp_comm_buffer(self._ccfg, self._ctx) 

673 

674 ins["CP KV Cache"] = self.mb(cp_memory.kv_cache_memory) 

675 ins["CP Attn Scores"] = self.mb(cp_memory.attention_scores_memory) 

676 ins["CP Softmax"] = self.mb(cp_memory.softmax_outputs_memory) 

677 ins["CP Comm Buffer"] = self.mb(cp_comm_buffer) 

678 ins["CP Reduction"] = self.mb(cp_memory.total_reduction) 

679 

680 ins["Node Log"] = sm["logs"][stage_id].node_compute_log 

681 # VERBOSE 

682 if verbose and spec_stage_id in (-1, stage_id): 

683 self.__verbose_insights(sm, stage_id, ins) 

684 ins["Static"] = self.mb(ins["Static"]) 

685 ins["Dynamic"] = self.mb(ins["Dynamic"]) 

686 insights += [ins] 

687 

688 def __verbose_insights(self, *args): 

689 """logging purpose""" 

690 sm = args[0] 

691 stage_id = args[1] 

692 ins = args[2] 

693 stat_i = max(1, ins["Static"]) 

694 dyn_i = max(1, ins["Dynamic"]) 

695 # logs_i = sm["logs"][stage_id] 

696 accu_i = sm["logs"][stage_id].accu_mem_type 

697 logger.info( 

698 "stage _%s : %s MB", 

699 stage_id, 

700 self.mb(ins["Static"] + ins["Dynamic"]), 

701 ) 

702 logger.info( 

703 "\tStatic\t%s\t" 

704 "ModelParam %s (%s%%), " 

705 "OptimStates %s (%s%%), " 

706 "Gradients %s (%s%%)", 

707 self.mb(ins["Static"]), 

708 ins["ModelParameters"], 

709 round(accu_i[MemType.MODEL_PARAM] / stat_i * 100), 

710 ins["OptimizerStates"], 

711 round(accu_i[MemType.OPTIM_STATE] / stat_i * 100), 

712 ins["AccumulGradients"], 

713 round(accu_i[MemType.ACCU_GRAD] / stat_i * 100), 

714 ) 

715 logger.info( 

716 "\tDynamic\t%s\t" 

717 "Attn %s (%d%%), " 

718 "FFn %s (%d%%), " 

719 "Norm %s (%d%%), " 

720 "AllGather Comm %s (%d%%), " 

721 "All2All Comm %s (%d%%), ", 

722 self.mb(ins["Dynamic"]), 

723 ins["Attn"], 

724 round(accu_i[MemType.ATTN_ACTIV] / dyn_i * 100), 

725 ins["FFn"], 

726 round(accu_i[MemType.FFN_ACTIV] / dyn_i * 100), 

727 ins["Norm"], 

728 round(accu_i[MemType.NORM_ACTIV] / dyn_i * 100), 

729 ins["AllGather Comm"], 

730 round(accu_i[MemType.AG_COMM] / dyn_i * 100), 

731 ins["All2All Comm"], 

732 round(accu_i[MemType.A2A_COMM] / dyn_i * 100), 

733 ) 

734 logger.info( 

735 "\tNode eval log : \n %s \n %s", 

736 "> Foreach: (stage_id,chunk_id,lay_id,name) -> (mem type,value)", 

737 pprint.pformat(ins["Node Log"], width=300), 

738 ) 

739 

740 def apply_hook(self, hook, ccfg=None, ctx=None): 

741 """apply hook on evaluator""" 

742 self._ccfg = ccfg if ccfg else self._ccfg 

743 self._ctx = ctx if ctx else self._ctx 

744 hook(self) 

745 

746 def set_layer_custom(self, _): 

747 """child implement""" 

748 pass # pylint: disable=unnecessary-pass 

749 

750 def is_regular_layer(self, _): 

751 """child implement""" 

752 return False 

753 

754 def import_eval_yaml(self): 

755 """child implement""" 

756 pass # pylint: disable=unnecessary-pass 

757 

758 def get_num_layers(self): 

759 """child implement""" 

760 return 0 

761 

762 def get_strategy(self): 

763 """child implement""" 

764 return {} 

765 

766 def get_max_device_memory(self): 

767 """child implement""" 

768 return 0 

769 

770 def print_stages(self, _): 

771 """child implement""" 

772 pass # pylint: disable=unnecessary-pass 

773 

774 def print_ctx(self): 

775 """child implement""" 

776 pass # pylint: disable=unnecessary-pass