File size: 20,440 Bytes
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
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
"""

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 schema ─────────────────────────────────────────────────────────────────
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

# ── 51 example stimuli organized by category ───────────────────────────────────
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}

# ── Model ──────────────────────────────────────────────────────────────────────
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] = {}  # sentence -> (56,) roi scores

def _load_model():
    global _model, _roi_names
    if _model is not None:
        return
    # 1. bundled alongside this script (self-contained zip)
    _HERE = Path(__file__).parent
    ckpt_path = _HERE / "text2roi_combined_v4.pt"
    # 2. CWD fallback
    if not ckpt_path.exists():
        ckpt_path = Path("text2roi_combined_v4.pt")
    # 3. HuggingFace (online fallback for HF Spaces / Colab)
    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
    # try bundled location first, then CWD
    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)  # (N, 2560)
    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()
    # Text-only mode: prepend zeros for the whisper portion (model trained with modality dropout)
    zeros = np.zeros((1, 1280), dtype=np.float32)
    inp = np.concatenate([zeros, qwen3_emb.reshape(1, -1)], axis=1)  # (1, 3840)
    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)


# ── Plotly chart ───────────────────────────────────────────────────────────────
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)


# ── Gradio callbacks ───────────────────────────────────────────────────────────
_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).", ""


# ── UI ─────────────────────────────────────────────────────────────────────────
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()