From 4803cd91b37dd8fbfafa6f3ce7dcebfc576cb368 Mon Sep 17 00:00:00 2001 From: Michael Gilbert Date: Fri, 25 Sep 2026 04:55:33 -0400 Subject: [PATCH 1/3] [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/3] [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/3] [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"