Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / nd / parallelize.py: 78%

348 statements  

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

1# Copyright 2024-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"""find parallelization""" 

16 

17import time 

18import copy 

19import multiprocessing as proc 

20import json 

21import os 

22import logging 

23 

24from hyper_parallel.auto_parallel.sapp_nd.memory_estimation.estimate_v2 import EvaluatorV2 

25from hyper_parallel.auto_parallel.sapp_nd.perf_estimation.estimate import estimate_performance 

26 

27from hyper_parallel.auto_parallel.sapp_nd.nd.global_config import GlobalConfig 

28from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger 

29import hyper_parallel.auto_parallel.sapp_nd.nd.dimensions as Dim 

30import hyper_parallel.auto_parallel.sapp_nd.nd.common.hardware as Hard 

31import hyper_parallel.auto_parallel.sapp_nd.nd.debug as Debug 

32from hyper_parallel.auto_parallel.sapp_nd.nd.dimensions import validate_cp_constraints 

33from hyper_parallel.auto_parallel.sapp_nd.nd.common.cost_model_preprocess import detect_attention_type 

34 

35# logger = proc.log_to_stderr() 

36# logger.setLevel(proc.SUBDEBUG) 

37 

38 

39class ParallelizeLayer: 

40 """Parallelize one layer type""" 

41 

42 def __init__( 

43 self, 

44 evaluator, 

45 machine, 

46 global_batch_size=None, 

47 dimensions=None, 

48 **extra_config, 

49 ): 

50 

51 self.enable_debug = logger.level < logging.CRITICAL 

52 self.machine = machine 

53 if "mppb" in extra_config: 

54 manual_ppb = extra_config.pop("mppb") 

55 else: 

56 manual_ppb = False 

57 

58 self.mem_eval = evaluator 

59 

60 self.model_name = self.mem_eval._ccfg.model_name 

61 logger.debug("model is %s", self.model_name) 

62 

63 if "mem_for_ppb" in extra_config: 

64 reserve_mem = extra_config.pop("mem_for_ppb") 

65 self.mem_eval._ccfg.device_capacity.decrease(reserve_mem) 

66 

67 if "max_mem" in extra_config: 

68 max_mem = extra_config.pop("max_mem") 

69 if max_mem is not None: 

70 self.mem_eval._ccfg.device_capacity.set(max_mem) 

71 

72 logger.debug("before global config init") 

73 

74 if "sub_model" in extra_config: 

75 sub_model = extra_config.pop("sub_model") 

76 if sub_model is not None: 

77 self.config = GlobalConfig( 

78 self.mem_eval._ccfg.mm_ccfgs[sub_model], 

79 dimensions, 

80 mppb=manual_ppb, 

81 ) 

82 else: 

83 self.config = GlobalConfig( 

84 self.mem_eval._ccfg, dimensions, mppb=manual_ppb 

85 ) 

86 else: 

87 self.config = GlobalConfig( 

88 self.mem_eval._ccfg, dimensions, mppb=manual_ppb 

89 ) 

90 

91 self.mem_eval.set_passes(**extra_config) 

92 

93 self.machine.update_num_if_none( 

94 self.config.ccfg.strategy_num_devices() 

95 ) 

96 

97 if global_batch_size: 

98 self.global_batch_size = global_batch_size 

99 else: 

100 self.global_batch_size = self.config.ccfg.gbs 

101 

102 self.bound_space() 

103 

104 def bound_space(self): 

105 """Set bounds for parallel dimensions""" 

106 vpp = ( 

107 1 

108 if Dim.VPP in self.config.dimensions 

109 else Dim.VPP.from_config(self.config.ccfg) 

110 ) 

111 pp_bound = min( 

112 self.machine.pipeline_bound(), 

113 self.config.total_layer_num() // vpp, 

114 self.global_batch_size, 

115 ) 

116 Dim.PP.set_bound(pp_bound) 

117 logger.info( 

118 "PP bound is %d, machine bound = %d, L = %d, VPP = %d, B = %d", 

119 pp_bound, 

120 self.machine.pipeline_bound(), 

121 self.config.total_layer_num(), 

122 vpp, 

123 self.global_batch_size, 

124 ) 

125 Dim.EP.set_bound(self.config.ccfg.n_exp) 

126 # if ( 

127 # self.config.dimensions.count(Dim.EP) > 0 

128 # and Dim.EP.from_config(self.config.ccfg) <= 1 

129 # ): 

130 # Dim.EP.set_bound(1) 

131 # self.config.dimensions.remove(Dim.EP) 

132 kv_heads = self.config.ccfg.n_kv 

133 if kv_heads: 

134 Dim.TP.set_bound(kv_heads) 

135 logger.warning( 

136 "Because of n_kv_heads, MP will be limited to %s", 

137 str(kv_heads), 

138 ) 

139 else: 

140 # num_head % (TP * UP) == 0. Add UP later 

141 Dim.TP.set_bound( 

142 Hard.highest_power_of_2_divisor(self.config.ccfg.a) 

143 ) 

144 

145 def filtered_out(self, _): 

146 """Manual conditions to remove config patterns""" 

147 # if parallel_config.has_dim(Dim.EP): 

148 # if self.config.dim_val(Dim.EP, parallel_config) < 8: 

149 # return True 

150 return False 

151 

152 def is_valid(self, parallel_config): 

153 """Check configuration validity""" 

154 if not parallel_config.is_valid(): 

155 logger.warning("configuration %s not valid", str(parallel_config)) 

156 return False 

157 if not self.config.moe_valid(parallel_config): 

158 logger.warning("expert parallel is higher than expert number") 

159 return False 

160 if hasattr(self.config, 'ep_constraints_valid') and not self.config.ep_constraints_valid(parallel_config): 

161 logger.warning("EP divisibility constraints not satisfied") 

162 return False 

163 if self.filtered_out(parallel_config): 

164 logger.warning("Config manually filtered out") 

165 return False 

166 

167 if hasattr(parallel_config, 'dims_val') and Dim.CP in parallel_config.dims_val: 

168 cp_degree = parallel_config.dims_val[Dim.CP] 

169 if cp_degree > 1: 

170 seq_len = self.config.ccfg.s 

171 tp_degree = parallel_config.dims_val.get(Dim.TP, 1) 

172 pp_degree = parallel_config.dims_val.get(Dim.PP, 1) 

173 device_per_node = self.machine.device.intra_node_num() 

174 total_devices = self.machine.number 

175 

176 attention_type = detect_attention_type(self.config.ccfg).name.lower() 

177 

178 bw_intra = self.config.ccfg.bw_intra 

179 bw_inter = self.config.ccfg.bw_inter 

180 

181 sp_enabled = bool(parallel_config.dims_val.get(Dim.SP, False)) 

182 

183 cp_result = validate_cp_constraints( 

184 seq_len=seq_len, 

185 cp_degree=cp_degree, 

186 tp_degree=tp_degree, 

187 pp_degree=pp_degree, 

188 device_per_node=device_per_node, 

189 attention_type_str=attention_type, 

190 bw_intra=bw_intra, 

191 bw_inter=bw_inter, 

192 total_devices=total_devices, 

193 sp_enabled=sp_enabled, 

194 cp_algo=getattr(self.config.ccfg, 'cp_algo', 'colossalai_cp'), 

195 attention_heads=self.config.ccfg.a, 

196 num_kv_heads=getattr(self.config.ccfg, 'n_kv', 0), 

197 ) 

198 

199 if not cp_result.is_valid: 

200 logger.warning("CP constraints violated: %s", cp_result.error_message) 

201 return False 

202 

203 if cp_result.warning_message: 

204 logger.info("CP warning: %s", cp_result.warning_message) 

205 

206 gbs = self.config.global_batch_size(parallel_config) 

207 if not gbs == self.global_batch_size: 

208 logger.error( 

209 "wrong global batch size: ccfg is %d, instead of %d", 

210 gbs, 

211 self.global_batch_size, 

212 ) 

213 return False 

214 return True 

215 

216 def memory_estim(self, debugger=None): 

217 """Whether the config fits memory""" 

218 logger.debug("estimate_peak") 

219 verbose = logger.level < logging.INFO 

220 self.mem_eval.set_config(self.config.ccfg) # = self.config.ccfg 

221 # self.mem_eval = EvaluatorV2(self.config) 

222 logger.debug("ccfg = %s", str(self.config.ccfg)) 

223 peak = self.mem_eval.estimate_peak( 

224 verbose=verbose 

225 ) # (logger.level>2)) 

226 logger.debug("peak memory = %d", peak) 

227 if debugger and debugger.is_enabled(): 

228 debugger.info[Debug.MemParts.TOTAL] = peak 

229 return peak 

230 

231 def generate_search_space(self, folder, threads_num): 

232 """Return a search space computed with memory estimation""" 

233 space = ({}, 0) 

234 configs = [] 

235 results = {} 

236 if threads_num: 

237 with proc.Pool(processes=threads_num) as pool: 

238 logger.debug("before loops") 

239 results, size = self.device_loops(space, pool) 

240 logger.debug("%d results", len(results)) 

241 for config, result in results.items(): 

242 logger.debug("result = %s", str(result)) 

243 logger.debug( 

244 "before get: is ready ? %s", str(result.ready()) 

245 ) 

246 logger.debug( 

247 "before get: is successful ? %s", 

248 str(result.successful()), 

249 ) 

250 # if result.successful(): 

251 peak_mem = result.get(1) 

252 logger.debug( 

253 "after get: is ready ? %s", str(result.ready()) 

254 ) 

255 logger.debug( 

256 "after get: is successful ? %s", 

257 str(result.successful()), 

258 ) 

259 logger.debug("peak_mem = %s", str(peak_mem)) 

260 if self.mem_eval.mem_fit(peak_mem): 

261 configs.append((config, peak_mem)) 

262 pool.close() 

263 pool.join() 

264 else: 

265 results, size = self.device_loops(space, None) 

266 for config, peak_mem in results.items(): 

267 if self.mem_eval.mem_fit(peak_mem): 

268 configs.append((config, peak_mem)) 

269 if folder: 

270 self.config.write(folder, config) 

271 logger.output("%d valid configurations generated", size) 

272 logger.output("%d configuration fitting memory to order", len(configs)) 

273 

274 return configs 

275 

276 def device_loops(self, space, pool): 

277 """Exploration loop nest level 0: parallel dimensions dividing devices""" 

278 for tp in self.config.space(Dim.TP, self.machine.number): 

279 for pp in self.config.space(Dim.PP, self.machine.number // tp): 

280 for cp in self.config.space( 

281 Dim.CP, self.machine.number // tp // pp 

282 ): 

283 logger.debug( 

284 "dp = %d / %d / %d / %d", 

285 self.machine.number, 

286 tp, 

287 cp, 

288 pp, 

289 ) 

290 dp = self.machine.number // tp // cp // pp 

291 if dp < 1: 

292 break 

293 space = self.batch_loops(space, pool, (dp, tp, pp, cp)) 

294 return space 

295 

296 def batch_loops(self, space, pool, dtpc_p): 

297 """Exploration loop nest level 1: dimensions dividing batch (except already processed DP)""" 

298 dp, _, pp, _ = dtpc_p 

299 # if pp > 1: 

300 for mbs in self.config.space( 

301 Dim.MBS, self.global_batch_size // pp // dp 

302 ): 

303 logger.debug("mbn= %d / %d / %d", self.global_batch_size, dp, mbs) 

304 mbn = self.global_batch_size // dp // mbs 

305 space = self.parallel_loops(space, pool, (dtpc_p, (mbs, mbn))) 

306 # else: 

307 # logger.debug("no pipeline so mbn = 1") 

308 # mbs = self.global_batch_size // dp 

309 # space = self.parallel_loops(space, pool, (dtpc_p, (mbs, 1))) 

310 return space 

311 

312 def parallel_loops(self, space, pool, dims): 

313 """Exploration loop nest level 2: dimensions dependent on others""" 

314 dtpc_p, mbsn = dims 

315 dp, tp, pp, _ = dtpc_p 

316 for ep in self.config.space(Dim.EP, dp * tp): 

317 for vpp in self.config.range_space( 

318 Dim.VPP, min(4, pp, self.config.total_layer_num() // pp) 

319 ): 

320 for op in self.config.space( 

321 Dim.OP, self.config.max_op(dp, tp, ep) 

322 ): 

323 for sp in self.config.bool_space(Dim.SP): 

324 space = self.inside_loop_nest( 

325 space, 

326 pool, 

327 (dtpc_p, mbsn, (ep, vpp, op, sp)), 

328 ) 

329 return space 

330 

331 def inside_loop_nest(self, space, pool, dims): 

332 """Exploration loop nest statements""" 

333 dtpc_p, mbsn, evos_p = dims 

334 configs, size = space 

335 parallel_config = self.config.make_parallel_config( 

336 dtpc_p, mbsn, evos_p 

337 ) 

338 logger.info("test config %d : %s", size, str(parallel_config)) 

339 size += 1 

340 

341 if self.is_valid(parallel_config) and self.config.set_parallel_config( 

342 parallel_config 

343 ): 

344 if pool is None: 

345 if self.enable_debug: 

346 mem_debugger = Debug.Debug( 

347 parallel_config, 

348 info_type=Debug.MemParts, 

349 enable=self.enable_debug, 

350 output_file="debug_mem.csv", 

351 ) 

352 # try: 

353 peak = self.memory_estim(mem_debugger) 

354 mem_debugger.write() 

355 else: 

356 peak = self.memory_estim() 

357 # except: 

358 # logger.error() 

359 # return (configs, size) 

360 else: 

361 # logger.debug("before evaluator copy") 

362 # evaluator = copy.deepcopy(self.mem_eval) 

363 logger.debug("before apply_async") 

364 peak = pool.apply_async( 

365 pool_estimate_memory, 

366 args=(copy.deepcopy(self.config.ccfg),), 

367 # args=(evaluator,), 

368 # self.memory_estim, 

369 ) 

370 logger.debug("after apply_async") 

371 configs[parallel_config] = peak 

372 

373 return (configs, size) 

374 

375 def order_search_space(self, space, threads_num, cache_file): 

376 """Sort the search space computed with performance estimation""" 

377 if not space: 

378 return ([], []) 

379 multiproc = False 

380 if threads_num and threads_num > 5 * len(space): 

381 multiproc = True 

382 scored_space = [] 

383 debug_parts = [] 

384 for config, mem in space: 

385 self.config.set_parallel_config(config) 

386 values = [] 

387 if multiproc: 

388 with proc.Pool(processes=threads_num) as pool: 

389 score = pool.apply_async( 

390 pool_estimate_performance, 

391 args=( 

392 copy.deepcopy(self.config), 

393 self.machine.device, 

394 cache_file, 

395 ), 

396 ) 

397 else: 

398 if self.enable_debug: 

399 debugger = Debug.Debug( 

400 config, 

401 info_type=Debug.PerfParts, 

402 enable=self.enable_debug, 

403 ) 

404 score = estimate_performance( 

405 self.config.ccfg, 

406 debugger=debugger, 

407 device_type=self.machine.device, 

408 memory=mem, 

409 cache_file=cache_file, 

410 ) 

411 debugger.write() 

412 debug_parts = list(debugger.info.keys()) 

413 values = list(debugger.info.values()) 

414 del values[-2:] 

415 del debug_parts[-2:] 

416 else: 

417 score = estimate_performance( 

418 self.config.ccfg, 

419 device_type=self.machine.device, 

420 memory=mem, 

421 ) 

422 scored_space.append((config, mem, score, values)) 

423 

424 logger.info("config %s has score %f", str(config), score) 

425 

426 if multiproc: 

427 new_scored_space = [] 

428 pool.close() 

429 pool.join() 

430 for config, mem, score, values in scored_space: 

431 new_scored_space.append((config, mem, score.get(), values)) 

432 else: 

433 new_scored_space = scored_space 

434 return (sorted(new_scored_space, key=lambda x: x[2]), debug_parts) 

435 

436 def order_space_test_comm_classified(self, space, order_by=2): 

437 """Order the given space with performance estimation""" 

438 scored_space = [] 

439 debug_parts = [] 

440 for config, real_time, real_comm_wait in space: 

441 debugger = Debug.Debug( 

442 config, info_type=Debug.PerfParts, enable=self.enable_debug 

443 ) 

444 self.config.set_parallel_config(config) 

445 peak_mem = self.memory_estim() 

446 score = estimate_performance( 

447 self.config.ccfg, 

448 debugger=debugger, 

449 device_type=self.machine.device, 

450 stage_focused=0, 

451 ) # , memory = mem) 

452 debugger.write() 

453 debug_parts = list(debugger.info.keys()) 

454 values = list(debugger.info.values()) 

455 del values[-2:] 

456 scored_space.append( 

457 (config, peak_mem, real_time, score, values, real_comm_wait) 

458 ) 

459 

460 logger.info("config %s has score %f", str(config), score) 

461 del debug_parts[-2:] 

462 return (sorted(scored_space, key=lambda x: x[order_by]), debug_parts) 

463 

464 def order_space_test(self, space, order_by=2): 

465 """Order the given space with performance estimation""" 

466 scored_space = [] 

467 debug_parts = [] 

468 for config, real_time in space: 

469 debugger = Debug.Debug( 

470 config, info_type=Debug.PerfParts, enable=self.enable_debug 

471 ) 

472 logger.info("Test config %s", str(config)) 

473 self.config.set_parallel_config(config) 

474 logger.debug(self.mem_eval.get_strategy()) 

475 peak_mem = self.memory_estim() 

476 score = estimate_performance( 

477 self.config.ccfg, 

478 debugger=debugger, 

479 device_type=self.machine.device, 

480 ) # , memory = mem) 

481 debugger.write() 

482 debug_parts = list(debugger.info.keys()) 

483 values = list(debugger.info.values()) 

484 del values[-2:] 

485 scored_space.append((config, peak_mem, real_time, score, values)) 

486 

487 logger.info("config %s has score %f", str(config), score) 

488 del debug_parts[-2:] 

489 return (sorted(scored_space, key=lambda x: x[order_by]), debug_parts) 

490 

491 def plot_title(self): 

492 """Generate plot title""" 

493 return ( 

494 f"{self.model_name} on {self.machine.number}" 

495 + f" {self.machine.device} with {self.global_batch_size} GBS" 

496 ) 

497 

498 def run_generation_to_ordering( 

499 self, yaml_folder, threads_num=None, top_num=None, cache_file=None 

500 ): 

501 """Test some functions""" 

502 start = time.time() 

503 space = self.generate_search_space(yaml_folder, threads_num) 

504 generation = time.time() 

505 scored_space, dbg = self.order_search_space( 

506 space, threads_num, cache_file=cache_file 

507 ) 

508 ordering = time.time() 

509 logger.output( 

510 space_to_string(scored_space, max_num=top_num, debug_parts=dbg) 

511 ) 

512 logger.output( 

513 "Space generation took %.2fs and ordering took %.2fs", 

514 generation - start, 

515 ordering - generation, 

516 ) 

517 is_not = " NOT" if not self.config.balancing.from_config else "" 

518 logger.output( 

519 "Offset & Recompute were%s computed from config info", is_not 

520 ) 

521 logger.output( 

522 "Device number is %d, global batch size is %d, dimensions are %s", 

523 self.machine.number, 

524 self.global_batch_size, 

525 str(self.config.dimensions), 

526 ) 

527 if self.enable_debug: 

528 file_path = os.path.dirname(os.path.realpath(__file__)) 

529 output_path = os.path.join(file_path, "output") 

530 if scored_space: 

531 Debug.plot_nd( 

532 scored_space, 

533 output_path, 

534 dbg, 

535 title=self.plot_title(), 

536 max_num=top_num, 

537 ) 

538 return scored_space 

539 

540 def to_ppb(self, scored_space, k, cfg_name): 

541 """Create an input file for pipeline balancing""" 

542 parallel_config = scored_space[k][0] 

543 self.config.set_parallel_config(parallel_config) 

544 self.mem_eval.update_config(self.config) 

545 m = cfg_name + "_nd_to_ppb_" + str(k) 

546 s = self.config.dim_val(Dim.PP, parallel_config) 

547 mb = self.config.dim_val(Dim.MBN, parallel_config) 

548 i = self.config.dim_val(Dim.VPP, parallel_config) 

549 mem = str(self.config.ccfg.device_capacity.to_mb) 

550 filename = ( 

551 os.path.dirname(os.path.realpath(__file__)) 

552 + "/../pipeline_balance/layers/" 

553 + m 

554 + ".json" 

555 ) 

556 with open(filename, "w+", encoding="utf-8") as fp: 

557 json.dump( 

558 self.mem_eval.estimate_layer_memory( 

559 device_type=self.machine.device 

560 ), 

561 fp, 

562 indent=4, 

563 ) 

564 logger.output( 

565 "To run pipeline balancing on configuration %s:" 

566 "\npython run_pipeline_balance.py " 

567 "-m %d -s %d -mb %d -i %d -mem %d", 

568 parallel_config, 

569 m, 

570 s, 

571 mb, 

572 i, 

573 mem, 

574 ) 

575 logger.output("Warning: currently select_recompute_memory \ 

576 should be removed & layer time need to be added") 

577 

578 def test_from_csv(self, csv_f, output_path=None): 

579 """Run estimation tests against a real run profiling in csv format""" 

580 configs, row_num = Debug.get_real_data(csv_f) 

581 configs_estimated, debug_parts = self.order_space_test( 

582 configs, order_by=2 

583 ) 

584 if output_path is not None: 

585 Debug.plot_vs_real( 

586 configs_estimated, 

587 csv_f, 

588 output_path, 

589 debug_parts, 

590 title=self.plot_title(), 

591 ) 

592 correl, topk = Debug.correlation_topk(configs_estimated, csv_f) 

593 return correl, topk, row_num 

594 

595 def test_from_csv_comm_classified( 

596 self, csv_f, output_path=None, plot_idle=False 

597 ): 

598 """Run test to compare estimation with detailed profiling""" 

599 configs = Debug.get_comm_classified_data(csv_f, plot_idle=plot_idle) 

600 configs_estimated, debug_parts = self.order_space_test_comm_classified( 

601 configs, order_by=2 

602 ) 

603 

604 if output_path is not None: 

605 Debug.plot_vs_real_comm_classified( 

606 configs_estimated, 

607 csv_f, 

608 output_path, 

609 debug_parts, 

610 title=self.plot_title(), 

611 plot_idle=plot_idle, 

612 ) 

613 

614 return Debug.correlation_with_classified_comms(configs_estimated) 

615 

616 

617class ParallelizeMultiModal(ParallelizeLayer): 

618 """Parallelize a MultiModel""" 

619 

620 def __init__( 

621 self, 

622 evaluator, 

623 machine, 

624 global_batch_size=None, 

625 dimensions=None, 

626 **extra_config, 

627 ): 

628 

629 super().__init__( 

630 evaluator, 

631 machine, 

632 global_batch_size=global_batch_size, 

633 dimensions=dimensions, 

634 sub_model="deepseekv3", 

635 **extra_config, 

636 ) 

637 

638 

639class Parallelize: # pylint: disable=R0903 

640 """Main class instantiated by one of the above two""" 

641 

642 def __init__( 

643 self, 

644 framework, 

645 config, 

646 machine, 

647 **extra_config, 

648 ): 

649 logger.debug("before evaluator init") 

650 if "model" in extra_config: 

651 model_name = extra_config.pop("model") 

652 mem_eval = EvaluatorV2( 

653 config, framework=framework, hook_cls=model_name, machine=machine 

654 ) 

655 else: 

656 mem_eval = EvaluatorV2(config, framework=framework, machine=machine) 

657 

658 if "global_batch_size" in extra_config: 

659 global_batch_size = extra_config.pop("global_batch_size") 

660 else: 

661 global_batch_size = None 

662 

663 if "dimensions" in extra_config: 

664 dimensions = extra_config.pop("dimensions") 

665 else: 

666 dimensions = None 

667 

668 if mem_eval.ccfg.multimodal: 

669 logger.debug("MultiModal is triggered") 

670 self.instance = ParallelizeMultiModal( 

671 mem_eval, 

672 machine, 

673 global_batch_size=global_batch_size, 

674 dimensions=dimensions, 

675 **extra_config, 

676 ) 

677 else: 

678 self.instance = ParallelizeLayer( 

679 mem_eval, 

680 machine, 

681 global_batch_size=global_batch_size, 

682 dimensions=dimensions, 

683 sub_model=None, 

684 **extra_config, 

685 ) 

686 

687 def __getattr__(self, name): 

688 return self.instance.__getattribute__(name) 

689 

690 

691def space_to_string(space, max_num=None, debug_parts=None): 

692 """Space printer""" 

693 i = 0 

694 s = "" 

695 if max_num is not None: 

696 s += "Top " + str(max_num) + " configurations:\n" 

697 else: 

698 s += "\n" 

699 if len(space) == 0: 

700 return s 

701 s += "\t" 

702 for d in space[0][0].all_dims: 

703 s += str(d) + " " * (6 - len(str(d))) 

704 s += "Memory Performance score " 

705 if debug_parts is not None: 

706 for dbg_part in debug_parts: 

707 s += "\t" + dbg_part.short_name() 

708 s += "\n" 

709 for config in space: 

710 if max_num is not None and max_num == i: 

711 break 

712 s += "\t" 

713 for v in config[0].values(): 

714 s += v + " " * (6 - len(v)) 

715 s += str(config[1]) + " MB " # + str(config[2]) 

716 s += f"{(config[2]):16.12e}" 

717 for v in config[3]: 

718 s += f"\t{(100*v/config[2]):.2f}%" 

719 s += "\n" 

720 i += 1 

721 return s 

722 

723 

724def pool_estimate_memory(config): 

725 """Calls memory estimation for multiprocessing""" 

726 logger.debug("estimate_peak") 

727 # print("estimate_peak") 

728 e = EvaluatorV2(config) 

729 return e.estimate_peak() 

730 

731 

732# def pool_estimate_memory(evaluator): 

733# """Calls memory estimation for multiprocessing""" 

734# logger.debug("estimate_peak") 

735# return evaluator.estimate_peak() 

736 

737 

738def pool_estimate_performance(config, device): 

739 """Calls performance estimation for multiprocessing""" 

740 return estimate_performance(config, device_type=device)