File size: 28,387 Bytes
1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf f2904ef 1dca9cf | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 | {
"nbformat": 4,
"nbformat_minor": 5,
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
"language_info": {
"name": "python",
"version": "3.10.0"
},
"colab": {
"provenance": [],
"gpuType": "T4"
},
"accelerator": "GPU"
},
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Amphora NeuroText v4 β Audio & Text β Brain ROI Activation\n",
"\n",
"Predict which brain regions activate in response to any audio or text stimulus.\n",
"\n",
"Trained on **real naturalistic fMRI data** from 1,600+ subjects across 4,480 sessions. Zero-shot: no brain scan needed at inference.\n",
"\n",
"**Audio model (Whisper v4) beats TRIBE v2** (Meta AI, Algonauts 2025 winner) **by +4.2%** \n",
"Holdout R = 0.257 vs TRIBE 0.215 Β· 23 held-out subjects Β· single shared model (no per-subject fine-tuning)\n",
"\n",
"**All v4 models use honest evaluation:** per-subject z-scoring + per-subject train/val split (no data leakage).\n",
"\n",
"**Runtime:** GPU recommended (T4 or better). The MLP projector runs on CPU instantly β only the encoder models (Whisper/Qwen3) need GPU.\n",
"\n",
"---\n",
"Model repo: [ffh92r32rm0/Amphora_NeuroText](https://huggingface.co/ffh92r32rm0/Amphora_NeuroText)"
],
"id": "markdown-intro"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": "# ββ Install dependencies ββββββββββββββββββββββββββββββββββββββββββββββββββββββ\n!pip install -q torch transformers huggingface_hub numpy matplotlib librosa\n!pip install -q nilearn mne # brain surface visualization\nprint('Done.')\n",
"id": "install"
},
{
"cell_type": "code",
"execution_count": null,
"id": "hf-auth",
"metadata": {},
"outputs": [],
"source": "# ββ HuggingFace authentication (required β repo is private) ββββββββββββββββββ\n# Option 1 (recommended): add your HF token to Colab Secrets as \"HF_TOKEN\"\n# Runtime β Secrets β Add new secret β Name: HF_TOKEN, Value: hf_...\n# Option 2: run notebook_login() interactively (prompts for token below)\nimport os\nfrom huggingface_hub import login\n\ntry:\n from google.colab import userdata\n _token = userdata.get('HF_TOKEN')\n if _token:\n login(token=_token, add_to_git_credential=False)\n print('Logged in via Colab Secrets (HF_TOKEN).')\n else:\n raise KeyError('HF_TOKEN not set in Colab Secrets')\nexcept Exception:\n print('HF_TOKEN not found in Colab Secrets β opening interactive login...')\n from huggingface_hub import notebook_login\n notebook_login()\n"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# ββ Choose modality βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ\n",
"# Set MODALITY to 'audio' or 'text'\n",
"MODALITY = 'audio' # 'audio' = Whisper v4 (holdout R=0.257, beats TRIBE v2 by +4.2%)\n",
" # 'text' = Combined v4 (val R=0.192, zero-padded whisper slot)\n",
"\n",
"MODEL_FILE = {'audio': 'text2roi_whisper_v4.pt', 'text': 'text2roi_combined_v4.pt'}[MODALITY]\n",
"IN_DIM = {'audio': 1280, 'text': 3840 }[MODALITY]\n",
"\n",
"print(f'Modality : {MODALITY}')\n",
"print(f'Model : {MODEL_FILE} (in_dim={IN_DIM})')"
],
"id": "choose-model"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": "# ββ Locate model files (local first, HuggingFace fallback) βββββββββββββββββββ\n#\n# Priority order:\n# 1. Files uploaded to this Colab session (e.g. from the self-contained zip)\n# 2. HuggingFace Hub download (requires internet; ~7-17 MB per model)\n#\n# To use the bundled zip offline:\n# - Extract NeuroText_v4_Demo.zip\n# - Upload the .pt files and predict.py to Colab via the Files panel\n# - Re-run this cell β it will find them locally and skip the HF download\n\nimport os, shutil\nfrom pathlib import Path\n\nREPO = 'ffh92r32rm0/Amphora_NeuroText'\n\ndef _resolve(filename, repo=REPO):\n \"\"\"Return local path to filename, downloading from HF if not found locally.\"\"\"\n local = Path(filename)\n if local.exists():\n print(f' {filename}: found locally')\n return str(local)\n print(f' {filename}: not found locally, downloading from HuggingFace...')\n from huggingface_hub import hf_hub_download\n path = hf_hub_download(repo, filename)\n print(f' {filename}: downloaded to {path}')\n return path\n\ndef _resolve_optional(filename, repo=REPO):\n \"\"\"Like _resolve but returns None on failure instead of raising.\"\"\"\n try:\n return _resolve(filename, repo)\n except Exception as e:\n print(f' {filename}: not available ({type(e).__name__}) β will be skipped')\n return None\n\nprint('Resolving files...')\nckpt_path = _resolve(MODEL_FILE)\npredict_path = _resolve('predict.py')\nexamples_path = _resolve_optional('examples_cache.npz')\n\n# make predict.py importable from CWD\nif predict_path != 'predict.py':\n shutil.copy(predict_path, 'predict.py')\n\nprint(f'\\nCheckpoint : {ckpt_path}')\nprint(f'predict.py : ready')\nprint(f'examples : {\"found\" if examples_path else \"not available (cached examples will be skipped)\"}')\nprint('Ready.')\n",
"id": "resolve-files"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# ββ Load model (tiny MLP, runs instantly on CPU) ββββββββββββββββββββββββββββββ\n",
"import torch, torch.nn as nn, numpy as np\n",
"\n",
"class Text2ROI(nn.Module):\n",
" def __init__(self, in_dim=1280, hidden=1024, out_dim=56, dropout=0.1):\n",
" super().__init__()\n",
" self.net = nn.Sequential(\n",
" nn.Linear(in_dim, hidden), nn.GELU(), nn.Dropout(dropout), nn.LayerNorm(hidden),\n",
" nn.Linear(hidden, hidden//2), nn.GELU(), nn.Dropout(dropout),\n",
" nn.Linear(hidden//2, out_dim),\n",
" )\n",
" def forward(self, x): return self.net(x)\n",
"\n",
"ckpt = torch.load(ckpt_path, map_location='cpu', weights_only=False)\n",
"model = Text2ROI(in_dim=ckpt['in_dim'], out_dim=ckpt.get('n_roi', 56))\n",
"model.load_state_dict(ckpt['state_dict'])\n",
"model.eval()\n",
"\n",
"ROI_NAMES = [\n",
" 'V1','V2','V3','V4','V3A','V3B','LO1','LO2','MT','MST','V7','IPS1',\n",
" 'FFA-1','FFA-2','PPA','RSC','OFA','EBA','IPS2','IPS3','IPS4','IPS5','SPL1',\n",
" 'hIP1','hIP2','hIP3','dlPFC','vlPFC','OFC','ACC','mPFC','FP1','FP2',\n",
" 'IFG','IFGorb','STG','STS','MTG','AG','PCC','mPFC_dmn','LP_L','LP_R',\n",
" 'HPC_L','HPC_R','AI','dACC','sgACC','vmPFC',\n",
" 'Amygdala_L','Amygdala_R','Caudate_L','Caudate_R','Putamen_L','Putamen_R','Thalamus',\n",
"]\n",
"\n",
"val_r = ckpt.get('best_val_r', 'N/A')\n",
"print(f'Model loaded in_dim={ckpt[\"in_dim\"]} val_R={val_r}')\n",
"print(f'Parameters: {sum(p.numel() for p in model.parameters()):,}')"
],
"id": "load-model"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": "# ββ Helper: run inference + plot ββββββββββββββββββββββββββββββββββββββββββββββ\nimport matplotlib.pyplot as plt\n\ndef predict_from_features(feat_1d, model, in_dim):\n \"\"\"Run model on a 1-D feature vector.\n\n For the combined model (in_dim=3840) with a 2560-d Qwen3 embedding,\n zero-pads the whisper slot [0:1280] automatically.\n Raises ValueError if the embedding cannot be adapted to in_dim.\n \"\"\"\n if in_dim == 3840 and feat_1d.shape[0] == 2560:\n feat_1d = np.concatenate([np.zeros(1280, dtype=np.float32), feat_1d])\n if feat_1d.shape[0] < in_dim:\n raise ValueError(\n f\"Embedding dim {feat_1d.shape[0]} < model in_dim {in_dim}. \"\n f\"Make sure you are using the correct model for this embedding type. \"\n f\"For text (Qwen3 2560-d) use the combined model (in_dim=3840); \"\n f\"for audio (Whisper 1280-d) use the audio model (in_dim=1280).\"\n )\n with torch.no_grad():\n pred = model(torch.from_numpy(feat_1d[:in_dim]).unsqueeze(0)).squeeze(0).numpy()\n return dict(zip(ROI_NAMES, pred.tolist()))\n\ndef plot_rois(roi_map, title='', top_n=20):\n items = sorted(roi_map.items(), key=lambda x: -x[1])[:top_n]\n names, vals = zip(*items)\n vmin, vmax = min(vals), max(vals)\n colors = plt.cm.RdYlGn([(v - vmin) / (vmax - vmin + 1e-9) for v in vals])\n fig, ax = plt.subplots(figsize=(9, 5))\n ax.barh(range(len(names)), vals, color=colors)\n ax.set_yticks(range(len(names))); ax.set_yticklabels(names, fontsize=9)\n ax.axvline(0, color='gray', linewidth=0.7)\n ax.set_xlabel('Predicted activation (z-score)')\n ax.set_title(title[:90], fontsize=10)\n ax.invert_yaxis(); plt.tight_layout(); plt.show()\n\nprint('Helpers ready.')\n",
"id": "helpers"
},
{
"cell_type": "code",
"execution_count": null,
"id": "brain-plotter-lib",
"metadata": {},
"outputs": [],
"source": "# ββ Brain map plotter (fsaverage5 / HCP-MMP1, same mesh as TRIBE v2) βββββββββ\n# Commercial-friendly: MNE + Nilearn only (BSD). Does NOT import tribev2.\n\nfrom __future__ import annotations\n\n!pip install -q mne nilearn Pillow\n\nfrom functools import lru_cache\nfrom pathlib import Path\nfrom typing import Dict, List, Mapping, Sequence, Tuple\nimport numpy as np\n\nFSAVERAGE5_VERTS_PER_HEMI = 10242\nFSAVERAGE5_VERTICES = FSAVERAGE5_VERTS_PER_HEMI * 2\n\n# ROI β HCP-MMP1 parcel mapping.\n# PPA/OFA entries are anatomical approximations (see comments).\nROI56_TO_HCP: Dict[str, List[Tuple[str, str]]] = {\n \"V1\": [(\"V1\",\"both\")], \"V2\": [(\"V2\",\"both\")], \"V3\": [(\"V3\",\"both\")], \"V4\": [(\"V4\",\"both\")],\n \"V3A\": [(\"V3A\",\"both\")], \"V3B\": [(\"V3B\",\"both\")], \"LO1\": [(\"LO1\",\"both\")], \"LO2\": [(\"LO2\",\"both\")],\n \"MT\": [(\"MT\",\"both\")], \"MST\": [(\"MST\",\"both\")], \"V7\": [(\"V7\",\"both\")], \"IPS1\": [(\"IPS1\",\"both\")],\n \"FFA-1\": [(\"FFC\",\"both\")], \"FFA-2\": [(\"FFC\",\"both\")],\n \"PPA\": [(\"PHA1\",\"both\"),(\"PHA2\",\"both\"),(\"PHA3\",\"both\")], # approx; PIT is posterior IT\n \"RSC\": [(\"POS1\",\"both\"),(\"7m\",\"both\")],\n \"OFA\": [(\"FFC\",\"both\"),(\"VVC\",\"both\")], # approx ventral temporal\n \"EBA\": [(\"FST\",\"both\")],\n \"IPS2\": [(\"IPS2\",\"both\")], \"IPS3\": [(\"IPS3\",\"both\")], \"IPS4\": [(\"IPS4\",\"both\")], \"IPS5\": [(\"IPS5\",\"both\")],\n \"SPL1\": [(\"SPL1\",\"both\")], \"hIP1\": [(\"7AL\",\"both\")], \"hIP2\": [(\"7PC\",\"both\")], \"hIP3\": [(\"7Am\",\"both\")],\n \"dlPFC\": [(\"9-46d\",\"both\"),(\"46\",\"both\")], \"vlPFC\": [(\"47l\",\"both\")], \"OFC\": [(\"11l\",\"both\")],\n \"ACC\": [(\"a24\",\"both\"),(\"p24\",\"both\")], \"mPFC\": [(\"9m\",\"both\")], \"FP1\": [(\"10d\",\"both\")], \"FP2\": [(\"10d\",\"both\")],\n \"IFG\": [(\"44\",\"both\"),(\"45a\",\"both\")], \"IFGorb\": [(\"47l\",\"both\")],\n \"STG\": [(\"STGa\",\"both\"),(\"STGr\",\"both\")], \"STS\": [(\"STSda\",\"both\")], \"MTG\": [(\"TE1a\",\"both\")],\n \"AG\": [(\"PFm\",\"both\"),(\"PGs\",\"both\")], \"PCC\": [(\"PCC\",\"both\")],\n \"mPFC_dmn\": [(\"9m\",\"both\"),(\"10r\",\"both\")],\n \"LP_L\": [(\"PFm\",\"left\"),(\"PGs\",\"left\")], \"LP_R\": [(\"PFm\",\"right\"),(\"PGs\",\"right\")],\n \"HPC_L\": [(\"Entorhinal\",\"left\")], \"HPC_R\": [(\"Entorhinal\",\"right\")],\n \"AI\": [(\"Ig\",\"both\"),(\"FOP4\",\"both\")], \"dACC\": [(\"a24\",\"both\")],\n \"sgACC\": [(\"s32\",\"both\")], \"vmPFC\": [(\"25\",\"both\")],\n \"Amygdala_L\": [], \"Amygdala_R\": [], \"Caudate_L\": [], \"Caudate_R\": [],\n \"Putamen_L\": [], \"Putamen_R\": [], \"Thalamus\": [],\n}\n\n@lru_cache(maxsize=1)\ndef _hcp_label_vertices(mesh: str = \"fsaverage5\") -> Dict[str, np.ndarray]:\n import mne, os\n if mesh != \"fsaverage5\":\n raise ValueError(\"Only fsaverage5 is supported\")\n max_v = FSAVERAGE5_VERTS_PER_HEMI\n\n # Fix 4/10: use fetch_fsaverage (~50 MB) instead of sample dataset (~1.7 GB)\n subjects_dir = Path(mne.datasets.fetch_fsaverage(verbose=False)).parent\n mne.datasets.fetch_hcp_mmp_parcellation(\n subjects_dir=subjects_dir, accept=True, verbose=False\n )\n\n # Fix 5: auto-detect annotation name (differs across MNE versions)\n annot_dir = subjects_dir / \"fsaverage\" / \"label\"\n annot_files = list(annot_dir.glob(\"lh.*.annot\"))\n if any(\"HCP-MMP1\" in f.name for f in annot_files):\n parc = \"HCP-MMP1\"\n else:\n parc = \"HCPMMP1\"\n\n out: Dict[str, List[np.ndarray]] = {}\n for hemi_code, offset in ((\"lh\", 0), (\"rh\", max_v)):\n labels = mne.read_labels_from_annot(\n \"fsaverage\", parc, hemi=hemi_code, subjects_dir=subjects_dir\n )\n for lab in labels:\n name = (lab.name\n .replace(\"_ROI-lh\",\"\").replace(\"_ROI-rh\",\"\")\n .replace(\"_ROI\",\"\").replace(\"-lh\",\"\").replace(\"-rh\",\"\"))\n verts = np.asarray(lab.vertices, dtype=np.int64)\n verts = verts[verts < max_v] + offset\n if verts.size:\n out.setdefault(name, []).append(verts)\n return {k: np.concatenate(v) for k, v in out.items()}\n\n\ndef get_hcp_roi_indices(rois, *, hemi=\"both\", mesh=\"fsaverage5\") -> np.ndarray:\n labels = _hcp_label_vertices(mesh)\n names = [rois] if isinstance(rois, str) else list(rois)\n selected = []\n for roi in names:\n if roi.endswith(\"*\"): selected.extend(k for k in labels if k.startswith(roi[:-1]))\n elif roi.startswith(\"*\"): selected.extend(k for k in labels if k.endswith(roi[1:]))\n elif roi in labels: selected.append(roi)\n else: raise ValueError(f\"ROI {roi!r} not found in HCP-MMP labels\")\n idx_parts = []\n for name in selected:\n verts = labels[name]\n if hemi == \"left\": idx_parts.append(verts[verts < FSAVERAGE5_VERTS_PER_HEMI])\n elif hemi == \"right\": idx_parts.append(verts[verts >= FSAVERAGE5_VERTS_PER_HEMI])\n else: idx_parts.append(verts)\n if not idx_parts:\n return np.array([], dtype=np.int64)\n # Fix 9: deduplicate vertices (wildcard matches can overlap)\n return np.unique(np.concatenate(idx_parts))\n\n\ndef roi_dict_to_fsaverage5(roi_values: Mapping[str, float]) -> Tuple[np.ndarray, List[str]]:\n acc = np.zeros(FSAVERAGE5_VERTICES, dtype=np.float64)\n counts = np.zeros(FSAVERAGE5_VERTICES, dtype=np.float64)\n skipped = []\n for roi_name, value in roi_values.items():\n targets = ROI56_TO_HCP.get(roi_name, [(roi_name, \"both\")])\n if not targets:\n skipped.append(roi_name); continue\n painted = False\n for hcp_name, hemi in targets:\n try:\n idx = get_hcp_roi_indices(hcp_name, hemi=hemi)\n except ValueError:\n continue\n if idx.size == 0:\n continue\n acc[idx] += float(value); counts[idx] += 1.0; painted = True\n if not painted:\n skipped.append(roi_name)\n mask = counts > 0\n acc[mask] /= counts[mask]\n return acc.astype(np.float32), skipped\n\n\ndef plot_brain(\n roi_map: Mapping[str, float],\n *,\n views: Sequence[str] = (\"left\", \"right\", \"dorsal\"),\n cmap: str = \"RdBu_r\",\n vmax: float | None = None,\n threshold: float = 0.02,\n title: str | None = None,\n):\n \"\"\"Paint roi_map on fsaverage5 and plot with Nilearn.\n\n views: any subset of (\"left\", \"right\", \"dorsal\", \"ventral\").\n MNE HCP-MMP1 parcellation downloads on first use (~50 MB).\n \"\"\"\n import matplotlib.pyplot as plt\n from nilearn.datasets import fetch_surf_fsaverage\n from nilearn.plotting import plot_surf_stat_map\n\n vertex_map, skipped = roi_dict_to_fsaverage5(roi_map)\n if skipped:\n print(f\" [brain plotter] skipped (no HCP mapping): {skipped}\")\n\n v = np.asarray(vertex_map, dtype=np.float32)\n # Fix 7: nanpercentile is safe when some vertices are zero/NaN\n if vmax is None:\n vmax = float(np.nanpercentile(np.abs(v), 99)) or 1.0\n\n fsa = fetch_surf_fsaverage(mesh=\"fsaverage5\")\n left = v[:FSAVERAGE5_VERTS_PER_HEMI]\n right = v[FSAVERAGE5_VERTS_PER_HEMI:]\n\n VIEW_SPEC = {\n \"left\": (\"left\", \"lateral\", \"infl_left\", \"sulc_left\"),\n \"right\": (\"right\", \"lateral\", \"infl_right\", \"sulc_right\"),\n \"dorsal\": (\"left\", \"dorsal\", \"infl_left\", \"sulc_left\"),\n \"ventral\": (\"left\", \"ventral\", \"infl_left\", \"sulc_left\"),\n }\n n = len(views)\n fig, axes = plt.subplots(1, n, figsize=(4.2*n, 3.8), subplot_kw={\"projection\": \"3d\"})\n if n == 1:\n axes = [axes]\n for ax, view in zip(axes, views):\n hemi, nv, infl_key, sulc_key = VIEW_SPEC.get(view, VIEW_SPEC[\"left\"])\n stat = left if hemi == \"left\" else right\n # Fix 3: surf_mesh first (nilearn keyword argument order)\n plot_surf_stat_map(\n surf_mesh=fsa[infl_key], stat_map=stat, bg_map=fsa[sulc_key],\n view=nv, axes=ax, cmap=cmap, vmin=-vmax, vmax=vmax,\n threshold=threshold, colorbar=False, symmetric_cbar=True,\n )\n ax.set_title(view, fontsize=9)\n sm = plt.cm.ScalarMappable(cmap=cmap, norm=plt.Normalize(vmin=-vmax, vmax=vmax))\n sm.set_array([])\n # Fix 6: ax=list(axes) avoids matplotlib deprecation warning\n fig.colorbar(sm, ax=list(axes), shrink=0.55, label=\"Predicted activation (a.u.)\")\n if title:\n fig.suptitle(title, fontsize=11)\n plt.tight_layout()\n plt.show()\n return fig\n\nprint(\"Brain plotter ready.\")\nprint(\" Usage: plot_brain(roi_map, title='My stimulus')\")\nprint(\" HCP-MMP1 parcellation downloads on first call (~50 MB via fetch_fsaverage).\")\n"
},
{
"cell_type": "code",
"execution_count": null,
"id": "brain-plotter-usage",
"metadata": {},
"outputs": [],
"source": "# ββ Plot predictions on the brain βββββββββββββββββββββββββββββββββββββββββββββ\n# Run any inference cell above first (Option A, B, or C), then run this cell.\n\n_map = (\n roi_map if 'roi_map' in dir() else\n roi_map_audio if 'roi_map_audio' in dir() else\n roi_map_text if 'roi_map_text' in dir() else\n None\n)\n\nif _map is None:\n print(\"Run an inference cell first (Option A, B, or C) to produce a roi_map.\")\nelse:\n # 4-view: left/right lateral + dorsal + ventral\n plot_brain(\n _map,\n views=(\"left\", \"right\", \"dorsal\", \"ventral\"),\n cmap=\"RdBu_r\",\n title=\"Predicted brain activation (fsaverage5, HCP-MMP1)\",\n )\n\n# ---- Compare two stimuli side-by-side ----------------------------------------\n# fig1 = plot_brain(roi_map_audio, title=\"Audio stimulus\")\n# fig2 = plot_brain(roi_map_text, title=\"Text stimulus\")\n\n# ---- Custom views (e.g. medial-only) -----------------------------------------\n# plot_brain(_map, views=(\"dorsal\", \"ventral\"), cmap=\"hot\", vmax=0.5)\n"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": "# ββ Option A: pre-cached text examples (CPU, instant, no encoder needed) βββββ\n# Uses examples_cache.npz bundled in the zip or downloaded from HuggingFace.\n# Skipped automatically if the file is unavailable.\n#\n# NOTE: The cache holds 2560-d Qwen3 embeddings β this cell always uses\n# text2roi_combined_v4.pt (in_dim=3840), regardless of the MODALITY setting above.\n\nif examples_path is None:\n print('examples_cache.npz not available β skip to Option B (audio) or Option C (text).')\nelse:\n # Load combined model (in_dim=3840) β required for Qwen3 2560-d embeddings\n _A_MODEL_FILE = 'text2roi_combined_v4.pt'\n _A_IN_DIM = 3840\n _a_ckpt_path = _resolve(_A_MODEL_FILE)\n _a_ckpt = torch.load(_a_ckpt_path, map_location='cpu', weights_only=False)\n _a_model = Text2ROI(in_dim=_a_ckpt['in_dim'], out_dim=_a_ckpt.get('n_roi', 56))\n _a_model.load_state_dict(_a_ckpt['state_dict'])\n _a_model.eval()\n print(f'Option A model: in_dim={_a_ckpt[\"in_dim\"]} val_R={_a_ckpt.get(\"best_val_r\",\"N/A\")}')\n\n cache = np.load(examples_path, allow_pickle=True)\n sentences = [s.decode() if isinstance(s, bytes) else str(s) for s in cache['sentences']]\n embeddings = cache['embeddings'].astype(np.float32) # (N, 2560)\n print(f'Loaded {len(sentences)} cached examples embedding_dim={embeddings.shape[1]}')\n\n # ββ pick a sentence to visualise ββ\n idx = 0 # change index or replace sentence below\n sentence = sentences[idx]\n roi_map = predict_from_features(embeddings[idx], _a_model, _A_IN_DIM)\n\n print(f'\\nInput: \"{sentence}\"')\n print('Top 10 ROIs:')\n for roi, val in sorted(roi_map.items(), key=lambda x: -x[1])[:10]:\n print(f' {roi:<16} {val:+.4f}')\n plot_rois(roi_map, title=sentence)\n",
"id": "cached-examples"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# ββ Option B: Audio inference with Whisper v4 (GPU recommended) ββββββββββββββ\n",
"# Upload a .wav file to Colab, or set AUDIO_PATH to a URL.\n",
"# Whisper-large-v3 downloads automatically (~3 GB, cached after first run).\n",
"\n",
"AUDIO_PATH = 'your_audio.wav' # <-- replace with your file\n",
"\n",
"# Uncomment to download a sample clip:\n",
"# import urllib.request\n",
"# urllib.request.urlretrieve('https://upload.wikimedia.org/wikipedia/commons/2/22/Beethoven_Moonlight_Sonata_First_Movement.ogg', 'sample.ogg')\n",
"# AUDIO_PATH = 'sample.ogg'\n",
"\n",
"import librosa\n",
"from transformers import WhisperProcessor, WhisperModel\n",
"\n",
"print('Loading Whisper-large-v3 (downloads ~3 GB on first run)...')\n",
"processor = WhisperProcessor.from_pretrained('openai/whisper-large-v3')\n",
"whisper_enc = WhisperModel.from_pretrained('openai/whisper-large-v3')\n",
"whisper_enc.eval()\n",
"print('Loaded.')\n",
"\n",
"audio, sr = librosa.load(AUDIO_PATH, sr=16000, mono=True)\n",
"inputs = processor(audio, sampling_rate=16000, return_tensors='pt')\n",
"with torch.no_grad():\n",
" enc_out = whisper_enc.encoder(inputs.input_features)\n",
"audio_emb = enc_out.last_hidden_state.mean(dim=1).squeeze(0).cpu().numpy() # (1280,)\n",
"print(f'Audio embedding: {audio_emb.shape}')\n",
"\n",
"roi_map_audio = predict_from_features(audio_emb, model, IN_DIM)\n",
"print('\\nTop 10 ROIs (audio):')\n",
"for roi, val in sorted(roi_map_audio.items(), key=lambda x: -x[1])[:10]:\n",
" print(f' {roi:<16} {val:+.4f}')\n",
"plot_rois(roi_map_audio, title=f'Audio: {AUDIO_PATH}')"
],
"id": "audio-inference"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": "# ββ Option C: Custom text inference with Qwen3 (GPU recommended, ~8 GB) ββββββ\n# NOTE: this cell always uses text2roi_combined_v4.pt (in_dim=3840), regardless\n# of the MODALITY setting above. The combined model was trained on both audio and\n# text; for text-only input the whisper slot is zero-padded automatically.\n# For best results on audio content, prefer Option B (actual audio file β Whisper).\nimport torch.nn.functional as F\nfrom transformers import AutoModel, AutoTokenizer\n\nYOUR_TEXT = 'listening to a symphony building to its climax'\n\n# ββ Load combined text model (always in_dim=3840) βββββββββββββββββββββββββββββ\n_TEXT_MODEL_FILE = 'text2roi_combined_v4.pt'\n_TEXT_IN_DIM = 3840\n\n_text_ckpt_path = _resolve(_TEXT_MODEL_FILE)\n_text_ckpt = torch.load(_text_ckpt_path, map_location='cpu', weights_only=False)\n_text_model = Text2ROI(in_dim=_text_ckpt['in_dim'], out_dim=_text_ckpt.get('n_roi', 56))\n_text_model.load_state_dict(_text_ckpt['state_dict'])\n_text_model.eval()\nprint(f'Text model loaded in_dim={_text_ckpt[\"in_dim\"]} val_R={_text_ckpt.get(\"best_val_r\",\"N/A\")}')\n\n# ββ Qwen3 embedding (last-token + L2 norm, matches training extraction) βββββββ\nprint('Loading Qwen3-Embedding-4B (downloads ~8 GB on first run)...')\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\ntok = AutoTokenizer.from_pretrained('Qwen/Qwen3-Embedding-4B', padding_side='left')\nqwen = AutoModel.from_pretrained(\n 'Qwen/Qwen3-Embedding-4B',\n torch_dtype=torch.bfloat16 if device != 'cpu' else torch.float32\n).to(device).eval()\n\nwith torch.no_grad():\n enc = tok([YOUR_TEXT], return_tensors='pt', padding=True, truncation=True, max_length=512).to(device)\n h = qwen(**enc).last_hidden_state[:, -1].float() # last token (decoder-style)\n emb = F.normalize(h, p=2, dim=1).cpu().numpy()[0].astype(np.float32) # (2560,)\n\ndel qwen; torch.cuda.empty_cache() if device == 'cuda' else None\nprint(f'Text embedding: {emb.shape} norm={float(np.linalg.norm(emb)):.4f}')\n\n# ββ Predict using the combined model βββββββββββββββββββββββββββββββββββββββββ\nroi_map_text = predict_from_features(emb, _text_model, _TEXT_IN_DIM)\nprint(f'\\nInput: \"{YOUR_TEXT}\"')\nprint('Top 10 ROIs:')\nfor roi, val in sorted(roi_map_text.items(), key=lambda x: -x[1])[:10]:\n print(f' {roi:<16} {val:+.4f}')\nplot_rois(roi_map_text, title=YOUR_TEXT)\n",
"id": "custom-text"
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# ββ Network-level summary βββββββββββββββββββββββββββββββββββββββββββββββββββββ\n",
"NETWORKS = {\n",
" 'Visual': ['V1','V2','V3','V4','V3A','V3B','LO1','LO2','MT','MST','V7','IPS1','FFA-1','FFA-2','PPA','RSC','OFA','EBA'],\n",
" 'Parietal': ['IPS2','IPS3','IPS4','IPS5','SPL1','hIP1','hIP2','hIP3'],\n",
" 'Frontal': ['dlPFC','vlPFC','OFC','ACC','mPFC','FP1','FP2'],\n",
" 'Language': ['IFG','IFGorb','STG','STS','MTG','AG'],\n",
" 'DefaultMode': ['PCC','mPFC_dmn','LP_L','LP_R','HPC_L','HPC_R'],\n",
" 'Salience': ['AI','dACC','sgACC','vmPFC','Amygdala_L','Amygdala_R'],\n",
" 'Subcortical': ['Caudate_L','Caudate_R','Putamen_L','Putamen_R','Thalamus'],\n",
"}\n",
"\n",
"# swap roi_map for whichever you ran above\n",
"active_map = roi_map if 'roi_map' in dir() else (roi_map_audio if 'roi_map_audio' in dir() else roi_map_text)\n",
"\n",
"nets = {net: float(np.mean([active_map[r] for r in rois if r in active_map]))\n",
" for net, rois in NETWORKS.items()}\n",
"\n",
"print('Network activations:')\n",
"for net, val in sorted(nets.items(), key=lambda x: -x[1]):\n",
" bar = '\\u2588' * max(0, int((val + 0.3) * 18))\n",
" print(f' {net:<14} {val:+.4f} {bar}')"
],
"id": "network-summary"
}
]
} |