"""Smoke test: loads the real checkpoints, stubs out EmbeddingGemma, and runs the full chain on CPU. Needs no HF token and downloads nothing. PHYSH_WEIGHTS_DIR=../physh_topic_supervised_classifier python test_local.py """ import os os.environ.setdefault( "PHYSH_WEIGHTS_DIR", os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "physh_topic_supervised_classifier"), ) import numpy as np import torch import app class FakeEmbedder: """Stands in for SentenceTransformer: deterministic unit-norm vectors.""" def to(self, device): return self def encode(self, text, prompt=None, convert_to_numpy=True): rng = np.random.default_rng(abs(hash(text)) % (2**32)) v = rng.normal(size=768).astype("float32") return v / np.linalg.norm(v) # EmbeddingGemma returns L2-normalised vectors app._EMBEDDER = FakeEmbedder() print("disciplines:", len(app.DISCIPLINE_LABELS), "concepts:", len(app.CONCEPT_LABELS)) print("remap is identity:", torch.equal(app._REMAP, torch.arange(18))) print("concept head in_features:", app.CONCEPT_MODEL.network[0].in_features) import sys print("real `spaces` package in use:", "spaces" in sys.modules) d, c, summary = app.classify(app.EXAMPLES[0], 0.5, app.DEFAULT_PROMPT, 8) assert len(d) == 8 and len(c) == 8, (len(d), len(c)) assert all(0.0 <= v <= 1.0 for v in {**d, **c}.values()) print("\ntop disciplines:", [f"{k} {v:.3f}" for k, v in d.items()][:3]) print("top concepts: ", [f"{k} {v:.3f}" for k, v in c.items()][:3]) # dropout must be off, so repeat calls agree d2, _, _ = app.classify(app.EXAMPLES[0], 0.5, app.DEFAULT_PROMPT, 8) assert d == d2, "eval() mode not applied — outputs are not deterministic" # empty input, and every prompt format assert app.classify(" ", 0.5, app.DEFAULT_PROMPT, 8)[0] == {} for p in app.PROMPT_TEMPLATES: app.classify("test abstract", 0.5, p, 5) # the GPU-decorated entry point returns plain lists (ZeroGPU pickles them) dl, cl = app.infer("test", app.PROMPT_TEMPLATES[app.DEFAULT_PROMPT]) assert isinstance(dl, list) and isinstance(cl, list) and len(dl) == 18 and len(cl) == 186 assert all(isinstance(x, float) for x in dl) # a missing embedder must surface as a readable error, not an AttributeError app._EMBEDDER, app._EMBEDDER_ERROR = None, "401 gated repo" try: app.infer("test", "{}") raise SystemExit("FAIL: expected a gr.Error when the embedder is missing") except Exception as exc: assert "HF_TOKEN" in str(exc), exc app._EMBEDDER = FakeEmbedder() print("\nOK: deterministic, empty input handled, all prompt formats run,") print(" infer() returns picklable lists, missing-token error is readable.") print("gradio Blocks built OK:", app.demo is not None)