File size: 5,924 Bytes
95456ed | 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 | import os
import sys
import json
import tempfile
from typing import Union, List, Dict, Any
import gradio as gr
import torch
from ScanDL2 import ScanDL2
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
_MODELS: Dict[str, ScanDL2] = {}
def get_model(text_type: str) -> ScanDL2:
"""Load and cache a ScanDL2 model for the given text type."""
if text_type not in _MODELS:
# `save=None` -> we handle saving ourselves in the Gradio app
_MODELS[text_type] = ScanDL2(text_type=text_type, save=None)
return _MODELS[text_type]
def predict(
text: str,
text_type: str,
progress=gr.Progress(track_tqdm=True),
) -> Dict[str, Any]:
"""
Run ScanDL2 on the input text and return a structured result.
"""
if text is None or text.strip() == "":
raise gr.Error("Please provide some input text.")
lines = [ln.strip() for ln in text.strip().split("\n") if ln.strip()]
if len(lines) == 0:
raise gr.Error("Input text is empty after cleaning.")
progress(0.05, desc="Loading ScanDL2 model...")
model = get_model(text_type)
progress(0.15, desc="Running ScanDL + FixDur modules...")
with torch.no_grad():
output = model(texts=lines)
return output
def format_output(output: Dict[str, Any]) -> str:
"""Pretty-print the ScanDL2 output for display in the UI."""
if output is None:
return ""
n = len(output.get("original_sn", []))
lines: List[str] = []
for i in range(n):
original_sn = output["original_sn"][i]
sp_words = output["predicted_sp_words"][i]
sp_ids = output["predicted_sp_ids"][i]
fix_durs = output["predicted_fix_durs"][i]
lines.append(f"### Example {i + 1}")
lines.append("")
lines.append(f"**Original text:** {' '.join(original_sn)}")
lines.append("")
# Build a readable scanpath table
lines.append("| # | Word | Word index | Fixation duration (ms) |")
lines.append("|---|------|-----------|------------------------|")
for step, (w, wid, dur) in enumerate(zip(sp_words, sp_ids, fix_durs), start=1):
lines.append(f"| {step} | {w} | {wid} | {dur} |")
lines.append("")
return "\n".join(lines)
def format_json(output: Dict[str, Any]) -> str:
"""Return the raw JSON string of the output."""
if output is None:
return "{}"
return json.dumps(output, indent=2, ensure_ascii=False)
DESCRIPTION = """
# ScanDL 2.0
**ScanDL 2.0** predicts human-like **eye-movement scanpaths** (which words are fixated, in what order)
and their **fixation durations** (in milliseconds) directly from text.
This Space wraps two jointly-trained modules:
1. **ScanDL module** — a discrete diffusion model that generates fixation *locations* (a scanpath) over the input text.
2. **FixDur module** — a sequence-to-sequence model that predicts the *duration* of each fixation.
### How to use
1. Paste your text in the box below.
- In **sentence** mode, each line is treated as a separate sentence.
- In **paragraph** mode, each line is treated as a separate paragraph.
2. Choose the text type (must match the model checkpoint you want to use).
3. Click **Run**.
### Output
- A human-readable scanpath table with predicted fixation durations per word.
- The raw JSON output (fixated words, word indices, and durations).
"""
EXAMPLES = [
[
"The quick brown fox jumps over the lazy dog.",
"sentence",
],
[
"Researchers have long been interested in how humans process written language.\n"
"Eye-tracking studies reveal where and for how long readers fixate on words.",
"paragraph",
],
]
def build_demo() -> gr.Blocks:
with gr.Blocks(
title="ScanDL 2.0 — Eye-Movement Scanpath Prediction",
theme=gr.themes.Soft(),
) as demo:
gr.Markdown(DESCRIPTION)
with gr.Row():
with gr.Column(scale=3):
text_in = gr.Textbox(
label="Input text",
placeholder="Paste a sentence or paragraph here...",
lines=8,
)
text_type_in = gr.Radio(
choices=["sentence", "paragraph"],
value="sentence",
label="Text type",
info="Must match an available ScanDL2 checkpoint.",
)
with gr.Row():
run_btn = gr.Button("Run", variant="primary")
clear_btn = gr.Button("Clear")
with gr.Column(scale=4):
table_out = gr.Markdown(
label="Predicted scanpath",
value="_Results will appear here._",
)
json_out = gr.Code(
label="Raw JSON output",
language="json",
value="{}",
)
gr.Examples(examples=EXAMPLES, inputs=[text_in, text_type_in])
def _run(text, text_type):
output = predict(text, text_type)
return format_output(output), format_json(output)
run_btn.click(
fn=_run,
inputs=[text_in, text_type_in],
outputs=[table_out, json_out],
)
clear_btn.click(
fn=lambda: ("", "sentence", "_Results will appear here._", "{}"),
inputs=None,
outputs=[text_in, text_type_in, table_out, json_out],
)
return demo
if __name__ == "__main__":
demo = build_demo()
demo.queue(max_size=16).launch(
server_name="0.0.0.0",
server_port=int(os.environ.get("PORT", 7860)),
show_error=True,
)
|