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
« 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"""
17import time
18import copy
19import multiprocessing as proc
20import json
21import os
22import logging
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
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
35# logger = proc.log_to_stderr()
36# logger.setLevel(proc.SUBDEBUG)
39class ParallelizeLayer:
40 """Parallelize one layer type"""
42 def __init__(
43 self,
44 evaluator,
45 machine,
46 global_batch_size=None,
47 dimensions=None,
48 **extra_config,
49 ):
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
58 self.mem_eval = evaluator
60 self.model_name = self.mem_eval._ccfg.model_name
61 logger.debug("model is %s", self.model_name)
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)
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)
72 logger.debug("before global config init")
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 )
91 self.mem_eval.set_passes(**extra_config)
93 self.machine.update_num_if_none(
94 self.config.ccfg.strategy_num_devices()
95 )
97 if global_batch_size:
98 self.global_batch_size = global_batch_size
99 else:
100 self.global_batch_size = self.config.ccfg.gbs
102 self.bound_space()
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 )
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
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
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
176 attention_type = detect_attention_type(self.config.ccfg).name.lower()
178 bw_intra = self.config.ccfg.bw_intra
179 bw_inter = self.config.ccfg.bw_inter
181 sp_enabled = bool(parallel_config.dims_val.get(Dim.SP, False))
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 )
199 if not cp_result.is_valid:
200 logger.warning("CP constraints violated: %s", cp_result.error_message)
201 return False
203 if cp_result.warning_message:
204 logger.info("CP warning: %s", cp_result.warning_message)
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
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
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))
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 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))
424 logger.info("config %s has score %f", str(config), score)
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)
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 )
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)
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))
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)
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 )
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
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")
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
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 )
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 )
614 return Debug.correlation_with_classified_comms(configs_estimated)
617class ParallelizeMultiModal(ParallelizeLayer):
618 """Parallelize a MultiModel"""
620 def __init__(
621 self,
622 evaluator,
623 machine,
624 global_batch_size=None,
625 dimensions=None,
626 **extra_config,
627 ):
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 )
639class Parallelize: # pylint: disable=R0903
640 """Main class instantiated by one of the above two"""
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)
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
663 if "dimensions" in extra_config:
664 dimensions = extra_config.pop("dimensions")
665 else:
666 dimensions = None
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 )
687 def __getattr__(self, name):
688 return self.instance.__getattribute__(name)
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
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()
732# def pool_estimate_memory(evaluator):
733# """Calls memory estimation for multiprocessing"""
734# logger.debug("estimate_peak")
735# return evaluator.estimate_peak()
738def pool_estimate_performance(config, device):
739 """Calls performance estimation for multiprocessing"""
740 return estimate_performance(config, device_type=device)