File size: 6,373 Bytes
01c4f63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d51e0f
01c4f63
 
 
 
 
 
 
5d51e0f
 
 
 
da139d3
 
 
 
 
5d51e0f
 
 
 
da139d3
 
5d51e0f
 
 
 
 
 
 
 
 
759c21f
5d51e0f
 
759c21f
5d51e0f
 
 
759c21f
5d51e0f
759c21f
5d51e0f
 
759c21f
 
 
5d51e0f
 
01c4f63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1c0d50b
 
 
01c4f63
1c0d50b
01c4f63
1c0d50b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
01c4f63
 
 
 
 
 
 
 
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
"""
Temporal Cell State Prediction β€” interactive demo (HF Space).

Frozen V-JEPA-2 ViT-L + a small attentive-pool head predicts per-cell cell-cycle
state (interphase / pre-mitosis / mitosis) directly from short single-cell clips,
replacing the classify stage of the segment->track->classify pipeline.

The Space serves PRE-COMPUTED results (no heavy model load on the free CPU tier):
a classification gallery (successes + failure modes), state-overlay + tracking videos,
the per-state counting probe, and a zero-shot VLM point-reasoning baseline.

Thesis of the study: in this small-data regime, *data scaling*, not model scaling,
is the binding constraint.
"""

from __future__ import annotations

import json
from pathlib import Path

import gradio as gr

ASSETS = Path(__file__).parent / "assets"


def _load_json(name: str, default):
    p = ASSETS / name
    return json.loads(p.read_text()) if p.exists() else default


METRICS = _load_json("metrics.json", {})
GALLERY = _load_json("gallery/gallery_manifest.json", [])


# ── Header / thesis ───────────────────────────────────────────────────────────
HEADER = """
# πŸ”¬ Temporal State Prediction
**Frozen V-JEPA-2 ViT-L + attentive-pool head β†’ per-cell cell-cycle state**, straight from a tracked clip β€”
no segment→track→**classify** hand-off.

> **Finding:** a frozen video foundation model is competitive with a purpose-built
> morphology–temporal baseline, and the binding constraint is **data scale, not model capacity.**
"""

LABELS3 = ["interphase", "pre-mitosis", "mitosis"]


def _selector_items():
    """Compact clip thumbnails for the bottom selector.

    Caption carries only a βœ…/❌ correct-vs-miss tag (not the predicted class), so you
    can browse all clips at a glance and choose a correct one or a misclassification.
    """
    items = []
    for g in GALLERY:
        img = ASSETS / "gallery" / g.get("raw_gif", g.get("gif", ""))
        if img.exists():
            tag = "βœ… correct" if g.get("correct") else "❌ miss"
            items.append((str(img), tag))
    return items


def _valid_indices():
    return [i for i, g in enumerate(GALLERY)
            if (ASSETS / "gallery" / g.get("raw_gif", g.get("gif", ""))).exists()]


def _clip_view(i: int):
    """Return (image_path, result_markdown) for gallery item i β€” predicted vs actual only."""
    g = GALLERY[i]
    img = str(ASSETS / "gallery" / g.get("raw_gif", g.get("gif", "")))
    icon = "βœ…" if g.get("correct") else "❌"
    note = ""
    if not g.get("correct"):
        if g["true"] == "pre-mitosis" and g["pred"] == "interphase":
            note = "\n\n*'Pre-mitosis' is a soft, lineage-defined window β€” this nucleus looks identical to interphase (not yet rounded/condensed).*"
        elif g["true"] == "mitosis":
            note = "\n\n*A missed division β€” the rare, high-value event.*"
        elif g["pred"] == "pre-mitosis":
            note = "\n\n*Over-called: interphase nucleus flagged as entering division.*"
    md = (f"### Predicted: **{g['pred']}**\n\n"
          f"### Actual: **{g['true']}**\n\n"
          f"{icon} {'correct' if g.get('correct') else 'incorrect'}{note}")
    return img, md


def _metrics_md() -> str:
    m = METRICS
    if not m:
        return "_metrics.json not found_"
    cm = m.get("confusion")
    lines = [
        "### Held-out HeLa (sequence 02, n=%s)" % m.get("n", "?"),
        "",
        "| model | macro-F1 | mitosis F1 | mitosis event P/R (Β±3fr) |",
        "|---|---|---|---|",
        f"| U-Net+BiLSTM baseline (3.8M) | {m.get('baseline_macro_f1','?')} | {m.get('baseline_mitosis_f1','?')} | {m.get('baseline_mit_pr','?')} |",
        f"| **frozen V-JEPA-2 head-only** | **{m.get('vjepa_macro_f1','?')}** | {m.get('vjepa_mitosis_f1','?')} | {m.get('vjepa_mit_pr','?')} |",
        "",
        f"*Data scaling (GOWT1β†’HeLa) lifts the same baseline +0.186 macro-F1; ~80Γ— model scaling adds only +0.046.*",
        f"*Seed band: {m.get('seed_band','0.635 Β± 0.098')} β€” single-seed gaps <0.08 are not significant.*",
    ]
    if cm:
        lines += [
            "",
            "**Confusion matrix** (rows = true, cols = pred):",
            "",
            "| true β§΅ pred | interphase | pre-mitosis | mitosis | recall |",
            "|---|---|---|---|---|",
        ]
        names = ["interphase", "pre-mitosis", "mitosis"]
        for i, nm in enumerate(names):
            row = cm[i]; rec = row[i] / max(sum(row), 1)
            lines.append(f"| **{nm}** | {row[0]} | {row[1]} | {row[2]} | {rec:.2f} |")
        lines.append("")
        lines.append("*Dominant error: pre-mitosis→interphase — a soft, lineage-defined 8-frame window with no sharp morphological boundary.*")
    return "\n".join(lines)


THEME = gr.themes.Soft(font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"])


def build():
    with gr.Blocks(title="Temporal State Prediction", theme=THEME) as demo:
        gr.Markdown(HEADER)
        gr.Markdown("**Select a single-cell clip below** β€” the model classifies that one clip into its cell-cycle state.")
        valid = _valid_indices()
        first = valid[0] if valid else 0
        img0, md0 = _clip_view(first) if valid else (None, "_no clips found_")
        with gr.Row():
            sel_clip = gr.Image(value=img0, label="selected clip", height=300)
            sel_md = gr.Markdown(md0)
        selector = gr.Gallery(value=_selector_items(), columns=10, height=170,
                              object_fit="cover", label="β–Ό pick a clip (βœ… correct Β· ❌ misclassified)",
                              allow_preview=False)

        def _on_select(evt: gr.SelectData):
            return _clip_view(evt.index)
        selector.select(_on_select, inputs=None, outputs=[sel_clip, sel_md])

        with gr.Accordion("Held-out test-set metrics (all 5,312 clips)", open=False):
            gr.Markdown(_metrics_md())

        gr.Markdown("---\nModels: `DnaRnaProteins/vjepa2-cell-cycle-vit-l`, `DnaRnaProteins/unet-bilstm-cell-cycle-baseline` Β· "
                    "Data: MICCAI Cell Tracking Challenge (Fluo-N2DL-HeLa). Labels derived from lineage trees (no manual annotation).")
    return demo


if __name__ == "__main__":
    build().launch()