Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_ppb / utils / compute_memory.py: 85%
267 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"""Derive per-layer memory parameters from a set of dry-run stage observations."""
16import numpy as np
18import hyper_parallel.auto_parallel.sapp_ppb.utils.recompute as Recompute
19from hyper_parallel.auto_parallel.sapp_ppb.utils.layer import Layer
20from hyper_parallel.auto_parallel.sapp_ppb.utils.logger import logger
21from hyper_parallel.auto_parallel.sapp_ppb.utils.stage import Stage, filter_stage_id
24class ComputeMemory:
25 """
26 ComputeMemory class to compute the different memories with stages information running (dry) log
28 stage{A|B} means stage with different configuration A and B
29 stage{1|2} means stage same configuration but different id (can be id other than 1 or 2)
31 number_of_stage_ (int): number of stages for the LLM
32 stagesA_ (list[Stage]): list of dry run stages information, with all the same configuration A,
33 required at least staged 0, 1, (n-2), (n-1)
34 Don't set directly stagesA_, but use set_stagesA
35 stagesB_ (list[Stage]): list of dry run stages information, with all the same configuration B,
36 different from config A required at least staged 0, 1, (n-2), (n-1)
37 Don't set directly stagesB_, but use set_stagesB
38 memory_parameter_ (float): memory_parameter_ of the BODY layer, memory required to run the layer
39 memory_activation_rec_ (dict[Recompute.TYPE, float]) activation memory per recompute types
40 recompute_considered_ (dict[Recompute.TYPE, bool]) recomputation types taken into consideration
41 memory_const_ (float): constant memory required for each stages
42 memory_head_ (float): memory required to run the head layer
43 memory_tail_ (float): memory required to run the tail layer
44 """
46 number_of_stage_: int
47 stages_a: list[Stage]
48 stages_b: list[Stage]
49 memory_parameter_: float
50 memory_activation_rec_: dict[Recompute.TYPE, float]
51 recompute_considered_: dict[Recompute.TYPE, bool]
52 memory_const_: float
53 memory_head_: float
54 memory_tail_: float
56 def __init__(self, number_of_stage: int, stages_a: list[Stage] = None,
57 stages_b: list[Stage] = None) -> None:
58 """Build a :class:`ComputeMemory` solver instance.
60 Args:
61 number_of_stage: Total number of pipeline stages in the target LLM.
62 stages_a: Dry-run observations with configuration A (at least stages ``0, i, j, n-1``).
63 stages_b: Dry-run observations with configuration B (must differ from A).
64 """
65 self.number_of_stage_ = number_of_stage
66 self.set_stages_a(stages_a)
67 self.set_stages_b(stages_b)
68 # number_of_stage != len(stages) can be true
69 self.memory_parameter_ = None
70 self.memory_activation_rec_ = {r: None for r in Recompute.TYPE}
71 self.find_recompute_considered()
72 self.memory_const_ = None
73 self.memory_head_ = None
74 self.memory_tail_ = None
76 def set_stages_a(self, stages: list[Stage]) -> None:
77 """Assign dry-run observations to configuration A after a consistency check."""
78 if stages is None:
79 self.stages_a = []
80 return
81 for stage1 in stages:
82 for stage2 in stages:
83 if not stage1.same_global_config(stage2):
84 logger.error(
85 "Cannot set stagesA, all elements don't have the same configuration",)
86 self.stages_a = []
87 return
88 self.stages_a = stages
90 def set_stages_b(self, stages: list[Stage]) -> None:
91 """Assign dry-run observations to configuration B (must differ from A)."""
92 if stages is None:
93 self.stages_b = []
94 return
95 for stage1 in stages:
96 for stage2 in stages:
97 if not stage1.same_global_config(stage2):
98 logger.error(
99 "Cannot set stagesB, all elements don't have the same configuration")
100 self.stages_b = []
101 return
102 for stage_a in self.stages_b:
103 if stage1.same_global_config(stage_a):
104 logger.error(
105 "Cannot set stagesB, an elements have the same configuration than stagesA")
106 self.stages_b = []
107 return
108 self.stages_b = stages
110 def find_recompute_considered(self) -> None:
111 """Populate :attr:`recompute_considered_` from the observed ``stages_a`` data."""
112 self.recompute_considered_ = {r: False for r in Recompute.TYPE}
113 self.recompute_considered_[Recompute.TYPE.NONE] = True
115 for stage in self.stages_a:
116 for rec in Recompute.TYPE:
117 if stage.nb_layer_rec_[rec] > 0:
118 self.recompute_considered_[rec] = True
120 def _compute_memory_parameter_local_(self, stage1: Stage, stage2: Stage) -> float:
121 """
122 Given 2 stages information with the same configuration, and different id,
123 Compute the memory_parameter
124 """
125 if stage1.same_config(stage2):
126 if stage1.id_ != stage2.id_:
127 res = stage1.memory_usage_ * (stage1.nb_stage_ - stage1.id_)
128 res -= stage2.memory_usage_ * (stage2.nb_stage_ - stage2.id_)
129 res /= stage1.id_ - stage2.id_
130 res = abs(res)
131 res /= stage1.nb_layer_
132 return res
133 logger.error(
134 "stage with same characteristic, BUT SAME ID too, cannot compute memory_parameter")
135 return 0
136 logger.error("stage with different characteristic, cannot compute memory_parameter")
137 return 0
139 def _compute_memory_parameter_(self, multi_run=False) -> float:
140 """Compute memory_parameter
141 With all available stages compute all combinations of memory parameter
142 and return the mean of all the memory_parameter found
143 BEWARE: can update memory_parameter_ & memory_activation_rec_
144 because of _compute_memories_layers_()
145 return: memory_parameter
146 """
147 if multi_run or (len(self.stages_a) < 5 and len(self.stages_b) < 5):
148 memory_parameter_list = []
149 for stage1 in self.stages_a:
150 if stage1.id_ in [0, (self.number_of_stage_ - 1)]:
151 continue
152 for stage2 in self.stages_a:
153 if stage2.id_ in [0, (self.number_of_stage_ - 1), stage1.id_]:
154 continue
155 mem_param = self._compute_memory_parameter_local_(stage1, stage2)
156 if mem_param != 0:
157 memory_parameter_list.append(mem_param)
158 for stage1 in self.stages_b:
159 if stage1.id_ not in [0, (self.number_of_stage_ - 1)]:
160 for stage2 in self.stages_b:
161 mem_param = self._compute_memory_parameter_local_(stage1, stage2)
162 if (stage2.id_ not in [0, (self.number_of_stage_ - 1),
163 stage1.id_] and mem_param != 0):
164 memory_parameter_list.append(mem_param)
165 return np.mean(memory_parameter_list)
166 if self._compute_memories_layers_():
167 return self.memory_parameter_
168 logger.error("Issue with _compute_memory_parameter_!!!")
169 return 0
171 def _compute_memory_activation_(self, rec, multi_run=False) -> float:
172 """
173 Compute memory_activation for a given recomputation type
174 return: memory_activation
175 """
176 if multi_run or (len(self.stages_a) < 5 and len(self.stages_b) < 5):
177 # look at solution 4 stages
178 logger.error("Not implemented yet!!!")
179 return 0
180 if self._compute_memories_layers_():
181 return self.memory_activation_rec_[rec]
182 logger.error("Issue with _compute_memory_activation_!!!")
183 return 0
185 def zero_offset(self) -> bool:
186 """Return ``True`` if every stage in ``stages_a`` hosts the same number of layers."""
187 nb_layer = self.stages_a[0].nb_layer_
188 for s in self.stages_a:
189 if s.nb_layer_ != nb_layer:
190 return False
191 return True
193 def _compute_memories_layers_(self) -> bool:
194 """check if enough stage number is provided"""
195 used_rec = Recompute.get_used_list(self.recompute_considered_)
196 used_rec_num = len(used_rec)
197 stage_num = len(self.stages_a)
198 if stage_num == used_rec_num + 3:
199 return self._compute_memories_layer_bodies_(False)
200 if stage_num >= used_rec_num + 4:
201 logger.info("Enabled const memory component because enough stages were given")
202 if self.zero_offset():
203 logger.error(
204 "The number of layer per stage cannot be the same for all stages "
205 "when const component is enabled. Some offset must be used"
206 )
207 return False
208 return self._compute_memories_layer_bodies_(True)
210 logger.error(
211 "%s stages found and (%s) recomputation considered"
212 "is not coherent. There should be 3 or 4 more stages than recomputation considered",
213 stage_num,
214 used_rec_num,
215 )
216 return False
218 def _compute_memories_layer_bodies_local_(
219 self, unused_rec: list[Recompute.TYPE],
220 stages: list[Stage]) -> tuple[float, float, float]:
221 """Compute memory_parameter & memory activation for all recomputation types
222 Require at least 3 Stages different from first and last stage
223 """
224 variable_factor_list = []
225 constant_memory_list = []
226 unused_rec.sort(reverse=True)
227 for stage in stages:
228 if stage.id_ not in [0, self.number_of_stage_ - 1]:
229 variable_factor_list.append(stage.get_index_memory_var())
230 for rec_i in unused_rec:
231 variable_factor_list[-1].pop(1 + rec_i)
232 constant_memory_list.append(stage.memory_usage_)
233 solution = list(
234 np.linalg.solve(np.array(variable_factor_list),
235 np.array(constant_memory_list)))
236 memory_param = solution.pop(0)
237 memory_act_rec = Recompute.assign_used(solution, unused_rec)
238 return (memory_param, memory_act_rec)
242 def _compute_memories_layer_bodies_local_with_fix_(
243 self, unused_rec: list[Recompute.TYPE],
244 stages: list[Stage]) -> tuple[float, float, float]:
245 """Compute memory_const, memory_parameter & memory activation for all recomputation types
246 Require at least 4 Stages different from first and last stage
247 """
248 variable_factor_list = []
249 constant_memory_list = []
250 unused_rec.sort(reverse=True)
251 for stage in stages:
252 if stage.id_ not in [0, self.number_of_stage_ - 1]:
253 variable_factor_list.append([1] + stage.get_index_memory_var())
254 for rec_i in unused_rec:
255 variable_factor_list[-1].pop(2 + rec_i)
256 constant_memory_list.append(stage.memory_usage_)
257 logger.debug(
258 "solve(\n %s, \n %s) ",
259 np.array(variable_factor_list),
260 np.array(constant_memory_list),
261 )
262 used_rec = Recompute.get_used_list(self.recompute_considered_)
263 used_rec_num = len(used_rec)
265 if len(stages) < used_rec_num + 4:
266 raise ValueError("Stages given are not enough to solve memory constraints")
267 if len(stages) == used_rec_num + 4:
268 solution = list(
269 np.linalg.solve(np.array(variable_factor_list),
270 np.array(constant_memory_list)))
271 else:
272 logger.warning("Stages given are more than needed, switch to least sqaures method")
273 solution = list(np.linalg.lstsq(np.array(variable_factor_list),
274 np.array(constant_memory_list), rcond=None)[0])
276 memory_const = solution.pop(0)
277 memory_param = solution.pop(0)
278 memory_act_rec = Recompute.assign_used(solution, unused_rec)
279 return (memory_const, memory_param, memory_act_rec)
281 def _compute_memories_layer_bodies_(self, with_fix: bool) -> bool:
282 """
283 Compute memory_parameter, memory_recompute, memory_activation
284 Require at least 3 Stages different from first and last stage
285 BEWARE: can update memory_parameter_, memory_recompute_, memory_activation_
286 return True if success to update memory_parameter_, memory_recompute_, memory_activation_
287 """
289 memory_const_a = None
290 memory_parameter_a = None
291 memory_recompute_a = {r: None for r in Recompute.TYPE}
293 memory_const_b = None
294 memory_parameter_b = None
295 memory_recompute_b = {r: None for r in Recompute.TYPE}
297 unused_rec = Recompute.get_unused_list(self.recompute_considered_)
298 logger.info("unused recomputation: %s", unused_rec)
300 if with_fix:
301 if len(self.stages_a) >= 5:
302 (memory_const_a,
303 memory_parameter_a,
304 memory_recompute_a) = (self._compute_memories_layer_bodies_local_with_fix_(
305 unused_rec, self.stages_a))
306 if len(self.stages_b) >= 5:
307 (memory_const_b,
308 memory_parameter_b,
309 memory_recompute_b) = (self._compute_memories_layer_bodies_local_with_fix_(
310 unused_rec, self.stages_b))
312 return self._average_if_needed_fix(
313 memory_const_a,
314 memory_parameter_a,
315 memory_recompute_a,
316 memory_const_b,
317 memory_parameter_b,
318 memory_recompute_b,
319 )
320 if len(self.stages_a) >= 5:
321 (memory_parameter_a,
322 memory_recompute_a) = (self._compute_memories_layer_bodies_local_(
323 unused_rec, self.stages_a))
324 if len(self.stages_b) >= 5:
325 (memory_parameter_b,
326 memory_recompute_b) = (self._compute_memories_layer_bodies_local_(
327 unused_rec, self.stages_b))
329 return self._average_if_needed(
330 memory_parameter_a,
331 memory_recompute_a,
332 memory_parameter_b,
333 memory_recompute_b,
334 )
336 def _average_if_needed_fix(
337 self,
338 memory_const_a,
339 memory_parameter_a,
340 memory_recompute_a,
341 memory_const_b,
342 memory_parameter_b,
343 memory_recompute_b,
344 ):
345 """check if average is needed"""
346 if memory_parameter_a is not None and memory_parameter_a != 0:
347 if memory_parameter_b is not None and memory_parameter_b != 0:
348 self.memory_const_ = (memory_const_a +
349 memory_const_b) / 2
350 self.memory_parameter_ = (memory_parameter_a +
351 memory_parameter_b) / 2
352 Recompute.average([memory_recompute_a, memory_recompute_b])
353 else:
354 self.memory_const_ = memory_const_a
355 self.memory_parameter_ = memory_parameter_a
356 self.memory_activation_rec_ = memory_recompute_a
358 elif memory_parameter_b is not None and memory_parameter_b != 0:
359 self.memory_const_ = memory_const_b
360 self.memory_parameter_ = memory_parameter_b
361 self.memory_activation_rec_ = memory_recompute_b
362 else:
363 logger.error("failed to compute memories")
364 return False
365 return True
367 def _average_if_needed(self, memory_parameter_a, memory_recompute_a, memory_parameter_b,
368 memory_recompute_b,):
369 """check if average is needed"""
370 if memory_parameter_a is not None and memory_parameter_a != 0:
371 if memory_parameter_b is not None and memory_parameter_b != 0:
372 self.memory_parameter_ = (memory_parameter_a + memory_parameter_b) / 2
373 Recompute.average([memory_recompute_a, memory_recompute_b])
374 else:
375 self.memory_parameter_ = memory_parameter_a
376 self.memory_activation_rec_ = memory_recompute_a
378 elif memory_parameter_b is not None and memory_parameter_b != 0:
379 self.memory_parameter_ = memory_parameter_b
380 self.memory_activation_rec_ = memory_recompute_b
381 else:
382 logger.error("failed to compute memories")
383 return False
384 return True
386 def _compute_memory_head_(self) -> float:
387 """compute the memory for the head"""
388 head_stages = filter_stage_id(self.stages_a, 0)
389 head_stages += filter_stage_id(self.stages_b, 0)
390 memory_head_list = []
391 mem_parameter = self.get_memory_parameter()
392 for head in head_stages:
393 head_memory = head.memory_usage_
394 for rec in Recompute.TYPE:
395 if self.recompute_considered_[rec] is True:
396 head_memory -= (head.nb_layer_rec_[rec] * self.get_memory_activation(
397 rec) * self.number_of_stage_)
398 head_memory -= (head.nb_layer_) * mem_parameter
399 memory_head_list.append(head_memory)
400 return np.mean(memory_head_list)
402 def _compute_memory_tail_(self) -> float:
403 """compute the memory for the tail"""
404 tail_stages = filter_stage_id(self.stages_a, self.number_of_stage_ - 1)
405 tail_stages += filter_stage_id(self.stages_b, self.number_of_stage_ - 1)
406 memory_tail_list = []
407 for tail in tail_stages:
408 tail_memory = tail.memory_usage_
409 for rec in Recompute.TYPE:
410 if self.recompute_considered_[rec] is True:
411 tail_memory -= (tail.nb_layer_rec_[rec] * self.get_memory_activation(rec) * 1)
412 tail_memory -= (tail.nb_layer_) * self.get_memory_parameter()
413 memory_tail_list.append(tail_memory)
414 return np.mean(memory_tail_list)
416 def get_memory_const(self) -> float:
417 """Return the solver-derived constant memory component per stage."""
418 return self.memory_const_
420 def get_memory_parameter(self, force_recompute: bool = False) -> float:
421 """Return the per-body-layer parameter memory, recomputing on demand."""
422 if force_recompute or self.memory_parameter_ is None:
423 self.memory_parameter_ = self._compute_memory_parameter_()
424 return self.memory_parameter_
426 def get_memory_activation(self, rec: Recompute.TYPE,
427 force_recompute: bool = False) -> float:
428 """Return the per-layer activation memory for a given recomputation type."""
429 if force_recompute or self.memory_activation_rec_[rec] is None:
430 self.memory_activation_rec_[rec] = self._compute_memory_activation_(rec)
431 return self.memory_activation_rec_[rec]
433 def get_memory_head(self, force_recompute: bool = False) -> float:
434 """Return the HEAD-layer memory, recomputing on demand."""
435 if force_recompute or self.memory_head_ is None:
436 self.memory_head_ = self._compute_memory_head_()
437 return self.memory_head_
439 def get_memory_tail(self, force_recompute: bool = False) -> float:
440 """Return the TAIL-layer memory, recomputing on demand."""
441 if force_recompute or self.memory_tail_ is None:
442 self.memory_tail_ = self._compute_memory_tail_()
443 return self.memory_tail_
446def compute_memories(layers: list[Layer], memory_folder: str, number_of_stage: int) -> list[Layer]:
447 """compute memories"""
448 filename = ""
449 # Put some meta information in a predefine .json file like layers info?
450 with open(memory_folder + filename, encoding="utf-8"):
451 pass
452 cm = ComputeMemory(number_of_stage=number_of_stage, stages_a=[], stages_b=[])
454 for layer in layers:
455 if layer.type_ == Layer.type_enum.HEAD:
456 layer.memory_parameter_ = cm.get_memory_head()
457 elif layer.type_ == Layer.type_enum.TAIL:
458 layer.memory_parameter_ = cm.get_memory_tail()
459 elif layer.type_ == Layer.type_enum.BODY:
460 layer.memory_parameter_ = cm.get_memory_parameter()
461 for rec in Recompute.TYPE:
462 layer.memory_activation_rec_[rec] = cm.get_memory_activation(rec)
463 return layers