CoolFace
Apppublic

rafmacalaba/data-use-annotate

sourceHugging Faceupdated 13d agoView on Hugging Face
0likes
build_gliner_queue.py218 linesDownload Raw Back to root
1#!/usr/bin/env python32"""Build the annotation queue from the gliner config of3rafmacalaba/datause-ner, with camp2 Luna verdicts mapped onto every4entity span and scores from the singlepass bundle5(rafmacalaba/gliner-datause-catchall-singlepass).6 7gliner rows: tokenized_text + ner (catch-all DATA_MENTION token spans)8+ spans[] traceability (key, luna_label, text, source). Join is positional:9ner[i] <-> spans[i] (keys are sequential per passage; verified).10 11ctx = ' '.join(tokenized_text); char offsets are computed on that grid, so12they are exact by construction. Per mention:13    luna            camp2 verdict, 1=keep / 0=drop14    head_score      probe_score from the singlepass infer head (MPS)15    extractor_score GLiNER proposer score @0.1 matched on the inference grid16    band            keep/confusion/drop via probe_labels.decide17 18Sampling: round-robin over origins, multi-mention passages first, span19budget --limit (default 520).20 21    uv run python human_labeling/build_gliner_queue.py [--limit 520] [--batch 8]22"""23 24import argparse25import json26import sys27from collections import defaultdict28from pathlib import Path29 30HERE = Path(__file__).resolve().parent31REPO = HERE.parent32sys.path.insert(0, str(REPO))33 34MIRROR = REPO / "hf_datause_ner"35OUT = HERE / "queue_gliner.json"36 37 38def token_char_offsets(tokens: list[str]) -> list[int]:39    offs, c = [], 040    for t in tokens:41        offs.append(c)42        c += len(t) + 143    return offs44 45 46def load_passages() -> list[dict]:47    rows = []48    for split in ("train", "val", "holdout"):49        for line in (MIRROR / f"gliner_{split}.jsonl").read_text().splitlines():50            if line.strip():51                r = json.loads(line)52                r["split"] = split53                rows.append(r)54    return rows55 56 57def main() -> None:58    ap = argparse.ArgumentParser()59    ap.add_argument("--model", default="rafmacalaba/gliner-datause-catchall-singlepass")60    ap.add_argument("--limit", type=int, default=520, help="span budget")61    ap.add_argument("--batch", type=int, default=8)62    a = ap.parse_args()63 64    import torch65    import torch.utils.data66    from training.singlepass_infer import default_device, load_bundle67    from training.probe_features_infer import char_to_infer_word68    from probe_labels import decide69 70    passages = load_passages()71    by_origin: dict[str, list[dict]] = defaultdict(list)72    for p in passages:73        if len(p.get("ner", [])) == len(p.get("spans", [])) and p["ner"]:74            by_origin[p["origin"]].append(p)75    origins = sorted(by_origin)76    for o in origins:77        by_origin[o].sort(key=lambda p: -min(len(p["ner"]), 2))78    ordered: list[dict] = []79    i = 080    while any(by_origin[o] for o in origins):81        o = origins[i % len(origins)]82        if by_origin[o]:83            ordered.append(by_origin[o].pop(0))84        i += 185 86    chosen, n_spans = [], 087    for p in ordered:88        if n_spans >= a.limit:89            break90        # dedupe identical spans (upstream extractor proposed the same span91        # twice under separate keys; 173 groups corpus-wide, 22 with92        # CONFLICTING luna verdicts). First key wins; conflicting duplicates93        # flag the survivor as luna_split for the UI.94        first, ner_u, spans_u = {}, [], []95        for i, (ner, s) in enumerate(zip(p["ner"], p["spans"])):96            k = (ner[0], ner[1])97            if k not in first:98                first[k] = len(spans_u)99                ner_u.append(ner)100                spans_u.append(dict(s))101            elif spans_u[first[k]].get("luna_label") != s.get("luna_label"):102                spans_u[first[k]]["luna_split"] = True103        q = dict(p, ner=ner_u, spans=spans_u)104        chosen.append(q)105        n_spans += len(ner_u)106    device = default_device()107    print(f"device={device} model={a.model} passages={len(chosen)} spans={n_spans}",108          flush=True)109    model, head, bundle = load_bundle(110        "rafmacalaba/gliner-datause-mentions-catch-all", a.model, device)111    thresholds = bundle.get("thresholds") or {}112    radius = bundle["radius"]113 114    INFER_LABELS = ["DATA_MENTION"]115    texts, char_maps = [], []116    for p in chosen:117        toks = p["tokenized_text"]118        texts.append(" ".join(toks))119        char_maps.append(token_char_offsets(toks))120    prepared = model.prepare_batch(texts, INFER_LABELS)121    collator = model.create_collator()122 123    def collate_fn(batch):124        return model.collate_batch(batch, prepared["entity_types"], collator)125 126    loader = torch.utils.data.DataLoader(127        prepared["input_x"], batch_size=a.batch, shuffle=False,128        collate_fn=collate_fn)129 130    v2o = prepared["valid_to_orig_idx"]131    n_probe = n_ext = n_skip = 0132    items: list[dict] = []133 134    def flush(p, ctx, char_offs, probes, extractors):135        mentions = []136        for i, (ner, s, probe, ext) in enumerate(137                zip(p["ner"], p["spans"], probes, extractors)):138            t0, t1 = ner[0], ner[1]139            start = char_offs[t0]140            end = char_offs[t1] + len(p["tokenized_text"][t1])141            mentions.append({142                "key": s["key"], "surface": " ".join(p["tokenized_text"][t0:t1 + 1]),143                "start": start, "end": end,144                "luna": s.get("luna_label"), "luna_split": s.get("luna_split", False),145                "head_score": probe, "extractor_score": ext,146                "band": decide(probe, p["origin"], thresholds),147            })148        bands = {m["band"] for m in mentions}149        pband = ("unscored" if "unscored" in bands else150                 "confusion" if "confusion" in bands else151                 "mixed" if len(bands) > 1 else bands.pop())152        items.append({153            "queue": "gliner", "origin": p["origin"], "split": p["split"],154            "ctx": ctx, "n": len(mentions), "band": pband,155            "mentions": mentions, "scored_by": a.model,156        })157 158    row = 0159    with torch.no_grad():160        for batch in loader:161            out = model.run_batch(batch, threshold=0.1, move_to_device=True)162            W = out.words_embedding.detach().float()163            mask = (out.mask.detach().cpu()164                    if getattr(out, "mask", None) is not None else None)165            decoded = model.decode_batch(out, batch, threshold=0.1,166                                         flat_ner=True, multi_label=False)167            B = W.shape[0]168            for bi in range(B):169                vi = row + bi170                oi = v2o[vi]171                p = chosen[oi]172                if oi not in set(v2o):  # unreachable; v2o IS the valid map173                    continue174                w = W[bi].to(device)175                L = int(mask[bi].sum()) if mask is not None else w.shape[0]176                starts = prepared["start_token_map"][vi]177                proposals = [(int(sp.start), int(sp.end), float(sp.score))178                             for sp in decoded[bi]]179                probes, extractors = [], []180                for ner in p["ner"]:181                    # token grid -> char grid -> inference word grid182                    t0, t1 = ner[0], ner[1]183                    char_offs = char_maps[oi]184                    cs = char_offs[t0]185                    ce = char_offs[t1] + len(p["tokenized_text"][t1])186                    g0, g1 = char_to_infer_word(starts, cs, ce)187                    probe = ext = None188                    if g1 < L and g0 < L:189                        idx = torch.arange(g0, g1 + 1, device=device)190                        parts = [w[g0], w[g1], w[idx].mean(dim=0)]191                        if radius > 0:192                            w0, w1 = max(0, g0 - radius), min(g1 + radius, L - 1)193                            parts.append(w[w0:w1 + 1].mean(dim=0))194                        probe = float(torch.sigmoid(195                            head(torch.cat(parts).unsqueeze(0))).item())196                        n_probe += 1197                        hit = [sc for (ps, pe, sc) in proposals198                               if ps == g0 and pe == g1]199                        if hit:200                            ext = hit[0]201                            n_ext += 1202                    probes.append(probe)203                    extractors.append(ext)204                flush(p, texts[oi], char_maps[oi], probes, extractors)205            row += B206 207    OUT.write_text("\n".join(json.dumps(it) for it in items) + "\n")208    from collections import Counter209    bands = Counter(m["band"] for it in items for m in it["mentions"])210    lu = Counter(m["luna"] for it in items for m in it["mentions"])211    print(f"queue: passages={len(items)} spans={sum(it['n'] for it in items)} "212          f"probe_scored={n_probe} extractor_matched={n_ext} "213          f"unscored={n_skip} bands={dict(bands)} luna={dict(lu)} -> {OUT}",214          flush=True)215 216 217if __name__ == "__main__":218    main()