Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_ppb / sapp / sapp_solver.py: 90%
877 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"""Solver Class"""
17import os
18from dataclasses import dataclass
19from enum import IntEnum
20from typing import Any, Dict, List, Optional
22import pulp as lpSolver
24import hyper_parallel.auto_parallel.sapp_ppb.utils.recompute as Recompute
25from hyper_parallel.auto_parallel.sapp_ppb.utils.layer import Layer
26from hyper_parallel.auto_parallel.sapp_ppb.utils.logger import logger
28# seqpipe const
29TENSOR_FLOAT_16 = 2
30TENSOR_FLOAT_32 = 4
31const_from_byte_to_mb = 1024 * 1024
32# llama intermideate_size
33LLAMA_INTERMEDIATE_SIZE = 11008
36@dataclass
37class PipelineMemoryConstraint:
38 """constraint struct"""
39 prob: Any
40 variables: Any
41 layers_sorted: dict[Any]
42 num_of_stage: int
43 num_of_interleave: int
44 micro_batch: int
45 memory_limit: int
48class SappSolver:
49 """solver for pipeline balance"""
51 BIG_M = 1000000
53 MEM_OVERHEAD_NAME = "memory_overhead"
54 TOTAL_SUM = "var_sum_FPi_BPi"
55 CHUNKS_SUM = "chunks_sum"
56 PREV_DIFF = "prev_diff"
57 NEXT_DIFF = "next_diff"
58 MAX_STAGE_TIME = "max_stage_time"
59 MAX_LAST_CHUNK = "max_last_chunk"
60 LAYER_FRONTIER = "layer_frontier"
61 REC_FRONTIER = "recompute_frontier"
62 PROP_PHASE = IntEnum("Propagation", ["FW", "BW"], start=0)
64 def __init__(
65 self,
66 num_of_stage: int,
67 num_of_interleave: int,
68 num_of_micro_batch: int,
69 max_memory: int,
70 layers: list[Layer],
71 layers_sorted: dict[Layer.type_enum, list[Layer]],
72 vpp_less_memory: bool = False,
73 # add dualpipe_v arg
74 dual: bool = False,
75 constant_memory: int = 0,
76 optimization_level: int = 1,
77 description: str = "Pipeline_execution_time_minimize",
78 extracted_training_params: dict[str, int] = None,
79 seq_split_num: int = 1,
80 ) -> None:
81 """Build the ILP variables and the empty problem skeleton.
83 Args:
84 num_of_stage: Number of physical pipeline stages.
85 num_of_interleave: Virtual-pipeline (VPP) chunk count.
86 num_of_micro_batch: Number of micro-batches.
87 max_memory: Per-device memory budget (MB).
88 layers: Flat list of :class:`Layer` descriptors covering the full model.
89 layers_sorted: ``layers`` indexed by HEAD / BODY / TAIL classification.
90 vpp_less_memory: Use the less-memory VPP scheduler variant.
91 dual: Enable dualpipe-V scheduling support.
92 constant_memory: Constant per-stage memory overhead (MB).
93 optimization_level: Solver optimization level (``0-2``).
94 description: Problem description used when exporting the LP model.
95 extracted_training_params: Optional training params for sequence-pipeline mode.
96 seq_split_num: Number of sequence splits (``>1`` enables sequence pipeline).
97 """
99 self.num_of_stage_ = num_of_stage
100 self.num_of_interleave_ = num_of_interleave
101 self.num_of_micro_batch_ = num_of_micro_batch
102 self.max_memory_ = max_memory
103 self.vpp_less_memory_ = vpp_less_memory
104 # Add dualpipe_v
105 self.dual_ = dual
106 self.constant_memory_ = constant_memory
107 self.optimization_level_ = optimization_level
108 self.layers_ = layers
109 self.layers_sorted_ = layers_sorted
111 self.recompute_considered_ = self.find_recompute_considered(
112 layers_sorted)
113 self.extracted_training_params_ = extracted_training_params
114 self.seq_split_num_ = seq_split_num
115 self.seq_pipe = self.seq_split_num_ > 1
116 if self.seq_pipe:
117 self._initialize_seq_pipe_layers()
119 self.variables_ = self._create_variables_to_solve_(
120 num_of_stage, num_of_interleave, layers_sorted)
121 self.problem_ = self._create_problem_(description)
123 def _initialize_seq_pipe_layers(self):
124 """Update memory and time metadata for sequence pipeline mode."""
125 self._update_seq_pipe_memory()
126 self.num_of_micro_batch_ *= self.seq_split_num_
127 self._update_seq_pipe_time()
129 def _update_seq_pipe_memory(self):
130 """Update layer memory values for sequence pipeline mode."""
131 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
132 self._update_body_seq_memory(layer)
133 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
134 self._update_head_seq_memory(head)
135 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
136 self._update_tail_seq_memory(tail)
138 def _update_body_seq_memory(self, layer):
139 """Update body layer memory values for sequence pipeline mode."""
140 if layer.memory_parameter_ is not None:
141 logger.info("Body Layer 1f1b Parameter Memory: %s", layer.memory_parameter_)
142 layer.memory_parameter_ = self.compute_seq_mem_parameter(
143 layer.memory_parameter_, self.extracted_training_params_)
144 logger.info("Body Layer Seq Parameter Memory: %s", layer.memory_parameter_)
145 for rec in Recompute.TYPE:
146 if not self.recompute_considered_[rec]:
147 continue
148 if rec.name == "FULL":
149 self.recompute_considered_[rec] = False
150 layer.recompute_considered_[rec] = False
151 layer.memory_activation_rec_[rec] = None
152 logger.error("Seqpipe doesn't support full recomputation, "
153 "recompute_activation is set as None for seqpp")
154 continue
155 logger.info(
156 "Body Layer 1f1b %s activation Memory: %s",
157 rec,
158 layer.memory_activation_rec_[rec],
159 )
160 layer.memory_activation_rec_[rec] = self.compute_seq_mem_activation(
161 layer.memory_activation_rec_[rec],
162 self.extracted_training_params_,
163 self.seq_split_num_
164 )
165 logger.info(
166 "Body Layer seq %s activation Memory: %s",
167 rec,
168 layer.memory_activation_rec_[rec],
169 )
171 def _update_head_seq_memory(self, head):
172 """Update head layer memory values for sequence pipeline mode."""
173 if head.memory_parameter_ is None:
174 return
175 logger.info("Head cost 1f1b: %s", head.memory_parameter_)
176 head.memory_parameter_ = self.compute_seq_mem_head_cost(
177 head.memory_parameter_, self.extracted_training_params_, self.seq_split_num_)
178 logger.info("Head cost Seq: %s", head.memory_parameter_)
180 def _update_tail_seq_memory(self, tail):
181 """Update tail layer memory values for sequence pipeline mode."""
182 if tail.memory_parameter_ is None:
183 return
184 logger.info("Tail cost 1f1b: %s", tail.memory_parameter_)
185 tail.memory_parameter_ = self.compute_seq_mem_tail_cost(
186 tail.memory_parameter_, self.extracted_training_params_, self.seq_split_num_)
187 logger.info("Tail cost seq: %s", tail.memory_parameter_)
189 def _update_seq_pipe_time(self):
190 """Update layer times for sequence pipeline mode."""
191 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
192 self._update_layer_seq_time(layer, "Body")
193 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
194 self._update_layer_seq_time(head, "Head")
195 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
196 self._update_layer_seq_time(tail, "Tail")
198 def _update_layer_seq_time(self, layer, layer_name):
199 """Scale one layer's time by the sequence split number."""
200 logger.info("%s Layer 1f1b fp time: %s", layer_name, layer.forward_time_)
201 logger.info("%s Layer 1f1b bp time:", layer_name)
202 for key, value in layer.backward_time_rec_.items():
203 logger.output("%s: %s", key, value)
204 layer.time_ = layer.time_ / self.seq_split_num_
205 layer.forward_time_ = layer.forward_time_ / self.seq_split_num_
206 layer.update_internal_time_for_seqpp()
207 logger.info("%s Layer seq fp time: %s", layer_name, layer.forward_time_)
208 logger.info("%s Layer seq bp time:", layer_name)
209 for key, value in layer.backward_time_rec_.items():
210 logger.output("%s: %s", key, value)
212 @staticmethod
213 def compute_forward_in_backward(num_of_stage: int,
214 micro_batch: int) -> list[int]:
215 """Computes the number of forward propagation happening after a backward"""
216 n = num_of_stage - 1
217 factors = []
218 for _ in range(num_of_stage):
219 factors.append(abs(n))
220 n -= 2
221 if micro_batch < 2 * num_of_stage:
222 for i in range(num_of_stage // 2):
223 factors[i] = 0
224 return factors
226 @staticmethod
227 def compute_lm_forward_in_backward(num_of_stage: int) -> list[int]:
228 """Function compute_forward_in_backward in less_memory schedule"""
229 return list(range(num_of_stage))
231 @staticmethod
232 def compute_activation_nums(num_of_stage: int, num_of_interleave: int,
233 micro_batch: int) -> list[list[int]]:
234 """compute the number of activation"""
235 activation_nums = []
237 if num_of_interleave > 1:
238 for i in range(num_of_interleave):
239 activation_nums.append([])
240 for _ in range(num_of_stage):
241 activation_nums[i].append(num_of_stage)
242 for s in range(num_of_stage):
243 activation_nums[0][s] += max(0, num_of_stage - 2 * s - 1)
244 for s in range(num_of_stage):
245 activation_nums[num_of_interleave - 1][s] += min(
246 0, num_of_stage - 2 * s - 1)
247 for i in range(num_of_interleave):
248 for s in range(num_of_stage):
249 activation_nums[i][s] = min(activation_nums[i][s],
250 micro_batch)
251 else:
252 for i in range(num_of_interleave):
253 activation_nums.append([])
254 for s in range(num_of_stage):
255 activation_nums[i].append(num_of_stage - s)
257 return activation_nums
259 @staticmethod
260 def compute_activation_nums_dual(num_of_stage: int, num_of_interleave: int,
261 micro_batch: int) -> list[list[int]]:
262 """compute the number of activation for dualpipe_v"""
263 activation_nums = []
265 for i in range(num_of_interleave):
266 activation_nums.append([])
267 for _ in range(num_of_stage):
268 activation_nums[i].append(0)
269 for s in range(num_of_stage):
270 activation_nums[0][s] += max(0, 2 * num_of_stage - s)
271 for s in range(num_of_stage):
272 activation_nums[num_of_interleave - 1][s] += max(
273 0, s + 1)
274 for i in range(num_of_interleave):
275 for s in range(num_of_stage):
276 activation_nums[i][s] = min(activation_nums[i][s],
277 micro_batch)
279 return activation_nums
281 @staticmethod
282 def compute_less_activation_nums(
283 num_of_stage: int, num_of_interleave: int) -> list[list[int]]:
284 """compute number of less_mem activation"""
285 activation_nums = []
286 if num_of_interleave > 1:
287 for i in range(num_of_interleave):
288 activation_nums.append([])
289 for _ in range(num_of_stage):
290 activation_nums[i].append(num_of_stage)
291 for s in range(num_of_stage):
292 activation_nums[num_of_interleave - 1][s] -= s
293 else:
294 for i in range(num_of_interleave):
295 activation_nums.append([])
296 for s in range(num_of_stage):
297 activation_nums[i].append(num_of_stage - s)
298 return activation_nums
300 #######################################################################
301 ## ##
302 ## SeqPipe ##
303 ## ##
304 #######################################################################
305 @staticmethod
306 def _compute_activation_seq_interleave(num_of_stage, num_of_interleave,
307 seq_split_num, micro_batch, act_gap):
308 """compute activation for seq chunks when num_of_interleave > 1."""
309 activation_nums = []
310 for i in range(num_of_interleave):
311 activation_nums.append([])
312 for _ in range(num_of_stage):
313 activation_nums[i].append(num_of_stage)
314 for s in range(num_of_stage):
315 activation_nums[num_of_interleave - 1][s] = seq_split_num
317 loop_index = 1
318 for stage_index in range(num_of_stage - 2, -1, -1):
319 flag_added = False
320 for chunk_index in range(num_of_interleave):
321 condition1 = activation_nums[chunk_index][stage_index + 1] % num_of_stage != 0
322 condition2 = activation_nums[chunk_index][stage_index + 1] // num_of_stage < loop_index
323 if condition1 or condition2:
324 for update in range(stage_index + 1):
325 activation_nums[chunk_index][update] += act_gap
326 flag_added = True
327 break
328 if not flag_added:
329 for update in range(stage_index + 1):
330 activation_nums[0][update] += act_gap
331 loop_index += 1
332 for i in range(num_of_interleave):
333 for s in range(num_of_stage):
334 activation_nums[i][s] = min(activation_nums[i][s], micro_batch)
335 return activation_nums
337 @staticmethod
338 def compute_activation_seq_nums(num_of_stage: int, num_of_interleave: int,
339 seq_split_num: int, micro_batch: int, less_memory: False) -> list[list[int]]:
340 """compute the number of activation for seq chunks"""
341 act_gap = 1 if less_memory else 2
342 if num_of_interleave > 1:
343 activation_nums = SappSolver._compute_activation_seq_interleave(
344 num_of_stage, num_of_interleave, seq_split_num, micro_batch, act_gap)
345 else:
346 activation_nums = []
347 for i in range(num_of_interleave):
348 activation_nums.append([])
349 for s in range(num_of_stage):
350 activation_nums[i].append(num_of_stage - s + seq_split_num - 1)
352 logger.output("compute_activation_seq_nums: %s", activation_nums)
353 return activation_nums
355 @staticmethod
356 def compute_seq_mem_activation(original_memory_activation: float,
357 extracted_training_params: dict[str, int],
358 seq_split_num: int) -> float:
359 """compute activation memory for seqpipe"""
360 # context parallel? cp?
361 batch_size = extracted_training_params['batch_size']
362 heads = extracted_training_params['num_heads']
363 seq_length = extracted_training_params['seq_length']
364 head_dim = extracted_training_params['head_dim']
365 mp = extracted_training_params['model_parallel']
366 # 2*Kv add
367 kv_update_mem_byte = 2 * ((TENSOR_FLOAT_16 * batch_size * heads * seq_length * head_dim) / (mp))
368 kv_update_mem = kv_update_mem_byte / const_from_byte_to_mb
369 # Attention Key,Value
370 # cp?
371 key_mem_byte = (TENSOR_FLOAT_16 * batch_size * heads * seq_length * head_dim) / (mp)
372 key_mem = key_mem_byte / const_from_byte_to_mb
373 # cp?
374 value_mem_byte = (TENSOR_FLOAT_16 * batch_size * heads * seq_length * head_dim) / (mp)
375 value_mem = value_mem_byte / const_from_byte_to_mb
377 seq_memory_activation = (original_memory_activation - key_mem - value_mem) / seq_split_num + kv_update_mem
378 return seq_memory_activation
380 @staticmethod
381 def compute_seq_mem_parameter(original_memory_parameter: float, extracted_training_params: dict[str, int]) -> float:
382 """compute layer parameter memory for seqpipe"""
383 # context parallel? cp?
384 batch_size = extracted_training_params['batch_size']
385 heads = extracted_training_params['num_heads']
386 seq_length = extracted_training_params['seq_length']
387 head_dim = extracted_training_params['head_dim']
388 mp = extracted_training_params['model_parallel']
389 kv_cache_parameter_mem_byte = 4 * (TENSOR_FLOAT_16 * batch_size * heads * seq_length * head_dim / (mp))
390 kv_cache_parameter_mem = kv_cache_parameter_mem_byte / const_from_byte_to_mb
391 seq_memory_parameter = original_memory_parameter + kv_cache_parameter_mem
392 return seq_memory_parameter
394 @staticmethod
395 def compute_seq_mem_head_cost(original_head_cost: float,
396 extracted_training_params: dict[str, int],
397 seq_split_num: int) -> float:
398 """compute head stage extra cost for seqpipe"""
399 batch_size = extracted_training_params['batch_size']
400 seq_length = extracted_training_params['seq_length']
401 hidden_size = extracted_training_params['hidden_size']
402 mp = extracted_training_params['model_parallel']
403 if mp > 1:
404 # comm operator Mem (recv+reduceScatter)
405 # cp?
406 comm_operator_mem_byte = 2 * (TENSOR_FLOAT_16 * batch_size * seq_length * hidden_size / (mp))
407 comm_operator_mem = comm_operator_mem_byte / const_from_byte_to_mb
408 # StridedSliceGrad Operator Mem
409 stridslice_operator_mem_byte = TENSOR_FLOAT_16 * batch_size * seq_length * hidden_size
410 stridslice_operator_mem = stridslice_operator_mem_byte / const_from_byte_to_mb
411 seq_head_cost = original_head_cost - (1 - 1 / seq_split_num) * (comm_operator_mem + stridslice_operator_mem)
412 else:
413 # comm operator Mem (recv)
414 # cp?
415 comm_operator_mem_byte = TENSOR_FLOAT_16 * batch_size * seq_length * hidden_size / (mp)
416 comm_operator_mem = comm_operator_mem_byte / const_from_byte_to_mb
417 # Grad/MatMul // Grad/Mul Operator Mem
418 # cp?
419 mul_operator_mem_byte = 1 * (TENSOR_FLOAT_16 * batch_size * seq_length * LLAMA_INTERMEDIATE_SIZE / (mp))
420 mul_operator_mem = mul_operator_mem_byte / const_from_byte_to_mb
421 seq_head_cost = original_head_cost - (1 - 1 / seq_split_num) * (comm_operator_mem + mul_operator_mem)
422 return seq_head_cost
424 @staticmethod
425 def compute_seq_mem_tail_cost(original_tail_cost: float,
426 extracted_training_params: dict[str, int],
427 seq_split_num: int) -> float:
428 """compute tail stage extra cost for seqpipe"""
429 batch_size = extracted_training_params['batch_size']
430 seq_length = extracted_training_params['seq_length']
431 vocab_size = extracted_training_params['vocab_size']
432 mp = extracted_training_params['model_parallel']
433 # Memory extra introduced by loss op:
434 loss_operator_mem_byte = TENSOR_FLOAT_32 * batch_size * seq_length * vocab_size / (mp)
435 loss_operator_mem = loss_operator_mem_byte / const_from_byte_to_mb
436 # New tail Cost = Old tail Cost - (3-3/k)M + (k-1)(M/k)
437 seq_tail_cost = original_tail_cost - (3 - 3 / seq_split_num) * loss_operator_mem + (
438 seq_split_num - 1) * (loss_operator_mem / seq_split_num)
439 return seq_tail_cost
441 def add_total_nb_layer_constraint(self, prob: Any, variables: Any,
442 sorted_layers: Dict[Layer.type_enum, list[Layer]]) -> Any:
443 """Enforce that the sum of assigned layers equals ``layer.nb_layer_`` per BODY layer."""
444 for layer in sorted_layers[Layer.type_enum.BODY]:
445 prob += (lpSolver.lpSum(
446 variables[layer.name_][rec] for rec in Recompute.TYPE
447 if self.recompute_considered_[rec]) == layer.nb_layer_)
448 return prob
450 def add_stage_nb_layer_constraint(self, prob: Any, variables: Any,
451 sorted_layers: Dict[Layer.type_enum, List[Layer]]) -> Any:
452 """Require each non-reserved ``(interleave, stage)`` cell to host at least one layer."""
453 layer_type_num = len(sorted_layers[Layer.type_enum.BODY])
454 reserved_positions = self._reserved_stage_positions()
455 for i in range(self.num_of_interleave_):
456 for s in range(self.num_of_stage_):
457 if (i, s) in reserved_positions:
458 continue
459 terms = []
460 for ll in range(layer_type_num):
461 body_layer = sorted_layers[Layer.type_enum.BODY][ll]
462 for rec in Recompute.TYPE:
463 if not self.recompute_considered_[rec]:
464 continue
466 terms.append(
467 variables[
468 body_layer.name_
469 ][rec][i][s]
470 )
472 prob += lpSolver.lpSum(terms) >= 1
473 return prob
475 def _reserved_stage_positions(self):
476 """Return stage positions reserved for head and tail layers."""
477 if self.dual_:
478 return {(0, 0), (1, 0)}
479 return {(0, 0), (self.num_of_interleave_ - 1, self.num_of_stage_ - 1)}
481 def add_multimodal_sequence_constraint(
482 self, prob: Any, variables: Any,
483 sorted_layers: Dict[Layer.type_enum, List[Layer]]) -> Any:
484 """Enforce a stage frontier between successive BODY layer types (multimodal models)."""
485 for frontier in range(1, len(sorted_layers[Layer.type_enum.BODY])):
486 layer = sorted_layers[Layer.type_enum.BODY][frontier].name_
487 for v in range(self.num_of_interleave_):
488 for s in range(self.num_of_stage_):
489 prob = self._add_frontier_lower_bound(prob, variables, layer, frontier, v, s)
490 return self._add_frontier_upper_bounds(prob, variables, sorted_layers)
492 def _add_frontier_lower_bound(self, prob, variables, layer, frontier, interleave, stage):
493 """Add the lower bound for one multimodal frontier variable."""
494 frontier_sum = self._frontier_layer_sum(variables, layer, interleave, stage)
495 if frontier_sum is None:
496 return prob
497 prob += (
498 variables[self.LAYER_FRONTIER][frontier - 1][interleave][stage]
499 >= frontier_sum / self.BIG_M
500 )
501 return prob
503 def _frontier_layer_sum(self, variables, layer, interleave, stage):
504 """Build the layer sum used by multimodal frontier constraints."""
505 if self.dual_:
506 return self._dual_frontier_layer_sum(variables, layer, interleave, stage)
507 return self._current_layer_sum(variables, layer, interleave, range(stage)) + (
508 self._previous_layer_sum(variables, layer, interleave)
509 )
511 def _dual_frontier_layer_sum(self, variables, layer, interleave, stage):
512 """Build the layer sum for dualpipe_v multimodal frontier constraints."""
513 if interleave == 0:
514 return self._current_layer_sum(variables, layer, interleave, range(stage))
515 if interleave == 1:
516 return self._current_layer_sum(variables, layer, interleave, range(stage, self.num_of_stage_)) + (
517 self._previous_layer_sum(variables, layer, interleave)
518 )
519 return None
521 def _current_layer_sum(self, variables, layer, interleave, stage_range):
522 """Sum current interleave variables over a stage range."""
523 terms = []
524 for rec in Recompute.TYPE:
525 if not self.recompute_considered_[rec]:
526 continue
528 for stage in stage_range:
529 terms.append(variables[layer][rec][interleave][stage])
531 return lpSolver.lpSum(terms)
533 def _previous_layer_sum(self, variables, layer, interleave):
534 """Sum variables from previous interleaves."""
535 terms = []
536 for rec in Recompute.TYPE:
537 if self.recompute_considered_[rec]:
538 for prev_interleave in range(interleave):
539 for stage in range(self.num_of_stage_):
540 terms.append(variables[layer][rec][prev_interleave][stage])
541 return lpSolver.lpSum(terms)
543 def _add_frontier_upper_bounds(self, prob, variables, sorted_layers):
544 """Prevent previous body layer types after each multimodal frontier."""
545 for frontier in range(1, len(sorted_layers[Layer.type_enum.BODY])):
546 layer = sorted_layers[Layer.type_enum.BODY][frontier - 1].name_
547 for stage in range(self.num_of_stage_):
548 for interleave in range(self.num_of_interleave_):
549 prob = self._add_frontier_upper_bound(prob, variables, layer, frontier, interleave, stage)
550 return prob
552 def _add_frontier_upper_bound(self, prob, variables, layer, frontier, interleave, stage):
553 """Add one upper bound constraint for a multimodal frontier."""
554 for rec in Recompute.TYPE:
555 if self.recompute_considered_[rec]:
556 prob += variables[layer][rec][interleave][stage] <= (
557 1 - variables[self.LAYER_FRONTIER][frontier - 1][interleave][stage]
558 ) * self.BIG_M
559 return prob
561 def add_multimodal_recompute_constraint(
562 self, prob: Any, variables: Any,
563 sorted_layers: Dict[Layer.type_enum, List[Layer]]) -> Any:
564 """Keep recomputation schemes consistent across BODY layer types (MindFormer constraint)."""
566 considered = Recompute.get_used_list(self.recompute_considered_)
567 if len(considered) > 2:
568 logger.error("Careful: MindFormer does not allow a fine recomputation scheme "
569 "for heterogeneous models. Pipeline balancing is currently unable to "
570 "comply with MF constraint for more than 1 recomputation type.")
571 return prob
573 if len(considered) < 2:
574 # this constraint is unnecessary if there is no recomputation
575 return prob
577 most_rec = max(considered)
578 layer_type_num = len(sorted_layers[Layer.type_enum.BODY])
579 for v in range(self.num_of_interleave_):
580 for s in range(self.num_of_stage_):
581 for rec in Recompute.TYPE:
582 if self.recompute_considered_[rec] and rec is not Recompute.TYPE.NONE:
583 for layer_idx in range(0, layer_type_num - 1):
584 prob += variables[self.REC_FRONTIER][v][s][layer_idx] >= (
585 lpSolver.lpSum(
586 variables[sorted_layers[Layer.type_enum.BODY][next_idx].name_][most_rec][v][s]
587 for next_idx in range(layer_idx + 1, layer_type_num))) / self.BIG_M
589 least_rec = min(considered)
590 for layer_idx in range(0, layer_type_num - 1):
591 layer_name = sorted_layers[Layer.type_enum.BODY][layer_idx].name_
592 for v in range(0, self.num_of_interleave_):
593 for s in range(0, self.num_of_stage_):
594 prob += variables[layer_name][least_rec][v][s] <= (
595 1 - variables[self.REC_FRONTIER][v][s][layer_idx]
596 ) * self.BIG_M
597 return prob
599 @staticmethod
600 def find_recompute_considered(
601 layers_sorted: Dict[Layer.type_enum, List[Layer]]) -> Dict[Recompute.TYPE, bool]:
602 """Return the recomputation-considered flags copied from the first BODY layer.
604 All BODY layers share the same recompute type mask (which types are
605 enabled); each layer may have different activation memory values for
606 the enabled types.
607 """
608 return dict(layers_sorted[Layer.type_enum.BODY][0].recompute_considered_)
610 def max_stage_micro_eq_stage(self, prob: Any,
611 layers_sorted: Dict[Layer.type_enum, List[Layer]]) -> Any:
612 """Apply additional VPP optimisations when ``pp == num_of_micro_batch``."""
613 last_chunk = self.num_of_interleave_ - 1
615 for i_stage in range(self.num_of_stage_):
616 for inter in range(last_chunk):
617 prob += self.variables_[self.MAX_STAGE_TIME] >= (
618 self._max_stage_bound_i_bp(layers_sorted, i_stage, inter) +
619 self._max_stage_bound_head_tail(layers_sorted, i_stage,
620 -1, inter))
622 if self.vpp_less_memory_:
623 factors = self.compute_lm_forward_in_backward(self.num_of_stage_)
624 else:
625 factors = self.compute_forward_in_backward(
626 self.num_of_stage_, self.num_of_micro_batch_)
628 for i_stage in range(self.num_of_stage_):
629 logger.debug(
630 "v=%s, s=%s: (BP + HT) + (%s / %s * FP)",
631 last_chunk,
632 i_stage,
633 factors[i_stage],
634 self.num_of_micro_batch_,
635 )
636 prob += self.variables_[self.MAX_LAST_CHUNK] >= (
637 self._max_stage_bound_i_bp(layers_sorted, i_stage, last_chunk) +
638 self._max_stage_bound_head_tail(layers_sorted, i_stage, last_chunk, last_chunk) +
639 (factors[i_stage] / self.num_of_micro_batch_) *
640 self._max_stage_bound_i_fp(layers_sorted, i_stage, last_chunk))
642 if self.optimization_level_ >= 2:
643 logger.debug("Approach 2a")
644 prob += self.variables_[self.MAX_STAGE_TIME] >= (
645 self.variables_[self.MAX_LAST_CHUNK])
647 return self.variables_[self.MAX_STAGE_TIME]
648 logger.debug("Approach 2b")
649 prob += self.variables_[self.MAX_LAST_CHUNK] >= (
650 self.variables_[self.MAX_STAGE_TIME])
652 return (self.variables_[self.MAX_STAGE_TIME] +
653 self.variables_[self.MAX_LAST_CHUNK])
655 def add_performance_constraint(self, prob: Any,
656 layers_sorted: Dict[Layer.type_enum, List[Layer]],
657 pipeline_total_time: Any) -> Any:
658 """Add the ``pipeline_total_time >= …`` performance constraints."""
659 max_stage_time = self.variables_[self.MAX_STAGE_TIME]
660 max_stage_time = self.add_max_stage_constraint(prob, layers_sorted, max_stage_time)
662 total_sum = self.variables_[self.TOTAL_SUM]
663 prob += total_sum >= self._total_sum(layers_sorted)
665 if self.optimization_level_ >= 2:
666 # approach A
667 for v in range(self.num_of_interleave_ - 1):
668 prob += self.variables_[self.PREV_DIFF][v] >= (
669 self._prev_diff_sum(layers_sorted, prob, v))
671 prob += self.variables_[self.CHUNKS_SUM][v] >= (
672 (self.num_of_interleave_ - v) / self.num_of_interleave_ *
673 self._chunks_sum(layers_sorted, v))
675 chunks_sum = lpSolver.lpSum(self.variables_[self.CHUNKS_SUM])
676 prev_diff = lpSolver.lpSum(self.variables_[self.PREV_DIFF])
678 next_diff = self.variables_[self.NEXT_DIFF]
679 prob += next_diff >= (
680 self._next_diff_sum(layers_sorted, prob))
682 prob += pipeline_total_time >= (
683 (total_sum + chunks_sum + prev_diff + next_diff)
684 / max(1, (self.num_of_interleave_ - 2))
685 + max_stage_time * (self.num_of_micro_batch_ - 2)
686 )
687 else:
688 # approach B
689 prob += pipeline_total_time >= max_stage_time
690 return prob
692 def add_max_stage_constraint(self, prob: Any,
693 layers_sorted: Dict[Layer.type_enum, List[Layer]],
694 max_stage_time: Any) -> Any:
695 """Add the ``max_stage_time`` lower-bound constraints over every ``(interleave, stage)``."""
696 if (self.num_of_interleave_ > 1 and self.optimization_level_ >= 1
697 and self.num_of_micro_batch_ == self.num_of_stage_):
698 max_stage_time = self.max_stage_micro_eq_stage(prob, layers_sorted)
699 else:
700 # Constraints on sub-main-part of a stage that it may take (for all stage)
701 for i_stage in range(self.num_of_stage_):
702 for inter_f in range(self.num_of_interleave_):
703 for inter_b in range(self.num_of_interleave_):
704 prob += max_stage_time >= (
705 self._max_stage_bound_i_fp(layers_sorted, i_stage, inter_f) +
706 self._max_stage_bound_i_bp(layers_sorted, i_stage, inter_b) +
707 self._max_stage_bound_head_tail(layers_sorted, i_stage,
708 inter_f, inter_b))
710 return max_stage_time
712 ############################################
713 # Memory Constraint #
714 ############################################
715 def _accumulate_body_param(self, variables, layers_sorted, stage_id, num_of_interleave):
716 """Accumulate BODY-layer parameter memory into an LP expression."""
717 bound = lpSolver.LpAffineExpression()
718 for inter_id in range(num_of_interleave):
719 for layer in layers_sorted[Layer.type_enum.BODY]:
720 for rec in Recompute.TYPE:
721 if self.recompute_considered_[rec]:
722 bound += (
723 variables[layer.name_][rec][inter_id][stage_id] *
724 layer.memory_parameter_)
725 return bound
727 def stage_param_memory(self, variables: Any,
728 layers_sorted: Dict[Layer.type_enum, List[Layer]],
729 stage_id: int, num_of_stage: int,
730 num_of_interleave: int) -> Any:
731 """Return an LP expression for the parameter memory of ``stage_id``."""
732 bound = self._accumulate_body_param(variables, layers_sorted, stage_id, num_of_interleave)
733 if stage_id == 0:
734 for head in layers_sorted[Layer.type_enum.HEAD]:
735 bound += head.memory_parameter_
736 if self.dual_:
737 for tail in layers_sorted[Layer.type_enum.TAIL]:
738 bound += tail.memory_parameter_
739 if not self.dual_ and stage_id == num_of_stage - 1:
740 for tail in layers_sorted[Layer.type_enum.TAIL]:
741 bound += tail.memory_parameter_
742 return bound
744 def stage_active_memory_per_micro(
745 self, variables: Any,
746 layers_sorted: Dict[Layer.type_enum, List[Layer]],
747 stage_id: int, inter_id: int) -> Any:
748 """Return an LP expression for the activation memory of ``stage_id`` per micro-batch."""
749 bound = lpSolver.LpAffineExpression()
750 for layer in layers_sorted[Layer.type_enum.BODY]:
751 for rec in Recompute.TYPE:
752 if self.recompute_considered_[rec]:
753 bound += (variables[layer.name_][rec][inter_id][stage_id] *
754 layer.memory_activation_rec_[rec])
755 return bound
757 def stage_active_memory(self, variables: Any,
758 layers_sorted: Dict[Layer.type_enum, List[Layer]],
759 stage_id: int, num_of_interleave: int,
760 activation_nums: List[List[int]]) -> Any:
761 """Return the total activation-memory LP expression for ``stage_id``."""
762 bound = lpSolver.LpAffineExpression()
763 for inter_id in range(num_of_interleave):
764 for layer in layers_sorted[Layer.type_enum.BODY]:
765 for rec in Recompute.TYPE:
766 if self.recompute_considered_[rec]:
767 bound += (
768 variables[layer.name_][rec][inter_id][stage_id] *
769 layer.memory_activation_rec_[rec] *
770 activation_nums[inter_id][stage_id])
771 return bound
773 def init_overhead_variables(self, variables: Any, s: int) -> Any:
774 """Compute the per-stage overhead LP expression used in the VPP memory constraint."""
775 bound = lpSolver.LpAffineExpression()
776 vf = self.num_of_interleave_ - 1
777 vb = self.num_of_interleave_ - 1
778 incr_f = True
779 if self.vpp_less_memory_:
780 for _ in range(self.num_of_interleave_ - 1):
781 if incr_f:
782 vf = (vf + 1) % self.num_of_interleave_
783 factor = abs(self.num_of_stage_ - s)
784 else:
785 vb = vb - 1
786 factor = s
787 incr_f = not incr_f
789 logger.debug("%s * (act(%s,%s) - act(%s,%s)", factor, vf, s, vb, s)
790 bound += factor * (
791 self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vf)
792 - self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vb))
793 else:
794 for _ in range(self.num_of_interleave_ - 1):
795 if incr_f:
796 vf = (vf + 1) % self.num_of_interleave_
797 logger.debug(
798 "%s * (act(%s,%s) - act(%s,%s)",
799 self.num_of_stage_ - abs(self.num_of_stage_ - 2 * s - 1),
800 vf,
801 s,
802 vb,
803 s,
804 )
805 bound += (self.num_of_stage_ - abs(self.num_of_stage_ - 2 * s - 1)) * (
806 self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vf)
807 - self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vb)
808 )
809 else:
810 vb = vb - 1
811 logger.debug(
812 "%s * (act(%s,%s) - act(%s,%s)",
813 max(self.num_of_stage_ - 2 * s - 1, 0),
814 vf + 1,
815 s,
816 vb + 1,
817 s,
818 )
819 bound += max(self.num_of_stage_ - 2 * s - 1, 0) * (
820 self.stage_active_memory_per_micro(variables,
821 self.layers_sorted_, s, vf + 1)
822 - self.stage_active_memory_per_micro(variables,
823 self.layers_sorted_, s, vb + 1)
824 )
825 logger.debug(
826 "%s * (act(%s,%s) - act(%s,%s)",
827 max(-(self.num_of_stage_ - 2 * s - 1), 0),
828 vf,
829 s,
830 vb,
831 s,
832 )
833 bound += max(-(self.num_of_stage_ - 2 * s - 1), 0) * (
834 self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vf)
835 - self.stage_active_memory_per_micro(variables, self.layers_sorted_, s, vb)
836 )
837 incr_f = not incr_f
839 return bound
841 def stage_overhead_memory(self, variables: Any, stage_id: int) -> Any:
842 """Return the stage-``stage_id`` memory overhead LP expression."""
843 bound = lpSolver.LpAffineExpression()
844 for v in range(self.num_of_interleave_ - 1):
845 bound += variables[self.MEM_OVERHEAD_NAME][stage_id][v]
846 return bound
848 def add_pipeline_memory_constraint(self,
849 constraint: PipelineMemoryConstraint) -> None:
850 """Add per-stage memory upper-bound constraints to the solver problem."""
851 prob = constraint.prob
852 variables = constraint.variables
853 layers_sorted = constraint.layers_sorted
854 num_of_stage = constraint.num_of_stage
855 num_of_interleave = constraint.num_of_interleave
856 micro_batch = constraint.micro_batch
857 memory_limit = constraint.memory_limit
859 if self.vpp_less_memory_:
860 if self.seq_pipe:
861 activation_nums = self.compute_activation_seq_nums(
862 num_of_stage, num_of_interleave, self.seq_split_num_, micro_batch, True)
863 else:
864 activation_nums = self.compute_less_activation_nums(
865 num_of_stage, num_of_interleave)
866 # Add if dual to decide whether dualpipe_v is used
867 elif self.dual_:
868 activation_nums = self.compute_activation_nums_dual(
869 num_of_stage, num_of_interleave, micro_batch)
871 else:
872 if self.seq_pipe:
873 activation_nums = self.compute_activation_seq_nums(
874 num_of_stage, num_of_interleave, self.seq_split_num_, micro_batch, False)
875 else:
876 activation_nums = self.compute_activation_nums(
877 num_of_stage, num_of_interleave, micro_batch)
878 logger.info("activation nums = %s", activation_nums)
880 if self.num_of_stage_ == self.num_of_micro_batch_:
881 for s in range(num_of_stage):
882 prob += memory_limit >= (
883 self.stage_param_memory(variables, layers_sorted, s,
884 num_of_stage, num_of_interleave) +
885 self.stage_active_memory(variables, layers_sorted, s,
886 num_of_interleave, activation_nums) +
887 self.constant_memory_)
888 else:
889 for s in range(num_of_stage):
890 prob += variables[self.MEM_OVERHEAD_NAME][s] >= (
891 self.init_overhead_variables(variables, s)
892 )
893 prob += memory_limit >= (
894 self.stage_param_memory(
895 variables, layers_sorted, s, num_of_stage, num_of_interleave
896 )
897 + self.stage_active_memory(
898 variables, layers_sorted, s, num_of_interleave, activation_nums
899 )
900 + variables[self.MEM_OVERHEAD_NAME][s]
901 + self.constant_memory_
902 )
904 def _stage_activation_memory(self, inter: int, stage: int) -> float:
905 """Compute activation memory for one ``(inter, stage)`` cell."""
906 memory_activation = 0
907 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
908 for rec in Recompute.TYPE:
909 if not self.recompute_considered_[rec]:
910 continue
911 var_value = self.variables_.get(layer.name_)[rec][inter][stage].varValue
912 memory_activation += var_value * layer.memory_activation_rec_[rec]
913 return memory_activation
915 def get_simulator_memory_activation(self) -> list[float]:
916 """Give the activation memory per stage for simulator."""
918 memory_active = []
919 if self.has_some_memory_info():
920 for inter in range(self.num_of_interleave_):
921 inter_list = []
922 for stage in range(self.num_of_stage_):
923 inter_list.append(self._stage_activation_memory(inter, stage))
924 memory_active.append(inter_list)
925 return memory_active
927 def get_simulator_memory_parameter(self) -> list[float]:
928 """Give the parameter memory per stage for simulator."""
929 memory_param_stage = [0] * self.num_of_stage_
930 if self.has_some_memory_info():
931 for inter in range(self.num_of_interleave_):
932 for stage in range(self.num_of_stage_):
933 memory_param_stage[stage] += self._get_stage_parameter_memory(inter, stage)
935 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
936 if head.memory_parameter_ is not None:
937 memory_param_stage[0] += head.memory_parameter_
938 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
939 if tail.memory_parameter_ is not None:
940 memory_param_stage[self.num_of_stage_ -
941 1] += tail.memory_parameter_
942 memory_param = [memory_param_stage] * self.num_of_interleave_
943 return memory_param
945 def _get_stage_parameter_memory(self, interleave, stage):
946 """Calculate BODY-layer parameter memory for one pipeline position."""
947 total = 0
948 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
949 if layer.memory_parameter_ is not None:
950 for rec in Recompute.TYPE:
951 if not self.recompute_considered_[rec]:
952 continue
954 var_value = self.variables_.get(layer.name_)[rec][interleave][stage].varValue
955 total += var_value * layer.memory_parameter_
956 return total
958 def get_simulator_time(self) -> list[float]:
959 """Give the time per stage for simulator."""
960 time = []
961 for i in range(self.num_of_interleave_):
962 time.append([])
963 for s in range(self.num_of_stage_):
964 time[i].append(0)
965 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
966 for rec in Recompute.TYPE:
967 if self.recompute_considered_[rec]:
968 time[i][s] += self.variables_.get(
969 layer.name_)[rec][i][s].varValue * (
970 layer.forward_time_ +
971 layer.backward_time_rec_[rec])
973 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
974 time[0][0] += head.forward_time_ + head.backward_time_rec_[Recompute.TYPE.NONE]
975 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
976 time[self.num_of_interleave_ - 1][self.num_of_stage_ -
977 1] += tail.forward_time_ + tail.backward_time_rec_[Recompute.TYPE.NONE]
978 return time
980 def get_simulator_forward_time(self) -> list[float]:
981 """Give the time per stage for simulator."""
982 time = []
983 for i in range(self.num_of_interleave_):
984 time.append([])
985 for s in range(self.num_of_stage_):
986 time[i].append(0)
987 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
988 for rec in Recompute.TYPE:
989 if self.recompute_considered_[rec]:
990 time[i][s] += self.variables_[layer.name_][rec][i][
991 s].varValue * (layer.forward_time_)
992 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
993 time[0][0] += head.forward_time_
994 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
995 time[self.num_of_interleave_ - 1][self.num_of_stage_ -
996 1] += tail.forward_time_
997 return time
999 def get_simulator_backward_time(self) -> list[float]:
1000 """Give the backward time per stage for simulator.
1002 Unlike :meth:`get_simulator_forward_time`, this returns the actual
1003 backward time derived from per-layer ``backward_time_rec_`` instead of
1004 relying on ``forward_time × backward_ratio``.
1005 """
1006 time = []
1007 for i in range(self.num_of_interleave_):
1008 time.append([])
1009 for s in range(self.num_of_stage_):
1010 time[i].append(0)
1011 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
1012 for rec in Recompute.TYPE:
1013 if self.recompute_considered_[rec]:
1014 time[i][s] += self.variables_[layer.name_][rec][i][
1015 s].varValue * layer.backward_time_rec_[rec]
1016 for head in self.layers_sorted_[Layer.type_enum.HEAD]:
1017 time[0][0] += head.backward_time_rec_[Recompute.TYPE.NONE]
1018 for tail in self.layers_sorted_[Layer.type_enum.TAIL]:
1019 time[self.num_of_interleave_ - 1][self.num_of_stage_ -
1020 1] += tail.backward_time_rec_[Recompute.TYPE.NONE]
1021 return time
1023 def get_simulator_recompute_time(self) -> list[float]:
1024 """Give the time per stage for simulator."""
1025 time_all_rec = []
1026 time_no_rec = []
1027 for i in range(self.num_of_interleave_):
1028 time_all_rec.append([])
1029 time_no_rec.append([])
1030 for s in range(self.num_of_stage_):
1031 time_all_rec[i].append(0)
1032 time_no_rec[i].append(0)
1033 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
1034 for rec in Recompute.TYPE:
1035 if self.recompute_considered_[rec]:
1036 time_all_rec[i][s] += self.variables_[
1037 layer.name_][rec][i][s].varValue * (
1038 layer.backward_time_rec_[rec])
1039 time_no_rec[i][s] += self.variables_[
1040 layer.name_][rec][i][s].varValue * (
1041 layer.backward_time_rec_[
1042 Recompute.TYPE.NONE])
1043 return [[r - n for r, n in zip(ar, nr)]
1044 for ar, nr in zip(time_all_rec, time_no_rec)]
1046 def has_some_memory_info(self) -> bool:
1047 """Check if there is some information for memory constraint."""
1048 return any(self.recompute_considered_.values())
1050 ############################################
1051 # General Constraint #
1052 ############################################
1053 def add_optional_recompute_constraint(
1054 self, prob: Any, variables: Any,
1055 sorted_layers: Dict[Layer.type_enum, List[Layer]]) -> None:
1056 """Pin unused recomputation variables to zero in the ILP."""
1057 for layer in sorted_layers[Layer.type_enum.BODY]:
1058 for rec in Recompute.TYPE:
1059 if not self.recompute_considered_[rec]:
1060 prob += lpSolver.lpSum(variables[layer.name_][rec]) == 0
1062 def dump_problem(self, folder: Optional[str] = None) -> None:
1063 """Serialize the pulp LP model to ``<folder>/<auto-generated-name>.lp``."""
1064 dump_name = "problem_" + str(self.layers_[0].model_name_)
1065 dump_name += "_" + str(self.max_memory_)
1066 dump_name += "_" + str(self.num_of_interleave_)
1067 dump_name += "_" + str(self.num_of_stage_)
1069 logger.info("dump_problem:out folder = %s", folder)
1070 if folder is not None:
1071 dump_name = os.path.join(folder, dump_name)
1072 dump_name += ".lp"
1073 logger.info("dump problem file: %s", dump_name)
1074 self.problem_.writeLP(dump_name)
1076 def _print_body_layer_assignments(self) -> None:
1077 """Log per-layer recompute assignments for each stage and interleave."""
1078 for body_layer in self.layers_sorted_[Layer.type_enum.BODY]:
1079 layer_name = body_layer.name_
1080 logger.output("For layer: %s", layer_name)
1081 logger.output("=========")
1082 logger.output(" Forward Prop time: %s", body_layer.forward_time_)
1083 for rec in Recompute.TYPE:
1084 if self.recompute_considered_[rec]:
1085 logger.output(" Backward Prop %s time: %s",
1086 Recompute.YAML_NAME[rec], body_layer.backward_time_rec_[rec])
1087 for inter in range(self.num_of_interleave_):
1088 for stage in range(self.num_of_stage_):
1089 parts = []
1090 for rec in Recompute.TYPE:
1091 if self.recompute_considered_[rec]:
1092 value = str(int(self.variables_[layer_name][rec][inter][stage].varValue))
1093 parts.append(value if rec is Recompute.TYPE.NONE else f"+ {value} {rec.name}")
1094 chunk = f" of chunk {inter}" if self.num_of_interleave_ != 1 else ""
1095 logger.output(" Assign %s: %s%s to stage %d",
1096 layer_name, " ".join(parts), chunk, stage)
1098 def _print_debug_variables(self) -> None:
1099 """Log debug-level variable values for the solver problem."""
1100 for s in range(self.num_of_stage_):
1101 logger.debug(
1102 "%s[%s] =%s",
1103 self.MEM_OVERHEAD_NAME,
1104 s,
1105 self.variables_[self.MEM_OVERHEAD_NAME][s].varValue,
1106 )
1108 for v in range(self.num_of_interleave_ - 1):
1109 logger.debug(
1110 "%s[%s] = %s",
1111 self.CHUNKS_SUM,
1112 v,
1113 self.variables_[self.CHUNKS_SUM][v].varValue,
1114 )
1116 for v in range(self.num_of_interleave_ - 1):
1117 logger.debug(
1118 "%s[%s] = %s",
1119 self.PREV_DIFF,
1120 v,
1121 self.variables_[self.PREV_DIFF][v].varValue,
1122 )
1124 logger.debug("%s = %s", self.NEXT_DIFF, self.variables_[self.NEXT_DIFF].varValue)
1125 logger.debug("%s = %s", self.TOTAL_SUM, self.variables_[self.TOTAL_SUM].varValue)
1126 logger.debug("%s = %s", self.MAX_STAGE_TIME, self.variables_[self.MAX_STAGE_TIME].varValue)
1127 logger.debug("%s = %s", self.MAX_LAST_CHUNK, self.variables_[self.MAX_LAST_CHUNK].varValue)
1129 for body_layer in range(len(self.layers_sorted_[Layer.type_enum.BODY]) - 1):
1130 for v in range(self.num_of_interleave_):
1131 for s in range(self.num_of_stage_):
1132 logger.info(
1133 "%s[%s][%s][%s] = %s",
1134 self.LAYER_FRONTIER,
1135 body_layer,
1136 v,
1137 s,
1138 self.variables_[self.LAYER_FRONTIER][body_layer][v][s].varValue,
1139 )
1141 def print_results(self) -> None:
1142 """Log the detailed per-layer solver assignment for the solved problem."""
1143 if self.has_some_memory_info():
1144 logger.output("For max memory %s", self.max_memory_)
1145 logger.output("==============")
1146 self._print_body_layer_assignments()
1147 self._print_debug_variables()
1149 def debug_print_solver_theoretical_memory(self) -> None:
1150 """Log the solver-implied per-stage theoretical memory (debug aid)."""
1151 logger.info("%s Solver Theoretical Memory Analysis %s", "=" * 20, "=" * 20)
1153 if self.vpp_less_memory_:
1154 if self.seq_pipe:
1155 activation_nums = self.compute_activation_seq_nums(
1156 self.num_of_stage_, self.num_of_interleave_, self.seq_split_num_, self.num_of_micro_batch_, True)
1157 else:
1158 activation_nums = self.compute_less_activation_nums(
1159 self.num_of_stage_, self.num_of_interleave_)
1160 else:
1161 if self.seq_pipe:
1162 activation_nums = self.compute_activation_seq_nums(
1163 self.num_of_stage_, self.num_of_interleave_, self.seq_split_num_, self.num_of_micro_batch_, False)
1164 else:
1165 activation_nums = self.compute_activation_nums(
1166 self.num_of_stage_, self.num_of_interleave_, self.num_of_micro_batch_)
1168 # compute theoretical value for each stage
1169 for s in range(self.num_of_stage_):
1170 param_mem = self.stage_param_memory(
1171 self.variables_,
1172 self.layers_sorted_,
1173 s,
1174 self.num_of_stage_,
1175 self.num_of_interleave_
1176 ).value()
1178 act_mem = self.stage_active_memory(
1179 self.variables_,
1180 self.layers_sorted_,
1181 s,
1182 self.num_of_interleave_,
1183 activation_nums
1184 ).value()
1186 overhead = 0
1187 total = param_mem + act_mem + overhead + self.constant_memory_
1189 logger.info("Stage %d Solver Memory Analysis:", s)
1190 logger.info("Parameter Memory: %.2f", param_mem)
1191 logger.info("Activation Memory: %.2f", act_mem)
1192 logger.info("Memory Overhead: %.2f", overhead)
1193 logger.info("Constant Memory: %.2f", self.constant_memory_)
1194 logger.info("Total Theoretical Memory: %.2f", total)
1197 def solve(self, time_limit: int = 90, dump_folder: Optional[str] = None) -> None:
1198 """Solve the ILP problem using PuLP's bundled CBC backend.
1200 Args:
1201 time_limit: Upper bound on solver wall-clock time in seconds.
1202 dump_folder: Directory to write the LP model to; ``None`` skips the dump.
1203 """
1204 logger.info("solve:out folder = %s", dump_folder)
1205 self.dump_problem(dump_folder)
1206 solver = lpSolver.getSolver("PULP_CBC_CMD", timeLimit=time_limit)
1207 self.problem_.solve(solver)
1209 self.print_results()
1211 self.debug_print_solver_theoretical_memory()
1213 for name, result in self.result().items():
1214 logger.output("%s %s %s", name, result, "\n")
1216 def result(self) -> dict[str, list[list[str]]]:
1217 """return schedule distribution for each layer (in the form of a dict)"""
1218 r = {}
1219 for layer in self.layers_sorted_[Layer.type_enum.BODY]:
1220 layer_name = layer.name_
1221 inter = []
1222 for i in range(self.num_of_interleave_):
1223 stage = []
1224 for s in range(self.num_of_stage_):
1225 for rec in Recompute.TYPE:
1226 if self.recompute_considered_[rec]:
1227 stage.append(
1228 str(
1229 self.variables_.get(layer_name)[rec][i]
1230 [s].varValue) + " + ")
1231 inter.append(stage)
1232 r[layer_name] = inter
1233 return r
1235 def _create_problem_(self, description: str) -> lpSolver.LpProblem:
1236 """create the problem"""
1237 prob = lpSolver.LpProblem(description, lpSolver.LpMinimize)
1238 layers_sorted = self.layers_sorted_
1239 num_of_stage = self.num_of_stage_
1240 num_of_interleave = self.num_of_interleave_
1241 num_of_micro_batch = self.num_of_micro_batch_
1242 max_memory = self.max_memory_
1243 # Local variable declaration
1244 # max time that a "main" stage have to take (var to minimize)
1245 pipeline_total_time = lpSolver.LpVariable("pipeline_total_time", 0,
1246 None, lpSolver.LpContinuous)
1248 # Var to Minimize
1249 prob += pipeline_total_time
1251 # Explicitly constrain unused recompute type variables to zero.
1252 # While current constraints filter by self.recompute_considered_[rec],
1253 # this prevents latent issues if future constraints forget to filter.
1254 for layer in layers_sorted[Layer.type_enum.BODY]:
1255 for rec in Recompute.TYPE:
1256 if not self.recompute_considered_[rec]:
1257 for inter in range(num_of_interleave):
1258 for stage in range(num_of_stage):
1259 prob += (
1260 self.variables_[layer.name_][rec][inter][stage] == 0
1261 )
1263 result = self.add_total_nb_layer_constraint(prob, self.variables_, layers_sorted)
1264 if result is None:
1265 raise RuntimeError("add_total_nb_layer_constraint() returned None.")
1266 # Add if dual to the original layer order constraint
1267 try:
1268 prob = self.add_stage_nb_layer_constraint(
1269 prob, self.variables_, layers_sorted
1270 )
1271 except Exception:
1272 logger.exception("Failed to add stage number layer constraint.")
1273 raise
1274 try:
1275 result = self.add_multimodal_sequence_constraint(prob, self.variables_, layers_sorted)
1276 except Exception:
1277 logger.exception("Failed to add multimodal sequence constraint.")
1278 raise
1280 #self.add_stage_nb_layer_constraint_dual(prob, self.variables_, layers_sorted)
1281 #self.add_multimodal_sequence_constraint_dual(prob, self.variables_, layers_sorted)
1282 try:
1283 result = self.add_multimodal_recompute_constraint(prob, self.variables_, layers_sorted)
1284 if result is None:
1285 raise RuntimeError("add_multimodal_recompute_constraint() returned None.")
1286 except Exception:
1287 logger.exception("Failed to add multimodal recompute constraint.")
1288 raise
1290 try:
1291 result = self.add_performance_constraint(prob, layers_sorted, pipeline_total_time)
1292 if result is None:
1293 raise RuntimeError("add_performance_constraint() returned None.")
1294 prob = result
1295 except Exception:
1296 logger.exception("Failed to add performance constraint.")
1297 raise
1299 constraint = PipelineMemoryConstraint(
1300 prob=prob,
1301 variables=self.variables_,
1302 layers_sorted=layers_sorted,
1303 num_of_stage=num_of_stage,
1304 num_of_interleave=num_of_interleave,
1305 micro_batch=num_of_micro_batch,
1306 memory_limit=max_memory,
1307 )
1308 if self.has_some_memory_info():
1309 self.add_pipeline_memory_constraint(constraint)
1310 return prob
1312 def _create_variables_to_solve_(
1313 self,
1314 num_of_stage: int,
1315 num_of_interleave: int,
1316 layers: dict[Layer.type_enum, list[Layer]],
1317 ):
1318 """create variables to solve"""
1319 variables = {}
1321 variables[self.TOTAL_SUM] = lpSolver.LpVariable(
1322 self.TOTAL_SUM, 0, None, lpSolver.LpContinuous)
1324 chunks_sum_dict = lpSolver.LpVariable.dicts(
1325 name=self.CHUNKS_SUM,
1326 indices=(range(0, self.num_of_interleave_ - 1)),
1327 lowBound=0,
1328 upBound=None,
1329 cat=lpSolver.LpContinuous
1330 )
1331 chunks_sum_list = list(chunks_sum_dict.values())
1332 variables[self.CHUNKS_SUM] = chunks_sum_list
1334 prev_diff_dict = lpSolver.LpVariable.dicts(
1335 name=self.PREV_DIFF,
1336 indices=(range(0, self.num_of_interleave_ - 1)),
1337 lowBound=0,
1338 upBound=None,
1339 cat=lpSolver.LpContinuous
1340 )
1341 prev_diff_list = list(prev_diff_dict.values())
1342 variables[self.PREV_DIFF] = prev_diff_list
1344 layer_frontier_dict = lpSolver.LpVariable.dicts(
1345 name=self.LAYER_FRONTIER,
1346 indices=(
1347 range(1, len(self.layers_sorted_[Layer.type_enum.BODY])),
1348 range(0, self.num_of_interleave_),
1349 range(0, self.num_of_stage_)),
1350 lowBound=0,
1351 upBound=1,
1352 cat=lpSolver.LpBinary
1353 )
1354 layer_frontier_list = list(layer_frontier_dict.values())
1355 variables[self.LAYER_FRONTIER] = layer_frontier_list
1357 rec_frontier_dict = lpSolver.LpVariable.dicts(
1358 name=self.REC_FRONTIER,
1359 indices=(
1360 range(0, self.num_of_interleave_),
1361 range(0, self.num_of_stage_),
1362 range(0, len(self.layers_sorted_[Layer.type_enum.BODY])-1)),
1363 lowBound=0,
1364 upBound=1,
1365 cat=lpSolver.LpBinary
1366 )
1367 rec_frontier_list = list(rec_frontier_dict.values())
1368 variables[self.REC_FRONTIER] = rec_frontier_list
1370 variables[self.NEXT_DIFF] = lpSolver.LpVariable(
1371 self.NEXT_DIFF, 0, None, lpSolver.LpContinuous)
1373 variables[self.MAX_STAGE_TIME] = lpSolver.LpVariable(
1374 self.MAX_STAGE_TIME, 0, None, lpSolver.LpContinuous)
1376 variables[self.MAX_LAST_CHUNK] = lpSolver.LpVariable(
1377 self.MAX_LAST_CHUNK, 0, None, lpSolver.LpContinuous)
1379 lp_variable_dict = lpSolver.LpVariable.dicts(
1380 name=self.MEM_OVERHEAD_NAME,
1381 indices=(range(0, self.num_of_stage_)),
1382 lowBound=0,
1383 upBound=None,
1384 cat=lpSolver.LpInteger,
1385 )
1386 variables_list = list(lp_variable_dict.values())
1387 variables[self.MEM_OVERHEAD_NAME] = variables_list
1389 for layer in layers[Layer.type_enum.BODY]:
1390 variable_dict = lpSolver.LpVariable.dicts(
1391 name=layer.name_,
1392 indices=(
1393 range(0, len(Recompute.TYPE)),
1394 range(0, num_of_interleave),
1395 range(0, num_of_stage),
1396 ),
1397 lowBound=0,
1398 upBound=None,
1399 cat=lpSolver.LpInteger,
1400 )
1401 variable_values = list(variable_dict.values())
1402 interleave_values = []
1403 for interleave in variable_values:
1404 interleave_value = list(interleave.values())
1405 interleave_values.append(interleave_value)
1406 variables[layer.name_] = interleave_values
1408 return variables
1410 ############################################
1411 # Time Constraint #
1412 ############################################
1413 def _max_stage_bound_i_fp(self, layers_sorted, stage_id, inter_f):
1414 bound = lpSolver.LpAffineExpression()
1415 for layer in layers_sorted[Layer.type_enum.BODY]:
1416 for rec in Recompute.TYPE:
1417 if self.recompute_considered_[rec]:
1418 bound += (self.variables_[layer.name_][rec][inter_f][stage_id] *
1419 layer.forward_time_)
1420 return bound
1422 def _max_stage_bound_i_bp(self, layers_sorted, stage_id, inter_b):
1423 bound = lpSolver.LpAffineExpression()
1424 for layer in layers_sorted[Layer.type_enum.BODY]:
1425 for rec in Recompute.TYPE:
1426 if self.recompute_considered_[rec]:
1427 bound += (self.variables_[layer.name_][rec][inter_b][stage_id] *
1428 layer.backward_time_rec_[rec])
1429 return bound
1431 def _max_stage_bound_head_tail(self, layers_sorted, stage_id, inter_f,
1432 inter_b):
1433 """maximize the stage bound of head and tail"""
1434 bound = lpSolver.LpAffineExpression()
1435 if stage_id == 0:
1436 if inter_f == 0:
1437 for head in layers_sorted[Layer.type_enum.HEAD]:
1438 bound += head.forward_time_
1439 if inter_b == 0:
1440 for head in layers_sorted[Layer.type_enum.HEAD]:
1441 bound += head.backward_time_rec_[Recompute.TYPE.NONE]
1442 if stage_id == self.num_of_stage_ - 1:
1443 if inter_f == self.num_of_interleave_ - 1:
1444 for tail in layers_sorted[Layer.type_enum.TAIL]:
1445 bound += tail.forward_time_
1446 if inter_b == self.num_of_interleave_ - 1:
1447 for tail in layers_sorted[Layer.type_enum.TAIL]:
1448 bound += tail.backward_time_rec_[Recompute.TYPE.NONE]
1449 return bound
1451 def _total_sum(self, layers_sorted):
1452 """sum up the layer time"""
1453 bound = lpSolver.LpAffineExpression()
1454 for layer in layers_sorted[Layer.type_enum.BODY]:
1455 for rec in Recompute.TYPE:
1456 if self.recompute_considered_[rec]:
1457 for inter in range(self.num_of_interleave_):
1458 for stage in range(self.num_of_stage_):
1459 bound += self.variables_[layer.name_][rec][inter][stage] * (
1460 layer.forward_time_ +
1461 layer.backward_time_rec_[rec])
1462 return bound
1464 def body_layer_time(self, prop: "SappSolver.PROP_PHASE", layer: Layer,
1465 inter: int, stage: int) -> Any:
1466 """Return a forward or backward time LP expression for ``layer`` at ``(inter, stage)``."""
1467 if prop == self.PROP_PHASE.FW:
1468 bound = lpSolver.lpSum(
1469 self.variables_[layer.name_][rec][inter][stage] * layer.forward_time_
1470 for rec in Recompute.TYPE if self.recompute_considered_[rec])
1471 else:
1472 bound = lpSolver.lpSum(
1473 self.variables_[layer.name_][rec][inter][stage] * layer.backward_time_rec_[rec]
1474 for rec in Recompute.TYPE if self.recompute_considered_[rec])
1476 return bound
1478 def _append_head_tail_time(self, prop, layers_sorted, inter, stage, bound):
1479 """Append head/tail time to the bound when at boundary stages."""
1480 if stage == 0 and inter == 0:
1481 for head in layers_sorted[Layer.type_enum.HEAD]:
1482 if prop == self.PROP_PHASE.FW:
1483 bound += head.forward_time_
1484 else:
1485 bound += head.backward_time_rec_[Recompute.TYPE.NONE]
1486 if stage == self.num_of_stage_ - 1 and inter == self.num_of_interleave_ - 1:
1487 for tail in layers_sorted[Layer.type_enum.TAIL]:
1488 if prop == self.PROP_PHASE.FW:
1489 bound += tail.forward_time_
1490 else:
1491 bound += tail.backward_time_rec_[Recompute.TYPE.NONE]
1492 return bound
1494 def micro_batch_time(self, prop: "SappSolver.PROP_PHASE",
1495 layers_sorted: Dict[Layer.type_enum, List[Layer]],
1496 inter: int, stage: int) -> Any:
1497 """Return the total micro-batch time LP expression at ``(inter, stage)``."""
1498 bound = lpSolver.LpAffineExpression()
1499 for layer in layers_sorted[Layer.type_enum.BODY]:
1500 bound += self.body_layer_time(prop, layer, inter, stage)
1501 bound = self._append_head_tail_time(prop, layers_sorted, inter, stage, bound)
1502 return bound
1504 def _chunks_sum(self, layers_sorted, v):
1505 """sum up the warm-up and cool-down time of a given chunk"""
1506 bound = lpSolver.LpAffineExpression()
1507 for stage in range(self.num_of_stage_):
1508 bound += self.micro_batch_time(self.PROP_PHASE.FW, layers_sorted, v, stage)
1509 bound += self.micro_batch_time(self.PROP_PHASE.BW, layers_sorted, v, stage)
1510 # normalize
1511 bound = bound / self.num_of_stage_
1512 return bound
1514 def _prev_diff_sum(self, layers_sorted, prob, v):
1515 """models bubble time for the first diagonal (forward, interleave 0)"""
1516 max_prev_stages = lpSolver.LpVariable.dicts(
1517 name="max_prev_stages_" + str(v),
1518 indices=(range(self.num_of_stage_)),
1519 lowBound=0,
1520 upBound=None,
1521 cat=lpSolver.LpContinuous,
1522 )
1524 diff_with_prev_stages = lpSolver.LpVariable.dicts(
1525 name="diff_with_prev_stages_" + str(v),
1526 indices=(range(self.num_of_stage_)),
1527 lowBound=0,
1528 upBound=None,
1529 cat=lpSolver.LpContinuous,
1530 )
1532 bound = lpSolver.LpAffineExpression()
1534 head_time = 0
1535 for head in layers_sorted[Layer.type_enum.HEAD]:
1536 head_time = head.forward_time_
1538 prob += max_prev_stages[0] >= (self.micro_batch_time(
1539 self.PROP_PHASE.FW, layers_sorted, v, 0)) - head_time
1541 for stage in range(1, self.num_of_stage_):
1542 prob += max_prev_stages[stage] >= max_prev_stages[stage - 1]
1543 prob += max_prev_stages[stage] >= (self.micro_batch_time(
1544 self.PROP_PHASE.FW, layers_sorted, v, stage))
1547 prob += diff_with_prev_stages[stage] >= (
1548 max_prev_stages[stage - 1] - self.micro_batch_time(
1549 self.PROP_PHASE.FW, layers_sorted, v, stage))
1551 bound += self.num_of_micro_batch_ * lpSolver.lpSum(
1552 diff_with_prev_stages[s] for s in range(1, self.num_of_stage_))
1553 return bound
1555 def _next_diff_sum(self, layers_sorted, prob):
1556 """models bubble time for the last diagonal (forward, last chunk)"""
1557 last_chunk = self.num_of_interleave_ - 1
1558 max_next_stages = lpSolver.LpVariable.dicts(
1559 name="max_next_stages",
1560 indices=(range(self.num_of_stage_)),
1561 lowBound=0,
1562 upBound=None,
1563 cat=lpSolver.LpContinuous,
1564 )
1566 diff_with_next_stages = lpSolver.LpVariable.dicts(
1567 name="diff_with_next_stages",
1568 indices=(range(self.num_of_stage_)),
1569 lowBound=0,
1570 upBound=None,
1571 cat=lpSolver.LpContinuous,
1572 )
1574 bound = lpSolver.LpAffineExpression()
1576 prob += max_next_stages[self.num_of_stage_ -
1577 1] >= (self.micro_batch_time(
1578 self.PROP_PHASE.FW, layers_sorted, last_chunk,
1579 self.num_of_stage_ - 1))
1581 for stage in reversed(range(0, self.num_of_stage_ - 1)):
1582 prob += max_next_stages[stage] >= max_next_stages[stage + 1]
1583 prob += max_next_stages[stage] >= (self.micro_batch_time(
1584 self.PROP_PHASE.FW, layers_sorted, last_chunk, stage))
1586 prob += diff_with_next_stages[stage] >= (
1587 max_next_stages[stage + 1] - self.micro_batch_time(
1588 self.PROP_PHASE.FW, layers_sorted, last_chunk, stage))
1590 bound += self.num_of_micro_batch_ * lpSolver.lpSum(
1591 diff_with_next_stages[s] for s in range(self.num_of_stage_ - 1))
1592 return bound