File size: 12,138 Bytes
4123863
49a8dc7
4123863
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49a8dc7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4123863
 
 
 
 
 
 
 
 
 
 
 
 
 
49a8dc7
 
4123863
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49a8dc7
 
 
 
 
 
 
4123863
 
49a8dc7
 
 
4123863
 
 
 
 
 
 
49a8dc7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4123863
 
 
 
 
49a8dc7
 
4123863
49a8dc7
 
4123863
 
 
 
49a8dc7
4123863
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
PhySH topic classifier β€” Gradio Space (ZeroGPU).

Pipeline:  text ──EmbeddingGemma-300m──> 768-d vector
                      β”‚
                      β”œβ”€β”€> discipline head (768 β†’ 1024 β†’ 512 β†’ 18), sigmoid
                      β”‚
                      └──> concept head ([768 + 18] β†’ 1024 β†’ 512 β†’ 186), sigmoid
                           conditioned on the discipline *probabilities*
                           (the checkpoint records use_logits = False)

Both heads are multi-label: each output is an independent sigmoid, so a text can
carry several disciplines and several concepts.
"""

from __future__ import annotations

import os

# `spaces` must be imported before torch β€” it patches CUDA init so the main
# process stays GPU-free until a @GPU function actually runs. ZeroGPU also scans
# for at least one decorated function at startup and stops the container if it
# finds none. The fallback keeps local runs and test_local.py working without it.
try:
    import spaces

    GPU = spaces.GPU
except ImportError:  # local development, or CPU hardware

    def GPU(*dargs, **dkwargs):
        if len(dargs) == 1 and callable(dargs[0]) and not dkwargs:
            return dargs[0]
        return lambda fn: fn


import gradio as gr
import torch
import torch.nn as nn
from huggingface_hub import hf_hub_download

# --------------------------------------------------------------------------- #
# Config
# --------------------------------------------------------------------------- #

MODEL_REPO = "LukeFP/physh_topic_supervised_classifier"
DISCIPLINE_CKPT = "discipline_classifier_gemma_20260130_140842.pt"
CONCEPT_CKPT = "concept_conditioned_gemma_20260130_140842.pt"
EMBED_MODEL = "google/embeddinggemma-300m"

# google/embeddinggemma-300m is gated: set HF_TOKEN as a Space *secret*, from an
# account that has accepted the Gemma license. Never commit the token itself.
HF_TOKEN = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")

# Set to a local directory to load the .pt files from disk instead of the Hub.
LOCAL_WEIGHTS_DIR = os.environ.get("PHYSH_WEIGHTS_DIR")

# EmbeddingGemma expects a task-specific prefix, and the prefix used at inference
# must match the one used to build the training embeddings β€” a mismatch degrades
# accuracy silently rather than erroring. Pick the one your training script used.
PROMPT_TEMPLATES = {
    "document β€” title: none | text: {}": "title: none | text: {}",
    "classification β€” task: classification | query: {}": "task: classification | query: {}",
    "none β€” raw text": "{}",
}
DEFAULT_PROMPT = "document β€” title: none | text: {}"

# --------------------------------------------------------------------------- #
# Model
# --------------------------------------------------------------------------- #


class MLPClassifier(nn.Module):
    """Linear/ReLU/Dropout stack. Layer indices line up with the checkpoints'
    `network.0`, `network.3`, `network.6` keys."""

    def __init__(self, input_dim: int, hidden_layers: list[int], output_dim: int, dropout: float):
        super().__init__()
        layers: list[nn.Module] = []
        prev = input_dim
        for width in hidden_layers:
            layers += [nn.Linear(prev, width), nn.ReLU(), nn.Dropout(dropout)]
            prev = width
        layers.append(nn.Linear(prev, output_dim))
        self.network = nn.Sequential(*layers)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.network(x)


def _weights_path(filename: str) -> str:
    if LOCAL_WEIGHTS_DIR:
        return os.path.join(LOCAL_WEIGHTS_DIR, filename)
    return hf_hub_download(MODEL_REPO, filename, token=HF_TOKEN)


def load_head(filename: str) -> tuple[MLPClassifier, dict]:
    ckpt = torch.load(_weights_path(filename), map_location="cpu", weights_only=False)
    cfg = ckpt["model_config"]
    # The discipline head records `input_dim`; the concept head records the two
    # halves of its input separately.
    input_dim = cfg.get("input_dim") or cfg["embedding_dim"] + cfg["discipline_dim"]
    model = MLPClassifier(input_dim, cfg["hidden_layers"], cfg["output_dim"], cfg["dropout"])
    model.load_state_dict(ckpt["model_state_dict"])
    model.eval()
    return model, ckpt


DISCIPLINE_MODEL, DISCIPLINE_CKPT_DATA = load_head(DISCIPLINE_CKPT)
CONCEPT_MODEL, CONCEPT_CKPT_DATA = load_head(CONCEPT_CKPT)

DISCIPLINE_LABELS = [d["label"] for d in DISCIPLINE_CKPT_DATA["class_labels"]]
CONCEPT_LABELS = [c["label"] for c in CONCEPT_CKPT_DATA["class_labels"]]

# The concept head was conditioned on the discipline vector in a specific order.
# Remap if the two checkpoints ever drift apart.
_CONDITION_ORDER = [d["discipline_id"] for d in CONCEPT_CKPT_DATA["discipline_labels"]]
_HEAD_ORDER = [d["discipline_id"] for d in DISCIPLINE_CKPT_DATA["class_labels"]]
_REMAP = torch.tensor([_HEAD_ORDER.index(i) for i in _CONDITION_ORDER], dtype=torch.long)

# EmbeddingGemma is loaded on CPU in the main process β€” under ZeroGPU nothing may
# touch CUDA outside a @GPU function, and the fork inherits this copy for free.
# A load failure is captured rather than raised so the Space still boots and can
# report the reason in the UI instead of crash-looping.
_EMBEDDER = None
_EMBEDDER_ERROR: str | None = None
try:
    from sentence_transformers import SentenceTransformer

    _EMBEDDER = SentenceTransformer(EMBED_MODEL, token=HF_TOKEN, device="cpu")
except Exception as exc:  # noqa: BLE001 β€” surfaced to the user verbatim
    _EMBEDDER_ERROR = f"{type(exc).__name__}: {exc}"


# --------------------------------------------------------------------------- #
# Inference
# --------------------------------------------------------------------------- #


@GPU(duration=60)
def infer(text: str, prompt_template: str) -> tuple[list[float], list[float]]:
    """Embed and run both heads. Returns plain lists β€” ZeroGPU pickles the return
    value across a process boundary, so nothing CUDA-resident may escape."""
    if _EMBEDDER is None:
        raise gr.Error(
            "EmbeddingGemma failed to load. It is a gated model, so the Space needs "
            "an HF_TOKEN secret from an account that has accepted the Gemma "
            f"license.\n\n{_EMBEDDER_ERROR}"
        )

    device = "cuda" if torch.cuda.is_available() else "cpu"
    embedder = _EMBEDDER.to(device)
    discipline_model = DISCIPLINE_MODEL.to(device)
    concept_model = CONCEPT_MODEL.to(device)
    remap = _REMAP.to(device)

    with torch.inference_mode():
        vector = embedder.encode(
            prompt_template.format(text),
            prompt="",  # stop ST applying the model's own default prefix on top
            convert_to_numpy=True,
        )
        embedding = torch.as_tensor(vector, dtype=torch.float32, device=device).unsqueeze(0)

        discipline_probs = torch.sigmoid(discipline_model(embedding))[0]
        conditioned = torch.cat([embedding, discipline_probs[remap].unsqueeze(0)], dim=1)
        concept_probs = torch.sigmoid(concept_model(conditioned))[0]

    return discipline_probs.float().cpu().tolist(), concept_probs.float().cpu().tolist()


def classify(text: str, threshold: float, prompt_choice: str, top_k: int):
    text = (text or "").strip()
    if not text:
        return {}, {}, "Paste some text β€” a title and abstract work best."

    template = PROMPT_TEMPLATES.get(prompt_choice, PROMPT_TEMPLATES[DEFAULT_PROMPT])
    discipline_scores, concept_scores = infer(text, template)

    disciplines = dict(zip(DISCIPLINE_LABELS, discipline_scores))
    concepts = dict(zip(CONCEPT_LABELS, concept_scores))

    return (
        dict(sorted(disciplines.items(), key=lambda kv: -kv[1])[:top_k]),
        dict(sorted(concepts.items(), key=lambda kv: -kv[1])[:top_k]),
        _summarize(disciplines, concepts, threshold),
    )


def _summarize(disciplines: dict, concepts: dict, threshold: float) -> str:
    def above(scores):
        hits = sorted((kv for kv in scores.items() if kv[1] >= threshold), key=lambda kv: -kv[1])
        return [f"**{name}** ({score:.2f})" for name, score in hits]

    d_hits, c_hits = above(disciplines), above(concepts)
    lines = [f"### Above threshold ({threshold:.2f})", ""]
    lines.append("**Disciplines** β€” " + (", ".join(d_hits) if d_hits else "_none_"))
    lines.append("")
    lines.append("**Concepts** β€” " + (", ".join(c_hits) if c_hits else "_none_"))
    if not d_hits and not c_hits:
        lines += ["", "_Nothing cleared the threshold. Lower it, or check that the "
                  "prompt format under Advanced matches your training setup._"]
    return "\n".join(lines)


# --------------------------------------------------------------------------- #
# UI
# --------------------------------------------------------------------------- #

EXAMPLES = [
    "We report the observation of a superconducting dome in magic-angle twisted "
    "bilayer graphene. Transport measurements below 1.7 K reveal a zero-resistance "
    "state whose critical temperature is tuned continuously by electrostatic gating, "
    "and the phase diagram closely tracks the filling of the flat moire bands.",
    "We present a measurement of the cosmic microwave background lensing power "
    "spectrum from four seasons of data. The reconstruction achieves a 40-sigma "
    "detection and, combined with baryon acoustic oscillation data, constrains the "
    "sum of the neutrino masses.",
    "A variational quantum eigensolver is used to compute ground-state energies of "
    "small molecular Hamiltonians on a superconducting processor. We introduce an "
    "error-mitigation scheme based on zero-noise extrapolation and show that it "
    "recovers chemical accuracy for LiH.",
]

with gr.Blocks(title="PhySH Topic Classifier") as demo:
    gr.Markdown(
        "# PhySH Topic Classifier\n"
        "Paste a physics title and abstract to get its **PhySH disciplines** and "
        "**top-level research-area concepts**. Both heads are multi-label, so several "
        "labels can fire at once.\n\n"
        f"Heads: [`{MODEL_REPO}`](https://huggingface.co/{MODEL_REPO}) Β· "
        f"Embeddings: [`{EMBED_MODEL}`](https://huggingface.co/{EMBED_MODEL})"
    )

    with gr.Row():
        with gr.Column(scale=3):
            text_input = gr.Textbox(
                label="Title + abstract",
                placeholder="Paste a paper title and abstract…",
                lines=12,
            )
            with gr.Row():
                submit = gr.Button("Classify", variant="primary")
                clear = gr.ClearButton(text_input, value="Clear")
            gr.Examples(examples=[[e] for e in EXAMPLES], inputs=[text_input], label="Try one")

        with gr.Column(scale=2):
            discipline_out = gr.Label(label="Disciplines (18)", num_top_classes=8)
            concept_out = gr.Label(label="Concepts (186)", num_top_classes=8)

    summary_out = gr.Markdown()

    with gr.Accordion("Advanced", open=False):
        threshold = gr.Slider(0.05, 0.95, value=0.5, step=0.05, label="Decision threshold")
        top_k = gr.Slider(3, 20, value=8, step=1, label="How many labels to show")
        prompt_choice = gr.Radio(
            choices=list(PROMPT_TEMPLATES),
            value=DEFAULT_PROMPT,
            label="EmbeddingGemma prompt format",
            info="Must match the prefix used to build the training embeddings. "
                 "If predictions look like noise, try the other options.",
        )
        gr.Markdown(
            f"Validation at training time β€” disciplines: micro-F1 "
            f"{DISCIPLINE_CKPT_DATA['metrics']['micro_f1']:.3f}, concepts: micro-F1 "
            f"{CONCEPT_CKPT_DATA['metrics']['micro_f1']:.3f}."
        )

    inputs = [text_input, threshold, prompt_choice, top_k]
    outputs = [discipline_out, concept_out, summary_out]
    submit.click(classify, inputs=inputs, outputs=outputs, api_name="classify")
    text_input.submit(classify, inputs=inputs, outputs=outputs)

if __name__ == "__main__":
    # Gradio 6 takes the theme on launch(), not on the Blocks constructor.
    demo.launch(theme=gr.themes.Soft())