| from __future__ import annotations |
|
|
| import re |
| from pathlib import Path |
| from typing import Iterable |
|
|
| import numpy as np |
| import pandas as pd |
| import plotly.graph_objects as go |
| from dash import Dash, Input, Output, dash_table, dcc, html |
| from plotly.subplots import make_subplots |
|
|
|
|
| DATA_DIR = Path(__file__).resolve().parent / "data" |
|
|
| PAPER_MODEL_ORDER = [ |
| "blend", |
| "blend_ndt", |
| "cebra", |
| "dnn", |
| "dpad", |
| "gpfa", |
| "gru", |
| "langevinflow_ccn", |
| "ldns", |
| "lfads_torch", |
| "lstm", |
| "marble", |
| "mint", |
| "neds", |
| "neds_pretrained", |
| "neuro_behavior_conditioning", |
| "pca", |
| "rnn", |
| "smc_rnns", |
| "svc", |
| "tndm", |
| "torchdfine", |
| "xg", |
| ] |
|
|
| DISPLAY_NAMES = { |
| "blend": "BLEND-LFADS", |
| "blend_ndt": "BLEND-NDT", |
| "cebra": "CEBRA", |
| "dnn": "DNN", |
| "dpad": "DPAD", |
| "gpfa": "GPFA", |
| "gru": "GRU", |
| "langevinflow_ccn": "LangevinFlow", |
| "ldns": "LDNS", |
| "lfads_torch": "AutoLFADS", |
| "lstm": "LSTM", |
| "marble": "MARBLE", |
| "mint": "MINT", |
| "neds": "NEDS", |
| "neds_pretrained": "NEDS-pt", |
| "neuro_behavior_conditioning": "mVAE", |
| "pca": "PCA", |
| "rnn": "RNN", |
| "smc_rnns": "SMC-RNN", |
| "svc": "SVC", |
| "tndm": "TNDM", |
| "torchdfine": "DFINE", |
| "xg": "XGBoost", |
| } |
|
|
| METHOD_FAMILY = { |
| "pca": "Linear latent", |
| "gpfa": "Linear latent", |
| "dnn": "Supervised decoder", |
| "gru": "Supervised decoder", |
| "lstm": "Supervised decoder", |
| "rnn": "Supervised decoder", |
| "svc": "Supervised decoder", |
| "xg": "Supervised decoder", |
| "lfads_torch": "Sequential latent", |
| "dpad": "Sequential latent", |
| "torchdfine": "Sequential latent", |
| "smc_rnns": "Sequential latent", |
| "tndm": "Sequential latent", |
| "langevinflow_ccn": "Sequential latent", |
| "ldns": "Sequential latent", |
| "neuro_behavior_conditioning": "Sequential latent", |
| "blend": "Distillation", |
| "blend_ndt": "Distillation", |
| "neds": "Foundation model", |
| "neds_pretrained": "Foundation model", |
| "cebra": "Contrastive", |
| "marble": "Geometric", |
| "mint": "Non-parametric", |
| } |
|
|
| METHOD_HARDWARE = { |
| "gpfa": "CPU", |
| "mint": "CPU", |
| "pca": "CPU", |
| "svc": "CPU", |
| "xg": "CPU", |
| } |
|
|
| TABLE_LABELS = { |
| "rank": "Rank", |
| "method": "Method", |
| "family": "Family", |
| "hardware": "Hardware", |
| "task_score": "Decoding score", |
| "score": "Decoding score", |
| "robustness_auc": "Robustness AUC", |
| "alignment_score": "Cross-session alignment", |
| "training_time_sec": "Training time (s)", |
| "inference_time_sec": "Evaluation time (s)", |
| "peak_ram_gb": "Peak memory (GB)", |
| "peak_vram_gb": "Peak GPU memory (GB)", |
| "reference_score": "Reference score", |
| "highest_noise_score": "Highest-noise score", |
| "average_noisy_score": "Average noisy score", |
| "latent_dim": "Latent dimensions", |
| "n_sessions": "Sessions", |
| "n_pairwise": "Session pairs", |
| "baseline_score": "Baseline score", |
| "full_model_score": "Full-model score", |
| "neuron_influence_auc": "Neuron influence AUC", |
| "trial_influence_auc": "Trial influence AUC", |
| "shap_mean_value": "Mean signed neuron contribution", |
| "shap_fraction_positive": "Fraction positive", |
| "n_train_trials": "Training trials", |
| "n_test_trials": "Test trials", |
| "n_neurons": "Neurons", |
| } |
|
|
| NUMERIC_COLUMNS = { |
| "task_score", |
| "score", |
| "robustness_auc", |
| "alignment_score", |
| "training_time_sec", |
| "inference_time_sec", |
| "peak_ram_gb", |
| "peak_vram_gb", |
| "reference_score", |
| "highest_noise_score", |
| "average_noisy_score", |
| "latent_dim", |
| "n_sessions", |
| "n_pairwise", |
| "baseline_score", |
| "full_model_score", |
| "neuron_influence_auc", |
| "trial_influence_auc", |
| "shap_mean_value", |
| "shap_fraction_positive", |
| "n_train_trials", |
| "n_test_trials", |
| "n_neurons", |
| } |
| RIGHT_ALIGNED_COLUMNS = NUMERIC_COLUMNS | {"rank"} |
|
|
| |
| TASK_COLOR = "#1565C0" |
| ROBUSTNESS_COLOR = "#2E7D32" |
| COMPUTE_COLOR = "#E65100" |
| INFLUENCE_COLOR = "#CC79A7" |
| ALIGNMENT_COLOR = "#0072B2" |
| SCORE_SCALE = [[0.0, "#EFF6FF"], [1.0, TASK_COLOR]] |
| ALIGNMENT_SCALE = [[0.0, "#F7FBF7"], [1.0, ROBUSTNESS_COLOR]] |
| DATASET_COLORS = { |
| "monkey": "#0072B2", |
| "allen_neuropixels": "#E69F00", |
| "speech": "#009E73", |
| "mc_pacman": "#CC79A7", |
| "ratinabox": "#D55E00", |
| } |
| CATEGORICAL_PALETTE = [ |
| "#E69F00", |
| "#56B4E9", |
| "#009E73", |
| "#F0E442", |
| "#0072B2", |
| "#D55E00", |
| "#CC79A7", |
| "#000000", |
| ] |
| DIRECTION_PALETTE = [ |
| "#B23AEE", |
| "#3B1C54", |
| "#2DD4F6", |
| "#289285", |
| "#E3D724", |
| "#00A65A", |
| "#5B8FF9", |
| "#F97316", |
| ] |
| SPEECH_PALETTE = { |
| "3": "#E69F00", |
| "2": "#56B4E9", |
| "4": "#009E73", |
| "7": "#999933", |
| "6": "#0072B2", |
| "1": "#D55E00", |
| "5": "#CC79A7", |
| "0": "#000000", |
| } |
| RATINABOX_SCALE = [ |
| [0.0, "#440154"], |
| [0.25, "#3B528B"], |
| [0.5, "#21918C"], |
| [0.75, "#5EC962"], |
| [1.0, "#FDE725"], |
| ] |
|
|
|
|
| def load_csv(name: str) -> pd.DataFrame: |
| path = DATA_DIR / name |
| if not path.exists(): |
| raise FileNotFoundError(f"Missing dashboard data: {path}") |
| return pd.read_csv(path) |
|
|
|
|
| prediction = load_csv("clean_prediction_summary.csv") |
| robustness = load_csv("robustness_summary.csv") |
| consistency = load_csv("consistency_summary.csv") |
| scalability = load_csv("scalability_summary.csv") |
| neuron_shap = load_csv("neuron_shap_summary.csv") |
| trial_shapley = load_csv("trial_shapley_summary.csv") |
| latent_samples = load_csv("latent_samples.csv") |
| latent_trajectories = load_csv("latent_trajectories.csv") |
|
|
|
|
| def present_rows(df: pd.DataFrame) -> pd.DataFrame: |
| if df.empty or "status" not in df.columns: |
| return df.copy() |
| return df[df["status"].fillna("") == "present"].copy() |
|
|
|
|
| def active_rows(df: pd.DataFrame) -> pd.DataFrame: |
| if df.empty or "is_active_model" not in df.columns: |
| return df.copy() |
| return df[df["is_active_model"].astype(str).str.lower() == "true"].copy() |
|
|
|
|
| def ordered_unique(values: Iterable[object]) -> list[str]: |
| seen: set[str] = set() |
| out: list[str] = [] |
| for value in values: |
| if pd.isna(value): |
| continue |
| text = str(value) |
| if text not in seen: |
| seen.add(text) |
| out.append(text) |
| return out |
|
|
|
|
| def build_dataset_labels() -> dict[str, str]: |
| pairs = ( |
| prediction[["dataset", "dataset_display"]] |
| .dropna(subset=["dataset"]) |
| .drop_duplicates(subset=["dataset"]) |
| ) |
| return {row.dataset: row.dataset_display for row in pairs.itertuples()} |
|
|
|
|
| DATASET_LABELS = build_dataset_labels() |
| DATASETS = ordered_unique(prediction.get("dataset", pd.Series(dtype=str))) |
| CONSISTENCY_DATASETS = set( |
| active_rows(consistency).get("dataset", pd.Series(dtype=str)).dropna().astype(str) |
| ) |
| MODEL_SET = set(prediction.get("model", pd.Series(dtype=str)).dropna().astype(str)) |
| MODELS = [model for model in PAPER_MODEL_ORDER if model in MODEL_SET] |
| MODELS += sorted(model for model in MODEL_SET if model not in set(MODELS)) |
| MODEL_RANK = {model: idx for idx, model in enumerate(MODELS)} |
|
|
| CONDITION_LABELS = { |
| "monkey": { |
| "0": "Up", |
| "1": "Up-right", |
| "2": "Right", |
| "3": "Down-right", |
| "4": "Down", |
| "5": "Down-left", |
| "6": "Left", |
| "7": "Up-left", |
| }, |
| "allen_neuropixels": { |
| str(i): f"{angle} deg" |
| for i, angle in enumerate([0, 45, 90, 135, 180, 225, 270, 315]) |
| }, |
| "speech": { |
| "0": "Do nothing", |
| "1": "ban", |
| "2": "choice", |
| "3": "day", |
| "4": "feel", |
| "5": "kite", |
| "6": "though", |
| "7": "were", |
| }, |
| } |
|
|
|
|
| def model_label(model: object) -> str: |
| if pd.isna(model): |
| return "" |
| text = str(model) |
| return DISPLAY_NAMES.get(text, text) |
|
|
|
|
| def model_sort_value(model: object) -> int: |
| return MODEL_RANK.get(str(model), len(MODEL_RANK)) |
|
|
|
|
| def hardware_label(model: object) -> str: |
| return METHOD_HARDWARE.get(str(model), "GPU") |
|
|
|
|
| def selected_models(models: list[str] | None) -> list[str]: |
| if not models: |
| return MODELS.copy() |
| return [model for model in MODELS if model in set(models)] |
|
|
|
|
| def filter_models(df: pd.DataFrame, models: list[str] | None) -> pd.DataFrame: |
| if df.empty or "model" not in df.columns: |
| return df.copy() |
| return df[df["model"].astype(str).isin(selected_models(models))].copy() |
|
|
|
|
| def condition_sort_key(value: object) -> tuple[int, float | str]: |
| try: |
| return (0, float(value)) |
| except (TypeError, ValueError): |
| return (1, str(value)) |
|
|
|
|
| def condition_label(dataset: str, condition: object) -> str: |
| if pd.isna(condition): |
| return "Unknown" |
| text = str(condition) |
| mapped = CONDITION_LABELS.get(dataset, {}).get(text) |
| if mapped is not None: |
| return mapped |
| if dataset == "ratinabox": |
| try: |
| idx = int(float(text)) |
| return f"x{idx % 10}, y{idx // 10}" |
| except ValueError: |
| return text |
| return text |
|
|
|
|
| def condition_axis_label(dataset: str, color_mode: str = "condition") -> str: |
| if dataset == "ratinabox": |
| if color_mode == "x": |
| return "X position bin" |
| if color_mode == "y": |
| return "Y position bin" |
| return "Position bin" |
| return { |
| "monkey": "Reach direction", |
| "allen_neuropixels": "Orientation", |
| "speech": "Cue", |
| }.get(dataset, "Condition") |
|
|
|
|
| def latent_color_label(dataset: str, color_mode: str, value: object) -> str: |
| if pd.isna(value): |
| return "Unknown" |
| if dataset == "ratinabox": |
| try: |
| idx = int(float(value)) |
| except ValueError: |
| return str(value) |
| if color_mode == "x": |
| return f"x{idx}" |
| if color_mode == "y": |
| return f"y{idx}" |
| return condition_label(dataset, value) |
|
|
|
|
| def add_latent_color_columns(df: pd.DataFrame, dataset: str, color_mode: str) -> pd.DataFrame: |
| out = df.copy() |
| condition_num = pd.to_numeric(out["condition"], errors="coerce") |
| if condition_num.isna().any() or (condition_num < 0).any(): |
| raise ValueError("Latent samples contain missing condition labels.") |
|
|
| if dataset == "ratinabox" and color_mode == "x": |
| values = (condition_num.astype(int) % 10).astype(str) |
| elif dataset == "ratinabox" and color_mode == "y": |
| values = (condition_num.astype(int) // 10).astype(str) |
| else: |
| values = condition_num.astype(int).astype(str) |
|
|
| out["color_value"] = values |
| out["color_num"] = pd.to_numeric(values, errors="coerce") |
| out["color_label"] = out["color_value"].map(lambda value: latent_color_label(dataset, color_mode, value)) |
| return out |
|
|
|
|
| def session_display_label(dataset: str, session: object) -> str: |
| text = "" if pd.isna(session) else str(session) |
| if dataset == "monkey": |
| match = re.search(r"sub-([A-Za-z])_ses-CO-(\d{4})(\d{2})(\d{2})", text) |
| if match: |
| monkey, year, month, day = match.groups() |
| return f"Monkey {monkey}, {year}-{month}-{day}" |
| if dataset == "ratinabox": |
| match = re.search(r"s(\d+)", text) |
| if match: |
| return f"Run {match.group(1)}" |
| if dataset == "speech": |
| return f"Session {text}" |
| if dataset == "allen_neuropixels": |
| return f"Session {text}" |
| return text or "Session" |
|
|
|
|
| def add_method_columns(df: pd.DataFrame) -> pd.DataFrame: |
| out = df.copy() |
| if "model" in out.columns: |
| out["method"] = out["model"].map(model_label) |
| out["model_order"] = out["model"].map(model_sort_value) |
| return out |
|
|
|
|
| def round_numeric(df: pd.DataFrame, columns: Iterable[str], digits: int = 3) -> pd.DataFrame: |
| out = df.copy() |
| for col in columns: |
| if col in out.columns: |
| out[col] = pd.to_numeric(out[col], errors="coerce").round(digits) |
| return out |
|
|
|
|
| def sort_table(df: pd.DataFrame, sort_by: list[dict] | None, default: list[tuple[str, bool]]) -> pd.DataFrame: |
| if sort_by: |
| sort_spec = [] |
| for item in sort_by: |
| col = item.get("column_id") |
| if col in df.columns: |
| sort_spec.append((col, item.get("direction") == "asc")) |
| if sort_spec: |
| return df.sort_values( |
| [col for col, _ in sort_spec], |
| ascending=[ascending for _, ascending in sort_spec], |
| na_position="last", |
| ) |
| return df.sort_values( |
| [col for col, _ in default], |
| ascending=[ascending for _, ascending in default], |
| na_position="last", |
| ) |
|
|
|
|
| def supported_sort(sort_by: list[dict] | None, columns: Iterable[str]) -> list[dict] | None: |
| allowed = set(columns) |
| if not sort_by: |
| return None |
| filtered = [item for item in sort_by if item.get("column_id") in allowed] |
| return filtered or None |
|
|
|
|
| def records(df: pd.DataFrame) -> list[dict]: |
| clean = df.astype(object).where(pd.notna(df), None) |
| return clean.to_dict("records") |
|
|
|
|
| def column_defs(columns: Iterable[str]) -> list[dict]: |
| out = [] |
| for col in columns: |
| item = {"name": TABLE_LABELS.get(col, col), "id": col} |
| if col in RIGHT_ALIGNED_COLUMNS: |
| item["type"] = "numeric" |
| out.append(item) |
| return out |
|
|
|
|
| def metric_text(value: object) -> str: |
| if pd.isna(value): |
| return "score" |
| return str(value) |
|
|
|
|
| def decoding_label(metric: object) -> str: |
| text = metric_text(metric) |
| if text.lower() == "r2": |
| return "Decoding R2" |
| if text.lower() == "accuracy": |
| return "Decoding accuracy" |
| return f"Decoding {text}" |
|
|
|
|
| def leaderboard_columns(columns: Iterable[str], metric: object) -> list[dict]: |
| defs = column_defs(columns) |
| label = decoding_label(metric) |
| for item in defs: |
| if item["id"] == "task_score": |
| item["name"] = label |
| return defs |
|
|
|
|
| def sort_status( |
| sort_by: list[dict] | None, |
| default: tuple[str, bool], |
| labels: dict[str, str] | None = None, |
| ) -> str: |
| column, ascending = default |
| if sort_by: |
| for item in sort_by: |
| candidate = item.get("column_id") |
| if candidate: |
| column = candidate |
| ascending = item.get("direction") == "asc" |
| break |
|
|
| merged_labels = dict(TABLE_LABELS) |
| if labels: |
| merged_labels.update(labels) |
| label = merged_labels.get(column, column) |
|
|
| if column in {"method", "family", "hardware"}: |
| direction = "A to Z" if ascending else "Z to A" |
| else: |
| direction = "low to high" if ascending else "high to low" |
| return f"Sorted by {label}: {direction}" |
|
|
|
|
| def value_text(value: object, digits: int = 3) -> str: |
| if value is None or pd.isna(value): |
| return "Not available" |
| number = float(value) |
| if abs(number) >= 1000: |
| return f"{number:,.0f}" |
| return f"{number:.{digits}f}" |
|
|
|
|
| def fig_layout(fig: go.Figure, *, height: int = 420) -> go.Figure: |
| fig.update_layout( |
| height=height, |
| paper_bgcolor="#ffffff", |
| plot_bgcolor="#ffffff", |
| margin=dict(l=22, r=22, t=52, b=38), |
| font=dict(family="Inter, Arial, sans-serif", size=13, color="#17202a"), |
| title=dict(font=dict(size=15, color="#17202a")), |
| legend=dict(orientation="h", yanchor="bottom", y=1.02, xanchor="left", x=0), |
| ) |
| fig.update_xaxes(showgrid=True, gridcolor="#edf1f4", zerolinecolor="#d7e0e6") |
| fig.update_yaxes(showgrid=True, gridcolor="#edf1f4", zerolinecolor="#d7e0e6") |
| return fig |
|
|
|
|
| def empty_figure(message: str) -> go.Figure: |
| fig = go.Figure() |
| fig.add_annotation( |
| text=message, |
| x=0.5, |
| y=0.5, |
| xref="paper", |
| yref="paper", |
| showarrow=False, |
| font=dict(size=15, color="#637381"), |
| ) |
| fig.update_xaxes(visible=False) |
| fig.update_yaxes(visible=False) |
| return fig_layout(fig, height=360) |
|
|
|
|
| def parse_float_list(value: object) -> list[float]: |
| if pd.isna(value): |
| return [] |
| out = [] |
| for part in str(value).split(";"): |
| part = part.strip() |
| if not part: |
| continue |
| try: |
| out.append(float(part)) |
| except ValueError: |
| continue |
| return out |
|
|
|
|
| def dataframe_table( |
| table_id: str, |
| *, |
| page_size: int = 8, |
| max_height: str = "520px", |
| sort_action: str = "native", |
| sort_by: list[dict] | None = None, |
| ) -> dash_table.DataTable: |
| return dash_table.DataTable( |
| id=table_id, |
| columns=[], |
| data=[], |
| page_size=page_size, |
| sort_action=sort_action, |
| sort_mode="single", |
| sort_by=sort_by or [], |
| cell_selectable=True, |
| style_as_list_view=True, |
| fixed_rows={"headers": True}, |
| style_table={"overflowX": "auto", "overflowY": "auto", "maxHeight": max_height}, |
| style_header={ |
| "backgroundColor": "#f3f6f8", |
| "fontWeight": "700", |
| "border": "0", |
| "borderBottom": "1px solid #cfd8df", |
| "color": "#26323f", |
| }, |
| style_cell={ |
| "fontFamily": "Inter, Arial, sans-serif", |
| "fontSize": "13px", |
| "padding": "10px 12px", |
| "textAlign": "left", |
| "minWidth": "90px", |
| "maxWidth": "260px", |
| "whiteSpace": "normal", |
| "height": "auto", |
| "border": "0", |
| "borderBottom": "1px solid #edf1f4", |
| }, |
| style_cell_conditional=[ |
| {"if": {"column_id": col}, "textAlign": "right"} for col in RIGHT_ALIGNED_COLUMNS |
| ], |
| style_data_conditional=[ |
| {"if": {"row_index": "odd"}, "backgroundColor": "#fbfcfd"}, |
| {"if": {"state": "active"}, "backgroundColor": "#e5f3f2", "border": "1px solid #4c908b"}, |
| ], |
| ) |
|
|
|
|
| def panel(title: str, *children, subtitle: str | None = None, className: str = "") -> html.Div: |
| heading = [html.H2(title)] |
| if subtitle: |
| heading.append(html.P(subtitle, className="panel-subtitle")) |
| classes = "panel" if not className else f"panel {className}" |
| return html.Div([html.Div(heading, className="panel-heading"), *children], className=classes) |
|
|
|
|
| def details_table(summary: str, table: dash_table.DataTable) -> html.Details: |
| return html.Details( |
| [html.Summary(summary), html.Div(table, className="details-body")], |
| className="details-table", |
| ) |
|
|
|
|
| def metric_card(label: str, value: str, detail: str | None = None) -> html.Div: |
| children = [ |
| html.Div(label, className="metric-label"), |
| html.Div(value, className="metric-value"), |
| ] |
| if detail: |
| children.append(html.Div(detail, className="metric-detail")) |
| return html.Div(children, className="metric-card") |
|
|
|
|
| def leaderboard_frame(dataset: str, models: list[str] | None) -> pd.DataFrame: |
| chosen = selected_models(models) |
| base = pd.DataFrame({"model": chosen}) |
| base = add_method_columns(base) |
|
|
| pred = filter_models(prediction, chosen) |
| pred = pred[pred["dataset"].astype(str) == str(dataset)].copy() |
| pred["score"] = pd.to_numeric(pred["score"], errors="coerce") |
| pred = pred[["model", "metric", "score", "n_train_trials", "n_test_trials", "n_neurons", "latent_dim"]] |
| pred = pred.rename(columns={"score": "task_score"}) |
|
|
| rob = filter_models(present_rows(robustness), chosen) |
| rob = rob[rob["dataset"].astype(str) == str(dataset)].copy() |
| rob["raw_auc"] = pd.to_numeric(rob["raw_auc"], errors="coerce") |
| rob = rob[["model", "raw_auc"]].rename(columns={"raw_auc": "robustness_auc"}) |
|
|
| cons = filter_models(active_rows(consistency), chosen) |
| cons = cons[cons["dataset"].astype(str) == str(dataset)].copy() |
| cons["mean_r2"] = pd.to_numeric(cons["mean_r2"], errors="coerce") |
| cons = cons[["model", "mean_r2"]].rename(columns={"mean_r2": "alignment_score"}) |
|
|
| scale = filter_models(present_rows(scalability), chosen) |
| scale = scale[scale["dataset"].astype(str) == str(dataset)].copy() |
| scale = scale[["model", "training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb"]] |
| for col in ["training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb"]: |
| scale[col] = pd.to_numeric(scale[col], errors="coerce") |
|
|
| df = base.merge(pred, on="model", how="left") |
| df = df.merge(rob, on="model", how="left") |
| df = df.merge(cons, on="model", how="left") |
| df = df.merge(scale, on="model", how="left") |
| available = df["task_score"].notna() |
| order = df.loc[available].sort_values( |
| ["task_score", "model_order"], ascending=[False, True] |
| ).index |
| df["rank"] = None |
| for rank, idx in enumerate(order, start=1): |
| df.at[idx, "rank"] = rank |
|
|
| df["id"] = df["model"] |
| return round_numeric(df, NUMERIC_COLUMNS) |
|
|
|
|
| def leaderboard_summary(dataset: str, table_df: pd.DataFrame) -> list[html.Div]: |
| label = DATASET_LABELS.get(dataset, dataset) |
| metric = metric_text(table_df["metric"].dropna().iloc[0]) if table_df["metric"].notna().any() else "score" |
| available = sort_table( |
| table_df, None, [("task_score", False), ("model_order", True)] |
| ).dropna(subset=["task_score"]).head(3) |
| cards = [metric_card("Dataset", label, f"Primary metric: {metric}")] |
| for _, row in available.iterrows(): |
| cards.append( |
| metric_card( |
| f"Rank {int(row['rank'])}", |
| str(row["method"]), |
| f"{value_text(row['task_score'])} {metric}", |
| ) |
| ) |
| if len(cards) == 1: |
| cards.append(metric_card("Top method", "Not available")) |
| return cards |
|
|
|
|
| def performance_heatmap(dataset: str, models: list[str] | None) -> go.Figure: |
| df = filter_models(prediction, models) |
| if df.empty: |
| return empty_figure("No decoding results are available.") |
| df = add_method_columns(df) |
| df["score"] = pd.to_numeric(df["score"], errors="coerce") |
| df["dataset_label"] = df["dataset"].map(DATASET_LABELS).fillna(df["dataset"]) |
| pivot = df.pivot_table( |
| index="method", columns="dataset_label", values="score", aggfunc="first" |
| ) |
|
|
| selected_label = DATASET_LABELS.get(dataset, dataset) |
| order_df = sort_table( |
| leaderboard_frame(dataset, models), |
| None, |
| [("task_score", False), ("model_order", True)], |
| ) |
| method_order = [m for m in order_df["method"] if m in set(pivot.index)] |
| pivot = pivot.reindex(method_order) |
| dataset_order = [DATASET_LABELS.get(ds, ds) for ds in DATASETS] |
| pivot = pivot.reindex(columns=[label for label in dataset_order if label in pivot.columns]) |
| text = pivot.map(lambda x: "" if pd.isna(x) else f"{x:.3f}") if not pivot.empty else pivot |
|
|
| zmin = min(0.0, float(np.nanmin(pivot.values))) if pivot.size and not np.isnan(pivot.values).all() else 0.0 |
| zmax = max(1.0, float(np.nanmax(pivot.values))) if pivot.size and not np.isnan(pivot.values).all() else 1.0 |
| fig = go.Figure( |
| go.Heatmap( |
| z=pivot.values if not pivot.empty else [[]], |
| x=list(pivot.columns), |
| y=list(pivot.index), |
| text=text.values if not pivot.empty else [[]], |
| texttemplate="%{text}", |
| colorscale=SCORE_SCALE, |
| zmin=zmin, |
| zmax=zmax, |
| colorbar=dict(title="Score", thickness=12), |
| hovertemplate="Method=%{y}<br>Dataset=%{x}<br>Score=%{z:.4f}<extra></extra>", |
| ) |
| ) |
| fig.update_layout(title=f"Decoding score matrix, sorted by {selected_label}") |
| return fig_layout(fig, height=max(430, 26 * len(pivot.index) + 150)) |
|
|
|
|
| def ranking_figure(dataset: str, table_df: pd.DataFrame) -> go.Figure: |
| rank_df = table_df.dropna(subset=["task_score"]).sort_values("task_score", ascending=True) |
| if rank_df.empty: |
| return empty_figure(f"No decoding results are available for {DATASET_LABELS.get(dataset, dataset)}.") |
| metric = metric_text(rank_df["metric"].dropna().iloc[0]) if rank_df["metric"].notna().any() else "score" |
| label = decoding_label(metric) |
| fig = go.Figure( |
| go.Bar( |
| x=rank_df["task_score"], |
| y=rank_df["method"], |
| orientation="h", |
| marker=dict(color=TASK_COLOR), |
| hovertemplate="Method=%{y}<br>Score=%{x:.4f}<extra></extra>", |
| ) |
| ) |
| fig.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} {label} ranking") |
| fig.update_xaxes(title=label) |
| fig.update_yaxes(title="") |
| return fig_layout(fig, height=max(420, 25 * len(rank_df) + 150)) |
|
|
|
|
| def consistency_frame(dataset: str, models: list[str] | None) -> pd.DataFrame: |
| chosen = selected_models(models) |
| base = add_method_columns(pd.DataFrame({"model": chosen})) |
|
|
| cons = filter_models(active_rows(consistency), chosen) |
| cons = cons[cons["dataset"].astype(str) == str(dataset)].copy() |
| for col in ["mean_r2", "latent_dim", "n_sessions", "n_pairwise"]: |
| cons[col] = pd.to_numeric(cons[col], errors="coerce") |
| cons = cons[["model", "mean_r2", "latent_dim", "n_sessions", "n_pairwise"]] |
| cons = cons.rename(columns={"mean_r2": "alignment_score"}) |
|
|
| df = base.merge(cons, on="model", how="left") |
| available = df["alignment_score"].notna() |
| order = df.loc[available].sort_values( |
| ["alignment_score", "model_order"], ascending=[False, True] |
| ).index |
| df["rank"] = None |
| for rank, idx in enumerate(order, start=1): |
| df.at[idx, "rank"] = rank |
| df["id"] = df["model"] |
| return round_numeric(df, NUMERIC_COLUMNS) |
|
|
|
|
| def selected_consistency_model(df: pd.DataFrame, active_cell: dict | None) -> str | None: |
| if active_cell and active_cell.get("row_id") in set(df["model"]): |
| return str(active_cell["row_id"]) |
| available = sort_table( |
| df, None, [("alignment_score", False), ("model_order", True)] |
| ).dropna(subset=["alignment_score"]) |
| if available.empty: |
| return None |
| return str(available.iloc[0]["model"]) |
|
|
|
|
| def consistency_bar_figure(dataset: str, df: pd.DataFrame) -> go.Figure: |
| bar_df = df.dropna(subset=["alignment_score"]).sort_values("alignment_score", ascending=True) |
| if bar_df.empty: |
| return empty_figure(f"No cross-session alignment results are available for {DATASET_LABELS.get(dataset, dataset)}.") |
| fig = go.Figure( |
| go.Bar( |
| x=bar_df["alignment_score"], |
| y=bar_df["method"], |
| orientation="h", |
| marker=dict(color=ALIGNMENT_COLOR), |
| hovertemplate="Method=%{y}<br>Alignment=%{x:.4f}<extra></extra>", |
| ) |
| ) |
| fig.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} cross-session alignment") |
| fig.update_xaxes(title="Alignment score") |
| fig.update_yaxes(title="") |
| return fig_layout(fig, height=max(360, 25 * len(bar_df) + 130)) |
|
|
|
|
| def consistency_heatmap(models: list[str] | None) -> go.Figure: |
| df = filter_models(active_rows(consistency), models) |
| if df.empty: |
| return empty_figure("No cross-session alignment results are available.") |
| df = add_method_columns(df) |
| df["mean_r2"] = pd.to_numeric(df["mean_r2"], errors="coerce") |
| df["dataset_label"] = df["dataset"].map(DATASET_LABELS).fillna(df["dataset"]) |
| pivot = df.pivot_table(index="method", columns="dataset_label", values="mean_r2", aggfunc="first") |
| if not pivot.empty: |
| pivot = pivot.loc[pivot.mean(axis=1, skipna=True).sort_values(ascending=False).index] |
| text = pivot.map(lambda x: "" if pd.isna(x) else f"{x:.2f}") if not pivot.empty else pivot |
| fig = go.Figure( |
| go.Heatmap( |
| z=pivot.values if not pivot.empty else [[]], |
| x=list(pivot.columns), |
| y=list(pivot.index), |
| text=text.values if not pivot.empty else [[]], |
| texttemplate="%{text}", |
| colorscale=ALIGNMENT_SCALE, |
| colorbar=dict(title="Score", thickness=12), |
| hovertemplate="Method=%{y}<br>Dataset=%{x}<br>Alignment=%{z:.4f}<extra></extra>", |
| ) |
| ) |
| fig.update_layout(title="Alignment across datasets") |
| return fig_layout(fig, height=max(360, 25 * len(pivot.index) + 135)) |
|
|
|
|
| def latent_space_figure(dataset: str, model: str | None, color_mode: str = "condition") -> go.Figure: |
| if not model: |
| return empty_figure("No latent-space view is available for this selection.") |
| if latent_samples.empty: |
| return empty_figure("Latent-space samples are unavailable in this view.") |
|
|
| plot_df = latent_samples[ |
| (latent_samples["dataset"].astype(str) == str(dataset)) |
| & (latent_samples["model"].astype(str) == str(model)) |
| ].copy() |
| if plot_df.empty: |
| return empty_figure(f"No latent-space samples are available for {model_label(model)} on {DATASET_LABELS.get(dataset, dataset)}.") |
|
|
| for col in ["x", "y", "z"]: |
| plot_df[col] = pd.to_numeric(plot_df[col], errors="coerce") |
| try: |
| plot_df = add_latent_color_columns(plot_df, dataset, color_mode) |
| except ValueError as exc: |
| return empty_figure(str(exc)) |
| plot_df["session_display"] = plot_df["session_label"].map(lambda value: session_display_label(dataset, value)) |
| plot_df = plot_df.dropna(subset=["x", "y", "z"]) |
| if plot_df.empty: |
| return empty_figure("No latent-space samples match this selection.") |
|
|
| trajectory_df = latent_trajectories[ |
| (latent_trajectories["dataset"].astype(str) == str(dataset)) |
| & (latent_trajectories["model"].astype(str) == str(model)) |
| ].copy() |
| for col in ["x", "y", "z"]: |
| if col in trajectory_df: |
| trajectory_df[col] = pd.to_numeric(trajectory_df[col], errors="coerce") |
| if not trajectory_df.empty: |
| try: |
| trajectory_df = add_latent_color_columns(trajectory_df, dataset, color_mode) |
| except ValueError: |
| trajectory_df = trajectory_df.iloc[0:0].copy() |
| trajectory_df["session_display"] = trajectory_df["session_label"].map(lambda value: session_display_label(dataset, value)) |
| trajectory_df = trajectory_df.dropna(subset=["x", "y", "z"]) |
| if "color_value" not in trajectory_df.columns: |
| trajectory_df["color_value"] = pd.Series(dtype=str) |
|
|
| sessions = ordered_unique(plot_df["session_label"]) |
| session_titles = [session_display_label(dataset, session) for session in sessions] |
| n_cols = 2 if len(sessions) > 1 else 1 |
| n_rows = int(np.ceil(len(sessions) / n_cols)) |
| specs = [[{"type": "scene"} for _ in range(n_cols)] for _ in range(n_rows)] |
| fig = make_subplots( |
| rows=n_rows, |
| cols=n_cols, |
| specs=specs, |
| subplot_titles=session_titles, |
| horizontal_spacing=0.045, |
| vertical_spacing=0.12, |
| ) |
|
|
| condition_values = sorted(plot_df["color_value"].astype(str).unique(), key=condition_sort_key) |
| use_categorical = len(condition_values) <= 12 and dataset != "ratinabox" |
| if dataset == "monkey": |
| condition_colors = { |
| condition: DIRECTION_PALETTE[int(float(condition)) % len(DIRECTION_PALETTE)] |
| for condition in condition_values |
| } |
| elif dataset == "speech": |
| condition_colors = { |
| condition: SPEECH_PALETTE.get(condition, CATEGORICAL_PALETTE[idx % len(CATEGORICAL_PALETTE)]) |
| for idx, condition in enumerate(condition_values) |
| } |
| else: |
| condition_colors = { |
| condition: CATEGORICAL_PALETTE[idx % len(CATEGORICAL_PALETTE)] |
| for idx, condition in enumerate(condition_values) |
| } |
| condition_name = condition_axis_label(dataset, color_mode) |
|
|
| for session_idx, session in enumerate(sessions): |
| session_df = plot_df[plot_df["session_label"].astype(str) == str(session)] |
| if session_df.empty: |
| continue |
| row = session_idx // n_cols + 1 |
| col = session_idx % n_cols + 1 |
| display_session = session_display_label(dataset, session) |
|
|
| if use_categorical: |
| for condition in condition_values: |
| cond_df = session_df[session_df["color_value"].astype(str) == condition] |
| if cond_df.empty: |
| continue |
| trace_name = latent_color_label(dataset, color_mode, condition) |
| session_traj = trajectory_df[ |
| (trajectory_df["session_label"].astype(str) == str(session)) |
| & (trajectory_df["color_value"].astype(str) == condition) |
| ].sort_values("time_index") |
| fig.add_trace( |
| go.Scatter3d( |
| x=cond_df["x"], |
| y=cond_df["y"], |
| z=cond_df["z"], |
| mode="markers", |
| name=trace_name, |
| legendgroup=condition, |
| showlegend=session_idx == 0, |
| marker=dict( |
| size=2.4 if not session_traj.empty else 3.0, |
| opacity=0.32 if not session_traj.empty else 0.78, |
| color=condition_colors[condition], |
| ), |
| customdata=np.stack( |
| [ |
| np.repeat(display_session, len(cond_df)), |
| cond_df["color_label"].astype(str), |
| cond_df["trial_index"].astype(str), |
| cond_df["time_index"].astype(str), |
| ], |
| axis=-1, |
| ), |
| hovertemplate=( |
| "Session=%{customdata[0]}<br>" |
| f"{condition_name}=%{{customdata[1]}}<br>" |
| "Trial=%{customdata[2]} time=%{customdata[3]}" |
| "<extra></extra>" |
| ), |
| ), |
| row=row, |
| col=col, |
| ) |
| if not session_traj.empty: |
| fig.add_trace( |
| go.Scatter3d( |
| x=session_traj["x"], |
| y=session_traj["y"], |
| z=session_traj["z"], |
| mode="lines", |
| name=trace_name, |
| legendgroup=condition, |
| showlegend=False, |
| line=dict(color=condition_colors[condition], width=5), |
| hovertemplate=( |
| f"{condition_name}={trace_name}<br>" |
| "Time=%{customdata}<extra></extra>" |
| ), |
| customdata=session_traj["time_index"], |
| ), |
| row=row, |
| col=col, |
| ) |
| else: |
| fig.add_trace( |
| go.Scatter3d( |
| x=session_df["x"], |
| y=session_df["y"], |
| z=session_df["z"], |
| mode="markers", |
| name=display_session, |
| showlegend=False, |
| marker=dict( |
| size=2.8, |
| opacity=0.72, |
| color=session_df["color_num"], |
| colorscale=RATINABOX_SCALE, |
| cmin=0, |
| cmax=9 if color_mode in {"x", "y"} else 99, |
| showscale=session_idx == 0, |
| colorbar=dict(title=condition_name, thickness=12), |
| ), |
| customdata=np.stack( |
| [ |
| np.repeat(display_session, len(session_df)), |
| session_df["color_label"].astype(str), |
| session_df["trial_index"].astype(str), |
| session_df["time_index"].astype(str), |
| ], |
| axis=-1, |
| ), |
| hovertemplate=( |
| "Session=%{customdata[0]}<br>" |
| f"{condition_name}=%{{customdata[1]}}<br>" |
| "Trial=%{customdata[2]} time=%{customdata[3]}" |
| "<extra></extra>" |
| ), |
| ), |
| row=row, |
| col=col, |
| ) |
|
|
| extent = max( |
| float(np.nanpercentile(np.abs(plot_df["x"]), 99)), |
| float(np.nanpercentile(np.abs(plot_df["y"]), 99)), |
| float(np.nanpercentile(np.abs(plot_df["z"]), 99)), |
| 1.0, |
| ) |
| lim = extent * 1.08 |
| for idx in range(len(sessions)): |
| scene_id = "scene" if idx == 0 else f"scene{idx + 1}" |
| fig.update_layout( |
| **{ |
| scene_id: dict( |
| xaxis=dict(title="", range=[-lim, lim], showgrid=False, zeroline=False, showticklabels=False), |
| yaxis=dict(title="", range=[-lim, lim], showgrid=False, zeroline=False, showticklabels=False), |
| zaxis=dict(title="", range=[-lim, lim], showgrid=False, zeroline=False, showticklabels=False), |
| aspectmode="cube", |
| bgcolor="#ffffff", |
| camera=dict(eye=dict(x=1.55, y=1.45, z=1.05)), |
| ) |
| } |
| ) |
|
|
| score_df = active_rows(consistency) |
| score_df = score_df[ |
| (score_df["dataset"].astype(str) == str(dataset)) |
| & (score_df["model"].astype(str) == str(model)) |
| ].copy() |
| score = None |
| if not score_df.empty: |
| score = pd.to_numeric(score_df["mean_r2"], errors="coerce").dropna() |
| score = float(score.iloc[0]) if not score.empty else None |
| score_suffix = "" if score is None else f" | alignment {score:.2f}" |
| fig.update_layout( |
| title=f"{model_label(model)} latent space on {DATASET_LABELS.get(dataset, dataset)}{score_suffix}", |
| height=760 if n_rows > 1 else 520, |
| paper_bgcolor="#ffffff", |
| plot_bgcolor="#ffffff", |
| margin=dict(l=8, r=8, t=78, b=92), |
| font=dict(family="Inter, Arial, sans-serif", size=13, color="#17202a"), |
| legend=dict(orientation="h", yanchor="top", y=-0.08, xanchor="left", x=0, title=condition_name), |
| ) |
| fig.for_each_annotation(lambda ann: ann.update(font=dict(size=12, color="#526171"))) |
| return fig |
|
|
|
|
| def robustness_figure(dataset: str, models: list[str] | None, hover_data: dict | None = None) -> go.Figure: |
| df = filter_models(present_rows(robustness), models) |
| df = df[df["dataset"].astype(str) == str(dataset)].copy() if not df.empty else df |
| if df.empty: |
| return empty_figure("No robustness results are available for this selection.") |
| df = add_method_columns(df).sort_values("model_order") |
| highlighted_trace = None |
| if hover_data and hover_data.get("points"): |
| highlighted_trace = hover_data["points"][0].get("curveNumber") |
|
|
| fig = go.Figure() |
| trace_idx = 0 |
| for idx, row in enumerate(df.itertuples()): |
| xs = parse_float_list(row.noise_levels) |
| ys = parse_float_list(row.scores) |
| if xs and len(xs) == len(ys): |
| is_highlighted = highlighted_trace is None or highlighted_trace == trace_idx |
| fig.add_trace( |
| go.Scatter( |
| x=xs, |
| y=ys, |
| mode="lines+markers", |
| name=row.method, |
| opacity=0.92 if is_highlighted else 0.16, |
| line=dict( |
| color=CATEGORICAL_PALETTE[idx % len(CATEGORICAL_PALETTE)], |
| width=4 if highlighted_trace == trace_idx else 2, |
| ), |
| marker=dict(size=8 if highlighted_trace == trace_idx else 5), |
| hovertemplate="Noise=%{x:.2f}<br>Score=%{y:.4f}<extra></extra>", |
| ) |
| ) |
| trace_idx += 1 |
| fig.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} robustness curves") |
| fig.update_layout(hovermode="closest") |
| fig.update_xaxes(title="Noise fraction") |
| metric = metric_text(df["metric"].dropna().iloc[0]) if df["metric"].notna().any() else "score" |
| fig.update_yaxes(title=decoding_label(metric)) |
| fig_layout(fig, height=560) |
| fig.update_layout( |
| margin=dict(l=58, r=170, t=64, b=48), |
| legend=dict( |
| orientation="v", |
| yanchor="top", |
| y=1, |
| xanchor="left", |
| x=1.02, |
| font=dict(size=11), |
| title=None, |
| ), |
| ) |
| return fig |
|
|
|
|
| def robustness_table_frame(dataset: str, models: list[str] | None) -> pd.DataFrame: |
| df = filter_models(present_rows(robustness), models) |
| df = df[df["dataset"].astype(str) == str(dataset)].copy() if not df.empty else df |
| if df.empty: |
| return pd.DataFrame(columns=["method", "reference_score", "highest_noise_score", "robustness_auc", "average_noisy_score"]) |
| df = add_method_columns(df) |
| df = df.rename( |
| columns={ |
| "score_at_noise0": "reference_score", |
| "score_at_max_noise": "highest_noise_score", |
| "raw_auc": "robustness_auc", |
| "mean_score": "average_noisy_score", |
| } |
| ) |
| cols = ["method", "reference_score", "highest_noise_score", "robustness_auc", "average_noisy_score"] |
| return round_numeric(df[cols].sort_values("robustness_auc", ascending=False), NUMERIC_COLUMNS) |
|
|
|
|
| def compute_figures(dataset: str, models: list[str] | None) -> tuple[go.Figure, go.Figure, pd.DataFrame]: |
| df = filter_models(present_rows(scalability), models) |
| df = df[df["dataset"].astype(str) == str(dataset)].copy() if not df.empty else df |
| if df.empty: |
| empty = pd.DataFrame(columns=["method", "hardware", "training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb"]) |
| return ( |
| empty_figure("No compute results are available for this selection."), |
| empty_figure("No memory results are available for this selection."), |
| empty, |
| ) |
|
|
| pred = present_rows(prediction)[["model", "dataset", "score", "metric"]].copy() |
| df = df.merge(pred, on=["model", "dataset"], how="left") |
| df = add_method_columns(df) |
| df["hardware"] = df["model"].map(hardware_label) |
| for col in ["training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb", "score"]: |
| df[col] = pd.to_numeric(df[col], errors="coerce") |
|
|
| scatter = go.Figure( |
| go.Scatter( |
| x=df["training_time_sec"], |
| y=df["score"], |
| mode="markers", |
| text=df["method"], |
| marker=dict( |
| size=np.clip(df["peak_ram_gb"].fillna(1.0) * 4, 8, 26), |
| color=df["hardware"].map({"CPU": "#E69F00", "GPU": TASK_COLOR}).fillna("#637381"), |
| opacity=0.82, |
| line=dict(color="#ffffff", width=1), |
| ), |
| customdata=np.stack( |
| [ |
| df["method"].astype(str), |
| df["hardware"].astype(str), |
| df["peak_ram_gb"].round(3).astype(str), |
| df["peak_vram_gb"].round(3).astype(str), |
| ], |
| axis=-1, |
| ), |
| hovertemplate=( |
| "Method=%{customdata[0]}<br>" |
| "Hardware=%{customdata[1]}<br>" |
| "Training time=%{x:.3f} s<br>" |
| "Score=%{y:.4f}<br>" |
| "Peak memory=%{customdata[2]} GB<br>" |
| "Peak GPU memory=%{customdata[3]} GB" |
| "<extra></extra>" |
| ), |
| ) |
| ) |
| scatter.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} performance and training time") |
| scatter.update_xaxes(title="Training time (s, log scale)", type="log") |
| scatter.update_yaxes(title=metric_text(df["metric"].dropna().iloc[0]) if df["metric"].notna().any() else "score") |
| fig_layout(scatter, height=500) |
|
|
| mem_df = df.sort_values("peak_ram_gb", ascending=True) |
| memory = go.Figure() |
| memory.add_trace(go.Bar(x=mem_df["peak_ram_gb"], y=mem_df["method"], orientation="h", name="RAM")) |
| memory.add_trace(go.Bar(x=mem_df["peak_vram_gb"], y=mem_df["method"], orientation="h", name="GPU memory")) |
| memory.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} peak memory", barmode="group") |
| memory.update_xaxes(title="GB") |
| memory.update_yaxes(title="") |
| fig_layout(memory, height=max(420, 26 * len(mem_df) + 150)) |
|
|
| table = df[["method", "hardware", "score", "training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb"]] |
| return scatter, memory, round_numeric(table.sort_values("training_time_sec"), NUMERIC_COLUMNS) |
|
|
|
|
| def influence_figures(dataset: str, models: list[str] | None) -> tuple[go.Figure, go.Figure, pd.DataFrame]: |
| nshap = filter_models(active_rows(neuron_shap), models) |
| nshap = nshap[nshap["dataset"].astype(str) == str(dataset)].copy() if not nshap.empty else nshap |
| if nshap.empty: |
| neuron_fig = empty_figure("No neuron-influence results are available for this dataset.") |
| table = pd.DataFrame(columns=["method", "baseline_score", "full_model_score", "neuron_influence_auc", "shap_mean_value", "shap_fraction_positive"]) |
| else: |
| nshap = add_method_columns(nshap) |
| nshap["auc"] = pd.to_numeric(nshap["auc"], errors="coerce") |
| bar_df = nshap.dropna(subset=["auc"]).sort_values("auc", ascending=True) |
| neuron_fig = go.Figure( |
| go.Bar( |
| x=bar_df["auc"], |
| y=bar_df["method"], |
| orientation="h", |
| marker=dict(color=INFLUENCE_COLOR), |
| hovertemplate="Method=%{y}<br>AUC=%{x:.4f}<extra></extra>", |
| ) |
| ) |
| neuron_fig.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} neuron influence") |
| neuron_fig.update_xaxes(title="Neuron influence AUC") |
| neuron_fig.update_yaxes(title="") |
| fig_layout(neuron_fig, height=max(420, 25 * len(bar_df) + 150)) |
| table = nshap.rename(columns={"auc": "neuron_influence_auc"}) |
| table = table[ |
| ["method", "baseline_score", "full_model_score", "neuron_influence_auc", "shap_mean_value", "shap_fraction_positive"] |
| ] |
|
|
| tshap = filter_models(active_rows(trial_shapley), models) |
| tshap = tshap[tshap["dataset"].astype(str) == str(dataset)].copy() if not tshap.empty else tshap |
| if tshap.empty: |
| trial_fig = empty_figure("No trial-influence results are available for this dataset.") |
| else: |
| tshap = add_method_columns(tshap) |
| tshap["perturbation_auc"] = pd.to_numeric(tshap["perturbation_auc"], errors="coerce") |
| trial_df = tshap.dropna(subset=["perturbation_auc"]).sort_values("perturbation_auc", ascending=True) |
| trial_fig = go.Figure( |
| go.Bar( |
| x=trial_df["perturbation_auc"], |
| y=trial_df["method"], |
| orientation="h", |
| marker=dict(color="#009E73"), |
| hovertemplate="Method=%{y}<br>AUC=%{x:.4f}<extra></extra>", |
| ) |
| ) |
| trial_fig.update_layout(title=f"{DATASET_LABELS.get(dataset, dataset)} trial influence") |
| trial_fig.update_xaxes(title="Trial influence AUC") |
| trial_fig.update_yaxes(title="") |
| fig_layout(trial_fig, height=max(420, 25 * len(trial_df) + 150)) |
|
|
| return neuron_fig, trial_fig, round_numeric(table, NUMERIC_COLUMNS) |
|
|
|
|
| def methods_frame(models: list[str] | None) -> pd.DataFrame: |
| chosen = selected_models(models) |
| rows = [] |
| for model in chosen: |
| rows.append( |
| { |
| "id": model, |
| "method": model_label(model), |
| "family": METHOD_FAMILY.get(model, "Model"), |
| "hardware": hardware_label(model), |
| } |
| ) |
| return pd.DataFrame(rows).sort_values("method") |
|
|
|
|
| app = Dash(__name__, title="Neural Model Benchmark") |
| server = app.server |
|
|
| app.layout = html.Div( |
| [ |
| html.Div( |
| [ |
| html.Div( |
| [ |
| html.Div("Tang Lab", className="eyebrow"), |
| html.H1("Neural Model Benchmark"), |
| html.P( |
| "Explore decoding performance, robustness, latent alignment, neuron and trial influence, and compute cost across 23 benchmarked methods.", |
| className="lede", |
| ), |
| ], |
| className="hero-copy", |
| ), |
| ], |
| className="hero", |
| ), |
| html.Div( |
| [ |
| html.Div( |
| [ |
| html.Label("Dataset"), |
| dcc.Dropdown( |
| id="dataset-filter", |
| options=[{"label": DATASET_LABELS.get(ds, ds), "value": ds} for ds in DATASETS], |
| value=DATASETS[0] if DATASETS else None, |
| clearable=False, |
| ), |
| ], |
| className="control", |
| ), |
| ], |
| className="toolbar", |
| ), |
| dcc.Tabs( |
| id="tabs", |
| value="leaderboard", |
| className="tabs", |
| children=[ |
| dcc.Tab( |
| label="Leaderboard", |
| value="leaderboard", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Leaderboard", |
| html.Div(id="leaderboard-top", className="top-strip"), |
| html.Div( |
| [ |
| html.Div( |
| [ |
| html.Div(id="leaderboard-sort-state", className="sort-state"), |
| dataframe_table("leaderboard-table", page_size=23, max_height="680px", sort_action="custom"), |
| ], |
| className="leaderboard-table-wrap", |
| ), |
| dcc.Graph(id="dataset-ranking", config={"displayModeBar": False}), |
| ], |
| className="leaderboard-grid", |
| ), |
| dcc.Graph(id="performance-heatmap", config={"displayModeBar": False}), |
| subtitle="Sorted by the selected dataset. Regression datasets use decoding R2; classification datasets use decoding accuracy. Higher is better.", |
| className="leaderboard-panel", |
| ) |
| ], |
| ), |
| dcc.Tab( |
| label="Consistency", |
| value="consistency", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Cross-session alignment", |
| html.Div( |
| [ |
| html.Div( |
| [ |
| html.Div(id="consistency-sort-state", className="sort-state"), |
| dataframe_table( |
| "consistency-table", |
| page_size=23, |
| max_height="520px", |
| sort_action="custom", |
| sort_by=[{"column_id": "alignment_score", "direction": "desc"}], |
| ), |
| ], |
| className="consistency-table-wrap", |
| ), |
| html.Div( |
| [ |
| html.Div( |
| [ |
| html.Label("Color by"), |
| dcc.Dropdown( |
| id="latent-color-mode", |
| clearable=False, |
| ), |
| ], |
| id="latent-color-control", |
| className="control latent-color-control", |
| ), |
| dcc.Graph( |
| id="latent-space", |
| config={ |
| "displayModeBar": "hover", |
| "toImageButtonOptions": { |
| "format": "png", |
| "filename": "benchdash_latent_space", |
| "height": 900, |
| "width": 1200, |
| "scale": 2, |
| }, |
| }, |
| ), |
| ], |
| className="latent-panel", |
| ), |
| ], |
| className="latent-grid", |
| ), |
| html.Div( |
| [ |
| dcc.Graph(id="consistency-bars", config={"displayModeBar": False}), |
| dcc.Graph(id="consistency-heatmap", config={"displayModeBar": False}), |
| ], |
| className="chart-grid two", |
| ), |
| subtitle="Each latent panel shows one recording session. Colors indicate task condition, stimulus, cue, or spatial bin.", |
| ) |
| ], |
| ), |
| dcc.Tab( |
| label="Robustness", |
| value="robustness", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Robustness", |
| dcc.Graph(id="robustness-curve", clear_on_unhover=True, config={"displayModeBar": False}), |
| details_table("View robustness rows", dataframe_table("robustness-table")), |
| subtitle="Hover a method to highlight its curve. Curves show how decoding performance changes as neural count noise increases.", |
| ) |
| ], |
| ), |
| dcc.Tab( |
| label="Influence", |
| value="influence", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Neuron and trial influence", |
| html.Div( |
| [ |
| dcc.Graph(id="neuron-influence-bars", config={"displayModeBar": False}), |
| dcc.Graph(id="trial-influence-bars", config={"displayModeBar": False}), |
| ], |
| className="chart-grid two", |
| ), |
| details_table("View neuron-influence rows", dataframe_table("influence-table")), |
| subtitle="Signed contribution summaries preserve whether neurons or trials helped or hurt prediction.", |
| ) |
| ], |
| ), |
| dcc.Tab( |
| label="Compute", |
| value="compute", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Performance and compute cost", |
| html.Div( |
| [ |
| dcc.Graph(id="compute-scatter", config={"displayModeBar": False}), |
| dcc.Graph(id="memory-bars", config={"displayModeBar": False}), |
| ], |
| className="chart-grid two", |
| ), |
| details_table("View compute rows", dataframe_table("compute-table")), |
| subtitle="Shows training time and peak memory for each method on the selected dataset.", |
| ) |
| ], |
| ), |
| dcc.Tab( |
| label="Methods", |
| value="methods", |
| className="tab", |
| selected_className="tab tab-selected", |
| children=[ |
| panel( |
| "Benchmarked methods", |
| dataframe_table("methods-table", page_size=23, max_height="600px"), |
| subtitle="Paper-facing method names and broad method families.", |
| ) |
| ], |
| ), |
| ], |
| ), |
| ], |
| className="app-shell", |
| ) |
|
|
|
|
| @app.callback( |
| Output("leaderboard-top", "children"), |
| Output("leaderboard-sort-state", "children"), |
| Output("leaderboard-table", "columns"), |
| Output("leaderboard-table", "data"), |
| Output("dataset-ranking", "figure"), |
| Output("performance-heatmap", "figure"), |
| Input("dataset-filter", "value"), |
| Input("leaderboard-table", "sort_by"), |
| ) |
| def update_leaderboard(dataset: str, sort_by: list[dict] | None): |
| df = leaderboard_frame(dataset, None) |
| visible_cols = [ |
| "method", |
| "task_score", |
| "robustness_auc", |
| "training_time_sec", |
| "peak_ram_gb", |
| "peak_vram_gb", |
| ] |
| if str(dataset) in CONSISTENCY_DATASETS: |
| visible_cols.insert(3, "alignment_score") |
| sort_by = supported_sort(sort_by, visible_cols) |
| sorted_df = sort_table(df, sort_by, [("task_score", False), ("model_order", True)]) |
| table_df = sorted_df[[c for c in visible_cols + ["id", "model"] if c in sorted_df.columns]] |
| metric = df["metric"].dropna().iloc[0] if df["metric"].notna().any() else "score" |
| return ( |
| leaderboard_summary(dataset, df), |
| sort_status(sort_by, ("task_score", False), {"task_score": decoding_label(metric)}), |
| leaderboard_columns([c for c in visible_cols if c in table_df.columns], metric), |
| records(table_df), |
| ranking_figure(dataset, df), |
| performance_heatmap(dataset, None), |
| ) |
|
|
|
|
| @app.callback( |
| Output("latent-color-mode", "options"), |
| Output("latent-color-mode", "value"), |
| Output("latent-color-control", "style"), |
| Input("dataset-filter", "value"), |
| ) |
| def update_latent_color_control(dataset: str): |
| if str(dataset) == "ratinabox": |
| return ( |
| [ |
| {"label": "Position bin", "value": "condition"}, |
| {"label": "X position", "value": "x"}, |
| {"label": "Y position", "value": "y"}, |
| ], |
| "condition", |
| {}, |
| ) |
| return ([{"label": condition_axis_label(str(dataset)), "value": "condition"}], "condition", {"display": "none"}) |
|
|
|
|
| @app.callback( |
| Output("consistency-table", "columns"), |
| Output("consistency-table", "data"), |
| Output("consistency-sort-state", "children"), |
| Output("latent-space", "figure"), |
| Output("consistency-bars", "figure"), |
| Output("consistency-heatmap", "figure"), |
| Input("dataset-filter", "value"), |
| Input("consistency-table", "active_cell"), |
| Input("consistency-table", "sort_by"), |
| Input("latent-color-mode", "value"), |
| ) |
| def update_consistency( |
| dataset: str, |
| active_cell: dict | None, |
| sort_by: list[dict] | None, |
| color_mode: str | None, |
| ): |
| if str(dataset) not in CONSISTENCY_DATASETS: |
| label = DATASET_LABELS.get(dataset, dataset) |
| message = f"Cross-session alignment is not defined for {label}." |
| return ( |
| [], |
| [], |
| message, |
| empty_figure(message), |
| empty_figure(message), |
| consistency_heatmap(None), |
| ) |
| df = consistency_frame(dataset, None) |
| model = selected_consistency_model(df, active_cell) |
| visible_cols = ["method", "alignment_score", "n_sessions", "latent_dim", "n_pairwise"] |
| sort_by = supported_sort(sort_by, visible_cols) |
| sorted_df = sort_table(df, sort_by, [("alignment_score", False), ("model_order", True)]) |
| table_df = sorted_df[[c for c in visible_cols + ["id", "model"] if c in sorted_df.columns]] |
| return ( |
| column_defs([c for c in visible_cols if c in table_df.columns]), |
| records(table_df), |
| sort_status(sort_by, ("alignment_score", False)), |
| latent_space_figure(dataset, model, color_mode or "condition"), |
| consistency_bar_figure(dataset, df), |
| consistency_heatmap(None), |
| ) |
|
|
|
|
| @app.callback( |
| Output("robustness-curve", "figure"), |
| Output("robustness-table", "columns"), |
| Output("robustness-table", "data"), |
| Input("dataset-filter", "value"), |
| Input("robustness-curve", "hoverData"), |
| ) |
| def update_robustness(dataset: str, hover_data: dict | None): |
| table = robustness_table_frame(dataset, None) |
| return ( |
| robustness_figure(dataset, None, hover_data), |
| column_defs(table.columns), |
| records(table), |
| ) |
|
|
|
|
| @app.callback( |
| Output("compute-scatter", "figure"), |
| Output("memory-bars", "figure"), |
| Output("compute-table", "columns"), |
| Output("compute-table", "data"), |
| Input("dataset-filter", "value"), |
| ) |
| def update_compute(dataset: str): |
| scatter, memory, table = compute_figures(dataset, None) |
| return scatter, memory, column_defs(table.columns), records(table) |
|
|
|
|
| @app.callback( |
| Output("neuron-influence-bars", "figure"), |
| Output("trial-influence-bars", "figure"), |
| Output("influence-table", "columns"), |
| Output("influence-table", "data"), |
| Input("dataset-filter", "value"), |
| ) |
| def update_influence(dataset: str): |
| neuron_fig, trial_fig, table = influence_figures(dataset, None) |
| return neuron_fig, trial_fig, column_defs(table.columns), records(table) |
|
|
|
|
| @app.callback( |
| Output("methods-table", "columns"), |
| Output("methods-table", "data"), |
| Input("dataset-filter", "value"), |
| ) |
| def update_methods(_dataset: str): |
| table = methods_frame(None) |
| return column_defs(table.columns), records(table) |
|
|
|
|
| if __name__ == "__main__": |
| app.run(host="0.0.0.0", port=7860, debug=False) |
|
|