| """
|
| Amphora NeuroText β Gradio Space Demo
|
|
|
| Predicts brain region activation from text using the text2roi_combined_v4.pt model.
|
| Pre-cached Qwen3 embeddings for 51 example stimuli let this run on CPU instantly.
|
| Custom text inference requires the full Qwen3-Embedding-4B model (GPU recommended).
|
|
|
| Model: text2roi_combined_v4.pt (val R=0.192, honest cross-subject holdout)
|
| Audio model: text2roi_whisper_v4.pt beats TRIBE v2 by +4.2% (R=0.257, 23 held-out subjects)
|
| """
|
| from __future__ import annotations
|
|
|
| import json
|
| from pathlib import Path
|
| from typing import Dict, List
|
|
|
| import gradio as gr
|
| import numpy as np
|
| import plotly.graph_objects as go
|
| import torch
|
| import torch.nn as nn
|
|
|
|
|
| ROI_NAMES: List[str] = [
|
| "V1","V2","V3","V4","V3A","V3B","LO1","LO2",
|
| "MT","MST","V7","IPS1","FFA-1","FFA-2","PPA","RSC",
|
| "OFA","EBA","IPS2","IPS3","IPS4","IPS5","SPL1",
|
| "hIP1","hIP2","hIP3","dlPFC","vlPFC","OFC","ACC",
|
| "mPFC","FP1","FP2","IFG","IFGorb","STG","STS",
|
| "MTG","AG","PCC","mPFC_dmn","LP_L","LP_R",
|
| "HPC_L","HPC_R","AI","dACC","sgACC","vmPFC",
|
| "Amygdala_L","Amygdala_R","Caudate_L","Caudate_R",
|
| "Putamen_L","Putamen_R","Thalamus",
|
| ]
|
|
|
| NETWORK_COLORS = {
|
| "Visual": ("#4B8BBE", ["V1","V2","V3","V4","V3A","V3B","LO1","LO2","MT","MST","V7","IPS1","FFA-1","FFA-2","PPA","RSC","OFA","EBA"]),
|
| "Parietal": ("#6AB187", ["IPS2","IPS3","IPS4","IPS5","SPL1","hIP1","hIP2","hIP3"]),
|
| "Frontal": ("#E07B39", ["dlPFC","vlPFC","OFC","ACC","mPFC","FP1","FP2"]),
|
| "Language": ("#9B59B6", ["IFG","IFGorb","STG","STS","MTG","AG"]),
|
| "Default Mode": ("#E74C3C", ["PCC","mPFC_dmn","LP_L","LP_R","HPC_L","HPC_R"]),
|
| "Salience": ("#F39C12", ["AI","dACC","sgACC","vmPFC","Amygdala_L","Amygdala_R"]),
|
| "Subcortical": ("#95A5A6", ["Caudate_L","Caudate_R","Putamen_L","Putamen_R","Thalamus"]),
|
| }
|
|
|
| ROI_TO_NET: Dict[str, str] = {}
|
| ROI_TO_COLOR: Dict[str, str] = {}
|
| for net, (col, rois) in NETWORK_COLORS.items():
|
| for r in rois:
|
| ROI_TO_NET[r] = net
|
| ROI_TO_COLOR[r] = col
|
|
|
|
|
| EXAMPLES = {
|
| "Face": [
|
| "the photograph showed a person raising an eyebrow in surprise",
|
| "she memorized the distinctive features of every face she met",
|
| "the newborn could already distinguish its mother's face from a stranger's",
|
| "an upside-down portrait makes it harder to recognize the person's identity",
|
| "identical twins are notoriously difficult to tell apart by facial features alone",
|
| ],
|
| "Scene": [
|
| "navigating the winding streets of an unfamiliar city neighbourhood",
|
| "the cabin sat in a dense forest clearing surrounded by tall pines",
|
| "she recognised the museum lobby from a single glimpse of its architecture",
|
| "the aerial view revealed a patchwork of farmland stretching to the horizon",
|
| "every corner of the childhood home was etched into their spatial memory",
|
| ],
|
| "Object": [
|
| "identifying the make and model of a vintage car from across the street",
|
| "the toolbox contained wrenches, pliers, and screwdrivers of every size",
|
| "grasping the difference between a cup and a bowl is trivial for humans",
|
| "the robotic arm picked up each item and sorted it into the correct bin",
|
| ],
|
| "Motor": [
|
| "the gymnast twisted her body into an impossible-looking backflip",
|
| "tying a shoelace is a motor skill that becomes automatic with practice",
|
| "the surgeon's hands moved with practised precision during the procedure",
|
| "drumming requires independent coordination of all four limbs simultaneously",
|
| ],
|
| "Language": [
|
| "the professor paused mid-sentence to choose a more precise word",
|
| "translating idioms between languages often loses the original meaning",
|
| "parsing a garden-path sentence requires revising your initial interpretation",
|
| "the radio announcer's voice was immediately recognisable to regular listeners",
|
| "metaphors allow us to understand abstract ideas through concrete comparisons",
|
| ],
|
| "Auditory": [
|
| "a sudden loud bang echoed through the empty corridor",
|
| "the melody of the piano piece lingered long after the concert ended",
|
| "distinguishing two similar vowel sounds is harder in a second language",
|
| ],
|
| "Math": [
|
| "estimating how many bricks it would take to fill the room",
|
| "the pattern of prime numbers has fascinated mathematicians for centuries",
|
| "keeping a running total while counting backwards from a hundred",
|
| ],
|
| "Attention":[
|
| "spotting the single red dot among hundreds of blue ones in a crowded display",
|
| "ignoring the conversation at the next table while trying to concentrate",
|
| ],
|
| "WM": [
|
| "holding seven random digits in mind while answering an unrelated question",
|
| "remembering the exact words of a sentence heard thirty seconds ago",
|
| ],
|
| "Fear": [
|
| "hearing an unexpected rustling sound in a dark forest at midnight",
|
| "the suspense built as the footsteps grew louder outside the locked door",
|
| "spotting a venomous spider sitting motionless on your pillow",
|
| ],
|
| "Disgust": [
|
| "the smell of rotting food was overwhelming as the bin had not been emptied for weeks",
|
| "the sight of the infected wound made his stomach turn",
|
| ],
|
| "Reward": [
|
| "the unexpected bonus triggered an immediate sense of pleasure and relief",
|
| "biting into a perfectly ripe piece of fruit on a hot summer day",
|
| "the slot machine paid out a jackpot after hours of near-misses",
|
| ],
|
| "Social": [
|
| "guessing what your friend is about to say before they finish the sentence",
|
| "recognising that someone is being sarcastic without explicit signals",
|
| "imagining how a stranger might feel after receiving devastating news",
|
| ],
|
| "Memory": [
|
| "replaying the exact sequence of events from a memorable birthday party",
|
| "the smell of cinnamon instantly transported her back to her grandmother's kitchen",
|
| "mentally walking through every room of a childhood home",
|
| "trying to recall whether you locked the door before leaving the house",
|
| ],
|
| "Pain": [
|
| "the throbbing headache made it impossible to focus on anything",
|
| "holding your breath underwater until your lungs burn",
|
| "the dentist's drill hitting a sensitive nerve sends a sharp jolt of pain",
|
| ],
|
| }
|
|
|
| ALL_EXAMPLES_FLAT = [s for sentences in EXAMPLES.values() for s in sentences]
|
| SENTENCE_TO_CAT = {s: cat for cat, sentences in EXAMPLES.items() for s in sentences}
|
|
|
|
|
| class Text2ROI(nn.Module):
|
| def __init__(self, in_dim=3840, hidden=1024, out_dim=56, dropout=0.1):
|
| super().__init__()
|
| self.net = nn.Sequential(
|
| nn.Linear(in_dim, hidden), nn.GELU(), nn.Dropout(dropout), nn.LayerNorm(hidden),
|
| nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout),
|
| nn.Linear(hidden // 2, out_dim),
|
| )
|
| def forward(self, x): return self.net(x)
|
|
|
|
|
| _model = None
|
| _roi_names = None
|
| _examples_cache: Dict[str, np.ndarray] = {}
|
|
|
| def _load_model():
|
| global _model, _roi_names
|
| if _model is not None:
|
| return
|
|
|
| _HERE = Path(__file__).parent
|
| ckpt_path = _HERE / "text2roi_combined_v4.pt"
|
|
|
| if not ckpt_path.exists():
|
| ckpt_path = Path("text2roi_combined_v4.pt")
|
|
|
| if not ckpt_path.exists():
|
| try:
|
| from huggingface_hub import hf_hub_download
|
| ckpt_path = Path(hf_hub_download("ffh92r32rm0/Amphora_NeuroText", "text2roi_combined_v4.pt"))
|
| except Exception as e:
|
| raise FileNotFoundError(
|
| "text2roi_combined_v4.pt not found locally or on HuggingFace. "
|
| "Make sure the .pt file is in the same folder as app.py."
|
| ) from e
|
| ckpt = torch.load(str(ckpt_path), map_location="cpu", weights_only=False)
|
| _roi_names = [s.decode() if isinstance(s, bytes) else str(s)
|
| for s in ckpt.get("roi_names", ROI_NAMES)]
|
| _model = Text2ROI(in_dim=ckpt["in_dim"], out_dim=ckpt["n_roi"])
|
| _model.load_state_dict(ckpt["state_dict"])
|
| _model.eval()
|
|
|
|
|
| def _load_examples():
|
| """Load pre-cached embeddings (NPZ with sentence index β qwen3 2560d vectors)."""
|
| _HERE = Path(__file__).parent
|
|
|
| cache = _HERE / "examples_cache.npz"
|
| if not cache.exists():
|
| cache = Path("examples_cache.npz")
|
| if not cache.exists():
|
| return False
|
| data = np.load(str(cache), allow_pickle=True)
|
| sentences = [s.decode() if isinstance(s, bytes) else str(s) for s in data["sentences"]]
|
| embs = data["embeddings"].astype(np.float32)
|
| for s, e in zip(sentences, embs):
|
| _examples_cache[s] = e
|
| return True
|
|
|
|
|
| def _run_inference(qwen3_emb: np.ndarray) -> Dict[str, float]:
|
| """Run combined projector on a (2560,) Qwen3 embedding in text-only mode."""
|
| _load_model()
|
|
|
| zeros = np.zeros((1, 1280), dtype=np.float32)
|
| inp = np.concatenate([zeros, qwen3_emb.reshape(1, -1)], axis=1)
|
| with torch.no_grad():
|
| pred = _model(torch.from_numpy(inp)).cpu().numpy()[0]
|
| return dict(zip(_roi_names or ROI_NAMES, pred.tolist()))
|
|
|
|
|
| def _embed_qwen3_live(text: str) -> np.ndarray:
|
| """Embed text with Qwen3-Embedding-4B (GPU strongly recommended)."""
|
| import torch.nn.functional as F
|
| from transformers import AutoModel, AutoTokenizer
|
| tok = AutoTokenizer.from_pretrained("Qwen/Qwen3-Embedding-4B", padding_side="left")
|
| device = "cuda" if torch.cuda.is_available() else "cpu"
|
| dtype = torch.bfloat16 if device != "cpu" else torch.float32
|
| mdl = AutoModel.from_pretrained("Qwen/Qwen3-Embedding-4B",
|
| torch_dtype=dtype).to(device).eval()
|
| with torch.no_grad():
|
| enc = tok([text], return_tensors="pt", padding=True,
|
| truncation=True, max_length=512).to(device)
|
| h = mdl(**enc).last_hidden_state[:, -1].float()
|
| emb = F.normalize(h, p=2, dim=1).cpu().numpy()[0]
|
| del mdl; torch.cuda.empty_cache()
|
| return emb.astype(np.float32)
|
|
|
|
|
|
|
| def _make_brain_chart(roi_map: Dict[str, float], title: str, category: str) -> go.Figure:
|
| names = list(roi_map.keys())
|
| scores = list(roi_map.values())
|
| colors = [ROI_TO_COLOR.get(n, "#95A5A6") for n in names]
|
| nets = [ROI_TO_NET.get(n, "Other") for n in names]
|
|
|
| order = sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)
|
| names = [names[i] for i in order]
|
| scores = [scores[i] for i in order]
|
| colors = [colors[i] for i in order]
|
| nets = [nets[i] for i in order]
|
|
|
| top15_names = names[:15]
|
| top15_scores = scores[:15]
|
| top15_colors = colors[:15]
|
| top15_nets = nets[:15]
|
|
|
| fig = go.Figure()
|
|
|
| fig.add_trace(go.Bar(
|
| x=top15_scores,
|
| y=top15_names,
|
| orientation="h",
|
| marker=dict(color=top15_colors, opacity=0.85,
|
| line=dict(color="rgba(255,255,255,0.3)", width=0.5)),
|
| customdata=top15_nets,
|
| hovertemplate="<b>%{y}</b><br>Activation: %{x:.3f}<br>Network: %{customdata}<extra></extra>",
|
| showlegend=False,
|
| ))
|
|
|
| fig.add_vline(x=0, line_color="rgba(150,150,150,0.5)", line_width=1)
|
|
|
| for net, (col, _) in NETWORK_COLORS.items():
|
| fig.add_trace(go.Bar(
|
| x=[None], y=[None],
|
| marker=dict(color=col),
|
| name=net,
|
| showlegend=True,
|
| ))
|
|
|
| fig.update_layout(
|
| title=dict(
|
| text=f"<b>Predicted brain activation</b><br><sub>{title[:80]}</sub>",
|
| font=dict(size=14),
|
| ),
|
| xaxis_title="Predicted activation (z-score)",
|
| yaxis=dict(autorange="reversed", tickfont=dict(size=11)),
|
| plot_bgcolor="rgba(20,20,30,0.95)",
|
| paper_bgcolor="rgba(20,20,30,0.0)",
|
| font=dict(color="#e0e0e0"),
|
| margin=dict(l=10, r=10, t=70, b=40),
|
| height=480,
|
| legend=dict(
|
| orientation="h",
|
| yanchor="bottom",
|
| y=-0.35,
|
| xanchor="center",
|
| x=0.5,
|
| font=dict(size=10),
|
| ),
|
| bargap=0.18,
|
| )
|
| return fig
|
|
|
|
|
| def _network_summary(roi_map: Dict[str, float]) -> str:
|
| net_scores: Dict[str, list] = {n: [] for n in NETWORK_COLORS}
|
| for roi, score in roi_map.items():
|
| net = ROI_TO_NET.get(roi)
|
| if net:
|
| net_scores[net].append(score)
|
| rows = []
|
| for net, scores in net_scores.items():
|
| if scores:
|
| mean = np.mean(scores)
|
| rows.append((net, mean))
|
| rows.sort(key=lambda x: -x[1])
|
| lines = [f"**Most activated network: {rows[0][0]}**\n"]
|
| for net, mean in rows:
|
| bar = "β" * max(0, int((mean + 0.5) * 12))
|
| lines.append(f"`{net:<14}` {mean:+.3f} {bar}")
|
| return "\n".join(lines)
|
|
|
|
|
|
|
| _cache_loaded = _load_examples()
|
|
|
| def predict_from_example(sentence: str):
|
| if not sentence:
|
| return None, "", ""
|
| cat = SENTENCE_TO_CAT.get(sentence, "Unknown")
|
| if sentence in _examples_cache:
|
| qwen3_emb = _examples_cache[sentence]
|
| roi_map = _run_inference(qwen3_emb)
|
| source = "Pre-cached embedding (instant)"
|
| else:
|
| return None, "β Embedding not cached. Use Custom Text tab.", ""
|
| fig = _make_brain_chart(roi_map, sentence, cat)
|
| summary = _network_summary(roi_map)
|
| top5 = "\n".join(f"**{i+1}. {r}** ({s:+.3f})"
|
| for i, (r, s) in enumerate(
|
| sorted(roi_map.items(), key=lambda x: -x[1])[:5]))
|
| return fig, f"**Category:** {cat} | {source}\n\n**Top 5 ROIs:**\n{top5}", summary
|
|
|
|
|
| def predict_from_custom(text: str):
|
| if not text or len(text.strip()) < 5:
|
| return None, "Please enter at least 5 characters.", ""
|
| try:
|
| qwen3_emb = _embed_qwen3_live(text.strip())
|
| roi_map = _run_inference(qwen3_emb)
|
| fig = _make_brain_chart(roi_map, text.strip(), "Custom")
|
| summary = _network_summary(roi_map)
|
| top5 = "\n".join(f"**{i+1}. {r}** ({s:+.3f})"
|
| for i, (r, s) in enumerate(
|
| sorted(roi_map.items(), key=lambda x: -x[1])[:5]))
|
| return fig, f"**Top 5 ROIs:**\n{top5}", summary
|
| except Exception as e:
|
| return None, f"β Error: {e}\n\nNote: Custom text requires Qwen3-Embedding-4B (GPU recommended).", ""
|
|
|
|
|
|
|
| CSS = """
|
| #title { text-align: center; }
|
| #subtitle { text-align: center; color: #aaa; margin-top: -10px; }
|
| .category-btn { font-size: 12px !important; }
|
| """
|
|
|
| DESCRIPTION = """
|
| **Amphora NeuroText** predicts which brain regions activate in response to any text stimulus β
|
| trained on **real fMRI data** from 1,600+ subjects across naturalistic experiments.
|
| No brain scanner needed at inference time.
|
|
|
| **Audio model beats TRIBE v2** (Meta AI, Algonauts 2025 winner) by **+4.2%** (R=0.257 vs 0.215).
|
| 10/10 cognitive category circuits correctly localized.
|
| """
|
|
|
| def build_interface():
|
| with gr.Blocks(css=CSS, title="Amphora NeuroText") as demo:
|
| gr.Markdown("# Amphora NeuroText", elem_id="title")
|
| gr.Markdown("### Text β Brain Region Activation", elem_id="subtitle")
|
| gr.Markdown(DESCRIPTION)
|
|
|
| with gr.Tabs():
|
| with gr.TabItem("Examples (instant, CPU)"):
|
| gr.Markdown("Select a cognitive category and example sentence:")
|
|
|
| with gr.Row():
|
| cat_dd = gr.Dropdown(
|
| label="Category",
|
| choices=list(EXAMPLES.keys()),
|
| value="Fear",
|
| scale=1,
|
| )
|
| sent_dd = gr.Dropdown(
|
| label="Example sentence",
|
| choices=EXAMPLES["Fear"],
|
| value=EXAMPLES["Fear"][2],
|
| scale=3,
|
| )
|
|
|
| predict_btn = gr.Button("Predict brain activation", variant="primary")
|
|
|
| with gr.Row():
|
| brain_plot = gr.Plot(label="Brain ROI Activations (top 15)")
|
|
|
| with gr.Row():
|
| result_md = gr.Markdown()
|
| network_md = gr.Markdown()
|
|
|
| cat_dd.change(
|
| fn=lambda cat: gr.Dropdown(choices=EXAMPLES[cat], value=EXAMPLES[cat][0]),
|
| inputs=cat_dd,
|
| outputs=sent_dd,
|
| )
|
| predict_btn.click(
|
| fn=predict_from_example,
|
| inputs=sent_dd,
|
| outputs=[brain_plot, result_md, network_md],
|
| )
|
| sent_dd.change(
|
| fn=predict_from_example,
|
| inputs=sent_dd,
|
| outputs=[brain_plot, result_md, network_md],
|
| )
|
|
|
| with gr.TabItem("Custom Text (GPU recommended)"):
|
| gr.Markdown("""
|
| Enter any text and see which brain regions the model predicts will activate.
|
|
|
| > **Note:** Custom text requires loading Qwen3-Embedding-4B (~8GB). This works on a GPU Space
|
| > but will be very slow on CPU. If you're running locally, install:
|
| > `pip install torch transformers`
|
| """)
|
| custom_text = gr.Textbox(
|
| label="Your text stimulus",
|
| placeholder='e.g. "hearing a jazz piano solo" or "solving a geometry puzzle"',
|
| lines=3,
|
| )
|
| custom_btn = gr.Button("Predict", variant="primary")
|
|
|
| with gr.Row():
|
| custom_plot = gr.Plot(label="Brain ROI Activations")
|
|
|
| with gr.Row():
|
| custom_result = gr.Markdown()
|
| custom_network = gr.Markdown()
|
|
|
| custom_btn.click(
|
| fn=predict_from_custom,
|
| inputs=custom_text,
|
| outputs=[custom_plot, custom_result, custom_network],
|
| )
|
|
|
| gr.Markdown("""
|
| ---
|
| **Model:** [Amphora_NeuroText](https://huggingface.co/ffh92r32rm0/Amphora_NeuroText) Β·
|
| **Training data:** Narratives, Little Prince, HCP-task, AOMIC, CNeuroMod, Clinical fMRI Β·
|
| **License:** MIT Β· **Contact:** hamiltonfrancesco5@gmail.com
|
| """)
|
|
|
| return demo
|
|
|
|
|
| if __name__ == "__main__":
|
| demo = build_interface()
|
| demo.launch()
|
|
|