From 4803cd91b37dd8fbfafa6f3ce7dcebfc576cb368 Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Fri, 25 Sep 2026 04:55:33 -0400 Subject: [PATCH 1/5] [frontend] Renames.einsums can be indexed by name; [basetypes] NameIndexableList may be used for classes that needs indexing by name but not Evalable --- accelforge/frontend/renames.py | 7 +- accelforge/util/_basetypes.py | 120 ++++++++++-------- .../test_renames.py | 11 +- 3 files changed, 74 insertions(+), 64 deletions(-) diff --git a/accelforge/frontend/renames.py b/accelforge/frontend/renames.py index 39fe20a2..0324bf10 100755 --- a/accelforge/frontend/renames.py +++ b/accelforge/frontend/renames.py @@ -4,6 +4,7 @@ EvalableList, EvalableModel, EvalsTo, + NameIndexableList, TryEvalTo, _PostCall, ) @@ -123,7 +124,7 @@ def __init__(self, *args, **kwargs) -> None: class Renames(EvalableModel): - einsums: list[EinsumRename] = list() + einsums: NameIndexableList[EinsumRename] = NameIndexableList() """ Renames for a workload. The Einsum list is a list of EinsumRename objects, and renames will be applied to Einsums whose names match the EinsumRename.name. If an @@ -151,7 +152,7 @@ def get_renames_for_einsum(self, einsum_name: EinsumName) -> EinsumRename: def _for_einsum(self, einsum_name: EinsumName) -> "Renames": """Return a copy of the renames with only the Einsum with the given name.""" new = self.model_copy(deep=False) - new.einsums = [ + new.einsums = NameIndexableList( e for e in new.einsums if e.name == einsum_name or e.name == "default" - ] + ) return new diff --git a/accelforge/util/_basetypes.py b/accelforge/util/_basetypes.py index 4728fac9..ed38ae7f 100755 --- a/accelforge/util/_basetypes.py +++ b/accelforge/util/_basetypes.py @@ -900,67 +900,19 @@ def get_validator(self, field: str) -> Type: return Any -class EvalableList(list[T], Evalable["EvalableList[T]"], Generic[T]): +class NameIndexableList(list[T], Generic[T]): """ - A list that can be evaluated from a string. EvalableList[T] means that a given string - can be evaluated, yielding a list of objects of type T. + A list that can be indexed by element name, in addition to the usual integer and + slice indexing. An element's name is its ``name`` attribute, or its ``"name"`` key + if it is a dict. It is not evaluated, so expressions in it are left as-is. """ - def get_validator(self, field: str) -> Type: - return T if self._validator is None else self._validator - - def _eval_expressions( - self, - symbol_table: dict[str, Any] = None, - order: tuple[str, ...] = (), - post_calls: tuple[_PostCall[T], ...] = (), - already_evaluated: dict[str, Any] | None = None, - **kwargs, - ) -> tuple["EvalableList[T]", dict[str, Any]]: - new = EvalableList[T](x for x in self) - symbol_table = symbol_table.copy() if symbol_table is not None else {} - order = order + tuple(x for x in range(len(new)) if x not in order) - return new._eval_expressions_final( - symbol_table, - order, - post_calls, - use_setattr=False, - already_evaluated=already_evaluated, - **kwargs, - ) - - def get_fields(self) -> list[str]: - return sorted(range(len(self))) - - @classmethod - def __get_pydantic_core_schema__( - cls, source_type: Any, handler: Callable - ) -> CoreSchema: - # Get the type parameter T from EvalableList[T] - type_args = get_args(source_type) - if not type_args: - raise TypeError( - f"EvalableList must be used with a type parameter, e.g. EvalableList[int]" - ) - item_type = type_args[0] - - # Get the schema for the item type - item_schema = handler(item_type) - - # Create a schema that validates lists of the item type - return chain_schema( - [ - list_schema(item_schema), - no_info_plain_validator_function(lambda x: cls(x)), - ] - ) - - def __getitem__(self, key: str | int | slice, _pretty_error: bool = True) -> T: + def __getitem__(self, key: str | int | slice, _pretty_error: bool = True): if isinstance(key, int): return super().__getitem__(key) # type: ignore elif isinstance(key, slice): - return EvalableList[T](super().__getitem__(key)) + return type(self)(super().__getitem__(key)) elif isinstance(key, str): found = None @@ -977,7 +929,7 @@ def __getitem__(self, key: str | int | slice, _pretty_error: bool = True) -> T: if found is not None: return found - fields = self.get_fields() + fields = list(range(len(self))) fields += [ ( x.name @@ -1000,10 +952,68 @@ def __contains__(self, item: Any) -> bool: except KeyError: return super().__contains__(item) + @classmethod + def __get_pydantic_core_schema__( + cls, source_type: Any, handler: Callable + ) -> CoreSchema: + # Get the type parameter T from cls[T] + type_args = get_args(source_type) + if not type_args: + raise TypeError( + f"{cls.__name__} must be used with a type parameter, e.g. " + f"{cls.__name__}[int]" + ) + item_type = type_args[0] + + # Get the schema for the item type + item_schema = handler(item_type) + + # Create a schema that validates lists of the item type + return chain_schema( + [ + list_schema(item_schema), + no_info_plain_validator_function(lambda x: cls(x)), + ] + ) + def __copy__(self) -> Self: return type(self)(x for x in self) +class EvalableList(NameIndexableList[T], Evalable["EvalableList[T]"], Generic[T]): + """ + A list that can be evaluated from a string. EvalableList[T] means that a given string + can be evaluated, yielding a list of objects of type T. It can also be indexed by + element name. + """ + + def get_validator(self, field: str) -> Type: + return T if self._validator is None else self._validator + + def _eval_expressions( + self, + symbol_table: dict[str, Any] = None, + order: tuple[str, ...] = (), + post_calls: tuple[_PostCall[T], ...] = (), + already_evaluated: dict[str, Any] | None = None, + **kwargs, + ) -> tuple["EvalableList[T]", dict[str, Any]]: + new = EvalableList[T](x for x in self) + symbol_table = symbol_table.copy() if symbol_table is not None else {} + order = order + tuple(x for x in range(len(new)) if x not in order) + return new._eval_expressions_final( + symbol_table, + order, + post_calls, + use_setattr=False, + already_evaluated=already_evaluated, + **kwargs, + ) + + def get_fields(self) -> list[str]: + return sorted(range(len(self))) + + class EvalableDict( dict[K, V], Evalable["EvalableDict[K, V]"], Generic[K, V], _FromYAMLAble ): diff --git a/tests/vibe_see_readme_in_this_dir/test_renames.py b/tests/vibe_see_readme_in_this_dir/test_renames.py index 42f2b739..431d060b 100644 --- a/tests/vibe_see_readme_in_this_dir/test_renames.py +++ b/tests/vibe_see_readme_in_this_dir/test_renames.py @@ -225,10 +225,9 @@ def test_get_renames_for_einsum_without_default(self): self.assertEqual(result.name, "SomeEinsum") self.assertEqual(len(result.tensor_accesses), 0) - def test_non_default_einsum_renames_applied_at_eval_time(self): - """Non-default einsum renames are only resolved during full spec - evaluation (name-based lookup requires EvalableList). Pre-evaluation, - get_renames_for_einsum only applies defaults.""" + def test_non_default_einsum_renames_found_before_eval(self): + """Renames.einsums is a NameIndexableList, so non-default einsum + renames are found by name without evaluating the spec.""" r = Renames( einsums=[ EinsumRename( @@ -239,10 +238,10 @@ def test_non_default_einsum_renames_applied_at_eval_time(self): ), ] ) - # Without evaluation, 'Matmul' is not found in the plain list, - # so a fresh EinsumRename is created with no tensor_accesses. result = r.get_renames_for_einsum("Matmul") self.assertEqual(result.name, "Matmul") + self.assertEqual(len(result.tensor_accesses), 1) + self.assertEqual(result.tensor_accesses["weight"].source, "W") def test_default_applied_when_no_specific_match(self): """When a specific einsum is not found, defaults are still applied.""" From 693ab783347c628f881c143ee7766ae6ca59b0cc Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Tue, 29 Sep 2026 09:21:46 -0400 Subject: [PATCH 2/5] [frontend] Fix global Einsum-specific rename not being used --- accelforge/frontend/workload.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/accelforge/frontend/workload.py b/accelforge/frontend/workload.py index f7f27601..77c445bd 100755 --- a/accelforge/frontend/workload.py +++ b/accelforge/frontend/workload.py @@ -752,14 +752,16 @@ def _eval_expressions(self, symbol_table: dict[str, Any], *args, **kwargs): self: Einsum = self.model_copy() self.renames = RenameList(self.renames) - # Grab the default renames and update the renames with more values - default_renames = renames.get_renames_for_einsum("default") - for tensor_rename in default_renames.tensor_accesses: - if tensor_rename.name not in self.renames: - self.renames.append(tensor_rename) - for rank_variable_rename in default_renames.rank_variables: - if rank_variable_rename.name not in self.renames: - self.renames.append(rank_variable_rename) + # Grab top-level Einsum-specific renames first, then load defaults that + # without overwriting + for rename_to_consider in [self.name, "default"]: + rename_to_consider = renames.get_renames_for_einsum(rename_to_consider) + for tensor_rename in rename_to_consider.tensor_accesses: + if tensor_rename.name not in self.renames: + self.renames.append(tensor_rename) + for rank_variable_rename in rename_to_consider.rank_variables: + if rank_variable_rename.name not in self.renames: + self.renames.append(rank_variable_rename) # Parse me! kwargs["musteval_tryeval_to"] = True From 95689adb3877f99321128835fb2ce9b174c19677 Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Tue, 29 Sep 2026 10:00:20 -0400 Subject: [PATCH 3/5] [FFM] Improve numerical stability when thresholding using OptimalityThresholder --- accelforge/mapper/FFM/_join_pmappings/join_pmappings.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py b/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py index 82b1a53f..d7b4b049 100755 --- a/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py +++ b/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py @@ -46,6 +46,9 @@ parallel, ) +# Small number for stability +EPS = 1e-5 + logger = logging.getLogger(__name__) @@ -111,7 +114,7 @@ def __init__( print(f"Filtering out pmappings worse than the following:") for i in chosen_indices.astype(int): - self.compare_to.append({c: compare_to[c].iloc[i] for c in compare_cols}) + self.compare_to.append({c: compare_to[c].iloc[i]*(1+EPS) for c in compare_cols}) if print_progress: print( "\t" From 7be6eb3560b2147badb8a957a8eb7645f4ba61c3 Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Mon, 5 Oct 2026 10:29:41 -0400 Subject: [PATCH 4/5] [FFM] Can now return pmapping statistics from join_pmappings --- accelforge/frontend/spec.py | 10 +- accelforge/mapper/FFM/__init__.py | 8 ++ .../FFM/_join_pmappings/join_pmappings.py | 125 +++++++++++++++--- accelforge/mapper/FFM/main.py | 32 ++++- 4 files changed, 148 insertions(+), 27 deletions(-) diff --git a/accelforge/frontend/spec.py b/accelforge/frontend/spec.py index 6f79884e..385c80c4 100755 --- a/accelforge/frontend/spec.py +++ b/accelforge/frontend/spec.py @@ -379,7 +379,8 @@ def map_workload_to_arch( print_progress: bool = True, print_number_of_pmappings: bool = False, _pmapping_row_filter_function: Callable[[pd.Series], bool] | None = None, - ) -> Mappings: + report_statistics: bool = False, + ) -> Mappings | tuple[Mappings, dict]: """ Maps the workload to the architecture using the AccelForge Fast and Fusiest Mapper (FFM). @@ -408,11 +409,15 @@ def map_workload_to_arch( A function that takes in a row of the pmapping dataframe and returns True if the row should be included in the final mappings, and False otherwise. If None, all rows will be included. + report_statistics: + If True, also return statistics about joining. See + `accelforge.mapper.FFM.join_pmappings`. Returns ------- Mappings - The mappings of the workload to the architecture. + The mappings of the workload to the architecture. If ``report_statistics`` + is True, a tuple of the mappings and the joining statistics. """ from accelforge.mapper.FFM.main import map_workload_to_arch @@ -423,6 +428,7 @@ def map_workload_to_arch( print_progress=print_progress, print_number_of_pmappings=print_number_of_pmappings, _pmapping_row_filter_function=_pmapping_row_filter_function, + report_statistics=report_statistics, ) diff --git a/accelforge/mapper/FFM/__init__.py b/accelforge/mapper/FFM/__init__.py index 878fa713..cac95f67 100755 --- a/accelforge/mapper/FFM/__init__.py +++ b/accelforge/mapper/FFM/__init__.py @@ -7,6 +7,11 @@ ) from accelforge.frontend.mapper.metrics import Metrics from accelforge.mapper.FFM._join_pmappings.pmapping_group import PmappingGroup +from accelforge.mapper.FFM._join_pmappings.join_pmappings import ( + JoinRunParameters, + JoinStatistics, + JoinStepStatistics, +) __all__ = [ "map_workload_to_arch", @@ -16,4 +21,7 @@ "Mappings", "Metrics", "PmappingGroup", + "JoinRunParameters", + "JoinStatistics", + "JoinStepStatistics", ] diff --git a/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py b/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py index d7b4b049..b74eb7aa 100755 --- a/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py +++ b/accelforge/mapper/FFM/_join_pmappings/join_pmappings.py @@ -4,10 +4,11 @@ CompatibilityDiff, ) from collections import defaultdict +from dataclasses import dataclass, field import itertools import logging import time -from typing import Any, Callable +from typing import Any, Callable, NamedTuple from accelforge._accelerated_imports import pd, np from accelforge.frontend.spec import Spec @@ -24,6 +25,7 @@ get_rank_variable_bounds_for_all_einsums, ) from accelforge.mapper.FFM._join_pmappings.pmapping_dataframe import ( + PmappingDataframe, row2pmappings, ) from accelforge.mapper.FFM._pareto_df.df_convention import ( @@ -52,6 +54,50 @@ logger = logging.getLogger(__name__) +@dataclass +class JoinStepStatistics: + """Statistics for joining one Einsum onto the partial mappings to its left.""" + + left_einsum: EinsumName + """The last Einsum in the partial mappings being joined from the left.""" + right_einsum: EinsumName + """The Einsum being joined from the right.""" + n_left_mappings: int + """Number of mappings on the left after grouping and pruning, before joining.""" + n_right_mappings: int + """Number of mappings on the right after grouping and pruning, before joining.""" + n_mappings_before_pruning: int + """Number of joined mappings before Pareto pruning (product of the sizes of + each pair of joined groups).""" + n_mappings: int + """Number of joined mappings after Pareto pruning.""" + group_compatibilities: list[Compatibility] + """The compatibility of each group of joined mappings.""" + + @property + def n_groups(self) -> int: + return len(self.group_compatibilities) + + @property + def mappings_per_group(self) -> float: + return self.n_mappings / self.n_groups + + +@dataclass +class JoinStatistics: + """Statistics for one call of join_pmappings.""" + + steps: list[JoinStepStatistics] = field(default_factory=list) + """Statistics for each Einsum joined, in joining order.""" + + +class JoinRunParameters(NamedTuple): + """Pruning tolerances used for one call of join_pmappings.""" + + objective_tolerance: float + resource_usage_tolerance: float + + class JoiningTimer: def __init__(self): self.prev_time = time.time() @@ -196,6 +242,7 @@ def join_strategy_2( for_model: bool, _pmapping_row_filter_function: Callable[[pd.DataFrame], np.ndarray] | None = None, resource_usage_tolerance: float = 0, + statistics: dict[JoinRunParameters, JoinStatistics] | None = None, ): thresholds = [1, 0] thresholds = [t for t in thresholds if t > spec.mapper.objective_tolerance] @@ -233,7 +280,12 @@ def join_strategy_2( _pmapping_row_filter_function=filter_func, print_progress=print_progress, metrics=metrics, + report_statistics=statistics is not None, ) + if statistics is not None: + joined, run_statistics = joined + run = JoinRunParameters(threshold, resource_usage_tolerance) + statistics[run] = run_statistics if i < len(thresholds) - 1: filter_func = OptimalityThresholder( joined, @@ -257,6 +309,7 @@ def multi_strategy_join( metrics: Metrics, for_model: bool, _pmapping_row_filter_function: Callable[[pd.DataFrame], np.ndarray] | None = None, + statistics: dict[JoinRunParameters, JoinStatistics] | None = None, ): for _, p in compressed.items(): for pg in p: @@ -264,13 +317,17 @@ def multi_strategy_join( # If it's for the model, just join things directly if for_model: - return join_pmappings( + joined = join_pmappings( deepcopy(compressed), spec, print_progress=print_progress, metrics=metrics, _pmapping_row_filter_function=_pmapping_row_filter_function, + report_statistics=statistics is not None, ) + if statistics is not None: + joined, statistics[JoinRunParameters(0, 0)] = joined + return joined if metrics & Metrics.RESOURCE_USAGE: return join_strategy_2( @@ -280,6 +337,7 @@ def multi_strategy_join( metrics, for_model, _pmapping_row_filter_function, + statistics=statistics, ) resource_usage_thresholds = [ @@ -309,6 +367,7 @@ def multi_strategy_join( for_model, _pmapping_row_filter_function, resource_usage_tolerance=threshold, + statistics=statistics, ) for c in joined.data.columns: if is_reservation_col(c): @@ -334,7 +393,13 @@ def clean_compress_and_join_pmappings( require_all_einsums: bool = True, _pmapping_row_filter_function: Callable[[pd.Series], bool] | None = None, print_progress: bool = True, -) -> Mappings: + report_statistics: bool = False, +) -> Mappings | tuple[Mappings, dict[JoinRunParameters, JoinStatistics]]: + """ + If report_statistics is True, returns a tuple of (mappings, statistics), where + statistics maps the parameters of each call of join_pmappings to its + JoinStatistics. + """ einsum2pmappings = pmappings.einsum2pmappings if not require_all_einsums: einsum2pmappings = { @@ -348,6 +413,7 @@ def clean_compress_and_join_pmappings( einsum2pmappings, print_progress ) + statistics = {} if report_statistics else None joined = multi_strategy_join( pmappings.spec, compressed, @@ -355,6 +421,7 @@ def clean_compress_and_join_pmappings( metrics, for_model, _pmapping_row_filter_function, + statistics=statistics, ) joined = decompress_pmappings(joined, decompress_data) @@ -387,7 +454,7 @@ def clean_compress_and_join_pmappings( # Fill nans with 0. We might get missing columns for some mapping entries if there # are energy entries for some pmappings but not others (e.g., one pmapping accesses # DRAM while another doesn't.) - return Mappings( + mappings = Mappings( pmappings.spec, list( x @@ -400,6 +467,9 @@ def clean_compress_and_join_pmappings( flattened_arches=pmappings.flattened_arches, evaluated_specs=pmappings.evaluated_specs, ) + if report_statistics: + return mappings, statistics + return mappings class PmappingsOneEinsum: @@ -503,8 +573,11 @@ def join_pmappings( metrics: Metrics = None, _pmapping_row_filter_function: Callable[[pd.Series], bool] | None = None, print_progress: bool = True, -): + report_statistics: bool = False, +) -> PmappingDataframe | tuple[PmappingDataframe, JoinStatistics]: """ + If report_statistics is True, returns a tuple of (mappings, JoinStatistics). + CONTRACT FOR MAPPINGS GETTING TO THIS POINT: - Reservations at a level include reservations at all levels above it. @@ -554,6 +627,7 @@ def join_pmappings( aliased_tensors = spec.workload.get_tensor_copies() runtime = {} + statistics = JoinStatistics() pmapping_groups = list(pmapping_groups.items()) @@ -754,6 +828,16 @@ def grab_einsum_pmappings() -> ( [s for lr in [left, right] for v in lr.values() for s, _ in v], live_tensors ) + if report_statistics: + # Groups may appear in multiple buckets (one per permutation), so dedupe + n_left_mappings, n_right_mappings = ( + sum( + len(s.mappings.data) + for s in {id(s): s for v in lr.values() for s, _ in v}.values() + ) + for lr in (left, right) + ) + DO_PRINT = False DELAY = True # ====================================================================== @@ -963,21 +1047,20 @@ def no_match_lookahead_error( # f"\tCombining {sum(len(s) for s in left.values())}({len(left)}) x {sum(len(s) for s in right.values())}({len(right)}) -> {len(combined)}" # ) - nmappings = sum(len(s.mappings.data) for s in combined) - for_einsum_text = f"for Einsum {right_einsum}" - # print(f"\tNumber of groups {for_einsum_text}: {len(combined)}") - # for c in combined: - # print(f"\t\t{c.compatibility}") - # print(f"\tNumber of mappings {for_einsum_text}: {nmappings}") - # print( - # f"\tMappings per group {for_einsum_text}: {nmappings / len(combined)}" - # ) - # logger.info( - # f"\tLargest left: {max(len(s2.mappings.data) for s in left.values() for s2, _ in s)}" - # ) - # logger.info( - # f"\tLargest right: {max(len(s2.mappings.data) for s in right.values() for s2, _ in s)}" - # ) + if report_statistics: + statistics.steps.append( + JoinStepStatistics( + left_einsum=left_einsum, + right_einsum=right_einsum, + n_left_mappings=n_left_mappings, + n_right_mappings=n_right_mappings, + n_mappings_before_pruning=sum( + s.n_pre_prune_mappings for s in combined + ), + n_mappings=sum(len(s.mappings.data) for s in combined), + group_compatibilities=[s.compatibility for s in combined], + ) + ) # ====================================================================== # Update left for the next iteration. @@ -1018,6 +1101,8 @@ def no_match_lookahead_error( # evaluations_tracker.n_mappings.update(n_mappings) # evaluations_tracker.runtime.update(runtime) + if report_statistics: + return mappings, statistics return mappings diff --git a/accelforge/mapper/FFM/main.py b/accelforge/mapper/FFM/main.py index de99a7e4..1ccf8f04 100755 --- a/accelforge/mapper/FFM/main.py +++ b/accelforge/mapper/FFM/main.py @@ -13,6 +13,8 @@ import accelforge.mapper.FFM._make_pmappings.make_pmappings as pmapper from accelforge.frontend.workload import EinsumName from accelforge.mapper.FFM._join_pmappings.join_pmappings import ( + JoinRunParameters, + JoinStatistics, clean_compress_and_join_pmappings, ) from accelforge._accelerated_imports import pd @@ -32,7 +34,8 @@ def map_workload_to_arch( print_number_of_pmappings: bool = False, eval_in_detail: bool = True, _pmapping_row_filter_function: Callable[[pd.Series], bool] | None = None, -) -> Mappings: + report_statistics: bool = False, +) -> Mappings | tuple[Mappings, dict[JoinRunParameters, JoinStatistics]]: """ Maps a workload to an architecture using the AccelForge Fast and Fusiest Mapper (FFM). @@ -63,6 +66,13 @@ def map_workload_to_arch( A function that takes in a row of the pmapping dataframe and returns True if the row should be included in the final mappings, and False otherwise. If None, all rows will be included. + report_statistics: + If True, also return statistics about joining. See `join_pmappings`. + + Returns + ------- + Mappings | tuple[Mappings, dict[JoinRunParameters, JoinStatistics]] + The mappings, and the joining statistics if ``report_statistics`` is True. """ from accelforge.model.main import evaluate_mapping @@ -86,10 +96,13 @@ def map_workload_to_arch( _pmapping_row_filter_function=_pmapping_row_filter_function, print_progress=print_progress, metrics=spec.mapper.metrics, + report_statistics=report_statistics, ) + if report_statistics: + mappings, statistics = mappings if not eval_in_detail: - return mappings + return (mappings, statistics) if report_statistics else mappings def eval_mapping(i, spec, mappings): local_spec = deepcopy(spec) @@ -146,6 +159,8 @@ def eval_mapping(i, spec, mappings): # print(f'\t{c}: {r[c]}') mappings.data = _fillna_and__numeric_cast(pd.concat(results), 0) + if report_statistics: + return mappings, statistics return mappings @@ -221,7 +236,8 @@ def join_pmappings( _skip_invalid: bool = True, _combine_reservations: bool = True, _runtime_log_file: str | None = None, -) -> Mappings: + report_statistics: bool = False, +) -> Mappings | tuple[Mappings, dict[JoinRunParameters, JoinStatistics]]: """ Joins pmappings into a full mappings for the entire workload. Pmappings can be generated using `make_pmappings`. @@ -247,10 +263,15 @@ def join_pmappings( If True, consolidate reservations to increase pruning effectiveness. _runtime_log_file: If set, append per-step runtime as JSON lines to this file. + report_statistics: + If True, also return statistics about joining. Joining may be run several + times with different pruning tolerances, so the statistics are a dictionary + mapping a `JoinRunParameters` for each run to that run's `JoinStatistics`. Returns ------- - Mappings - A Mappings object containing all valid, optimal mappings for the workload. + Mappings | tuple[Mappings, dict[JoinRunParameters, JoinStatistics]] + A Mappings object containing all valid, optimal mappings for the workload, and + the joining statistics if ``report_statistics`` is True. """ spec = pmappings.spec if _skip_invalid is not True: @@ -266,6 +287,7 @@ def join_pmappings( _pmapping_row_filter_function=_pmapping_row_filter_function, print_progress=print_progress, for_model=False, + report_statistics=report_statistics, ) From bfc12161157912bea8bdb0cab58d15cee1323431 Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Mon, 5 Oct 2026 10:38:57 -0400 Subject: [PATCH 5/5] [util] set_n_parallel_jobs now checks input type --- accelforge/util/parallel.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/accelforge/util/parallel.py b/accelforge/util/parallel.py index 564716c5..08544387 100755 --- a/accelforge/util/parallel.py +++ b/accelforge/util/parallel.py @@ -75,6 +75,8 @@ def set_n_parallel_jobs(n_jobs: int, print_message: bool = False) -> None: print_message : bool, optional Whether to print a message when the number of parallel jobs is set. """ + if not isinstance(n_jobs, int): + raise TypeError(f"n_jobs must be an integer, got {n_jobs} with type {type(n_jobs)}") global N_PARALLEL_PROCESSES N_PARALLEL_PROCESSES = n_jobs global PARALLELIZE