data-use-annotate / build_gliner_queue.py
rafmacalaba's picture
annotation review app (per-user queues, Hub-backed rulings, static-safe direct commit)
53ea208 verified
Raw History Blame Contribute Delete
8.8 kB
#!/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()