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