#!/usr/bin/env python3 """Build the annotation queue from the gliner config of rafmacalaba/datause-ner, with camp2 Luna verdicts mapped onto every entity span and scores from the singlepass bundle (rafmacalaba/gliner-datause-catchall-singlepass). gliner rows: tokenized_text + ner (catch-all DATA_MENTION token spans) + spans[] traceability (key, luna_label, text, source). Join is positional: ner[i] <-> spans[i] (keys are sequential per passage; verified). ctx = ' '.join(tokenized_text); char offsets are computed on that grid, so they are exact by construction. Per mention: luna camp2 verdict, 1=keep / 0=drop head_score probe_score from the singlepass infer head (MPS) extractor_score GLiNER proposer score @0.1 matched on the inference grid band keep/confusion/drop via probe_labels.decide Sampling: round-robin over origins, multi-mention passages first, span budget --limit (default 520). uv run python human_labeling/build_gliner_queue.py [--limit 520] [--batch 8] """ import argparse import json import sys from collections import defaultdict from pathlib import Path HERE = Path(__file__).resolve().parent REPO = HERE.parent sys.path.insert(0, str(REPO)) MIRROR = REPO / "hf_datause_ner" OUT = HERE / "queue_gliner.json" def token_char_offsets(tokens: list[str]) -> list[int]: offs, c = [], 0 for t in tokens: offs.append(c) c += len(t) + 1 return offs def load_passages() -> list[dict]: rows = [] for split in ("train", "val", "holdout"): for line in (MIRROR / f"gliner_{split}.jsonl").read_text().splitlines(): if line.strip(): r = json.loads(line) r["split"] = split rows.append(r) return rows def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--model", default="rafmacalaba/gliner-datause-catchall-singlepass") ap.add_argument("--limit", type=int, default=520, help="span budget") ap.add_argument("--batch", type=int, default=8) a = ap.parse_args() import torch import torch.utils.data from training.singlepass_infer import default_device, load_bundle from training.probe_features_infer import char_to_infer_word from probe_labels import decide passages = load_passages() by_origin: dict[str, list[dict]] = defaultdict(list) for p in passages: if len(p.get("ner", [])) == len(p.get("spans", [])) and p["ner"]: by_origin[p["origin"]].append(p) origins = sorted(by_origin) for o in origins: by_origin[o].sort(key=lambda p: -min(len(p["ner"]), 2)) ordered: list[dict] = [] i = 0 while any(by_origin[o] for o in origins): o = origins[i % len(origins)] if by_origin[o]: ordered.append(by_origin[o].pop(0)) i += 1 chosen, n_spans = [], 0 for p in ordered: if n_spans >= a.limit: break # dedupe identical spans (upstream extractor proposed the same span # twice under separate keys; 173 groups corpus-wide, 22 with # CONFLICTING luna verdicts). First key wins; conflicting duplicates # flag the survivor as luna_split for the UI. first, ner_u, spans_u = {}, [], [] for i, (ner, s) in enumerate(zip(p["ner"], p["spans"])): k = (ner[0], ner[1]) if k not in first: first[k] = len(spans_u) ner_u.append(ner) spans_u.append(dict(s)) elif spans_u[first[k]].get("luna_label") != s.get("luna_label"): spans_u[first[k]]["luna_split"] = True q = dict(p, ner=ner_u, spans=spans_u) chosen.append(q) n_spans += len(ner_u) device = default_device() print(f"device={device} model={a.model} passages={len(chosen)} spans={n_spans}", flush=True) model, head, bundle = load_bundle( "rafmacalaba/gliner-datause-mentions-catch-all", a.model, device) thresholds = bundle.get("thresholds") or {} radius = bundle["radius"] INFER_LABELS = ["DATA_MENTION"] texts, char_maps = [], [] for p in chosen: toks = p["tokenized_text"] texts.append(" ".join(toks)) char_maps.append(token_char_offsets(toks)) prepared = model.prepare_batch(texts, INFER_LABELS) collator = model.create_collator() def collate_fn(batch): return model.collate_batch(batch, prepared["entity_types"], collator) loader = torch.utils.data.DataLoader( prepared["input_x"], batch_size=a.batch, shuffle=False, collate_fn=collate_fn) v2o = prepared["valid_to_orig_idx"] n_probe = n_ext = n_skip = 0 items: list[dict] = [] def flush(p, ctx, char_offs, probes, extractors): mentions = [] for i, (ner, s, probe, ext) in enumerate( zip(p["ner"], p["spans"], probes, extractors)): t0, t1 = ner[0], ner[1] start = char_offs[t0] end = char_offs[t1] + len(p["tokenized_text"][t1]) mentions.append({ "key": s["key"], "surface": " ".join(p["tokenized_text"][t0:t1 + 1]), "start": start, "end": end, "luna": s.get("luna_label"), "luna_split": s.get("luna_split", False), "head_score": probe, "extractor_score": ext, "band": decide(probe, p["origin"], thresholds), }) bands = {m["band"] for m in mentions} pband = ("unscored" if "unscored" in bands else "confusion" if "confusion" in bands else "mixed" if len(bands) > 1 else bands.pop()) items.append({ "queue": "gliner", "origin": p["origin"], "split": p["split"], "ctx": ctx, "n": len(mentions), "band": pband, "mentions": mentions, "scored_by": a.model, }) row = 0 with torch.no_grad(): for batch in loader: out = model.run_batch(batch, threshold=0.1, move_to_device=True) W = out.words_embedding.detach().float() mask = (out.mask.detach().cpu() if getattr(out, "mask", None) is not None else None) decoded = model.decode_batch(out, batch, threshold=0.1, flat_ner=True, multi_label=False) B = W.shape[0] for bi in range(B): vi = row + bi oi = v2o[vi] p = chosen[oi] if oi not in set(v2o): # unreachable; v2o IS the valid map continue w = W[bi].to(device) L = int(mask[bi].sum()) if mask is not None else w.shape[0] starts = prepared["start_token_map"][vi] proposals = [(int(sp.start), int(sp.end), float(sp.score)) for sp in decoded[bi]] probes, extractors = [], [] for ner in p["ner"]: # token grid -> char grid -> inference word grid t0, t1 = ner[0], ner[1] char_offs = char_maps[oi] cs = char_offs[t0] ce = char_offs[t1] + len(p["tokenized_text"][t1]) g0, g1 = char_to_infer_word(starts, cs, ce) probe = ext = None if g1 < L and g0 < L: idx = torch.arange(g0, g1 + 1, device=device) parts = [w[g0], w[g1], w[idx].mean(dim=0)] if radius > 0: w0, w1 = max(0, g0 - radius), min(g1 + radius, L - 1) parts.append(w[w0:w1 + 1].mean(dim=0)) probe = float(torch.sigmoid( head(torch.cat(parts).unsqueeze(0))).item()) n_probe += 1 hit = [sc for (ps, pe, sc) in proposals if ps == g0 and pe == g1] if hit: ext = hit[0] n_ext += 1 probes.append(probe) extractors.append(ext) flush(p, texts[oi], char_maps[oi], probes, extractors) row += B OUT.write_text("\n".join(json.dumps(it) for it in items) + "\n") from collections import Counter bands = Counter(m["band"] for it in items for m in it["mentions"]) lu = Counter(m["luna"] for it in items for m in it["mentions"]) print(f"queue: passages={len(items)} spans={sum(it['n'] for it in items)} " f"probe_scored={n_probe} extractor_matched={n_ext} " f"unscored={n_skip} bands={dict(bands)} luna={dict(lu)} -> {OUT}", flush=True) if __name__ == "__main__": main()