File size: 2,834 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
#!/usr/bin/env python3
"""Group scored flat pool items into passage examples (multi-mention).

Reads the flat rescored queue (build_pool_queue.py → rescore_singlepass.py)
and emits one JSON object per distinct (origin, ctx):

    {queue, origin, ctx, n, band, scored_by,
     mentions: [{key, surface, start, end, head_score, band}...]}   # by start

Passages with one mention stay as-is; multi stays multi — the UI drives
them as one example with a per-mention ruling target. The passage-level
band is the most-informative mention state:

    unscored  any unscored
    confusion any confusion
    mixed     keep + drop, no confusion
    else      the single uniform band

    uv run python human_labeling/build_passage_queue.py \
        [--input human_labeling/queue.json] [--out human_labeling/queue_passages.json]
"""

import argparse
import json
from collections import defaultdict
from pathlib import Path

HERE = Path(__file__).resolve().parent


def passage_band(mentions: list[dict]) -> str:
    bands = {m.get("band") for m in mentions}
    if "unscored" in bands:
        return "unscored"
    if "confusion" in bands:
        return "confusion"
    if len(bands) == 1:
        return bands.pop()
    return "mixed"


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--input", default=str(HERE / "queue.json"))
    ap.add_argument("--out", default=str(HERE / "queue_passages.json"))
    a = ap.parse_args()

    groups: dict[tuple, list[dict]] = defaultdict(list)
    for line in Path(a.input).read_text().splitlines():
        if not line.strip():
            continue
        r = json.loads(line)
        if r.get("ctx"):
            groups[(r.get("origin"), r["ctx"])].append(r)

    passages = []
    for (origin, ctx), rows in groups.items():
        mentions = sorted(
            ({"key": r["key"], "surface": r.get("surface"),
              "start": r.get("start"), "end": r.get("end"),
              "head_score": r.get("head_score"), "band": r.get("band")}
             for r in rows),
            key=lambda m: (m["start"] if isinstance(m["start"], int) else 0,
                           m["surface"] or ""))
        passages.append({
            "queue": "probe_candidates/passages",
            "origin": origin, "ctx": ctx, "n": len(mentions),
            "band": passage_band(mentions), "mentions": mentions,
            "scored_by": rows[0].get("scored_by"),
        })

    passages.sort(key=lambda p: (p["origin"] or "", p["ctx"] or ""))
    Path(a.out).write_text(
        "\n".join(json.dumps(p) for p in passages) + "\n")
    n_multi = sum(1 for p in passages if p["n"] > 1)
    from collections import Counter
    print(f"passages={len(passages)} multi={n_multi} "
          f"{Counter(p['band'] for p in passages)} -> {a.out}")


if __name__ == "__main__":
    main()