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

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 

17 

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 

22 

23 

24class ComputeMemory: 

25 """ 

26 ComputeMemory class to compute the different memories with stages information running (dry) log 

27 

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) 

30 

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 """ 

45 

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 

55 

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. 

59 

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 

75 

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 

89 

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 

109 

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 

114 

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 

119 

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 

138 

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 

170 

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 

184 

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 

192 

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) 

209 

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 

217 

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) 

239 

240 

241 

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) 

264 

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]) 

275 

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) 

280 

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 """ 

288 

289 memory_const_a = None 

290 memory_parameter_a = None 

291 memory_recompute_a = {r: None for r in Recompute.TYPE} 

292 

293 memory_const_b = None 

294 memory_parameter_b = None 

295 memory_recompute_b = {r: None for r in Recompute.TYPE} 

296 

297 unused_rec = Recompute.get_unused_list(self.recompute_considered_) 

298 logger.info("unused recomputation: %s", unused_rec) 

299 

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)) 

311 

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)) 

328 

329 return self._average_if_needed( 

330 memory_parameter_a, 

331 memory_recompute_a, 

332 memory_parameter_b, 

333 memory_recompute_b, 

334 ) 

335 

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 

357 

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 

366 

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 

377 

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 

385 

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) 

401 

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) 

415 

416 def get_memory_const(self) -> float: 

417 """Return the solver-derived constant memory component per stage.""" 

418 return self.memory_const_ 

419 

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_ 

425 

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] 

432 

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_ 

438 

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_ 

444 

445 

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=[]) 

453 

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