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

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

16 

17import os 

18from dataclasses import dataclass 

19from enum import IntEnum 

20from typing import Any, Dict, List, Optional 

21 

22import pulp as lpSolver 

23 

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 

27 

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 

34 

35 

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 

46 

47 

48class SappSolver: 

49 """solver for pipeline balance""" 

50 

51 BIG_M = 1000000 

52 

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) 

63 

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. 

82 

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

98 

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 

110 

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

118 

119 self.variables_ = self._create_variables_to_solve_( 

120 num_of_stage, num_of_interleave, layers_sorted) 

121 self.problem_ = self._create_problem_(description) 

122 

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

128 

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) 

137 

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 ) 

170 

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

179 

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

188 

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

197 

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) 

211 

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 

225 

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

230 

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

236 

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) 

256 

257 return activation_nums 

258 

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

264 

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) 

278 

279 return activation_nums 

280 

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 

299 

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 

316 

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 

336 

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) 

351 

352 logger.output("compute_activation_seq_nums: %s", activation_nums) 

353 return activation_nums 

354 

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 

376 

377 seq_memory_activation = (original_memory_activation - key_mem - value_mem) / seq_split_num + kv_update_mem 

378 return seq_memory_activation 

379 

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 

393 

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 

423 

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 

440 

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 

449 

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 

465 

466 terms.append( 

467 variables[ 

468 body_layer.name_ 

469 ][rec][i][s] 

470 ) 

471 

472 prob += lpSolver.lpSum(terms) >= 1 

473 return prob 

474 

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

480 

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) 

491 

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 

502 

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 ) 

510 

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 

520 

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 

527 

528 for stage in stage_range: 

529 terms.append(variables[layer][rec][interleave][stage]) 

530 

531 return lpSolver.lpSum(terms) 

532 

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) 

542 

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 

551 

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 

560 

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

565 

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 

572 

573 if len(considered) < 2: 

574 # this constraint is unnecessary if there is no recomputation 

575 return prob 

576 

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 

588 

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 

598 

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. 

603 

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

609 

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 

614 

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

621 

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

627 

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

641 

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

646 

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

651 

652 return (self.variables_[self.MAX_STAGE_TIME] + 

653 self.variables_[self.MAX_LAST_CHUNK]) 

654 

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) 

661 

662 total_sum = self.variables_[self.TOTAL_SUM] 

663 prob += total_sum >= self._total_sum(layers_sorted) 

664 

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

670 

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

674 

675 chunks_sum = lpSolver.lpSum(self.variables_[self.CHUNKS_SUM]) 

676 prev_diff = lpSolver.lpSum(self.variables_[self.PREV_DIFF]) 

677 

678 next_diff = self.variables_[self.NEXT_DIFF] 

679 prob += next_diff >= ( 

680 self._next_diff_sum(layers_sorted, prob)) 

681 

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 

691 

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

709 

710 return max_stage_time 

711 

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 

726 

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 

743 

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 

756 

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 

772 

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 

788 

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 

838 

839 return bound 

840 

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 

847 

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 

858 

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) 

870 

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) 

879 

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 ) 

903 

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 

914 

915 def get_simulator_memory_activation(self) -> list[float]: 

916 """Give the activation memory per stage for simulator.""" 

917 

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 

926 

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) 

934 

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 

944 

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 

953 

954 var_value = self.variables_.get(layer.name_)[rec][interleave][stage].varValue 

955 total += var_value * layer.memory_parameter_ 

956 return total 

957 

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

972 

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 

979 

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 

998 

999 def get_simulator_backward_time(self) -> list[float]: 

1000 """Give the backward time per stage for simulator. 

1001 

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 

1022 

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

1045 

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

1049 

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 

1061 

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

1068 

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) 

1075 

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) 

1097 

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 ) 

1107 

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 ) 

1115 

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 ) 

1123 

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) 

1128 

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 ) 

1140 

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

1148 

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) 

1152 

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

1167 

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

1177 

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

1185 

1186 overhead = 0 

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

1188 

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) 

1195 

1196 

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. 

1199 

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) 

1208 

1209 self.print_results() 

1210 

1211 self.debug_print_solver_theoretical_memory() 

1212 

1213 for name, result in self.result().items(): 

1214 logger.output("%s %s %s", name, result, "\n") 

1215 

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 

1234 

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) 

1247 

1248 # Var to Minimize 

1249 prob += pipeline_total_time 

1250 

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 ) 

1262 

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 

1279 

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 

1289 

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 

1298 

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 

1311 

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 = {} 

1320 

1321 variables[self.TOTAL_SUM] = lpSolver.LpVariable( 

1322 self.TOTAL_SUM, 0, None, lpSolver.LpContinuous) 

1323 

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 

1333 

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 

1343 

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 

1356 

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 

1369 

1370 variables[self.NEXT_DIFF] = lpSolver.LpVariable( 

1371 self.NEXT_DIFF, 0, None, lpSolver.LpContinuous) 

1372 

1373 variables[self.MAX_STAGE_TIME] = lpSolver.LpVariable( 

1374 self.MAX_STAGE_TIME, 0, None, lpSolver.LpContinuous) 

1375 

1376 variables[self.MAX_LAST_CHUNK] = lpSolver.LpVariable( 

1377 self.MAX_LAST_CHUNK, 0, None, lpSolver.LpContinuous) 

1378 

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 

1388 

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 

1407 

1408 return variables 

1409 

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 

1421 

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 

1430 

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 

1450 

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 

1463 

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

1475 

1476 return bound 

1477 

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 

1493 

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 

1503 

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 

1513 

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 ) 

1523 

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 ) 

1531 

1532 bound = lpSolver.LpAffineExpression() 

1533 

1534 head_time = 0 

1535 for head in layers_sorted[Layer.type_enum.HEAD]: 

1536 head_time = head.forward_time_ 

1537 

1538 prob += max_prev_stages[0] >= (self.micro_batch_time( 

1539 self.PROP_PHASE.FW, layers_sorted, v, 0)) - head_time 

1540 

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

1545 

1546 

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

1550 

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 

1554 

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 ) 

1565 

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 ) 

1573 

1574 bound = lpSolver.LpAffineExpression() 

1575 

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

1580 

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

1585 

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

1589 

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