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

350 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-21 04:29 +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 

17from contextlib import nullcontext 

18import time 

19import copy 

20import multiprocessing as proc 

21import json 

22import os 

23import logging 

24from typing import Optional 

25 

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

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

28 

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

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

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

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

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

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

35from hyper_parallel.auto_parallel.sapp_nd.nd.common.cost_model_preprocess import ( 

36 CostModelConfig, 

37 detect_attention_type, 

38) 

39 

40# logger = proc.log_to_stderr() 

41# logger.setLevel(proc.SUBDEBUG) 

42 

43 

44class ParallelizeLayer: 

45 """Parallelize one layer type""" 

46 

47 def __init__( 

48 self, 

49 evaluator, 

50 machine, 

51 global_batch_size=None, 

52 dimensions=None, 

53 **extra_config, 

54 ): 

55 

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

57 self.machine = machine 

58 if "mppb" in extra_config: 

59 manual_ppb = extra_config.pop("mppb") 

60 else: 

61 manual_ppb = False 

62 

63 self.mem_eval = evaluator 

64 

65 self.model_name = self.mem_eval._ccfg.model_name 

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

67 

68 if "mem_for_ppb" in extra_config: 

69 reserve_mem = extra_config.pop("mem_for_ppb") 

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

71 

72 if "max_mem" in extra_config: 

73 max_mem = extra_config.pop("max_mem") 

74 if max_mem is not None: 

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

76 

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

78 

79 if "sub_model" in extra_config: 

80 sub_model = extra_config.pop("sub_model") 

81 if sub_model is not None: 

82 self.config = GlobalConfig( 

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

84 dimensions, 

85 mppb=manual_ppb, 

86 ) 

87 else: 

88 self.config = GlobalConfig( 

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

90 ) 

91 else: 

92 self.config = GlobalConfig( 

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

94 ) 

95 

96 self.mem_eval.set_passes(**extra_config) 

97 

98 self.machine.update_num_if_none( 

99 self.config.ccfg.strategy_num_devices() 

100 ) 

101 

102 if global_batch_size: 

103 self.global_batch_size = global_batch_size 

104 else: 

105 self.global_batch_size = self.config.ccfg.gbs 

106 

107 self.bound_space() 

108 

109 def bound_space(self): 

110 """Set bounds for parallel dimensions""" 

111 vpp = ( 

112 1 

113 if Dim.VPP in self.config.dimensions 

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

115 ) 

116 pp_bound = min( 

117 self.machine.pipeline_bound(), 

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

119 self.global_batch_size, 

120 ) 

121 Dim.PP.set_bound(pp_bound) 

122 logger.info( 

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

124 pp_bound, 

125 self.machine.pipeline_bound(), 

126 self.config.total_layer_num(), 

127 vpp, 

128 self.global_batch_size, 

129 ) 

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

131 # if ( 

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

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

134 # ): 

135 # Dim.EP.set_bound(1) 

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

137 kv_heads = self.config.ccfg.n_kv 

138 if kv_heads: 

139 Dim.TP.set_bound(kv_heads) 

140 logger.warning( 

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

142 str(kv_heads), 

143 ) 

144 else: 

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

146 Dim.TP.set_bound( 

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

148 ) 

149 

150 def filtered_out(self, _): 

151 """Manual conditions to remove config patterns""" 

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

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

154 # return True 

155 return False 

156 

157 def is_valid(self, parallel_config): 

158 """Check configuration validity""" 

159 if not parallel_config.is_valid(): 

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

161 return False 

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

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

164 return False 

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

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

167 return False 

168 if self.filtered_out(parallel_config): 

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

170 return False 

171 

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

173 cp_degree = parallel_config.dims_val[Dim.CP] 

174 if cp_degree > 1: 

175 seq_len = self.config.ccfg.s 

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

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

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

179 total_devices = self.machine.number 

180 

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

182 

183 bw_intra = self.config.ccfg.bw_intra 

184 bw_inter = self.config.ccfg.bw_inter 

185 

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

187 

188 cp_result = validate_cp_constraints( 

189 seq_len=seq_len, 

190 cp_degree=cp_degree, 

191 tp_degree=tp_degree, 

192 pp_degree=pp_degree, 

193 device_per_node=device_per_node, 

194 attention_type_str=attention_type, 

195 bw_intra=bw_intra, 

196 bw_inter=bw_inter, 

197 total_devices=total_devices, 

198 sp_enabled=sp_enabled, 

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

200 attention_heads=self.config.ccfg.a, 

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

202 ) 

203 

204 if not cp_result.is_valid: 

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

206 return False 

207 

208 if cp_result.warning_message: 

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

210 

211 gbs = self.config.global_batch_size(parallel_config) 

212 if not gbs == self.global_batch_size: 

213 logger.error( 

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

215 gbs, 

216 self.global_batch_size, 

217 ) 

218 return False 

219 return True 

220 

221 def memory_estim(self, debugger=None): 

222 """Whether the config fits memory""" 

223 logger.debug("estimate_peak") 

224 verbose = logger.level < logging.INFO 

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

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

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

228 peak = self.mem_eval.estimate_peak( 

229 verbose=verbose 

230 ) # (logger.level>2)) 

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

232 if debugger and debugger.is_enabled(): 

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

234 return peak 

235 

236 def generate_search_space(self, folder, threads_num): 

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

238 space = ({}, 0) 

239 configs = [] 

240 results = {} 

241 if threads_num: 

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

243 logger.debug("before loops") 

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

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

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

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

248 logger.debug( 

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

250 ) 

251 peak_mem = result.get() 

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

385 proc.Pool(processes=threads_num) 

386 if multiproc 

387 else nullcontext() 

388 ) as pool: 

389 for config, mem in space: 

390 self.config.set_parallel_config(config) 

391 values = [] 

392 if multiproc: 

393 score = pool.apply_async( 

394 pool_estimate_performance, 

395 args=( 

396 copy.deepcopy(self.config.ccfg), 

397 self.machine.device, 

398 mem, 

399 cache_file, 

400 ), 

401 ) 

402 else: 

403 if self.enable_debug: 

404 debugger = Debug.Debug( 

405 config, 

406 info_type=Debug.PerfParts, 

407 enable=self.enable_debug, 

408 ) 

409 score = estimate_performance( 

410 self.config.ccfg, 

411 debugger=debugger, 

412 device_type=self.machine.device, 

413 memory=mem, 

414 cache_file=cache_file, 

415 ) 

416 debugger.write() 

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

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

419 del values[-2:] 

420 del debug_parts[-2:] 

421 else: 

422 score = estimate_performance( 

423 self.config.ccfg, 

424 device_type=self.machine.device, 

425 memory=mem, 

426 ) 

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

428 

429 if not multiproc: 

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

431 

432 if multiproc: 

433 new_scored_space = [] 

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

435 score_value = score.get() 

436 logger.info( 

437 "config %s has score %f", str(config), score_value 

438 ) 

439 new_scored_space.append( 

440 (config, mem, score_value, values) 

441 ) 

442 else: 

443 new_scored_space = scored_space 

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

445 

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

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

448 scored_space = [] 

449 debug_parts = [] 

450 for config, real_time, real_comm_wait in space: 

451 debugger = Debug.Debug( 

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

453 ) 

454 self.config.set_parallel_config(config) 

455 peak_mem = self.memory_estim() 

456 score = estimate_performance( 

457 self.config.ccfg, 

458 debugger=debugger, 

459 device_type=self.machine.device, 

460 stage_focused=0, 

461 ) # , memory = mem) 

462 debugger.write() 

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

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

465 del values[-2:] 

466 scored_space.append( 

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

468 ) 

469 

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

471 del debug_parts[-2:] 

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

473 

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

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

476 scored_space = [] 

477 debug_parts = [] 

478 for config, real_time in space: 

479 debugger = Debug.Debug( 

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

481 ) 

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

483 self.config.set_parallel_config(config) 

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

485 peak_mem = self.memory_estim() 

486 score = estimate_performance( 

487 self.config.ccfg, 

488 debugger=debugger, 

489 device_type=self.machine.device, 

490 ) # , memory = mem) 

491 debugger.write() 

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

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

494 del values[-2:] 

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

496 

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

498 del debug_parts[-2:] 

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

500 

501 def plot_title(self): 

502 """Generate plot title""" 

503 return ( 

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

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

506 ) 

507 

508 def run_generation_to_ordering( 

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

510 ): 

511 """Test some functions""" 

512 start = time.time() 

513 space = self.generate_search_space(yaml_folder, threads_num) 

514 generation = time.time() 

515 scored_space, dbg = self.order_search_space( 

516 space, threads_num, cache_file=cache_file 

517 ) 

518 ordering = time.time() 

519 logger.output( 

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

521 ) 

522 logger.output( 

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

524 generation - start, 

525 ordering - generation, 

526 ) 

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

528 logger.output( 

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

530 ) 

531 logger.output( 

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

533 self.machine.number, 

534 self.global_batch_size, 

535 str(self.config.dimensions), 

536 ) 

537 if self.enable_debug: 

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

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

540 if scored_space: 

541 Debug.plot_nd( 

542 scored_space, 

543 output_path, 

544 dbg, 

545 title=self.plot_title(), 

546 max_num=top_num, 

547 ) 

548 return scored_space 

549 

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

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

552 parallel_config = scored_space[k][0] 

553 self.config.set_parallel_config(parallel_config) 

554 self.mem_eval.update_config(self.config) 

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

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

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

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

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

560 filename = ( 

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

562 + "/../pipeline_balance/layers/" 

563 + m 

564 + ".json" 

565 ) 

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

567 json.dump( 

568 self.mem_eval.estimate_layer_memory( 

569 device_type=self.machine.device 

570 ), 

571 fp, 

572 indent=4, 

573 ) 

574 logger.output( 

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

576 "\npython run_pipeline_balance.py " 

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

578 parallel_config, 

579 m, 

580 s, 

581 mb, 

582 i, 

583 mem, 

584 ) 

585 logger.output("Warning: currently select_recompute_memory \ 

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

587 

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

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

590 configs, row_num = Debug.get_real_data(csv_f) 

591 configs_estimated, debug_parts = self.order_space_test( 

592 configs, order_by=2 

593 ) 

594 if output_path is not None: 

595 Debug.plot_vs_real( 

596 configs_estimated, 

597 csv_f, 

598 output_path, 

599 debug_parts, 

600 title=self.plot_title(), 

601 ) 

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

603 return correl, topk, row_num 

604 

605 def test_from_csv_comm_classified( 

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

607 ): 

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

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

610 configs_estimated, debug_parts = self.order_space_test_comm_classified( 

611 configs, order_by=2 

612 ) 

613 

614 if output_path is not None: 

615 Debug.plot_vs_real_comm_classified( 

616 configs_estimated, 

617 csv_f, 

618 output_path, 

619 debug_parts, 

620 title=self.plot_title(), 

621 plot_idle=plot_idle, 

622 ) 

623 

624 return Debug.correlation_with_classified_comms(configs_estimated) 

625 

626 

627class ParallelizeMultiModal(ParallelizeLayer): 

628 """Parallelize a MultiModel""" 

629 

630 def __init__( 

631 self, 

632 evaluator, 

633 machine, 

634 global_batch_size=None, 

635 dimensions=None, 

636 **extra_config, 

637 ): 

638 

639 super().__init__( 

640 evaluator, 

641 machine, 

642 global_batch_size=global_batch_size, 

643 dimensions=dimensions, 

644 sub_model="deepseekv3", 

645 **extra_config, 

646 ) 

647 

648 

649class Parallelize: # pylint: disable=R0903 

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

651 

652 def __init__( 

653 self, 

654 framework, 

655 config, 

656 machine, 

657 **extra_config, 

658 ): 

659 logger.debug("before evaluator init") 

660 if "model" in extra_config: 

661 model_name = extra_config.pop("model") 

662 mem_eval = EvaluatorV2( 

663 config, framework=framework, hook_cls=model_name, machine=machine 

664 ) 

665 else: 

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

667 

668 if "global_batch_size" in extra_config: 

669 global_batch_size = extra_config.pop("global_batch_size") 

670 else: 

671 global_batch_size = None 

672 

673 if "dimensions" in extra_config: 

674 dimensions = extra_config.pop("dimensions") 

675 else: 

676 dimensions = None 

677 

678 if mem_eval.ccfg.multimodal: 

679 logger.debug("MultiModal is triggered") 

680 self.instance = ParallelizeMultiModal( 

681 mem_eval, 

682 machine, 

683 global_batch_size=global_batch_size, 

684 dimensions=dimensions, 

685 **extra_config, 

686 ) 

687 else: 

688 self.instance = ParallelizeLayer( 

689 mem_eval, 

690 machine, 

691 global_batch_size=global_batch_size, 

692 dimensions=dimensions, 

693 sub_model=None, 

694 **extra_config, 

695 ) 

696 

697 def __getattr__(self, name): 

698 return self.instance.__getattribute__(name) 

699 

700 

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

702 """Space printer""" 

703 i = 0 

704 s = "" 

705 if max_num is not None: 

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

707 else: 

708 s += "\n" 

709 if len(space) == 0: 

710 return s 

711 s += "\t" 

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

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

714 s += "Memory Performance score " 

715 if debug_parts is not None: 

716 for dbg_part in debug_parts: 

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

718 s += "\n" 

719 for config in space: 

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

721 break 

722 s += "\t" 

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

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

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

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

727 for v in config[3]: 

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

729 s += "\n" 

730 i += 1 

731 return s 

732 

733 

734def pool_estimate_memory(config: CostModelConfig) -> float: 

735 """Calls memory estimation for multiprocessing""" 

736 logger.debug("estimate_peak") 

737 # print("estimate_peak") 

738 e = EvaluatorV2(None, ccfg=config) 

739 return e.estimate_peak() 

740 

741 

742# def pool_estimate_memory(evaluator): 

743# """Calls memory estimation for multiprocessing""" 

744# logger.debug("estimate_peak") 

745# return evaluator.estimate_peak() 

746 

747 

748def pool_estimate_performance( 

749 config: CostModelConfig, 

750 device: Hard.Type, 

751 memory: Optional[float] = None, 

752 cache_file: Optional[str] = None, 

753) -> float: 

754 """Calls performance estimation for multiprocessing""" 

755 return estimate_performance( 

756 config, 

757 device_type=device, 

758 memory=memory, 

759 cache_file=cache_file, 

760 )