Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_ppb / sapp / sapp_pipeline.py: 91%
426 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
1# Copyright 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"""High-level orchestrator around :class:`SappSolver`: build, solve, simulate, export YAML."""
16import os
17import sys
18from typing import Any, Dict, List, Optional, Union
20import matplotlib.pyplot as plt
21import yaml
23import hyper_parallel.auto_parallel.sapp_ppb.simulator.pp_simulator as sim
24import hyper_parallel.auto_parallel.sapp_ppb.utils.recompute as Recompute
25from hyper_parallel.auto_parallel.sapp_ppb.sapp.sapp_solver import SappSolver
26from hyper_parallel.auto_parallel.sapp_ppb.utils.check_rules import check_yaml_depth_before_loading
27from hyper_parallel.auto_parallel.sapp_ppb.utils.layer import Layer, filter_layer_type
28from hyper_parallel.auto_parallel.sapp_ppb.utils.logger import logger
31class SappPipeline:
32 """pipeline balancer"""
34 def __init__(
35 self,
36 model_name: str,
37 num_of_stage: int,
38 num_of_micro_batch: int,
39 max_memory: int,
40 layers: List[Layer],
41 vpp_less_memory: bool = False,
42 # Add arg dual
43 dual: bool = False,
44 num_of_interleave: int = 1,
45 constant_memory: int = 0,
46 optimization_level: int = 1,
47 extracted_training_params: Optional[Dict[str, int]] = None,
48 seq_split_num: int = 1,
49 use_backward_time: bool = False,
50 ) -> None:
51 """Cache pipeline parameters and index the input ``layers`` by HEAD / BODY / TAIL.
53 Args:
54 model_name (str): Model identifier, used for dump filenames and log prefixes.
55 num_of_stage (int): Number of physical pipeline stages.
56 num_of_micro_batch (int): Number of micro-batches scheduled per iteration.
57 max_memory (int): Per-device memory budget in MB.
58 layers (List[Layer]): Ordered list of layer descriptors covering HEAD/BODY/TAIL.
59 vpp_less_memory (bool, optional): If ``True``, use the less-memory VPP scheduler variant.
60 Default: ``False``.
61 dual (bool, optional): Enable dualpipe-V scheduling support. Default: ``False``.
62 num_of_interleave (int, optional): Virtual-pipeline (VPP) chunk count. Default: ``1``.
63 constant_memory (int, optional): Constant per-stage memory overhead (MB). Default: ``0``.
64 optimization_level (int, optional): Solver optimization level (``0-2``). Default: ``1``.
65 extracted_training_params (Optional[Dict[str, int]], optional): Optional training-config parameters for
66 seqpp. Default: ``None``.
67 seq_split_num (int, optional): Number of sequence splits; ``>1`` enables sequence pipeline.
68 Default: ``1``.
69 """
70 self.model_name_ = model_name
71 self.num_of_stage_ = num_of_stage
72 self.num_of_micro_batch_ = num_of_micro_batch
73 self.num_of_interleave_ = num_of_interleave
74 self.max_memory_ = max_memory
75 self.vpp_less_memory_ = vpp_less_memory
76 # Add arg dual_
77 self.dual_ = dual
78 self.constant_memory_ = constant_memory
79 self.optimization_level = optimization_level
80 self.extracted_training_params_ = extracted_training_params
81 self.seq_split_num_ = seq_split_num
82 self.use_backward_time_ = use_backward_time
83 self.seqpipe_ = self.seq_split_num_ > 1
84 # logger.output("seq chunk: %s",self.seq_split_num_)
86 self.problem_ = None
87 self.layers_ = layers
88 self.layers_sorted_ = {
89 Layer.type_enum.HEAD: filter_layer_type(layers,
90 Layer.type_enum.HEAD),
91 Layer.type_enum.BODY: filter_layer_type(layers,
92 Layer.type_enum.BODY),
93 Layer.type_enum.TAIL: filter_layer_type(layers,
94 Layer.type_enum.TAIL),
95 }
97 @property
98 def simulator(self):
99 """Pipeline simulator instance (available after :meth:`simulate`)."""
100 return self._simulator
102 def has_some_memory_info(self) -> bool:
103 """Check if there is all information for memory constraint."""
104 return self.problem_.has_some_memory_info()
106 def construct_problem(self, solver: str = "pulp") -> None:
107 """Construct the underlying ILP problem using the requested solver backend."""
108 if solver == "pulp":
109 self.problem_ = self._construct_problem_pulp_()
110 elif solver == "other":
111 logger.warning(
112 "No other solver available..., automatically switch to pulp!!!"
113 )
114 self.problem_ = self._construct_problem_pulp_()
115 else:
116 logger.warning(
117 "No other solver available..., automatically switch to pulp!!!"
118 )
119 self.problem_ = self._construct_problem_pulp_()
121 def solve_problem(self, time_limit: int = 90, dump_folder: Optional[str] = None) -> None:
122 """Solve the ILP, optionally dumping the LP model into ``dump_folder``."""
123 self.problem_.solve(time_limit, dump_folder)
125 def get_result(self) -> dict[str, list[list[str]]]:
126 """Get result distribution of the solution (compact form)."""
127 return self.problem_.result()
129 def get_memory_activation(self) -> list[float]:
130 """Get the activation memory per stage for simulator."""
131 return self.problem_.get_simulator_memory_activation()
133 def get_memory_parameter(self) -> list[float]:
134 """Get the parameter memory per stage for simulator."""
135 return self.problem_.get_simulator_memory_parameter()
137 def get_fw_time(self) -> list[float]:
138 """Get the forward time per stage for simulator."""
139 time = self.problem_.get_simulator_forward_time()
140 return time
142 def get_recompute_time(self) -> list[float]:
143 """Get the recompute time per stage for simulator."""
144 time = self.problem_.get_simulator_recompute_time()
145 return time
147 def get_time(self) -> list[float]:
148 """Get the time per stage for simulator."""
149 return self.problem_.get_simulator_time()
151 def naive_layer_per_stage(self,
152 layer_num: int,
153 num_of_interleave: int = 1) -> List[List[int]]:
154 """Return the naive layer-to-stage assignment (``layer_num`` evenly split)."""
155 logger.output("layer_num = %s", layer_num)
156 layer_count = layer_num // (self.num_of_stage_ * num_of_interleave)
157 return [[layer_count] * self.num_of_stage_ for _ in range(num_of_interleave)]
159 def print_yaml_results(self) -> None:
160 """Log the solver output in the MindFormers YAML schema."""
162 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
163 nass = self.naive_layer_per_stage(layer.nb_layer_,
164 self.num_of_interleave_)
165 yaml_format = Recompute.yaml_from_internal(
166 self.num_of_interleave_,
167 self.num_of_stage_,
168 self.problem_.variables_[layer.name_],
169 nass,
170 )
171 logger.output("layer-to-stage assignment baseline is \n\t%s", nass)
172 yaml_results = "\nTo put in yaml configuration:"
173 for y, v in yaml_format.items():
174 yaml_results += f"\n\t{y}: {v}"
175 logger.output(yaml_results)
177 def get_manual_memory_activation(
178 self,
179 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
180 interleave_num: int = 1) -> List[List[float]]:
181 """Return the per-stage activation memory for a user-supplied layer assignment."""
182 memory_active = []
183 if self.has_some_memory_info():
184 for inter in range(interleave_num):
185 memory_active.append([])
186 for stage in range(self.num_of_stage_):
187 memory_activation = 0
188 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
189 memory_activation += self._get_layer_memory_activation(
190 each_layer_per_recompute, layer, inter, stage
191 )
192 memory_active[inter].append(memory_activation)
193 return memory_active
195 @staticmethod
196 def _get_layer_memory_activation(each_layer_per_recompute, layer, interleave, stage):
197 """Calculate activation memory for one layer at one pipeline position."""
198 memory_activation = 0
199 unused_recompute_list = Recompute.get_unused_list(each_layer_per_recompute[layer])
200 for rec in Recompute.TYPE:
201 if rec in unused_recompute_list:
202 continue
203 value = each_layer_per_recompute[layer][rec][interleave][stage]
204 if value > 0:
205 memory_activation += value * layer.memory_activation_rec_[rec]
206 return memory_activation
208 def get_manual_memory_parameter(
209 self,
210 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
211 interleave_num: int = 1) -> List[List[float]]:
212 """Return the per-stage parameter memory for a user-supplied layer assignment."""
213 memory_param_stage = [0] * self.num_of_stage_
214 for inter in range(interleave_num):
215 for stage in range(self.num_of_stage_):
216 for rec in Recompute.TYPE:
217 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
218 if layer.memory_parameter_ is None:
219 continue
221 if rec in Recompute.get_unused_list(each_layer_per_recompute[layer]):
222 continue
224 value = each_layer_per_recompute[layer][rec][inter][stage]
225 if value <= 0:
226 continue
228 memory_param_stage[stage] += value * layer.memory_parameter_
229 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
230 if head.memory_parameter_ is not None:
231 memory_param_stage[0] += head.memory_parameter_
232 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
233 if tail.memory_parameter_ is not None:
234 memory_param_stage[self.num_of_stage_ -
235 1] += tail.memory_parameter_
236 memory_param = [memory_param_stage] * interleave_num
237 return memory_param
239 def get_manual_time(
240 self,
241 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
242 interleave_num: int = 1) -> List[List[float]]:
243 """Return the per-stage execution time for a user-supplied layer assignment."""
244 time = []
245 for i in range(interleave_num):
246 time.append([])
247 for s in range(self.num_of_stage_):
248 time[i].append(0)
249 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
250 for r in Recompute.TYPE:
251 if each_layer_per_recompute[layer][r][i][s] > 0:
252 time[i][s] += each_layer_per_recompute[layer][r][i][s] * (
253 layer.forward_time_ +
254 layer.backward_time_rec_[r])
256 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
257 time[0][0] += head.forward_time_ + head.backward_time_rec_[Recompute.TYPE.NONE]
258 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
259 time[interleave_num - 1][self.num_of_stage_ - 1] += (
260 tail.forward_time_
261 + tail.backward_time_rec_[Recompute.TYPE.NONE]
262 )
263 return time
265 def get_manual_fw_time(
266 self,
267 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
268 interleave_num: int = 1) -> List[List[float]]:
269 """Return the per-stage forward time for a user-supplied layer assignment."""
270 time = []
271 for i in range(interleave_num):
272 time.append([])
273 for s in range(self.num_of_stage_):
274 time[i].append(0)
275 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
276 for r in Recompute.TYPE:
277 if (r not in Recompute.get_unused_list(each_layer_per_recompute[layer])
278 and each_layer_per_recompute[layer][r][i][s] > 0):
279 time[i][s] += each_layer_per_recompute[layer][r][i][s] * (
280 layer.forward_time_)
281 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
282 time[0][0] += head.forward_time_
283 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
284 time[interleave_num - 1][self.num_of_stage_ - 1] += tail.forward_time_
285 return time
287 def get_manual_backward_time(
288 self,
289 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
290 interleave_num: int = 1) -> List[List[float]]:
291 """Return the per-stage backward time for a user-supplied layer assignment."""
292 time = []
293 for i in range(interleave_num):
294 time.append([])
295 for s in range(self.num_of_stage_):
296 time[i].append(0)
297 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
298 for r in Recompute.TYPE:
299 if (r not in Recompute.get_unused_list(each_layer_per_recompute[layer])
300 and each_layer_per_recompute[layer][r][i][s] > 0):
301 time[i][s] += each_layer_per_recompute[layer][r][i][s] * (
302 layer.backward_time_rec_[r])
303 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
304 time[0][0] += head.backward_time_rec_[Recompute.TYPE.NONE]
305 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
306 time[interleave_num - 1][self.num_of_stage_ - 1] += (
307 tail.backward_time_rec_[Recompute.TYPE.NONE]
308 )
309 return time
311 def get_manual_recompute_time(
312 self,
313 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
314 interleave_num: int = 1) -> List[List[float]]:
315 """Return the per-stage recompute-only time for a user-supplied layer assignment."""
316 logger.output("each_layer_per_recompute = %s", each_layer_per_recompute)
317 time_all_rec = []
318 time_no_rec = []
319 for i in range(interleave_num):
320 time_all_rec.append([])
321 time_no_rec.append([])
322 for s in range(self.num_of_stage_):
323 time_all_rec[i].append(0)
324 time_no_rec[i].append(0)
325 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
326 self._add_manual_recompute_time(
327 each_layer_per_recompute, layer, i, s, time_all_rec, time_no_rec)
329 return [[r - n for r, n in zip(ar, nr)]
330 for ar, nr in zip(time_all_rec, time_no_rec)]
332 def _add_manual_recompute_time(self, each_layer_per_recompute, layer, interleave, stage,
333 time_all_rec, time_no_rec):
334 """Accumulate recompute time for a single layer and stage."""
335 logger.output("backward_time_rec_(%s) = %s", layer, layer.backward_time_rec_)
336 unused_rec = Recompute.get_unused_list(each_layer_per_recompute[layer])
337 for rec in Recompute.TYPE:
338 layer_num = each_layer_per_recompute[layer][rec][interleave][stage]
339 if rec in unused_rec or layer_num <= 0:
340 continue
341 if layer.backward_time_rec_[rec] is None:
342 raise ValueError("No backward tme is specified for this "
343 "recomputation. Recomputation "
344 f"'{Recompute.YAML_NAME[rec]}' is likely not considered")
345 logger.output("r = %s; i = %s; s = %s", rec, interleave, stage)
346 time_all_rec[interleave][stage] += layer_num * layer.backward_time_rec_[rec]
347 time_no_rec[interleave][stage] += layer_num * layer.backward_time_rec_[Recompute.TYPE.NONE]
349 def simulate(self, show: bool = True, file_name: Optional[str] = None,
350 sub_fig: Optional[plt.Figure] = None, comm_time: float = 0.0) -> float:
351 """Run the simulator on the solved schedule and return its estimated total time."""
352 forward_time = self.get_fw_time()
353 recompute_overhead = self.get_recompute_time()
354 backward_time = self.problem_.get_simulator_backward_time() if self.use_backward_time_ else 0
355 stage_mem_par = 0
356 stage_mem_act = 0
357 if self.has_some_memory_info():
358 stage_mem_par = self.get_memory_parameter()
359 stage_mem_act = self.get_memory_activation()
361 return self.simulation(
362 forward_time,
363 recompute_overhead,
364 stage_mem_par,
365 stage_mem_act,
366 self.constant_memory_,
367 backward_time=backward_time,
368 show=show,
369 file_name=file_name,
370 sub_fig=sub_fig,
371 comm_time=comm_time,
372 )
374 def simulate_naive(self, layers: List[Layer], output_folder: str) -> None:
375 """Simulate the naive (even) layer-to-stage assignments for sanity comparison."""
376 num_layers = 0
377 rec_considered = {}
378 for layer in layers:
379 if layer.type_ == Layer.type_enum.BODY:
380 num_layers = layer.nb_layer_
381 rec_considered = layer.recompute_considered_
383 all_recomp = {"offset": 0}
384 no_recomp = {"offset": 0}
385 for rec in [Recompute.TYPE.FULL, Recompute.TYPE.SLCT, Recompute.TYPE.COMM]:
386 if rec_considered.get(rec, False):
387 all_recomp[Recompute.YAML_NAME[rec]] = True
388 no_recomp[Recompute.YAML_NAME[rec]] = False
390 self.simulate_yaml(
391 yaml_format=all_recomp,
392 show=True,
393 interleave_num=self.num_of_interleave_,
394 file_name=os.path.join(output_folder,
395 "result_naive_all_recomp.svg"),
396 )
398 if num_layers % self.num_of_stage_ == 0:
399 self.simulate_yaml(
400 yaml_format=no_recomp,
401 show=True,
402 interleave_num=self.num_of_interleave_,
403 file_name=os.path.join(output_folder,
404 "result_naive_no_recomp.svg"),
405 )
406 else:
407 logger.warning("num layer cannot be divided by num stage")
409 def simulate_comparison(self, manual_config_file: str, output_folder: str) -> None:
410 """Render side-by-side automatic vs manual simulations for every entry in the YAML."""
411 with open(manual_config_file, encoding="utf-8") as fp:
412 check_yaml_depth_before_loading(fp)
413 fp.seek(0)
414 data = yaml.safe_load(fp)
415 yaml_data = {}
416 for manual in data.values():
417 yaml_data[Recompute.OFFSET] = manual.get(Recompute.OFFSET)
418 if isinstance(yaml_data[Recompute.OFFSET], list) and all(
419 isinstance(item, int) for item in yaml_data[Recompute.OFFSET]):
420 yaml_data[Recompute.OFFSET] = [yaml_data[Recompute.OFFSET]]
422 for rec in Recompute.YAML_NAME.values():
423 yaml_data[rec] = manual.get(rec)
424 if isinstance(yaml_data[rec], list) and all(
425 isinstance(item, int) for item in yaml_data[rec]):
426 yaml_data[rec] = [yaml_data[rec]]
427 interleave_num = manual.get("interleave_num",
428 self.num_of_interleave_)
429 show = manual.get("show", False)
430 file_name = manual.get("file_name")
431 full_file_name = os.path.join(output_folder,
432 file_name) if (file_name) else None
434 fig = plt.figure(figsize=(24, 8))
435 sub_figs = fig.subfigures(1, 2, wspace=0.07)
436 sub_figs[0].suptitle('Automatic', fontsize='x-large')
437 try:
438 simulate_result = self.simulate(
439 show=False,
440 file_name=os.path.join(output_folder, "Auto_" + file_name),
441 sub_fig=sub_figs[0],
442 )
443 except Exception:
444 logger.exception("Failed to simulate auto pipeline.")
445 raise
447 if simulate_result is None:
448 raise RuntimeError("simulate() returned None.")
450 sub_figs[1].suptitle('Manual', fontsize='x-large')
451 self.simulate_yaml(yaml_data, False, interleave_num, full_file_name, sub_figs[1])
452 plt.savefig(os.path.join(output_folder, "Comparison_" + file_name))
453 if show:
454 plt.show()
456 def simulate_only_manual(self, manual_config_file: str, output_folder: str) -> None:
457 """Render only the manual simulation for every entry in ``manual_config_file``."""
458 with open(manual_config_file, encoding="utf-8") as fp:
459 check_yaml_depth_before_loading(fp)
460 fp.seek(0)
461 data = yaml.safe_load(fp)
462 yaml_data = {}
463 for manual in data.values():
464 yaml_data[Recompute.OFFSET] = manual.get(Recompute.OFFSET)
465 if isinstance(yaml_data[Recompute.OFFSET], list) and all(
466 isinstance(item, int) for item in yaml_data[Recompute.OFFSET]):
467 yaml_data[Recompute.OFFSET] = [yaml_data[Recompute.OFFSET]]
469 for rec in Recompute.YAML_NAME.values():
470 yaml_data[rec] = manual.get(rec)
471 if isinstance(yaml_data[rec], list) and all(
472 isinstance(item, int) for item in yaml_data[rec]):
473 yaml_data[rec] = [yaml_data[rec]]
474 interleave_num = manual.get("interleave_num",
475 self.num_of_interleave_)
476 show = manual.get("show", False)
477 file_name = manual.get("file_name")
478 full_file_name = os.path.join(output_folder,
479 file_name) if (file_name) else None
481 fig = plt.figure(figsize=(12, 8))
482 self.simulate_yaml(yaml_data, False, interleave_num, full_file_name, fig)
483 plt.savefig(os.path.join(output_folder, "manual_file_" + file_name))
484 if show:
485 plt.show()
487 def simulate_yaml(self, yaml_format: Dict[str, Any], show: bool = True,
488 interleave_num: int = 1,
489 file_name: Optional[str] = None,
490 sub_fig: Optional[plt.Figure] = None) -> float:
491 """Simulate a manual pipeline configuration encoded as a YAML-compatible dict."""
492 layer_num = 0
493 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
494 layer_num += layer.nb_layer_
495 nass = self.naive_layer_per_stage(layer_num,
496 num_of_interleave=interleave_num)
497 layer_per_recompute = Recompute.internal_from_yaml(
498 interleave_num, self.num_of_stage_, yaml_format, nass)
499 each_layer_per_recompute = self.split_layer_per_recompute(layer_per_recompute)
500 return self.simulate_manual(
501 each_layer_per_recompute,
502 show,
503 interleave_num=interleave_num,
504 file_name=file_name,
505 sub_fig=sub_fig
506 )
508 #######################################################################
509 ## ##
510 ## Print Solver Model ##
511 ## ##
512 #######################################################################
513 def _calculate_activation_memory(self, each_layer_per_recompute, v, s):
514 """Calculate activation memory for next and current stage"""
515 act_mem_next = 0
516 act_mem_curr = 0
518 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
519 for rec in Recompute.TYPE:
520 if self.problem_.recompute_considered_[rec]:
521 if each_layer_per_recompute[layer][rec][v + 1][s] > 0: # next
522 act_mem_next += (each_layer_per_recompute[layer][rec][v + 1][s] *
523 layer.memory_activation_rec_[rec])
524 if each_layer_per_recompute[layer][rec][v][s] > 0: # current
525 act_mem_curr += (each_layer_per_recompute[layer][rec][v][s] *
526 layer.memory_activation_rec_[rec])
528 return act_mem_next, act_mem_curr
530 def _compute_parameter_memory_manually_solver(self, each_layer_per_recompute, s, interleave_num=1):
531 """Solver memory model: parameter memory"""
532 param_mem = 0
533 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
534 if layer.memory_parameter_ is not None:
535 param_mem += self._calculate_layer_parameter_memory(
536 layer, each_layer_per_recompute[layer], s, interleave_num)
537 return param_mem
539 def _calculate_layer_parameter_memory(self, layer, layer_per_recompute, s, interleave_num):
540 """Calculate parameter memory for a single layer"""
541 layer_mem = 0
542 for inter in range(interleave_num):
543 for rec in Recompute.TYPE:
544 if self.problem_.recompute_considered_[rec]:
545 if layer_per_recompute[rec][inter][s] > 0:
546 layer_mem += layer_per_recompute[rec][inter][s] * layer.memory_parameter_
547 return layer_mem
549 def _calculate_activation_memory_solver(self, each_layer_per_recompute, s, interleave_num, activation_nums):
550 """Calculate activation memory for a given stage"""
551 act_mem = 0
552 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
553 for inter in range(interleave_num):
554 for rec in Recompute.TYPE:
555 if self.problem_.recompute_considered_[rec]:
556 if each_layer_per_recompute[layer][rec][inter][s] > 0:
557 act_mem += (each_layer_per_recompute[layer][rec][inter][s] *
558 layer.memory_activation_rec_[rec] *
559 activation_nums[inter][s])
560 return act_mem
563 def debug_print_manual_theoretical_memory(
564 self,
565 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]],
566 interleave_num: int = 1) -> None:
567 """Log the per-stage theoretical memory implied by the solver model (debug aid)."""
568 logger.info("%s Manual Theoretical Memory Analysis %s", "=" * 20, "=" * 20)
570 if self.vpp_less_memory_:
571 if self.seqpipe_:
572 activation_nums = self.problem_.compute_activation_seq_nums(
573 self.num_of_stage_, interleave_num, self.seq_split_num_, self.num_of_micro_batch_, True)
574 else:
575 activation_nums = self.problem_.compute_less_activation_nums(
576 self.num_of_stage_, interleave_num)
577 # Add if dual to decide whether dualpipe_v is used
578 elif self.dual_:
579 activation_nums = self.problem_.compute_activation_nums_dual(
580 self.num_of_stage_, interleave_num, self.num_of_micro_batch_)
581 else:
582 if self.seqpipe_:
583 activation_nums = self.problem_.compute_activation_seq_nums(
584 self.num_of_stage_, interleave_num, self.seq_split_num_, self.num_of_micro_batch_, False)
585 else:
586 activation_nums = self.problem_.compute_activation_nums(
587 self.num_of_stage_, interleave_num, self.num_of_micro_batch_)
589 logger.info("Activation nums = %s", activation_nums)
591 # compute for each stage
592 for s in range(self.num_of_stage_):
594 # parameter memory
595 param_mem = self._compute_parameter_memory_manually_solver(each_layer_per_recompute, s, interleave_num)
597 # head memory
598 if s == 0:
599 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
600 if head.memory_parameter_ is not None:
601 param_mem += head.memory_parameter_
603 # tail memory
604 if s == self.num_of_stage_ - 1:
605 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
606 if tail.memory_parameter_ is not None:
607 param_mem += tail.memory_parameter_
609 # act memory
610 act_mem = self._calculate_activation_memory_solver(each_layer_per_recompute, s,
611 interleave_num, activation_nums)
613 # overhead
614 overhead = 0
616 total = param_mem + act_mem + overhead + self.constant_memory_
618 logger.info("Stage %d Manual Memory Analysis:", s)
619 logger.info("Parameter Memory: %.2f", param_mem)
620 logger.info("Activation Memory: %.2f", act_mem)
621 logger.info("Memory Overhead: %.2f", overhead)
622 logger.info("Constant Memory: %.2f", self.constant_memory_)
623 logger.info("Total Theoretical Memory: %.2f", total)
625 def split_layer_per_recompute(
626 self,
627 layer_per_recompute: Dict[Recompute.TYPE, List[List[int]]]
628 ) -> Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]]:
629 """Split aggregate per-recompute layer counts into counts per BODY layer."""
630 each_layer_per_recompute = {}
631 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
632 rest = layer.nb_layer_
633 each_layer_per_recompute[layer] = {r: [] for r in Recompute.TYPE}
634 for rec in Recompute.TYPE:
635 for i in range(self.num_of_interleave_):
636 each_layer_per_recompute[layer][rec].append([0]*self.num_of_stage_)
637 for s in range(self.num_of_stage_):
638 subtract = min(layer_per_recompute[rec][i][s], rest)
639 layer_per_recompute[rec][i][s] -= subtract
640 rest -= subtract
641 each_layer_per_recompute[layer][rec][i][s] += subtract
642 return each_layer_per_recompute
644 def fuse_layer_per_recompute(
645 self,
646 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]]
647 ) -> Dict[Recompute.TYPE, List[List[int]]]:
648 """Fuse per-layer recompute counts back into aggregate per-recompute-type totals."""
649 all_layers_per_recompute = {r: [] for r in Recompute.TYPE}
650 for rec in Recompute.TYPE:
651 for i in range(self.num_of_interleave_):
652 all_layers_per_recompute[rec].append([])
653 for s in range(self.num_of_stage_):
654 all_layers_per_recompute[rec][i].append(sum(
655 each_layer_per_recompute[layer][rec][i][s]
656 for layer in self.layers_sorted_[Layer.type_enum.BODY]
657 ))
658 return all_layers_per_recompute
661 def simulate_manual(
662 self,
663 each_layer_per_recompute: Optional[Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]]] = None,
664 show: bool = True,
665 interleave_num: int = 1,
666 file_name: Optional[str] = None,
667 sub_fig: Optional[plt.Figure] = None) -> float:
668 """Run the simulator on a user-supplied per-layer recompute strategy."""
669 logger.output("Simulating given strategy: %s", each_layer_per_recompute)
671 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
672 for rec in Recompute.TYPE:
673 if len(each_layer_per_recompute[layer][rec]) != interleave_num:
674 logger.error(
675 "For layer %s with recompute %s, %s does not match interleave number %s",
676 layer,
677 rec,
678 len(each_layer_per_recompute[layer][rec]),
679 interleave_num,
680 )
681 return sys.maxsize
683 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
684 for rec in Recompute.TYPE:
685 if any(x < 0 for sublist in each_layer_per_recompute[layer][rec]
686 for x in sublist):
687 raise ValueError(
688 f"for {rec}, there is strategy less than 0 in "
689 f"{each_layer_per_recompute[layer][rec]}"
690 )
692 forward_time = self.get_manual_fw_time(each_layer_per_recompute,
693 interleave_num)
694 recompute_overhead = self.get_manual_recompute_time(
695 each_layer_per_recompute, interleave_num)
696 backward_time = (
697 self.get_manual_backward_time(
698 each_layer_per_recompute, interleave_num)
699 if self.use_backward_time_
700 else 0
701 )
702 stage_mem_par = 0
703 stage_mem_act = 0
704 if self.has_some_memory_info():
705 stage_mem_par = self.get_manual_memory_parameter(
706 each_layer_per_recompute, interleave_num=interleave_num)
707 stage_mem_act = self.get_manual_memory_activation(
708 each_layer_per_recompute, interleave_num=interleave_num)
710 self.debug_print_manual_theoretical_memory(each_layer_per_recompute, interleave_num)
712 return self.simulation(
713 forward_time,
714 recompute_overhead,
715 stage_mem_par,
716 stage_mem_act,
717 constant_mem=self.constant_memory_,
718 backward_time=backward_time,
719 show=show,
720 file_name=file_name,
721 sub_fig=sub_fig,
722 comm_time=0.0,
723 )
725 def simulation(
726 self,
727 forward_time: List[List[float]],
728 recompute_overhead: Union[int, List[List[float]]] = 0,
729 stage_mem_par: Union[int, List[List[float]]] = 0,
730 stage_mem_act: Union[int, List[List[float]]] = 0,
731 constant_mem: int = 0,
732 backward_time: Union[int, List[List[float]]] = 0,
733 show: bool = True,
734 file_name: Optional[str] = None,
735 sub_fig: Optional[plt.Figure] = None,
736 comm_time: float = 0.0,
737 ) -> float:
738 """Run the low-level :class:`PipelineSimulator` and return its reported end time."""
739 use_comm = comm_time > 0.0
740 if self.has_some_memory_info():
741 logger.output(
742 "PipelineSimulator(\n\t%s, %s,"
743 "\n\tblock_mem_act=%s,"
744 "\n\tblock_mem_par=%s,"
745 "\n\tlayer_recompute=%s,"
746 "\n\tbackward_time=%s,"
747 "\n\tless_memory=%s )",
748 forward_time,
749 self.num_of_micro_batch_,
750 stage_mem_act,
751 stage_mem_par,
752 recompute_overhead,
753 backward_time,
754 self.vpp_less_memory_,
755 )
757 sim_method = "vpp2" if self.vpp_less_memory_ else "vpp"
758 simulator = sim.PipelineSimulator(
759 forward_time,
760 self.num_of_micro_batch_,
761 comm_time=comm_time,
762 block_mem=stage_mem_act,
763 block_mem_par=stage_mem_par,
764 constant_mem=constant_mem,
765 layer_recompute=recompute_overhead,
766 backward_time=backward_time,
767 method=sim_method,
768 sub_fig=sub_fig
769 )
770 else:
771 logger.output(
772 "PipelineSimulator(\n\t%s, %s,"
773 "\n\tlayer_recompute=%s,"
774 "\n\tbackward_time=%s,"
775 "\n\tless_memory=%s )",
776 forward_time,
777 self.num_of_micro_batch_,
778 recompute_overhead,
779 backward_time,
780 self.vpp_less_memory_,
781 )
782 simulator = sim.PipelineSimulator(
783 forward_time,
784 self.num_of_micro_batch_,
785 comm_time=comm_time,
786 layer_recompute=recompute_overhead,
787 backward_time=backward_time,
788 less_memory=self.vpp_less_memory_,
789 sub_fig=sub_fig
790 )
792 simulator.run(comm=use_comm)
793 self._simulator = simulator
794 if file_name:
795 simulator.save(file_name)
796 if show:
797 simulator.show()
798 return simulator.end_time
800 def _construct_problem_pulp_(self) -> SappSolver:
801 """construct the problem using pulp"""
802 prob = SappSolver(
803 num_of_stage=self.num_of_stage_,
804 num_of_micro_batch=self.num_of_micro_batch_,
805 num_of_interleave=self.num_of_interleave_,
806 max_memory=self.max_memory_,
807 vpp_less_memory=self.vpp_less_memory_,
808 # Add arg dual
809 dual = self.dual_,
810 constant_memory=self.constant_memory_,
811 layers=self.layers_,
812 layers_sorted=self.layers_sorted_,
813 optimization_level=self.optimization_level,
814 extracted_training_params=self.extracted_training_params_,
815 seq_split_num=self.seq_split_num_
816 )
817 return prob
819 def _recompute_considered(self):
820 return self.problem_.recompute_considered_
823def choose_interleave(
824 model_name: str,
825 number_of_stage: int,
826 number_of_micro_batch: int,
827 max_memory: int,
828 layers: list[Layer],
829) -> tuple[int, int, dict[str, list[list[str]]]]:
830 """Simulates different interleaves and returns the best."""
831 max_inter = 4
832 best_time = int(sys.maxsize)
833 best_inter = 1
834 best_distribution = {}
836 for inter in range(1, max_inter + 1):
837 pipe = SappPipeline(
838 model_name=model_name,
839 num_of_stage=number_of_stage,
840 num_of_micro_batch=number_of_micro_batch,
841 max_memory=max_memory,
842 layers=layers,
843 num_of_interleave=inter,
844 )
846 pipe.construct_problem(solver="pulp")
847 pipe.solve_problem()
848 time = pipe.simulate(show=False)
849 logger.output("for interleave %s, time = %s", inter, time)
850 if time < best_time:
851 best_time = time
852 best_inter = inter
853 best_distribution = pipe.get_result()
855 return (best_inter, best_time, best_distribution)
858def flatten(inter_stage_list: List[List[float]]) -> List[float]:
859 """Collapse an ``[interleave][stage]`` matrix into a per-stage list via summation."""
860 stage_list = [0] * len(inter_stage_list[0])
861 for inter, _ in enumerate(inter_stage_list):
862 for stage, _ in enumerate(inter_stage_list[inter]):
863 stage_list[stage] += inter_stage_list[inter][stage]
864 return stage_list