Spaces:
Running on Zero
Running on Zero
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())
|