LukeFP's picture
Run on ZeroGPU: add @spaces.GPU entry point
49a8dc7
Raw History Blame Contribute Delete
12.1 kB
"""
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())