Spaces:
Paused
Paused
| """UNI Task PSD Explorer: EC / EO / SM condition comparison.""" | |
| import numpy as np | |
| import pandas as pd | |
| import plotly.graph_objects as go | |
| from scipy import signal | |
| import gradio as gr | |
| import lcmv_xtra as lx | |
| import logging | |
| from pathlib import Path | |
| from typing import Dict, List, Tuple | |
| logging.basicConfig(level=logging.INFO) | |
| logger = logging.getLogger(__name__) | |
| # ============================================================================= | |
| # 1. CONFIGURATION & CONSTANTS | |
| # ============================================================================= | |
| TENSOR_DIR = Path("./data") | |
| CONDITION_LABELS = { | |
| "ec": "Eyes Closed", | |
| "eo": "Eyes Open", | |
| "sm": "Motor Task", | |
| } | |
| CONDITION_COLORS = { | |
| "ec": "#1F77B4", # Blue | |
| "eo": "#2CA02C", # Green | |
| "sm": "#D62728", # Red | |
| } | |
| PSD_WINDOW_SECONDS: float = 4.0 | |
| PSD_OVERLAP_FRACTION: float = 0.75 | |
| PSD_EPSILON: float = 1e-15 | |
| REFERENCE_BAND_HZ: Tuple[float, float] = (1.0, 4.0) | |
| FREQ_MAX_PLOT_HZ: float = 40.0 | |
| PLOT_BANDS = [ | |
| (1, 4, 'Delta', '#90B3F9'), | |
| (4, 8, 'Theta', '#FFF9B2'), | |
| (8, 13, 'Alpha', '#AAFCD2'), | |
| (13, 20, 'Low Beta', '#97C2F9'), | |
| (20, 30, 'High Beta', '#90BEF5'), | |
| ] | |
| BAND_OPTIONS = { | |
| 'Delta (1-4 Hz)': (1, 4), | |
| 'Theta (4-8 Hz)': (4, 8), | |
| 'Alpha (8-13 Hz)': (8, 13), | |
| 'Low Beta (13-20 Hz)': (13, 20), | |
| 'High Beta (20-30 Hz)': (20, 30), | |
| 'Low Gamma (30-50 Hz)': (30, 50), | |
| } | |
| # ============================================================================= | |
| # 2. DATA LOADING & ATLAS MANAGEMENT | |
| # ============================================================================= | |
| def load_psd_cache() -> dict: | |
| """Load single precomputed PSD cache file.""" | |
| cache_path = TENSOR_DIR / "psd_cache.npz" | |
| cache = np.load(cache_path, allow_pickle=True) | |
| logger.info( | |
| f"Loaded PSD cache: {len(cache['entries'])} entries × " | |
| f"{cache['n_rois']} ROIs × {len(cache['freqs'])} freq bins" | |
| ) | |
| return cache | |
| def build_cascading_roi_map(atlas_df: pd.DataFrame) -> Tuple[Dict[str, List[str]], Dict[str, int]]: | |
| """Parse CIMT atlas DataFrame into cascading dropdown structures.""" | |
| required_cols = ['index', 'region_full_name', 'hemisphere', 'functional_system'] | |
| assert all(col in atlas_df.columns for col in required_cols), \ | |
| f"Atlas missing required columns: {set(required_cols) - set(atlas_df.columns)}" | |
| atlas_df = atlas_df.copy() | |
| atlas_df['display_label'] = atlas_df['region_full_name'] + " (" + atlas_df['hemisphere'].str[0] + ")" | |
| system_to_rois: Dict[str, List[str]] = {} | |
| for system in sorted(atlas_df['functional_system'].unique()): | |
| labels = atlas_df[atlas_df['functional_system'] == system]['display_label'].tolist() | |
| system_to_rois[system] = sorted(labels) | |
| label_to_index: Dict[str, int] = dict( | |
| zip(atlas_df['display_label'], atlas_df['index'].astype(int)) | |
| ) | |
| logger.info(f"Built cascading map: {len(system_to_rois)} systems, {len(label_to_index)} ROIs") | |
| return system_to_rois, label_to_index | |
| def get_default_roi_state( | |
| system_to_rois: Dict[str, List[str]], | |
| label_to_index: Dict[str, int] | |
| ) -> Tuple[str, str, int]: | |
| """Return (default_system, default_roi_label, default_roi_index).""" | |
| systems = sorted(system_to_rois.keys()) | |
| assert len(systems) > 0, "No functional systems found in atlas" | |
| default_system = systems[0] | |
| rois = system_to_rois[default_system] | |
| assert len(rois) > 0, f"No ROIs found in system '{default_system}'" | |
| default_roi = rois[0] | |
| default_index = label_to_index[default_roi] | |
| return default_system, default_roi, default_index | |
| # ============================================================================= | |
| # 3. CORE COMPUTATION (REMOVED — now served from cache) | |
| # ============================================================================= | |
| # compute_aligned_psd is no longer needed at runtime. | |
| # All PSDs are precomputed in psd_cache.npz. | |
| # ============================================================================= | |
| # 4. VISUALIZATION | |
| # ============================================================================= | |
| def build_psd_figure( | |
| freqs: np.ndarray, | |
| psd_ec_db: np.ndarray, | |
| psd_eo_db: np.ndarray, | |
| psd_sm_db: np.ndarray, | |
| roi_label: str, | |
| freq_max: float = FREQ_MAX_PLOT_HZ | |
| ) -> go.Figure: | |
| """Construct PSD Plotly figure with band shading (original visual style).""" | |
| fig = go.Figure() | |
| # Band shading with annotations (identical to original) | |
| for f_lo, f_hi, name, color in PLOT_BANDS: | |
| if f_hi <= freq_max: | |
| fig.add_vrect(x0=f_lo, x1=f_hi, fillcolor=color, opacity=0.08, layer="below", line_width=0) | |
| fig.add_annotation( | |
| x=(f_lo + f_hi) / 2, y=0.97, xref="x", yref="paper", | |
| text=f"<b>{name}</b>", showarrow=False, | |
| font=dict(size=10, color='#1E3A5F'), opacity=0.8 | |
| ) | |
| mask = freqs <= freq_max | |
| traces = [ | |
| (CONDITION_LABELS["ec"], psd_ec_db, CONDITION_COLORS["ec"]), | |
| (CONDITION_LABELS["eo"], psd_eo_db, CONDITION_COLORS["eo"]), | |
| (CONDITION_LABELS["sm"], psd_sm_db, CONDITION_COLORS["sm"]), | |
| ] | |
| for label, psd_db, color in traces: | |
| fig.add_trace(go.Scatter( | |
| x=freqs[mask], y=psd_db[mask], mode='lines', | |
| name=label, line=dict(color=color, width=2.5), | |
| )) | |
| # Identical layout to original | |
| fig.update_layout( | |
| legend=dict(yanchor="top", y=0.99, xanchor="right", x=0.99, font=dict(size=12)), | |
| template='plotly_white', margin=dict(t=80, b=60, l=70, r=30), height=500, | |
| ) | |
| fig.update_xaxes(showgrid=True, gridwidth=1, gridcolor='rgba(0,0,0,0.08)') | |
| fig.update_yaxes(showgrid=True, gridwidth=1, gridcolor='rgba(0,0,0,0.08)') | |
| return fig | |
| def build_ratio_figure( | |
| roi_label: str, | |
| band_label: str, | |
| ec_mean: float, | |
| eo_mean: float, | |
| sm_mean: float, | |
| eps: float = PSD_EPSILON | |
| ) -> go.Figure: | |
| """Horizontal bar chart: modulation index relative to EC baseline (original visual style).""" | |
| comparisons = ['Motor Task', 'Eyes Open'] | |
| ratio_sm = ((sm_mean - ec_mean) / (sm_mean + ec_mean + eps)) * 100 | |
| ratio_eo = ((eo_mean - ec_mean) / (eo_mean + ec_mean + eps)) * 100 | |
| values = [ratio_sm, ratio_eo] | |
| colors = [ | |
| CONDITION_COLORS["sm"] if ratio_sm >= 0 else CONDITION_COLORS["ec"], | |
| CONDITION_COLORS["eo"] if ratio_eo >= 0 else CONDITION_COLORS["ec"], | |
| ] | |
| fig = go.Figure() | |
| fig.add_trace(go.Bar( | |
| y=comparisons, x=values, orientation='h', marker_color=colors, | |
| text=[f'{v:+.1f}%' for v in values], textposition='inside', | |
| textfont=dict(size=12, family='monospace', color='white'), insidetextanchor='middle', | |
| )) | |
| # Baseline reference annotations (mirrors original Drug annotation style) | |
| fig.add_annotation(x=1.02, y='MT', xref='paper', yref='y', | |
| text='<b>EC</b>', showarrow=False, font=dict(size=11, color='#333'), xanchor='left') | |
| fig.add_annotation(x=1.02, y='EO', xref='paper', yref='y', | |
| text='<b>EC</b>', showarrow=False, font=dict(size=11, color='#333'), xanchor='left') | |
| # Identical axis/layout styling to original | |
| fig.update_layout( | |
| xaxis=dict(tickfont=dict(size=10), zeroline=True, zerolinewidth=1, | |
| zerolinecolor='#999', showgrid=True, gridwidth=1, gridcolor='rgba(0,0,0,0.06)'), | |
| yaxis=dict(tickfont=dict(size=11, weight='bold'), showgrid=False, zeroline=False, side='left'), | |
| template='plotly_white', height=200, margin=dict(t=50, b=30, l=80, r=60), | |
| ) | |
| return fig | |
| # ============================================================================= | |
| # 5. GRADIO CALLBACKS (ZERO COMPUTATION — pure cache lookup) | |
| # ============================================================================= | |
| def update_psd(roi_label, entry_name, label_to_index, cache): | |
| """Callback for PSD plot update from precomputed cache.""" | |
| if roi_label not in label_to_index: | |
| raise ValueError(f"ROI label '{roi_label}' not found in index map") | |
| idx = label_to_index[roi_label] | |
| entry_idx = np.where(cache['entries'] == entry_name)[0][0] | |
| freqs = cache['freqs'] | |
| ec_db = np.nan_to_num(cache['ec_db'][entry_idx, idx, :], nan=0.0) | |
| eo_db = np.nan_to_num(cache['eo_db'][entry_idx, idx, :], nan=0.0) | |
| sm_db = np.nan_to_num(cache['sm_db'][entry_idx, idx, :], nan=0.0) | |
| return build_psd_figure(freqs, ec_db, eo_db, sm_db, roi_label) | |
| def update_ratio(roi_label, band_label, entry_name, label_to_index, cache): | |
| """Callback for ratio plot from precomputed cache.""" | |
| if roi_label not in label_to_index: | |
| raise ValueError(f"ROI label '{roi_label}' not found in index map") | |
| if band_label not in BAND_OPTIONS: | |
| raise ValueError(f"Band label '{band_label}' not found in BAND_OPTIONS") | |
| idx = label_to_index[roi_label] | |
| f_lo, f_hi = BAND_OPTIONS[band_label] | |
| entry_idx = np.where(cache['entries'] == entry_name)[0][0] | |
| freqs = cache['freqs'] | |
| ec_db = np.nan_to_num(cache['ec_db'][entry_idx, idx, :], nan=0.0) | |
| eo_db = np.nan_to_num(cache['eo_db'][entry_idx, idx, :], nan=0.0) | |
| sm_db = np.nan_to_num(cache['sm_db'][entry_idx, idx, :], nan=0.0) | |
| band_mask = (freqs >= f_lo) & (freqs <= f_hi) | |
| ec_mean = float(np.mean(10 ** (ec_db[band_mask] / 10))) | |
| eo_mean = float(np.mean(10 ** (eo_db[band_mask] / 10))) | |
| sm_mean = float(np.mean(10 ** (sm_db[band_mask] / 10))) | |
| return build_ratio_figure(roi_label, band_label, ec_mean, eo_mean, sm_mean) | |
| def on_system_change(system, system_to_rois): | |
| """Update ROI dropdown choices when functional system changes.""" | |
| rois = system_to_rois.get(system, []) | |
| new_default = rois[0] if rois else None | |
| return gr.update(choices=rois, value=new_default) | |
| # ============================================================================= | |
| # 6. APP INITIALIZATION | |
| # ============================================================================= | |
| def create_app(): | |
| """Build and return the Gradio Blocks app. Importable entry point.""" | |
| cache = load_psd_cache() | |
| # Load CIMT labels from bundled atlas | |
| import lcmv_xtra | |
| labels_path = Path(lcmv_xtra.__file__).parent / 'data' / 'cimt_atlas' / 'cimt_atlas_labels.csv' | |
| atlas_df = pd.read_csv(labels_path) | |
| SYSTEM_TO_ROIS, LABEL_TO_INDEX = build_cascading_roi_map(atlas_df) | |
| DEFAULT_SYS, DEFAULT_ROI, _ = get_default_roi_state(SYSTEM_TO_ROIS, LABEL_TO_INDEX) | |
| entries = list(cache['entries']) # ["Group Average", "sub-01", "sub-02", ...] | |
| initial_fig = update_psd(DEFAULT_ROI, "Group Average", LABEL_TO_INDEX, cache) | |
| initial_ratio = update_ratio(DEFAULT_ROI, 'Alpha (8-13 Hz)', "Group Average", LABEL_TO_INDEX, cache) | |
| with gr.Blocks(title="UNI Task Atlas Explorer") as app: | |
| gr.Markdown( | |
| "# UNI Task: Full Atlas PSD Explorer\n" | |
| "Interactive delta-aligned PSD analysis across Eyes Closed / Eyes Open / Motor Task conditions" | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| subject_dropdown = gr.Dropdown( | |
| choices=entries, | |
| value="Group Average", | |
| label="Subject", | |
| info="Select individual subject or group average" | |
| ) | |
| sys_dropdown = gr.Dropdown( | |
| choices=sorted(SYSTEM_TO_ROIS.keys()), | |
| value=DEFAULT_SYS, | |
| label="Functional System", | |
| info="Select brain network to filter ROIs" | |
| ) | |
| roi_dropdown = gr.Dropdown( | |
| choices=SYSTEM_TO_ROIS[DEFAULT_SYS], | |
| value=DEFAULT_ROI, | |
| label="Region of Interest", | |
| info="Select specific anatomical region" | |
| ) | |
| band_dropdown = gr.Dropdown( | |
| choices=list(BAND_OPTIONS.keys()), | |
| value='Alpha (8-13 Hz)', | |
| label="Frequency Band", | |
| info="Band-averaged power comparison" | |
| ) | |
| ratio_output = gr.Plot(label="Condition Modulation", value=initial_ratio) | |
| with gr.Column(scale=2): | |
| psd_plot = gr.Plot(label="Delta-Aligned PSD", value=initial_fig) | |
| sys_dropdown.change( | |
| fn=lambda s: on_system_change(s, SYSTEM_TO_ROIS), | |
| inputs=sys_dropdown, outputs=roi_dropdown | |
| ) | |
| roi_dropdown.change( | |
| fn=lambda r, subj: update_psd(r, subj, LABEL_TO_INDEX, cache), | |
| inputs=[roi_dropdown, subject_dropdown], outputs=psd_plot | |
| ) | |
| subject_dropdown.change( | |
| fn=lambda r, subj: update_psd(r, subj, LABEL_TO_INDEX, cache), | |
| inputs=[roi_dropdown, subject_dropdown], outputs=psd_plot | |
| ) | |
| roi_dropdown.change( | |
| fn=lambda r, b, subj: update_ratio(r, b, subj, LABEL_TO_INDEX, cache), | |
| inputs=[roi_dropdown, band_dropdown, subject_dropdown], outputs=ratio_output | |
| ) | |
| band_dropdown.change( | |
| fn=lambda r, b, subj: update_ratio(r, b, subj, LABEL_TO_INDEX, cache), | |
| inputs=[roi_dropdown, band_dropdown, subject_dropdown], outputs=ratio_output | |
| ) | |
| subject_dropdown.change( | |
| fn=lambda r, b, subj: update_ratio(r, b, subj, LABEL_TO_INDEX, cache), | |
| inputs=[roi_dropdown, band_dropdown, subject_dropdown], outputs=ratio_output | |
| ) | |
| return app | |
| if __name__ == "__main__": | |
| app = create_app() | |
| app.launch(theme=gr.themes.Soft(), css=".gradio-container { max-width: 1200px !important; }") |