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
10 changes: 8 additions & 2 deletions accelforge/frontend/spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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

Expand All @@ -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,
)


Expand Down
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
8 changes: 8 additions & 0 deletions accelforge/mapper/FFM/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -16,4 +21,7 @@
"Mappings",
"Metrics",
"PmappingGroup",
"JoinRunParameters",
"JoinStatistics",
"JoinStepStatistics",
]
130 changes: 109 additions & 21 deletions accelforge/mapper/FFM/_join_pmappings/join_pmappings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 (
Expand All @@ -46,9 +48,56 @@
parallel,
)

# Small number for stability
EPS = 1e-5

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()
Expand Down Expand Up @@ -111,7 +160,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 Expand Up @@ -193,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]
Expand Down Expand Up @@ -230,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,
Expand All @@ -254,20 +309,25 @@ 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:
pg.mappings.drop_valid_reservations = not (Metrics.RESOURCE_USAGE & metrics)

# 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(
Expand All @@ -277,6 +337,7 @@ def multi_strategy_join(
metrics,
for_model,
_pmapping_row_filter_function,
statistics=statistics,
)

resource_usage_thresholds = [
Expand Down Expand Up @@ -306,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):
Expand All @@ -331,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 = {
Expand All @@ -345,13 +413,15 @@ def clean_compress_and_join_pmappings(
einsum2pmappings, print_progress
)

statistics = {} if report_statistics else None
joined = multi_strategy_join(
pmappings.spec,
compressed,
print_progress,
metrics,
for_model,
_pmapping_row_filter_function,
statistics=statistics,
)

joined = decompress_pmappings(joined, decompress_data)
Expand Down Expand Up @@ -384,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
Expand All @@ -397,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:
Expand Down Expand Up @@ -500,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.
Expand Down Expand Up @@ -551,6 +627,7 @@ def join_pmappings(
aliased_tensors = spec.workload.get_tensor_copies()

runtime = {}
statistics = JoinStatistics()

pmapping_groups = list(pmapping_groups.items())

Expand Down Expand Up @@ -751,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
# ======================================================================
Expand Down Expand Up @@ -960,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.
Expand Down Expand Up @@ -1015,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


Expand Down
Loading
Loading