""" 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())