data-use-annotate / rescore_singlepass.py
rafmacalaba's picture
annotation review app (per-user queues, Hub-backed rulings, static-safe direct commit)
53ea208 verified
Raw History Blame Contribute Delete
3.66 kB
#!/usr/bin/env python3
"""Rescore labeling queues with the singlepass infer head (local MPS).
Same weights the H100 fit (probe_head.pt in the singlepass bundle), applied
to fixed queue spans via training/probe_features_infer.encode_split_infer:
one prompt-conditioned forward per passage batch, score =
sigmoid(head([start; end; mean; +-64 window])).
Scoring happens only for items without scored_by (or --force). Every queue
item — scored here or carrying authoritative H100 scores from
build_gold_queue.py — is then tagged with the singlepass decision rule
(human_labeling/probe_labels.py):
keep score >= per-origin best-F1 threshold (thresholds.json)
drop score <= 0.05
confusion between threshold and the drop floor (was "mid")
unscored score missing (out-of-grid span)
uv run python human_labeling/rescore_singlepass.py [--force] [--limit N]
"""
import argparse
import json
import sys
from pathlib import Path
REPO = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO))
MODEL = "rafmacalaba/gliner-datause-catchall-singlepass"
QUEUES = ["human_labeling/queue.json", "human_labeling/queue_human473.json"]
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--model", default=MODEL)
ap.add_argument("--force", action="store_true")
ap.add_argument("--limit", type=int, default=0)
ap.add_argument("--batch", type=int, default=8)
a = ap.parse_args()
import torch
from training.singlepass_infer import default_device, load_bundle
from training.probe_features_infer import encode_split_infer
from probe_labels import decide, load_thresholds
device = default_device()
print(f"device={device} model={a.model}", flush=True)
model, head, bundle = load_bundle(
"rafmacalaba/gliner-datause-mentions-catch-all", a.model, device)
thresholds = bundle.get("thresholds") or load_thresholds()
args = argparse.Namespace(encode_batch_size=a.batch, context_radius=64)
for qname in QUEUES:
p = REPO / qname
if not p.exists():
continue
items = [json.loads(l) for l in p.read_text().splitlines() if l.strip()]
todo = [(i, r) for i, r in enumerate(items) if r.get("ctx")]
to_score = [(i, r) for i, r in todo if a.force or not r.get("scored_by")]
if a.limit:
to_score = to_score[:a.limit]
n_scored = n_skip = skipped = 0
if to_score:
texts = [r["ctx"] for _, r in to_score]
spans = [(k, r["start"], r["end"], 0, r["key"])
for k, (_, r) in enumerate(to_score)]
feats, dim, skipped = encode_split_infer(model, texts, spans, args, device)
F = torch.from_numpy(feats).to(device)
with torch.no_grad():
probs = torch.sigmoid(head(F)).cpu().numpy().ravel()
for (i, r), frow, ps in zip(to_score, feats, probs):
if not frow.any(): # out-of-grid: zero row, not a real score
r["head_score"] = None
n_skip += 1
else:
r["head_score"] = float(ps)
n_scored += 1
r["scored_by"] = a.model
n_tag = 0
for i, r in todo:
r["band"] = decide(r.get("head_score"), r.get("origin"), thresholds)
n_tag += 1
if to_score or n_tag:
p.write_text("\n".join(json.dumps(r) for r in items) + "\n")
print(f"{qname}: rescored={n_scored} unscored={n_skip} "
f"tagged={n_tag} skipped_align={skipped}", flush=True)
if __name__ == "__main__":
main()