#!/usr/bin/env python3 """Minimal example: load the SRT-Adapter v8a checkpoint, score a passage, and print the four semiotic readouts. Usage: cd examples pip install -r ../requirements.txt python load_and_score.py --text "Vaccine mandates are an obvious public health win." First run downloads Qwen/Qwen2.5-7B (~15 GB) from HuggingFace. """ from __future__ import annotations import argparse import json import sys from pathlib import Path import torch from transformers import AutoTokenizer HERE = Path(__file__).resolve().parent sys.path.insert(0, str((HERE.parent / "src").resolve())) from srt.adapter import SRTAdapter # noqa: E402 from srt.config import ( # noqa: E402 SRTConfig, MAHConfig, RRMConfig, BENConfig, CommunityConfig, LossConfig, ) def build_config(config_path: Path) -> SRTConfig: raw = json.loads(config_path.read_text()) return SRTConfig( backbone_id=raw["backbone_id"], backbone_dtype=raw["backbone_dtype"], mah_layer_indices=list(raw["mah_layer_indices"]), rrm_inject_indices=list(raw["rrm_inject_indices"]), community_layer_idx=raw["community_layer_idx"], num_mah_layers=raw["num_mah_layers"], mah=MAHConfig(**raw["mah"]), rrm=RRMConfig(**raw["rrm"]), ben=BENConfig(**raw["ben"]), community=CommunityConfig(**raw["community"]), loss=LossConfig(**{ k: v for k, v in raw["loss"].items() if k in LossConfig.__dataclass_fields__ }), ) def main() -> None: ap = argparse.ArgumentParser() default_adapter = HERE.parent / "adapter.safetensors" if not default_adapter.exists(): default_adapter = HERE.parent / "adapter.pt" ap.add_argument("--adapter", default=str(default_adapter), help="Path to adapter.safetensors (preferred) or adapter.pt.") ap.add_argument("--config", default=str(HERE.parent / "config.json")) ap.add_argument("--text", required=True, help="Passage to score.") ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") ap.add_argument("--max-seq-len", type=int, default=512) args = ap.parse_args() print(f"[load] config: {args.config}") config = build_config(Path(args.config)) print(f"[load] backbone: {config.backbone_id} ({config.backbone_dtype})") print(f"[load] adapter: {args.adapter}") model = SRTAdapter(config).to(args.device) if args.adapter.endswith(".safetensors"): from safetensors.torch import load_file state = load_file(args.adapter, device=args.device) else: state = torch.load(args.adapter, map_location=args.device) missing, unexpected = model.load_state_dict(state, strict=False) print(f"[load] missing={len(missing)} unexpected={len(unexpected)}") model.eval() tok = AutoTokenizer.from_pretrained(config.backbone_id) enc = tok(args.text, return_tensors="pt", truncation=True, max_length=args.max_seq_len).to(args.device) with torch.no_grad(): out = model(input_ids=enc.input_ids, attention_mask=enc.attention_mask) print("\n=== SRT-Adapter readouts ===") print(f"input tokens: {enc.input_ids.shape[1]}") print(f"backbone vocab logits shape: {tuple(out.logits.shape)}") if out.community_output is not None: cv = out.community_output.vector[0] # (d_community,) print(f"community vector ({cv.shape[0]}-D): " f"norm={cv.norm().item():.3f} " f"first 5 dims={[round(x, 3) for x in cv[:5].tolist()]}") for i, d in enumerate(out.divergences): mean_norm = d.norm(dim=-1).mean().item() print(f"divergence layer {i} mean ||d||: {mean_norm:.3f}") if out.ben_output is not None: r_hat = out.ben_output.r_hat[0] regime_prob_super = torch.softmax(out.ben_output.regime_logits[0], dim=-1)[:, 1] print(f"reflexivity r_hat: " f"mean={r_hat.mean().item():+.3f} " f"min={r_hat.min().item():+.3f} " f"max={r_hat.max().item():+.3f}") print(f"P(supercritical): " f"mean={regime_prob_super.mean().item():.3f} " f"max={regime_prob_super.max().item():.3f}") print("\nSee paper.pdf §3 for what each readout means and §5 for headline numbers.") if __name__ == "__main__": main()