Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_ppb / sapp / sapp_pipeline.py: 91%

426 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-21 04:29 +0800

1# Copyright 2026 Huawei Technologies Co., Ltd 

2# 

3# Licensed under the Apache License, Version 2.0 (the "License"); 

4# you may not use this file except in compliance with the License. 

5# You may obtain a copy of the License at 

6# 

7# http://www.apache.org/licenses/LICENSE-2.0 

8# 

9# Unless required by applicable law or agreed to in writing, software 

10# distributed under the License is distributed on an "AS IS" BASIS, 

11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. 

12# See the License for the specific language governing permissions and 

13# limitations under the License. 

14# ============================================================================ 

15"""High-level orchestrator around :class:`SappSolver`: build, solve, simulate, export YAML.""" 

16import os 

17import sys 

18from typing import Any, Dict, List, Optional, Union 

19 

20import matplotlib.pyplot as plt 

21import yaml 

22 

23import hyper_parallel.auto_parallel.sapp_ppb.simulator.pp_simulator as sim 

24import hyper_parallel.auto_parallel.sapp_ppb.utils.recompute as Recompute 

25from hyper_parallel.auto_parallel.sapp_ppb.sapp.sapp_solver import SappSolver 

26from hyper_parallel.auto_parallel.sapp_ppb.utils.check_rules import check_yaml_depth_before_loading 

27from hyper_parallel.auto_parallel.sapp_ppb.utils.layer import Layer, filter_layer_type 

28from hyper_parallel.auto_parallel.sapp_ppb.utils.logger import logger 

29 

30 

31class SappPipeline: 

32 """pipeline balancer""" 

33 

34 def __init__( 

35 self, 

36 model_name: str, 

37 num_of_stage: int, 

38 num_of_micro_batch: int, 

39 max_memory: int, 

40 layers: List[Layer], 

41 vpp_less_memory: bool = False, 

42 # Add arg dual 

43 dual: bool = False, 

44 num_of_interleave: int = 1, 

45 constant_memory: int = 0, 

46 optimization_level: int = 1, 

47 extracted_training_params: Optional[Dict[str, int]] = None, 

48 seq_split_num: int = 1, 

49 use_backward_time: bool = False, 

50 ) -> None: 

51 """Cache pipeline parameters and index the input ``layers`` by HEAD / BODY / TAIL. 

52 

53 Args: 

54 model_name (str): Model identifier, used for dump filenames and log prefixes. 

55 num_of_stage (int): Number of physical pipeline stages. 

56 num_of_micro_batch (int): Number of micro-batches scheduled per iteration. 

57 max_memory (int): Per-device memory budget in MB. 

58 layers (List[Layer]): Ordered list of layer descriptors covering HEAD/BODY/TAIL. 

59 vpp_less_memory (bool, optional): If ``True``, use the less-memory VPP scheduler variant. 

60 Default: ``False``. 

61 dual (bool, optional): Enable dualpipe-V scheduling support. Default: ``False``. 

62 num_of_interleave (int, optional): Virtual-pipeline (VPP) chunk count. Default: ``1``. 

63 constant_memory (int, optional): Constant per-stage memory overhead (MB). Default: ``0``. 

64 optimization_level (int, optional): Solver optimization level (``0-2``). Default: ``1``. 

65 extracted_training_params (Optional[Dict[str, int]], optional): Optional training-config parameters for 

66 seqpp. Default: ``None``. 

67 seq_split_num (int, optional): Number of sequence splits; ``>1`` enables sequence pipeline. 

68 Default: ``1``. 

69 """ 

70 self.model_name_ = model_name 

71 self.num_of_stage_ = num_of_stage 

72 self.num_of_micro_batch_ = num_of_micro_batch 

73 self.num_of_interleave_ = num_of_interleave 

74 self.max_memory_ = max_memory 

75 self.vpp_less_memory_ = vpp_less_memory 

76 # Add arg dual_ 

77 self.dual_ = dual 

78 self.constant_memory_ = constant_memory 

79 self.optimization_level = optimization_level 

80 self.extracted_training_params_ = extracted_training_params 

81 self.seq_split_num_ = seq_split_num 

82 self.use_backward_time_ = use_backward_time 

83 self.seqpipe_ = self.seq_split_num_ > 1 

84 # logger.output("seq chunk: %s",self.seq_split_num_) 

85 

86 self.problem_ = None 

87 self.layers_ = layers 

88 self.layers_sorted_ = { 

89 Layer.type_enum.HEAD: filter_layer_type(layers, 

90 Layer.type_enum.HEAD), 

91 Layer.type_enum.BODY: filter_layer_type(layers, 

92 Layer.type_enum.BODY), 

93 Layer.type_enum.TAIL: filter_layer_type(layers, 

94 Layer.type_enum.TAIL), 

95 } 

96 

97 @property 

98 def simulator(self): 

99 """Pipeline simulator instance (available after :meth:`simulate`).""" 

100 return self._simulator 

101 

102 def has_some_memory_info(self) -> bool: 

103 """Check if there is all information for memory constraint.""" 

104 return self.problem_.has_some_memory_info() 

105 

106 def construct_problem(self, solver: str = "pulp") -> None: 

107 """Construct the underlying ILP problem using the requested solver backend.""" 

108 if solver == "pulp": 

109 self.problem_ = self._construct_problem_pulp_() 

110 elif solver == "other": 

111 logger.warning( 

112 "No other solver available..., automatically switch to pulp!!!" 

113 ) 

114 self.problem_ = self._construct_problem_pulp_() 

115 else: 

116 logger.warning( 

117 "No other solver available..., automatically switch to pulp!!!" 

118 ) 

119 self.problem_ = self._construct_problem_pulp_() 

120 

121 def solve_problem(self, time_limit: int = 90, dump_folder: Optional[str] = None) -> None: 

122 """Solve the ILP, optionally dumping the LP model into ``dump_folder``.""" 

123 self.problem_.solve(time_limit, dump_folder) 

124 

125 def get_result(self) -> dict[str, list[list[str]]]: 

126 """Get result distribution of the solution (compact form).""" 

127 return self.problem_.result() 

128 

129 def get_memory_activation(self) -> list[float]: 

130 """Get the activation memory per stage for simulator.""" 

131 return self.problem_.get_simulator_memory_activation() 

132 

133 def get_memory_parameter(self) -> list[float]: 

134 """Get the parameter memory per stage for simulator.""" 

135 return self.problem_.get_simulator_memory_parameter() 

136 

137 def get_fw_time(self) -> list[float]: 

138 """Get the forward time per stage for simulator.""" 

139 time = self.problem_.get_simulator_forward_time() 

140 return time 

141 

142 def get_recompute_time(self) -> list[float]: 

143 """Get the recompute time per stage for simulator.""" 

144 time = self.problem_.get_simulator_recompute_time() 

145 return time 

146 

147 def get_time(self) -> list[float]: 

148 """Get the time per stage for simulator.""" 

149 return self.problem_.get_simulator_time() 

150 

151 def naive_layer_per_stage(self, 

152 layer_num: int, 

153 num_of_interleave: int = 1) -> List[List[int]]: 

154 """Return the naive layer-to-stage assignment (``layer_num`` evenly split).""" 

155 logger.output("layer_num = %s", layer_num) 

156 layer_count = layer_num // (self.num_of_stage_ * num_of_interleave) 

157 return [[layer_count] * self.num_of_stage_ for _ in range(num_of_interleave)] 

158 

159 def print_yaml_results(self) -> None: 

160 """Log the solver output in the MindFormers YAML schema.""" 

161 

162 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

163 nass = self.naive_layer_per_stage(layer.nb_layer_, 

164 self.num_of_interleave_) 

165 yaml_format = Recompute.yaml_from_internal( 

166 self.num_of_interleave_, 

167 self.num_of_stage_, 

168 self.problem_.variables_[layer.name_], 

169 nass, 

170 ) 

171 logger.output("layer-to-stage assignment baseline is \n\t%s", nass) 

172 yaml_results = "\nTo put in yaml configuration:" 

173 for y, v in yaml_format.items(): 

174 yaml_results += f"\n\t{y}: {v}" 

175 logger.output(yaml_results) 

176 

177 def get_manual_memory_activation( 

178 self, 

179 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

180 interleave_num: int = 1) -> List[List[float]]: 

181 """Return the per-stage activation memory for a user-supplied layer assignment.""" 

182 memory_active = [] 

183 if self.has_some_memory_info(): 

184 for inter in range(interleave_num): 

185 memory_active.append([]) 

186 for stage in range(self.num_of_stage_): 

187 memory_activation = 0 

188 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

189 memory_activation += self._get_layer_memory_activation( 

190 each_layer_per_recompute, layer, inter, stage 

191 ) 

192 memory_active[inter].append(memory_activation) 

193 return memory_active 

194 

195 @staticmethod 

196 def _get_layer_memory_activation(each_layer_per_recompute, layer, interleave, stage): 

197 """Calculate activation memory for one layer at one pipeline position.""" 

198 memory_activation = 0 

199 unused_recompute_list = Recompute.get_unused_list(each_layer_per_recompute[layer]) 

200 for rec in Recompute.TYPE: 

201 if rec in unused_recompute_list: 

202 continue 

203 value = each_layer_per_recompute[layer][rec][interleave][stage] 

204 if value > 0: 

205 memory_activation += value * layer.memory_activation_rec_[rec] 

206 return memory_activation 

207 

208 def get_manual_memory_parameter( 

209 self, 

210 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

211 interleave_num: int = 1) -> List[List[float]]: 

212 """Return the per-stage parameter memory for a user-supplied layer assignment.""" 

213 memory_param_stage = [0] * self.num_of_stage_ 

214 for inter in range(interleave_num): 

215 for stage in range(self.num_of_stage_): 

216 for rec in Recompute.TYPE: 

217 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

218 if layer.memory_parameter_ is None: 

219 continue 

220 

221 if rec in Recompute.get_unused_list(each_layer_per_recompute[layer]): 

222 continue 

223 

224 value = each_layer_per_recompute[layer][rec][inter][stage] 

225 if value <= 0: 

226 continue 

227 

228 memory_param_stage[stage] += value * layer.memory_parameter_ 

229 for head in self.layers_sorted_[Layer.type_enum.HEAD]: 

230 if head.memory_parameter_ is not None: 

231 memory_param_stage[0] += head.memory_parameter_ 

232 for tail in self.layers_sorted_[Layer.type_enum.TAIL]: 

233 if tail.memory_parameter_ is not None: 

234 memory_param_stage[self.num_of_stage_ - 

235 1] += tail.memory_parameter_ 

236 memory_param = [memory_param_stage] * interleave_num 

237 return memory_param 

238 

239 def get_manual_time( 

240 self, 

241 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

242 interleave_num: int = 1) -> List[List[float]]: 

243 """Return the per-stage execution time for a user-supplied layer assignment.""" 

244 time = [] 

245 for i in range(interleave_num): 

246 time.append([]) 

247 for s in range(self.num_of_stage_): 

248 time[i].append(0) 

249 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

250 for r in Recompute.TYPE: 

251 if each_layer_per_recompute[layer][r][i][s] > 0: 

252 time[i][s] += each_layer_per_recompute[layer][r][i][s] * ( 

253 layer.forward_time_ + 

254 layer.backward_time_rec_[r]) 

255 

256 for head in self.layers_sorted_[Layer.type_enum.HEAD]: 

257 time[0][0] += head.forward_time_ + head.backward_time_rec_[Recompute.TYPE.NONE] 

258 for tail in self.layers_sorted_[Layer.type_enum.TAIL]: 

259 time[interleave_num - 1][self.num_of_stage_ - 1] += ( 

260 tail.forward_time_ 

261 + tail.backward_time_rec_[Recompute.TYPE.NONE] 

262 ) 

263 return time 

264 

265 def get_manual_fw_time( 

266 self, 

267 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

268 interleave_num: int = 1) -> List[List[float]]: 

269 """Return the per-stage forward time for a user-supplied layer assignment.""" 

270 time = [] 

271 for i in range(interleave_num): 

272 time.append([]) 

273 for s in range(self.num_of_stage_): 

274 time[i].append(0) 

275 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

276 for r in Recompute.TYPE: 

277 if (r not in Recompute.get_unused_list(each_layer_per_recompute[layer]) 

278 and each_layer_per_recompute[layer][r][i][s] > 0): 

279 time[i][s] += each_layer_per_recompute[layer][r][i][s] * ( 

280 layer.forward_time_) 

281 for head in self.layers_sorted_[Layer.type_enum.HEAD]: 

282 time[0][0] += head.forward_time_ 

283 for tail in self.layers_sorted_[Layer.type_enum.TAIL]: 

284 time[interleave_num - 1][self.num_of_stage_ - 1] += tail.forward_time_ 

285 return time 

286 

287 def get_manual_backward_time( 

288 self, 

289 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

290 interleave_num: int = 1) -> List[List[float]]: 

291 """Return the per-stage backward time for a user-supplied layer assignment.""" 

292 time = [] 

293 for i in range(interleave_num): 

294 time.append([]) 

295 for s in range(self.num_of_stage_): 

296 time[i].append(0) 

297 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

298 for r in Recompute.TYPE: 

299 if (r not in Recompute.get_unused_list(each_layer_per_recompute[layer]) 

300 and each_layer_per_recompute[layer][r][i][s] > 0): 

301 time[i][s] += each_layer_per_recompute[layer][r][i][s] * ( 

302 layer.backward_time_rec_[r]) 

303 for head in self.layers_sorted_[Layer.type_enum.HEAD]: 

304 time[0][0] += head.backward_time_rec_[Recompute.TYPE.NONE] 

305 for tail in self.layers_sorted_[Layer.type_enum.TAIL]: 

306 time[interleave_num - 1][self.num_of_stage_ - 1] += ( 

307 tail.backward_time_rec_[Recompute.TYPE.NONE] 

308 ) 

309 return time 

310 

311 def get_manual_recompute_time( 

312 self, 

313 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

314 interleave_num: int = 1) -> List[List[float]]: 

315 """Return the per-stage recompute-only time for a user-supplied layer assignment.""" 

316 logger.output("each_layer_per_recompute = %s", each_layer_per_recompute) 

317 time_all_rec = [] 

318 time_no_rec = [] 

319 for i in range(interleave_num): 

320 time_all_rec.append([]) 

321 time_no_rec.append([]) 

322 for s in range(self.num_of_stage_): 

323 time_all_rec[i].append(0) 

324 time_no_rec[i].append(0) 

325 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

326 self._add_manual_recompute_time( 

327 each_layer_per_recompute, layer, i, s, time_all_rec, time_no_rec) 

328 

329 return [[r - n for r, n in zip(ar, nr)] 

330 for ar, nr in zip(time_all_rec, time_no_rec)] 

331 

332 def _add_manual_recompute_time(self, each_layer_per_recompute, layer, interleave, stage, 

333 time_all_rec, time_no_rec): 

334 """Accumulate recompute time for a single layer and stage.""" 

335 logger.output("backward_time_rec_(%s) = %s", layer, layer.backward_time_rec_) 

336 unused_rec = Recompute.get_unused_list(each_layer_per_recompute[layer]) 

337 for rec in Recompute.TYPE: 

338 layer_num = each_layer_per_recompute[layer][rec][interleave][stage] 

339 if rec in unused_rec or layer_num <= 0: 

340 continue 

341 if layer.backward_time_rec_[rec] is None: 

342 raise ValueError("No backward tme is specified for this " 

343 "recomputation. Recomputation " 

344 f"'{Recompute.YAML_NAME[rec]}' is likely not considered") 

345 logger.output("r = %s; i = %s; s = %s", rec, interleave, stage) 

346 time_all_rec[interleave][stage] += layer_num * layer.backward_time_rec_[rec] 

347 time_no_rec[interleave][stage] += layer_num * layer.backward_time_rec_[Recompute.TYPE.NONE] 

348 

349 def simulate(self, show: bool = True, file_name: Optional[str] = None, 

350 sub_fig: Optional[plt.Figure] = None, comm_time: float = 0.0) -> float: 

351 """Run the simulator on the solved schedule and return its estimated total time.""" 

352 forward_time = self.get_fw_time() 

353 recompute_overhead = self.get_recompute_time() 

354 backward_time = self.problem_.get_simulator_backward_time() if self.use_backward_time_ else 0 

355 stage_mem_par = 0 

356 stage_mem_act = 0 

357 if self.has_some_memory_info(): 

358 stage_mem_par = self.get_memory_parameter() 

359 stage_mem_act = self.get_memory_activation() 

360 

361 return self.simulation( 

362 forward_time, 

363 recompute_overhead, 

364 stage_mem_par, 

365 stage_mem_act, 

366 self.constant_memory_, 

367 backward_time=backward_time, 

368 show=show, 

369 file_name=file_name, 

370 sub_fig=sub_fig, 

371 comm_time=comm_time, 

372 ) 

373 

374 def simulate_naive(self, layers: List[Layer], output_folder: str) -> None: 

375 """Simulate the naive (even) layer-to-stage assignments for sanity comparison.""" 

376 num_layers = 0 

377 rec_considered = {} 

378 for layer in layers: 

379 if layer.type_ == Layer.type_enum.BODY: 

380 num_layers = layer.nb_layer_ 

381 rec_considered = layer.recompute_considered_ 

382 

383 all_recomp = {"offset": 0} 

384 no_recomp = {"offset": 0} 

385 for rec in [Recompute.TYPE.FULL, Recompute.TYPE.SLCT, Recompute.TYPE.COMM]: 

386 if rec_considered.get(rec, False): 

387 all_recomp[Recompute.YAML_NAME[rec]] = True 

388 no_recomp[Recompute.YAML_NAME[rec]] = False 

389 

390 self.simulate_yaml( 

391 yaml_format=all_recomp, 

392 show=True, 

393 interleave_num=self.num_of_interleave_, 

394 file_name=os.path.join(output_folder, 

395 "result_naive_all_recomp.svg"), 

396 ) 

397 

398 if num_layers % self.num_of_stage_ == 0: 

399 self.simulate_yaml( 

400 yaml_format=no_recomp, 

401 show=True, 

402 interleave_num=self.num_of_interleave_, 

403 file_name=os.path.join(output_folder, 

404 "result_naive_no_recomp.svg"), 

405 ) 

406 else: 

407 logger.warning("num layer cannot be divided by num stage") 

408 

409 def simulate_comparison(self, manual_config_file: str, output_folder: str) -> None: 

410 """Render side-by-side automatic vs manual simulations for every entry in the YAML.""" 

411 with open(manual_config_file, encoding="utf-8") as fp: 

412 check_yaml_depth_before_loading(fp) 

413 fp.seek(0) 

414 data = yaml.safe_load(fp) 

415 yaml_data = {} 

416 for manual in data.values(): 

417 yaml_data[Recompute.OFFSET] = manual.get(Recompute.OFFSET) 

418 if isinstance(yaml_data[Recompute.OFFSET], list) and all( 

419 isinstance(item, int) for item in yaml_data[Recompute.OFFSET]): 

420 yaml_data[Recompute.OFFSET] = [yaml_data[Recompute.OFFSET]] 

421 

422 for rec in Recompute.YAML_NAME.values(): 

423 yaml_data[rec] = manual.get(rec) 

424 if isinstance(yaml_data[rec], list) and all( 

425 isinstance(item, int) for item in yaml_data[rec]): 

426 yaml_data[rec] = [yaml_data[rec]] 

427 interleave_num = manual.get("interleave_num", 

428 self.num_of_interleave_) 

429 show = manual.get("show", False) 

430 file_name = manual.get("file_name") 

431 full_file_name = os.path.join(output_folder, 

432 file_name) if (file_name) else None 

433 

434 fig = plt.figure(figsize=(24, 8)) 

435 sub_figs = fig.subfigures(1, 2, wspace=0.07) 

436 sub_figs[0].suptitle('Automatic', fontsize='x-large') 

437 try: 

438 simulate_result = self.simulate( 

439 show=False, 

440 file_name=os.path.join(output_folder, "Auto_" + file_name), 

441 sub_fig=sub_figs[0], 

442 ) 

443 except Exception: 

444 logger.exception("Failed to simulate auto pipeline.") 

445 raise 

446 

447 if simulate_result is None: 

448 raise RuntimeError("simulate() returned None.") 

449 

450 sub_figs[1].suptitle('Manual', fontsize='x-large') 

451 self.simulate_yaml(yaml_data, False, interleave_num, full_file_name, sub_figs[1]) 

452 plt.savefig(os.path.join(output_folder, "Comparison_" + file_name)) 

453 if show: 

454 plt.show() 

455 

456 def simulate_only_manual(self, manual_config_file: str, output_folder: str) -> None: 

457 """Render only the manual simulation for every entry in ``manual_config_file``.""" 

458 with open(manual_config_file, encoding="utf-8") as fp: 

459 check_yaml_depth_before_loading(fp) 

460 fp.seek(0) 

461 data = yaml.safe_load(fp) 

462 yaml_data = {} 

463 for manual in data.values(): 

464 yaml_data[Recompute.OFFSET] = manual.get(Recompute.OFFSET) 

465 if isinstance(yaml_data[Recompute.OFFSET], list) and all( 

466 isinstance(item, int) for item in yaml_data[Recompute.OFFSET]): 

467 yaml_data[Recompute.OFFSET] = [yaml_data[Recompute.OFFSET]] 

468 

469 for rec in Recompute.YAML_NAME.values(): 

470 yaml_data[rec] = manual.get(rec) 

471 if isinstance(yaml_data[rec], list) and all( 

472 isinstance(item, int) for item in yaml_data[rec]): 

473 yaml_data[rec] = [yaml_data[rec]] 

474 interleave_num = manual.get("interleave_num", 

475 self.num_of_interleave_) 

476 show = manual.get("show", False) 

477 file_name = manual.get("file_name") 

478 full_file_name = os.path.join(output_folder, 

479 file_name) if (file_name) else None 

480 

481 fig = plt.figure(figsize=(12, 8)) 

482 self.simulate_yaml(yaml_data, False, interleave_num, full_file_name, fig) 

483 plt.savefig(os.path.join(output_folder, "manual_file_" + file_name)) 

484 if show: 

485 plt.show() 

486 

487 def simulate_yaml(self, yaml_format: Dict[str, Any], show: bool = True, 

488 interleave_num: int = 1, 

489 file_name: Optional[str] = None, 

490 sub_fig: Optional[plt.Figure] = None) -> float: 

491 """Simulate a manual pipeline configuration encoded as a YAML-compatible dict.""" 

492 layer_num = 0 

493 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

494 layer_num += layer.nb_layer_ 

495 nass = self.naive_layer_per_stage(layer_num, 

496 num_of_interleave=interleave_num) 

497 layer_per_recompute = Recompute.internal_from_yaml( 

498 interleave_num, self.num_of_stage_, yaml_format, nass) 

499 each_layer_per_recompute = self.split_layer_per_recompute(layer_per_recompute) 

500 return self.simulate_manual( 

501 each_layer_per_recompute, 

502 show, 

503 interleave_num=interleave_num, 

504 file_name=file_name, 

505 sub_fig=sub_fig 

506 ) 

507 

508 ####################################################################### 

509 ## ## 

510 ## Print Solver Model ## 

511 ## ## 

512 ####################################################################### 

513 def _calculate_activation_memory(self, each_layer_per_recompute, v, s): 

514 """Calculate activation memory for next and current stage""" 

515 act_mem_next = 0 

516 act_mem_curr = 0 

517 

518 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

519 for rec in Recompute.TYPE: 

520 if self.problem_.recompute_considered_[rec]: 

521 if each_layer_per_recompute[layer][rec][v + 1][s] > 0: # next 

522 act_mem_next += (each_layer_per_recompute[layer][rec][v + 1][s] * 

523 layer.memory_activation_rec_[rec]) 

524 if each_layer_per_recompute[layer][rec][v][s] > 0: # current 

525 act_mem_curr += (each_layer_per_recompute[layer][rec][v][s] * 

526 layer.memory_activation_rec_[rec]) 

527 

528 return act_mem_next, act_mem_curr 

529 

530 def _compute_parameter_memory_manually_solver(self, each_layer_per_recompute, s, interleave_num=1): 

531 """Solver memory model: parameter memory""" 

532 param_mem = 0 

533 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

534 if layer.memory_parameter_ is not None: 

535 param_mem += self._calculate_layer_parameter_memory( 

536 layer, each_layer_per_recompute[layer], s, interleave_num) 

537 return param_mem 

538 

539 def _calculate_layer_parameter_memory(self, layer, layer_per_recompute, s, interleave_num): 

540 """Calculate parameter memory for a single layer""" 

541 layer_mem = 0 

542 for inter in range(interleave_num): 

543 for rec in Recompute.TYPE: 

544 if self.problem_.recompute_considered_[rec]: 

545 if layer_per_recompute[rec][inter][s] > 0: 

546 layer_mem += layer_per_recompute[rec][inter][s] * layer.memory_parameter_ 

547 return layer_mem 

548 

549 def _calculate_activation_memory_solver(self, each_layer_per_recompute, s, interleave_num, activation_nums): 

550 """Calculate activation memory for a given stage""" 

551 act_mem = 0 

552 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

553 for inter in range(interleave_num): 

554 for rec in Recompute.TYPE: 

555 if self.problem_.recompute_considered_[rec]: 

556 if each_layer_per_recompute[layer][rec][inter][s] > 0: 

557 act_mem += (each_layer_per_recompute[layer][rec][inter][s] * 

558 layer.memory_activation_rec_[rec] * 

559 activation_nums[inter][s]) 

560 return act_mem 

561 

562 

563 def debug_print_manual_theoretical_memory( 

564 self, 

565 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]], 

566 interleave_num: int = 1) -> None: 

567 """Log the per-stage theoretical memory implied by the solver model (debug aid).""" 

568 logger.info("%s Manual Theoretical Memory Analysis %s", "=" * 20, "=" * 20) 

569 

570 if self.vpp_less_memory_: 

571 if self.seqpipe_: 

572 activation_nums = self.problem_.compute_activation_seq_nums( 

573 self.num_of_stage_, interleave_num, self.seq_split_num_, self.num_of_micro_batch_, True) 

574 else: 

575 activation_nums = self.problem_.compute_less_activation_nums( 

576 self.num_of_stage_, interleave_num) 

577 # Add if dual to decide whether dualpipe_v is used 

578 elif self.dual_: 

579 activation_nums = self.problem_.compute_activation_nums_dual( 

580 self.num_of_stage_, interleave_num, self.num_of_micro_batch_) 

581 else: 

582 if self.seqpipe_: 

583 activation_nums = self.problem_.compute_activation_seq_nums( 

584 self.num_of_stage_, interleave_num, self.seq_split_num_, self.num_of_micro_batch_, False) 

585 else: 

586 activation_nums = self.problem_.compute_activation_nums( 

587 self.num_of_stage_, interleave_num, self.num_of_micro_batch_) 

588 

589 logger.info("Activation nums = %s", activation_nums) 

590 

591 # compute for each stage 

592 for s in range(self.num_of_stage_): 

593 

594 # parameter memory 

595 param_mem = self._compute_parameter_memory_manually_solver(each_layer_per_recompute, s, interleave_num) 

596 

597 # head memory 

598 if s == 0: 

599 for head in self.layers_sorted_[Layer.type_enum.HEAD]: 

600 if head.memory_parameter_ is not None: 

601 param_mem += head.memory_parameter_ 

602 

603 # tail memory 

604 if s == self.num_of_stage_ - 1: 

605 for tail in self.layers_sorted_[Layer.type_enum.TAIL]: 

606 if tail.memory_parameter_ is not None: 

607 param_mem += tail.memory_parameter_ 

608 

609 # act memory 

610 act_mem = self._calculate_activation_memory_solver(each_layer_per_recompute, s, 

611 interleave_num, activation_nums) 

612 

613 # overhead 

614 overhead = 0 

615 

616 total = param_mem + act_mem + overhead + self.constant_memory_ 

617 

618 logger.info("Stage %d Manual Memory Analysis:", s) 

619 logger.info("Parameter Memory: %.2f", param_mem) 

620 logger.info("Activation Memory: %.2f", act_mem) 

621 logger.info("Memory Overhead: %.2f", overhead) 

622 logger.info("Constant Memory: %.2f", self.constant_memory_) 

623 logger.info("Total Theoretical Memory: %.2f", total) 

624 

625 def split_layer_per_recompute( 

626 self, 

627 layer_per_recompute: Dict[Recompute.TYPE, List[List[int]]] 

628 ) -> Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]]: 

629 """Split aggregate per-recompute layer counts into counts per BODY layer.""" 

630 each_layer_per_recompute = {} 

631 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

632 rest = layer.nb_layer_ 

633 each_layer_per_recompute[layer] = {r: [] for r in Recompute.TYPE} 

634 for rec in Recompute.TYPE: 

635 for i in range(self.num_of_interleave_): 

636 each_layer_per_recompute[layer][rec].append([0]*self.num_of_stage_) 

637 for s in range(self.num_of_stage_): 

638 subtract = min(layer_per_recompute[rec][i][s], rest) 

639 layer_per_recompute[rec][i][s] -= subtract 

640 rest -= subtract 

641 each_layer_per_recompute[layer][rec][i][s] += subtract 

642 return each_layer_per_recompute 

643 

644 def fuse_layer_per_recompute( 

645 self, 

646 each_layer_per_recompute: Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]] 

647 ) -> Dict[Recompute.TYPE, List[List[int]]]: 

648 """Fuse per-layer recompute counts back into aggregate per-recompute-type totals.""" 

649 all_layers_per_recompute = {r: [] for r in Recompute.TYPE} 

650 for rec in Recompute.TYPE: 

651 for i in range(self.num_of_interleave_): 

652 all_layers_per_recompute[rec].append([]) 

653 for s in range(self.num_of_stage_): 

654 all_layers_per_recompute[rec][i].append(sum( 

655 each_layer_per_recompute[layer][rec][i][s] 

656 for layer in self.layers_sorted_[Layer.type_enum.BODY] 

657 )) 

658 return all_layers_per_recompute 

659 

660 

661 def simulate_manual( 

662 self, 

663 each_layer_per_recompute: Optional[Dict[Layer, Dict[Recompute.TYPE, List[List[int]]]]] = None, 

664 show: bool = True, 

665 interleave_num: int = 1, 

666 file_name: Optional[str] = None, 

667 sub_fig: Optional[plt.Figure] = None) -> float: 

668 """Run the simulator on a user-supplied per-layer recompute strategy.""" 

669 logger.output("Simulating given strategy: %s", each_layer_per_recompute) 

670 

671 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

672 for rec in Recompute.TYPE: 

673 if len(each_layer_per_recompute[layer][rec]) != interleave_num: 

674 logger.error( 

675 "For layer %s with recompute %s, %s does not match interleave number %s", 

676 layer, 

677 rec, 

678 len(each_layer_per_recompute[layer][rec]), 

679 interleave_num, 

680 ) 

681 return sys.maxsize 

682 

683 for layer in self.layers_sorted_[Layer.type_enum.BODY]: 

684 for rec in Recompute.TYPE: 

685 if any(x < 0 for sublist in each_layer_per_recompute[layer][rec] 

686 for x in sublist): 

687 raise ValueError( 

688 f"for {rec}, there is strategy less than 0 in " 

689 f"{each_layer_per_recompute[layer][rec]}" 

690 ) 

691 

692 forward_time = self.get_manual_fw_time(each_layer_per_recompute, 

693 interleave_num) 

694 recompute_overhead = self.get_manual_recompute_time( 

695 each_layer_per_recompute, interleave_num) 

696 backward_time = ( 

697 self.get_manual_backward_time( 

698 each_layer_per_recompute, interleave_num) 

699 if self.use_backward_time_ 

700 else 0 

701 ) 

702 stage_mem_par = 0 

703 stage_mem_act = 0 

704 if self.has_some_memory_info(): 

705 stage_mem_par = self.get_manual_memory_parameter( 

706 each_layer_per_recompute, interleave_num=interleave_num) 

707 stage_mem_act = self.get_manual_memory_activation( 

708 each_layer_per_recompute, interleave_num=interleave_num) 

709 

710 self.debug_print_manual_theoretical_memory(each_layer_per_recompute, interleave_num) 

711 

712 return self.simulation( 

713 forward_time, 

714 recompute_overhead, 

715 stage_mem_par, 

716 stage_mem_act, 

717 constant_mem=self.constant_memory_, 

718 backward_time=backward_time, 

719 show=show, 

720 file_name=file_name, 

721 sub_fig=sub_fig, 

722 comm_time=0.0, 

723 ) 

724 

725 def simulation( 

726 self, 

727 forward_time: List[List[float]], 

728 recompute_overhead: Union[int, List[List[float]]] = 0, 

729 stage_mem_par: Union[int, List[List[float]]] = 0, 

730 stage_mem_act: Union[int, List[List[float]]] = 0, 

731 constant_mem: int = 0, 

732 backward_time: Union[int, List[List[float]]] = 0, 

733 show: bool = True, 

734 file_name: Optional[str] = None, 

735 sub_fig: Optional[plt.Figure] = None, 

736 comm_time: float = 0.0, 

737 ) -> float: 

738 """Run the low-level :class:`PipelineSimulator` and return its reported end time.""" 

739 use_comm = comm_time > 0.0 

740 if self.has_some_memory_info(): 

741 logger.output( 

742 "PipelineSimulator(\n\t%s, %s," 

743 "\n\tblock_mem_act=%s," 

744 "\n\tblock_mem_par=%s," 

745 "\n\tlayer_recompute=%s," 

746 "\n\tbackward_time=%s," 

747 "\n\tless_memory=%s )", 

748 forward_time, 

749 self.num_of_micro_batch_, 

750 stage_mem_act, 

751 stage_mem_par, 

752 recompute_overhead, 

753 backward_time, 

754 self.vpp_less_memory_, 

755 ) 

756 

757 sim_method = "vpp2" if self.vpp_less_memory_ else "vpp" 

758 simulator = sim.PipelineSimulator( 

759 forward_time, 

760 self.num_of_micro_batch_, 

761 comm_time=comm_time, 

762 block_mem=stage_mem_act, 

763 block_mem_par=stage_mem_par, 

764 constant_mem=constant_mem, 

765 layer_recompute=recompute_overhead, 

766 backward_time=backward_time, 

767 method=sim_method, 

768 sub_fig=sub_fig 

769 ) 

770 else: 

771 logger.output( 

772 "PipelineSimulator(\n\t%s, %s," 

773 "\n\tlayer_recompute=%s," 

774 "\n\tbackward_time=%s," 

775 "\n\tless_memory=%s )", 

776 forward_time, 

777 self.num_of_micro_batch_, 

778 recompute_overhead, 

779 backward_time, 

780 self.vpp_less_memory_, 

781 ) 

782 simulator = sim.PipelineSimulator( 

783 forward_time, 

784 self.num_of_micro_batch_, 

785 comm_time=comm_time, 

786 layer_recompute=recompute_overhead, 

787 backward_time=backward_time, 

788 less_memory=self.vpp_less_memory_, 

789 sub_fig=sub_fig 

790 ) 

791 

792 simulator.run(comm=use_comm) 

793 self._simulator = simulator 

794 if file_name: 

795 simulator.save(file_name) 

796 if show: 

797 simulator.show() 

798 return simulator.end_time 

799 

800 def _construct_problem_pulp_(self) -> SappSolver: 

801 """construct the problem using pulp""" 

802 prob = SappSolver( 

803 num_of_stage=self.num_of_stage_, 

804 num_of_micro_batch=self.num_of_micro_batch_, 

805 num_of_interleave=self.num_of_interleave_, 

806 max_memory=self.max_memory_, 

807 vpp_less_memory=self.vpp_less_memory_, 

808 # Add arg dual 

809 dual = self.dual_, 

810 constant_memory=self.constant_memory_, 

811 layers=self.layers_, 

812 layers_sorted=self.layers_sorted_, 

813 optimization_level=self.optimization_level, 

814 extracted_training_params=self.extracted_training_params_, 

815 seq_split_num=self.seq_split_num_ 

816 ) 

817 return prob 

818 

819 def _recompute_considered(self): 

820 return self.problem_.recompute_considered_ 

821 

822 

823def choose_interleave( 

824 model_name: str, 

825 number_of_stage: int, 

826 number_of_micro_batch: int, 

827 max_memory: int, 

828 layers: list[Layer], 

829) -> tuple[int, int, dict[str, list[list[str]]]]: 

830 """Simulates different interleaves and returns the best.""" 

831 max_inter = 4 

832 best_time = int(sys.maxsize) 

833 best_inter = 1 

834 best_distribution = {} 

835 

836 for inter in range(1, max_inter + 1): 

837 pipe = SappPipeline( 

838 model_name=model_name, 

839 num_of_stage=number_of_stage, 

840 num_of_micro_batch=number_of_micro_batch, 

841 max_memory=max_memory, 

842 layers=layers, 

843 num_of_interleave=inter, 

844 ) 

845 

846 pipe.construct_problem(solver="pulp") 

847 pipe.solve_problem() 

848 time = pipe.simulate(show=False) 

849 logger.output("for interleave %s, time = %s", inter, time) 

850 if time < best_time: 

851 best_time = time 

852 best_inter = inter 

853 best_distribution = pipe.get_result() 

854 

855 return (best_inter, best_time, best_distribution) 

856 

857 

858def flatten(inter_stage_list: List[List[float]]) -> List[float]: 

859 """Collapse an ``[interleave][stage]`` matrix into a per-stage list via summation.""" 

860 stage_list = [0] * len(inter_stage_list[0]) 

861 for inter, _ in enumerate(inter_stage_list): 

862 for stage, _ in enumerate(inter_stage_list[inter]): 

863 stage_list[stage] += inter_stage_list[inter][stage] 

864 return stage_list