Spaces:
Running on Zero
Running on Zero
Download app.py from LukeFP/Physh_Classification: direct link, hf CLI and curl.
- Browser
- Download file 12.1 kB
-
https://huggingface.co/spaces/LukeFP/Physh_Classification/resolve/main/app.py
- Command line
-
hf download hf://spaces/LukeFP/Physh_Classification/app.py
-
curl -L -o app.py https://huggingface.co/spaces/LukeFP/Physh_Classification/resolve/main/app.py
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 | |
| # --------------------------------------------------------------------------- # | |
| 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()) | |