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
« 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"""
17from contextlib import nullcontext
18import time
19import copy
20import multiprocessing as proc
21import json
22import os
23import logging
24from typing import Optional
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
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)
40# logger = proc.log_to_stderr()
41# logger.setLevel(proc.SUBDEBUG)
44class ParallelizeLayer:
45 """Parallelize one layer type"""
47 def __init__(
48 self,
49 evaluator,
50 machine,
51 global_batch_size=None,
52 dimensions=None,
53 **extra_config,
54 ):
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
63 self.mem_eval = evaluator
65 self.model_name = self.mem_eval._ccfg.model_name
66 logger.debug("model is %s", self.model_name)
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)
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)
77 logger.debug("before global config init")
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 )
96 self.mem_eval.set_passes(**extra_config)
98 self.machine.update_num_if_none(
99 self.config.ccfg.strategy_num_devices()
100 )
102 if global_batch_size:
103 self.global_batch_size = global_batch_size
104 else:
105 self.global_batch_size = self.config.ccfg.gbs
107 self.bound_space()
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 )
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
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
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
181 attention_type = detect_attention_type(self.config.ccfg).name.lower()
183 bw_intra = self.config.ccfg.bw_intra
184 bw_inter = self.config.ccfg.bw_inter
186 sp_enabled = bool(parallel_config.dims_val.get(Dim.SP, False))
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 )
204 if not cp_result.is_valid:
205 logger.warning("CP constraints violated: %s", cp_result.error_message)
206 return False
208 if cp_result.warning_message:
209 logger.info("CP warning: %s", cp_result.warning_message)
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
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
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))
274 return configs
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
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
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
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
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
373 return (configs, size)
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))
429 if not multiproc:
430 logger.info("config %s has score %f", str(config), score)
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)
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 )
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)
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))
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)
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 )
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
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")
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
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 )
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 )
624 return Debug.correlation_with_classified_comms(configs_estimated)
627class ParallelizeMultiModal(ParallelizeLayer):
628 """Parallelize a MultiModel"""
630 def __init__(
631 self,
632 evaluator,
633 machine,
634 global_batch_size=None,
635 dimensions=None,
636 **extra_config,
637 ):
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 )
649class Parallelize: # pylint: disable=R0903
650 """Main class instantiated by one of the above two"""
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)
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
673 if "dimensions" in extra_config:
674 dimensions = extra_config.pop("dimensions")
675 else:
676 dimensions = None
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 )
697 def __getattr__(self, name):
698 return self.instance.__getattribute__(name)
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
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()
742# def pool_estimate_memory(evaluator):
743# """Calls memory estimation for multiprocessing"""
744# logger.debug("estimate_peak")
745# return evaluator.estimate_peak()
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 )