Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions accelforge/frontend/renames.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
EvalableList,
EvalableModel,
EvalsTo,
NameIndexableList,
TryEvalTo,
_PostCall,
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
18 changes: 10 additions & 8 deletions accelforge/frontend/workload.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion accelforge/mapper/FFM/_join_pmappings/join_pmappings.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,9 @@
parallel,
)

# Small number for stability
EPS = 1e-5

logger = logging.getLogger(__name__)


Expand Down Expand Up @@ -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"
Expand Down
120 changes: 65 additions & 55 deletions accelforge/util/_basetypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
):
Expand Down
11 changes: 5 additions & 6 deletions tests/vibe_see_readme_in_this_dir/test_renames.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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."""
Expand Down
Loading