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
252 changes: 145 additions & 107 deletions ml_peg/app/build_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,12 +36,17 @@
build_page_loading_spinner,
build_weight_components,
)
from ml_peg.app.utils.build_frameworks import (
build_framework_page_layout,
build_framework_summary_tables,
build_framework_views,
)
from ml_peg.app.utils.onboarding import (
build_onboarding_modal,
register_onboarding_callbacks,
)
from ml_peg.app.utils.register_callbacks import (
register_benchmark_to_category_callback,
register_benchmark_to_group_callback,
register_filter_loading_callback,
register_filter_tables_callback,
)
Expand All @@ -51,6 +56,7 @@
)
from ml_peg.app.utils.utils import (
build_level_of_theory_warnings,
framework_sort_key,
get_framework_config,
get_mlip_column_width,
load_model_registry_configs,
Expand Down Expand Up @@ -543,7 +549,9 @@ def build_category(

# Register callback for all benchmark tables -> category table
# Category summary table columns add "Score" to name for clarity
register_benchmark_to_category_callback(all_tables, category_to_title)
register_benchmark_to_group_callback(
all_tables, category_to_title, table_id_suffix="-summary-table"
)

return category_views, category_tables, category_weights, framework_ids

Expand Down Expand Up @@ -604,100 +612,12 @@ def build_category_page_layout(
)


def build_framework_views(
category_views: dict[str, dict[str, object]],
framework_ids: set[str],
) -> dict[str, dict[str, object]]:
"""
Build extra framework-focused page metadata for non-default frameworks.

Parameters
----------
category_views
Category metadata including benchmark layout components.
framework_ids
All framework IDs discovered from benchmark apps.

Returns
-------
dict[str, dict[str, object]]
Mapping of framework ID to grouped benchmark layouts by category.
"""
framework_views: dict[str, dict[str, object]] = {}
for framework_id in sorted(framework_ids):
if framework_id == "ml_peg":
continue

category_groups = []
for category_name, category_view in category_views.items():
tests = [
test["layout"]
for test in category_view["tests"]
if framework_id in test["framework_ids"]
]
if tests:
category_groups.append({"category": category_name, "tests": tests})

if category_groups:
framework_views[framework_id] = {
"framework_id": framework_id,
"label": get_framework_config(framework_id)["label"],
"category_groups": category_groups,
}

return framework_views


def build_framework_page_layout(framework_view: dict[str, object]) -> Div:
"""
Build a framework-focused page containing benchmark sections only.

Parameters
----------
framework_view
Framework page metadata with grouped benchmark layouts by category.

Returns
-------
Div
Framework page layout.
"""
framework_label = framework_view["label"]
category_groups = framework_view["category_groups"]

sections = []
for group in category_groups:
sections.append(H3(group["category"], style={"marginTop": "26px"}))
sections.append(Div(group["tests"], style={"display": "grid", "gap": "24px"}))

return Div(
[
H1(f"{framework_label} Benchmarks"),
Div(
(
"These benchmarks also remain in their category pages. "
"Framework pages omit the category summary table and reuse the "
"same benchmark controls, so weight and threshold edits stay in "
"sync across both views."
),
style={
"fontSize": "13px",
"fontStyle": "italic",
"color": "#64748b",
"marginTop": "8px",
"marginBottom": "8px",
},
),
*sections,
]
)


def build_summary_table(
tables: dict[str, DataTable],
table_id: str = "summary-table",
description: str | None = None,
weights: dict[str, float] | None = None,
header_labels: dict[str, str] | None = None,
) -> DataTable:
"""
Build summary table from a set of tables.
Expand All @@ -712,6 +632,9 @@ def build_summary_table(
Description of summary table. Default is None.
weights
Weights for each column. Default is `None`, which sets all weights to 1.
header_labels
Optional mapping of table key to display header text. Column ids are
unchanged, so only the rendered header differs. Default is None.

Returns
-------
Expand Down Expand Up @@ -749,19 +672,27 @@ def build_summary_table(
data = calc_table_scores(data, weights=weights)

columns_headers = ("MLIP", "Score") + tuple(key + " Score" for key in tables)
if table_id == "summary-table":
display_headers = {
header: (
header
if header in {"MLIP", "Score"} or not header.endswith(" Score")
else "\n".join([*header.removesuffix(" Score").split(), "Score"])

# Map each column id to its visible header text (labels + line wrapping)
display_headers = {}
for column_id in columns_headers:
header = column_id
if (
header not in {"MLIP", "Score"}
and header.endswith(" Score")
and header_labels
):
key = header.removesuffix(" Score")
if key in header_labels:
header = header_labels[key] + " Score"
if table_id != "summary-table":
display_headers[column_id] = _format_summary_column_header(header)
elif header in {"MLIP", "Score"} or not header.endswith(" Score"):
display_headers[column_id] = header
else:
display_headers[column_id] = "\n".join(
[*header.removesuffix(" Score").split(), "Score"]
)
for header in columns_headers
}
else:
display_headers = {
header: _format_summary_column_header(header) for header in columns_headers
}

columns = [
{"name": display_headers[header], "id": header} for header in columns_headers
Expand Down Expand Up @@ -902,6 +833,8 @@ def build_nav(
summary_table: DataTable,
weight_components: Div,
all_apps: dict[str, Dash],
combined_framework_table: DataTable | None = None,
framework_weight_components: Div | None = None,
) -> None:
"""
Build page layouts and sidebar navigation.
Expand All @@ -920,15 +853,18 @@ def build_nav(
Weight sliders, text boxes and reset button.
all_apps
Dictionary of all test apps.
combined_framework_table
Frameworks summary table shown on the home page, or None when there are
no external frameworks.
framework_weight_components
Weight controls for the combined framework summary table, or None.
"""
category_paths = {
category_name: _category_to_path(category_name)
for category_name in sorted(category_views)
}
framework_order = sorted(
framework_views,
key=lambda framework_id: framework_views[framework_id]["label"],
)
# Frameworks first, then papers, alphabetical by label within each
framework_order = sorted(framework_views, key=framework_sort_key)
framework_paths = {
framework_id: _framework_to_path(framework_id)
for framework_id in framework_order
Expand Down Expand Up @@ -1011,6 +947,27 @@ def build_nav(
]
)

# Computed + weight stores for each per-framework summary table
framework_state_stores = []
for framework_view in framework_views.values():
fw_table = framework_view.get("summary_table")
if fw_table is None:
continue
framework_state_stores.extend(
[
Store(
id=f"{fw_table.id}-computed-store",
storage_type="session",
data=fw_table.data,
),
Store(
id=f"{fw_table.id}-weight-store",
storage_type="session",
data=_default_weight_store_data(fw_table),
),
]
)

test_state_stores = []
for app in all_apps.values():
test_state_stores.extend(app.stores)
Expand All @@ -1023,8 +980,17 @@ def build_nav(
),
Store(id="cmap-store", storage_type="local", data="viridis_r"),
*category_state_stores,
*framework_state_stores,
*test_state_stores,
]
if combined_framework_table is not None:
global_state_stores.append(
Store(
id="framework-summary-table-weight-store",
storage_type="session",
data=_default_weight_store_data(combined_framework_table),
)
)

full_layout = [
# Start-up mask: covers the page and tutorial with a spinner until the
Expand Down Expand Up @@ -1092,6 +1058,11 @@ def build_nav(
id="summary-table-scores-store",
storage_type="session",
),
*(
[Store(id="framework-summary-table-scores-store", storage_type="session")]
if combined_framework_table is not None
else []
),
Div(global_state_stores, style={"display": "none"}),
Div(
[
Expand Down Expand Up @@ -1159,6 +1130,17 @@ def build_nav(
storage_type="session",
data=summary_table.data,
),
*(
[
Store(
id="framework-summary-table-computed-store",
storage_type="session",
data=combined_framework_table.data,
)
]
if combined_framework_table is not None
else []
),
Store(
id="filter-recompute-done",
storage_type="memory",
Expand Down Expand Up @@ -1348,6 +1330,7 @@ def select_page(
if pathname in (None, "", "/", "/summary"):
summary_counts = (
f"{len(category_views)} categories · {len(all_apps)} benchmarks"
f" · {len(framework_views)} frameworks"
)
return Div(
[
Expand Down Expand Up @@ -1389,6 +1372,29 @@ def select_page(
],
style={"width": "fit-content"},
),
*(
[
H1(
"Frameworks Summary",
style={"marginTop": "32px"},
),
Div(
[
build_download_controls(
combined_framework_table.id, row=True
),
build_loading_summary_table(
combined_framework_table
),
Br(),
framework_weight_components,
],
style={"width": "fit-content"},
),
]
if combined_framework_table is not None
else []
),
build_faqs(),
]
), sidebar_children
Expand Down Expand Up @@ -1439,6 +1445,36 @@ def build_full_app(full_app: Dash, category: str = "*", test: str = "*") -> None
all_layouts, all_tables, all_frameworks
)
framework_views = build_framework_views(cat_views, framework_ids)
# Build per-framework summary tables and the combined framework summary
framework_tables, framework_grouping = build_framework_summary_tables(
all_tables, all_frameworks, framework_views
)
combined_framework_table = None
framework_weight_components = None
if framework_tables:
combined_framework_table = build_summary_table(
dict(
sorted(
framework_tables.items(),
key=lambda item: framework_sort_key(item[0]),
)
),
table_id="framework-summary-table",
header_labels={
fid: str(framework_views[fid]["label"]) for fid in framework_tables
},
)
framework_weight_components = build_weight_components(
header="Weights",
table=combined_framework_table,
include_download_controls=False,
column_widths=combined_framework_table.column_widths,
)
register_benchmark_to_group_callback(
framework_grouping,
{fid: fid for fid in framework_grouping},
table_id_suffix="-framework-summary-table",
)
# Build overall summary table
summary_table = build_summary_table(
dict(sorted(cat_tables.items())), weights=cat_weights
Expand All @@ -1457,5 +1493,7 @@ def build_full_app(full_app: Dash, category: str = "*", test: str = "*") -> None
summary_table,
weight_components,
all_apps,
combined_framework_table,
framework_weight_components,
)
register_onboarding_callbacks()
Loading
Loading