Spaces:
Sleeping
Sleeping
| """Custom Gradio UI for showing the trained model's predictions. | |
| Mounts as a second tab next to the OpenEnv playground when ``ENABLE_WEB_INTERFACE`` | |
| is on. Lets a viewer: | |
| * pick a subject / condition / run / seed, | |
| * run one full 20-step rollout under the trained policy (or the static | |
| fallback when the Space has no GPU), | |
| * see what the model predicted -- which parcels, in what order, and how the | |
| cumulative R^2 climbs as a function of selection step, | |
| * optionally see those parcels rendered on a glass brain (via nilearn's | |
| Schaefer-2018 atlas, lazy-loaded so a CPU-only Space still serves the | |
| numerical prediction view). | |
| This module is purely a visualization layer over :func:`server.app.run_brain_rollout` | |
| and the existing ``/rollout`` HTTP endpoint, so anything the UI shows is | |
| reproducible by an external client. | |
| """ | |
| from __future__ import annotations | |
| import io | |
| import json | |
| import os | |
| import threading | |
| from pathlib import Path | |
| from typing import Any, Callable, Iterable | |
| import gradio as gr | |
| # --------------------------------------------------------------------------- | |
| # Lookups (subjects, runs, conditions) -- read once at build_prediction_ui() | |
| # --------------------------------------------------------------------------- | |
| CONDITIONS = ["single_m", "single_f", "mixed_m", "mixed_f"] | |
| def _participant_info_path() -> Path | None: | |
| """Locate ``participant_run_info.json`` from configured data dirs.""" | |
| config_dir = os.getenv("BRAINRL_CONFIG_DIR") | |
| if config_dir: | |
| candidate = Path(config_dir) / "participant_run_info.json" | |
| if candidate.exists(): | |
| return candidate | |
| # Fall back to the bundled configs dir shipped with the Space image. | |
| bundled = Path(__file__).resolve().parent.parent / "configs" / "participant_run_info.json" | |
| if bundled.exists(): | |
| return bundled | |
| return None | |
| def _participant_choices() -> tuple[list[str], dict[str, dict[str, str]]]: | |
| """Return (sorted subject ids, raw subject -> {run: condition} mapping).""" | |
| info_path = _participant_info_path() | |
| if info_path is None: | |
| return ["sub-01"], {} | |
| try: | |
| info = json.loads(info_path.read_text()) | |
| except Exception: | |
| return ["sub-01"], {} | |
| return sorted(info.keys()), info | |
| # --------------------------------------------------------------------------- | |
| # Schaefer-2018 atlas glass-brain plot -- lazy because nilearn is heavy | |
| # --------------------------------------------------------------------------- | |
| _ATLAS_LOCK = threading.Lock() | |
| _ATLAS_CACHE: dict[str, Any] = {"loaded": False, "labels_img": None, "label_to_index": None, "error": None} | |
| def _load_schaefer_atlas() -> dict[str, Any]: | |
| """Fetch the Schaefer-2018 atlas once and cache the result. | |
| The atlas (NIfTI) is ~10MB and nilearn caches it on disk, so this only | |
| pays the download cost on the very first /web request after a Space | |
| cold start. Subsequent calls reuse the in-memory ``labels_img``. | |
| """ | |
| if _ATLAS_CACHE["loaded"]: | |
| return _ATLAS_CACHE | |
| with _ATLAS_LOCK: | |
| if _ATLAS_CACHE["loaded"]: | |
| return _ATLAS_CACHE | |
| try: | |
| from nilearn import datasets as nl_datasets | |
| from nilearn import image as nl_image | |
| except ImportError as exc: # pragma: no cover - graceful degrade | |
| _ATLAS_CACHE["loaded"] = True | |
| _ATLAS_CACHE["error"] = ( | |
| "nilearn not installed in this Space; install the [atlas] extra " | |
| f"to enable the glass-brain view ({exc})." | |
| ) | |
| return _ATLAS_CACHE | |
| try: | |
| atlas = nl_datasets.fetch_atlas_schaefer_2018( | |
| n_rois=200, yeo_networks=7, resolution_mm=2, | |
| ) | |
| labels_img = nl_image.load_img(atlas.maps) | |
| raw_labels: list[str] = [] | |
| for entry in atlas.labels: | |
| raw_labels.append(entry.decode("utf-8") if isinstance(entry, bytes) else str(entry)) | |
| label_to_index = {lbl: idx + 1 for idx, lbl in enumerate(raw_labels)} | |
| except Exception as exc: # pragma: no cover - graceful degrade | |
| _ATLAS_CACHE["loaded"] = True | |
| _ATLAS_CACHE["error"] = f"Failed to fetch Schaefer atlas: {exc}" | |
| return _ATLAS_CACHE | |
| _ATLAS_CACHE.update( | |
| loaded=True, | |
| labels_img=labels_img, | |
| label_to_index=label_to_index, | |
| error=None, | |
| ) | |
| return _ATLAS_CACHE | |
| def _glass_brain_image(steps: Iterable[dict[str, Any]]) -> tuple[Any | None, str | None]: | |
| """Render a glass-brain overlay where each parcel is colored by selection order. | |
| The first selected parcel gets the lowest priority value, the last gets the | |
| highest, and unselected voxels are masked out. Using a divergent colormap | |
| on a numeric "selection rank" makes the time-ordering of the model's | |
| predictions readable at a glance. | |
| """ | |
| atlas = _load_schaefer_atlas() | |
| if atlas.get("error"): | |
| return None, atlas["error"] | |
| labels_img = atlas["labels_img"] | |
| label_to_index = atlas["label_to_index"] or {} | |
| if labels_img is None or not label_to_index: | |
| return None, "Atlas labels not available." | |
| try: | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| from nilearn import image as nl_image | |
| from nilearn import plotting as nl_plotting | |
| except ImportError as exc: # pragma: no cover - graceful degrade | |
| return None, f"Plot dependencies missing: {exc}" | |
| label_data = np.asarray(labels_img.dataobj).astype(int) | |
| overlay = np.zeros_like(label_data, dtype=np.float32) | |
| rendered = 0 | |
| for rank, step in enumerate(steps, start=1): | |
| label = step.get("region_label") | |
| if not label: | |
| continue | |
| idx = label_to_index.get(label) | |
| if not idx: | |
| continue | |
| overlay[label_data == idx] = float(rank) | |
| rendered += 1 | |
| if rendered == 0: | |
| return None, "Model picked no parcels that match the Schaefer atlas labels." | |
| overlay_img = nl_image.new_img_like(labels_img, overlay) | |
| fig = plt.figure(figsize=(9, 4.0)) | |
| nl_plotting.plot_glass_brain( | |
| overlay_img, | |
| figure=fig, | |
| colorbar=True, | |
| cmap="viridis", | |
| plot_abs=False, | |
| threshold=0.5, | |
| title=f"Predicted parcels (selection order, n={rendered})", | |
| ) | |
| return fig, None | |
| # --------------------------------------------------------------------------- | |
| # 2D summary plots -- always available because matplotlib ships in [plots] | |
| # --------------------------------------------------------------------------- | |
| def _r2_curve_figure(steps: list[dict[str, Any]]) -> Any | None: | |
| if not steps: | |
| return None | |
| try: | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| except ImportError: # pragma: no cover | |
| return None | |
| timesteps = [s["timestep"] for s in steps] | |
| cumulative = [s["current_r2"] for s in steps] | |
| rewards = [s["reward"] for s in steps] | |
| fallback = [bool(s["fallback_used"]) for s in steps] | |
| fig, axes = plt.subplots(1, 2, figsize=(11, 4.0)) | |
| axes[0].plot([0, *timesteps], [0.0, *cumulative], color="#3b82f6", linewidth=2) | |
| axes[0].scatter( | |
| timesteps, | |
| cumulative, | |
| c=["#ef4444" if fb else "#3b82f6" for fb in fallback], | |
| zorder=3, | |
| s=40, | |
| ) | |
| axes[0].set_title("Predicted cumulative R\u00b2 over selection budget") | |
| axes[0].set_xlabel("Selection step") | |
| axes[0].set_ylabel("Cumulative R\u00b2") | |
| axes[0].grid(linestyle="--", alpha=0.4) | |
| bar_colors = ["#ef4444" if fb else "#10b981" for fb in fallback] | |
| axes[1].bar(timesteps, rewards, color=bar_colors) | |
| axes[1].set_title("Per-step reward (red = fallback used)") | |
| axes[1].set_xlabel("Selection step") | |
| axes[1].set_ylabel("Reward") | |
| axes[1].grid(axis="y", linestyle="--", alpha=0.4) | |
| fig.suptitle("Trained policy: predictions on a single episode") | |
| fig.tight_layout() | |
| return fig | |
| def _network_distribution_figure(steps: list[dict[str, Any]]) -> Any | None: | |
| """Bar chart of selected parcels grouped by network x hemisphere. | |
| Useful sanity check: a sensible policy should bias toward language / | |
| auditory networks (Default, Cont, SalVentAttn) on auditory conditions | |
| rather than spreading uniformly across all 7 Yeo networks. | |
| """ | |
| if not steps: | |
| return None | |
| try: | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| except ImportError: # pragma: no cover | |
| return None | |
| counts: dict[str, dict[str, int]] = {} | |
| for step in steps: | |
| network = step.get("network") or "Other" | |
| hemi = step.get("hemisphere") or "n/a" | |
| counts.setdefault(network, {"left": 0, "right": 0, "other": 0}) | |
| counts[network][hemi if hemi in ("left", "right") else "other"] += 1 | |
| if not counts: | |
| return None | |
| networks = sorted(counts.keys()) | |
| left = [counts[n]["left"] for n in networks] | |
| right = [counts[n]["right"] for n in networks] | |
| other = [counts[n]["other"] for n in networks] | |
| fig, ax = plt.subplots(figsize=(9, 3.5)) | |
| x = list(range(len(networks))) | |
| ax.bar(x, left, color="#3b82f6", label="left hemisphere") | |
| ax.bar(x, right, bottom=left, color="#10b981", label="right hemisphere") | |
| bottoms = [l + r for l, r in zip(left, right)] | |
| if any(other): | |
| ax.bar(x, other, bottom=bottoms, color="#94a3b8", label="other") | |
| ax.set_xticks(x) | |
| ax.set_xticklabels(networks, rotation=20) | |
| ax.set_ylabel("Number of selected parcels") | |
| ax.set_title("Predicted parcels by Yeo-7 network") | |
| ax.legend(loc="upper right") | |
| ax.grid(axis="y", linestyle="--", alpha=0.4) | |
| fig.tight_layout() | |
| return fig | |
| # --------------------------------------------------------------------------- | |
| # Markdown helpers | |
| # --------------------------------------------------------------------------- | |
| def _stimulus_markdown(rollout: dict[str, Any]) -> str: | |
| stimulus = rollout.get("stimulus") or {} | |
| if not stimulus: | |
| return "_No stimulus context attached to this episode._" | |
| parts = [ | |
| f"**Condition**: `{rollout.get('condition')}`", | |
| f"**Subject / run**: `{rollout.get('subject_id')}` / `{rollout.get('run_id')}`", | |
| ] | |
| window = stimulus.get("window") if isinstance(stimulus, dict) else None | |
| if window: | |
| parts.append(f"**Window**: `{window}`") | |
| onset = stimulus.get("start_time_s") if isinstance(stimulus, dict) else None | |
| offset = stimulus.get("end_time_s") if isinstance(stimulus, dict) else None | |
| if onset is not None and offset is not None: | |
| parts.append(f"**Time**: `{onset:.1f}s` -> `{offset:.1f}s`") | |
| pos = stimulus.get("dominant_pos") if isinstance(stimulus, dict) else None | |
| if pos: | |
| parts.append(f"**Dominant POS**: `{pos}`") | |
| text = stimulus.get("text") if isinstance(stimulus, dict) else None | |
| if text: | |
| parts.append(f"\n> {text}") | |
| return "\n\n".join(parts) | |
| def _status_markdown(rollout: dict[str, Any]) -> str: | |
| policy = rollout.get("policy") or {} | |
| steps = rollout.get("steps") or [] | |
| fallback_count = sum(1 for s in steps if s.get("fallback_used")) | |
| final_r2 = float(rollout.get("final_r2") or 0.0) | |
| budget = int(rollout.get("selection_budget") or 0) | |
| if policy.get("loaded"): | |
| model_line = f"**Model**: `{policy.get('model_repo')}` (device=`{policy.get('device')}`)" | |
| elif policy.get("enabled"): | |
| err = policy.get("load_error") or "loading" | |
| model_line = f"**Model**: `{policy.get('model_repo')}` (status=`{err}`)" | |
| else: | |
| model_line = "**Model**: static fallback policy (no `BRAINRL_POLICY_MODEL_REPO`)." | |
| return ( | |
| f"{model_line}\n\n" | |
| f"**Predicted final R\u00b2**: `{final_r2:.4f}` over `{len(steps)}/{budget}` steps. " | |
| f"Fallback used on `{fallback_count}/{len(steps)}` steps." | |
| ) | |
| def _steps_table(steps: list[dict[str, Any]]) -> list[list[Any]]: | |
| rows: list[list[Any]] = [] | |
| for s in steps: | |
| rows.append( | |
| [ | |
| s.get("timestep"), | |
| s.get("region_id"), | |
| s.get("region_label"), | |
| s.get("network"), | |
| s.get("hemisphere"), | |
| round(float(s.get("current_r2", 0.0)), 4), | |
| round(float(s.get("delta_r2", 0.0)), 4), | |
| round(float(s.get("reward", 0.0)), 4), | |
| "yes" if s.get("fallback_used") else "no", | |
| ] | |
| ) | |
| return rows | |
| # --------------------------------------------------------------------------- | |
| # Top-level builder | |
| # --------------------------------------------------------------------------- | |
| def build_prediction_ui( | |
| run_rollout: Callable[..., dict[str, Any]], | |
| policy: Any, | |
| ) -> Callable[..., gr.Blocks]: | |
| """Return an OpenEnv-compatible ``gradio_builder``. | |
| ``run_rollout`` and ``policy`` are bound up front so the Gradio handlers | |
| can call back into the trained-policy + environment without re-importing | |
| them at module load time (which would create a circular import with | |
| ``server.app``). | |
| """ | |
| subjects, info = _participant_choices() | |
| default_subject = subjects[0] if subjects else "sub-01" | |
| def _ui_run_rollout( | |
| subject: str, | |
| condition: str, | |
| run_id: str, | |
| seed: int, | |
| stimulus_window: int, | |
| ): | |
| try: | |
| rollout = run_rollout( | |
| seed=int(seed), | |
| subject_id=subject or None, | |
| run_id=run_id or None, | |
| condition=condition or None, | |
| stimulus_window=int(stimulus_window) if stimulus_window >= 0 else None, | |
| ) | |
| except Exception as exc: | |
| err = f"Rollout failed: {type(exc).__name__}: {exc}" | |
| return ( | |
| err, | |
| None, | |
| None, | |
| None, | |
| "_(no rollout)_", | |
| [], | |
| "{}", | |
| "Atlas plot unavailable: rollout failed.", | |
| ) | |
| steps = rollout.get("steps") or [] | |
| status = _status_markdown(rollout) | |
| r2_fig = _r2_curve_figure(steps) | |
| net_fig = _network_distribution_figure(steps) | |
| glass_fig, glass_err = _glass_brain_image(steps) | |
| stimulus_md = _stimulus_markdown(rollout) | |
| table_rows = _steps_table(steps) | |
| raw_json = json.dumps(rollout, indent=2) | |
| atlas_status = ( | |
| glass_err if glass_err else f"Glass brain rendered for {len(steps)} predicted parcels." | |
| ) | |
| return ( | |
| status, | |
| r2_fig, | |
| net_fig, | |
| glass_fig, | |
| stimulus_md, | |
| table_rows, | |
| raw_json, | |
| atlas_status, | |
| ) | |
| def _suggest_run_for_condition(subject: str, condition: str) -> str: | |
| if not subject or not condition: | |
| return "" | |
| subject_info = info.get(subject, {}) | |
| for run, cond in subject_info.items(): | |
| if cond == condition: | |
| # Translate "run1" -> "run-1" because the env uses both forms | |
| # in different code paths; the dash form is what the notebook | |
| # demos and the verifier all default to. | |
| if run.startswith("run") and len(run) > 3 and run[3].isdigit(): | |
| return f"run-{run[3:]}" | |
| return run | |
| return "" | |
| # Required signature for OpenEnv `gradio_builder`: | |
| # (web_manager, action_fields, metadata, is_chat_env, title, quick_start_md) -> gr.Blocks | |
| # We ignore most of those because the prediction tab does not drive the | |
| # OpenEnv state machine -- it spins up its own env per request. | |
| def _builder( | |
| web_manager: Any, | |
| action_fields: list[Any], | |
| metadata: Any, | |
| is_chat_env: bool, | |
| title: str, | |
| quick_start_md: str | None, | |
| ) -> gr.Blocks: | |
| with gr.Blocks(title="BrainRL prediction demo") as demo: | |
| gr.Markdown( | |
| "## BrainRL prediction demo\n\n" | |
| "Pick a subject + condition and run a 20-step model rollout. " | |
| "Each step shows which parcel the trained policy predicted, " | |
| "its cumulative R\u00b2, and where it sits on the brain. " | |
| "Same call as `POST /rollout`, just visualised." | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| subject_dd = gr.Dropdown( | |
| choices=subjects or [default_subject], | |
| value=default_subject, | |
| label="Subject", | |
| ) | |
| condition_dd = gr.Dropdown( | |
| choices=CONDITIONS, | |
| value="single_m", | |
| label="Condition", | |
| ) | |
| run_tb = gr.Textbox( | |
| value=_suggest_run_for_condition(default_subject, "single_m") or "run-1", | |
| label="Run ID", | |
| info="Override if you want a specific run; otherwise defaults to " | |
| "the run that matches the condition for this subject.", | |
| ) | |
| seed_sl = gr.Slider( | |
| minimum=0, maximum=1024, step=1, value=42, label="Seed", | |
| ) | |
| stim_sl = gr.Slider( | |
| minimum=-1, | |
| maximum=64, | |
| step=1, | |
| value=-1, | |
| label="Stimulus window (-1 = auto)", | |
| ) | |
| run_btn = gr.Button("Run model prediction", variant="primary") | |
| status_md = gr.Markdown( | |
| value=("**Model**: not yet queried. Click " | |
| "**Run model prediction** to invoke the policy.") | |
| ) | |
| with gr.Column(scale=2): | |
| with gr.Tabs(): | |
| with gr.Tab("R\u00b2 curve & rewards"): | |
| r2_plot = gr.Plot(label="Predicted R\u00b2 / step reward") | |
| with gr.Tab("Brain map"): | |
| atlas_status_md = gr.Markdown( | |
| value=("Glass-brain map is rendered after the first " | |
| "rollout; the Schaefer atlas is fetched on demand.") | |
| ) | |
| glass_plot = gr.Plot(label="Predicted parcels (glass brain)") | |
| with gr.Tab("Network distribution"): | |
| net_plot = gr.Plot(label="Selected parcels by Yeo-7 network") | |
| with gr.Tab("Stimulus context"): | |
| stimulus_md = gr.Markdown( | |
| value="_Run a prediction to see the stimulus window._" | |
| ) | |
| with gr.Tab("Per-step predictions"): | |
| steps_table = gr.Dataframe( | |
| headers=[ | |
| "step", | |
| "region_id", | |
| "label", | |
| "network", | |
| "hemi", | |
| "cum_R\u00b2", | |
| "delta_R\u00b2", | |
| "reward", | |
| "fallback", | |
| ], | |
| value=[], | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| with gr.Tab("Raw JSON"): | |
| raw_json_box = gr.Code( | |
| label="Rollout response", language="json", interactive=False, | |
| ) | |
| # Auto-suggest a sensible run-id when subject/condition change. | |
| def _refresh_run(subject: str, condition: str) -> str: | |
| return _suggest_run_for_condition(subject, condition) or "run-1" | |
| subject_dd.change(_refresh_run, inputs=[subject_dd, condition_dd], outputs=[run_tb]) | |
| condition_dd.change(_refresh_run, inputs=[subject_dd, condition_dd], outputs=[run_tb]) | |
| run_btn.click( | |
| fn=_ui_run_rollout, | |
| inputs=[subject_dd, condition_dd, run_tb, seed_sl, stim_sl], | |
| outputs=[ | |
| status_md, | |
| r2_plot, | |
| net_plot, | |
| glass_plot, | |
| stimulus_md, | |
| steps_table, | |
| raw_json_box, | |
| atlas_status_md, | |
| ], | |
| ) | |
| return demo | |
| return _builder | |