File size: 3,661 Bytes
53ea208
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
#!/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()