File size: 3,693 Bytes
e4805d0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Join teacher labels onto forms and write train/val/test sets + an MLM corpus.

Usage: python build_data.py REPS_DIR OUT_DIR PATTERN [PATTERN ...] --labels DIR [--labels DIR ...]
  REPS_DIR/reps.jsonl.gz (from select.py over the same PATTERNs) fixes each signature's split/representative;
  teacher labels are matched by signature across all --labels dirs.
  Every harvested copy of a labeled signature inherits its labels (max 3 copies per signature).
"""
import glob
import gzip
import json
import os
import random
import sys
from collections import Counter, defaultdict

sys.path.insert(0, "/job/work")
from label.common import load_forms  # noqa: E402
from label.select import labeled_sigs  # noqa: E402
from label.teacher import FIELD_LABELS, FORM_LABELS  # noqa: E402

MAX_COPIES = 3


def main(label_dir, out_dir, patterns, label_dirs):
    os.makedirs(out_dir, exist_ok=True)
    reps = {}
    with gzip.open(os.path.join(label_dir, "reps.jsonl.gz"), "rt") as f:
        for line in f:
            r = json.loads(line)
            reps[r["sig"]] = {"split": r["split"], "uid": r["uid"], "prio": r["prio"]}
    by_sig = {sig: lab for sig, lab in labeled_sigs(label_dirs).items() if sig in reps}
    print(f"{len(reps)} signatures, {len(by_sig)} labeled signatures", flush=True)

    outs = {s: open(os.path.join(out_dir, f"{s}.jsonl"), "w") for s in ("train", "val", "test")}
    mlm = gzip.open(os.path.join(out_dir, "mlm.jsonl.gz"), "wt")
    n_copies = defaultdict(int)
    seen_uid = set()
    stats = Counter()
    field_c, form_c = Counter(), Counter()
    for pat in patterns:
        for r in load_forms(pat):
            if r["uid"] in seen_uid:
                continue
            seen_uid.add(r["uid"])
            rep = reps.get(r["sig"])
            if rep is None:
                continue
            split = rep["split"]
            # held-out splits keep only the representative (exact copies would inflate the test set)
            if split != "train" and r["uid"] != rep["uid"]:
                continue
            if n_copies[r["sig"]] >= MAX_COPIES:
                continue
            n_copies[r["sig"]] += 1
            if split == "train" and n_copies[r["sig"]] == 1:
                mlm.write(json.dumps({"ctx": r["ctx"], "fields": r["fields"]}, ensure_ascii=False) + "\n")
                stats["mlm"] += 1
            lab = by_sig.get(r["sig"])
            if lab is None or len(lab["fields"]) != len(r["fields"]):
                continue
            row = {"uid": r["uid"], "sig": r["sig"], "url": r["url"], "host": r["host"], "lang": r.get("lang", ""),
                   "ctx": r["ctx"], "fields": r["fields"], "form_label": lab["form"], "field_labels": lab["fields"],
                   "prio": rep["prio"], "rep": r["uid"] == rep["uid"]}
            outs[split].write(json.dumps(row, ensure_ascii=False) + "\n")
            stats[split] += 1
            if split == "train":
                form_c[lab["form"]] += 1
                field_c.update(lab["fields"])
    for f in outs.values():
        f.close()
    mlm.close()
    summary = {"stats": stats, "train_form_labels": form_c, "train_field_labels": field_c,
               "field_labels": FIELD_LABELS, "form_labels": FORM_LABELS}
    json.dump(summary, open(os.path.join(out_dir, "data_stats.json"), "w"), indent=1)
    print(json.dumps(summary, indent=1), flush=True)


if __name__ == "__main__":
    import argparse
    ap = argparse.ArgumentParser()
    ap.add_argument("reps_dir")
    ap.add_argument("out")
    ap.add_argument("patterns", nargs="+")
    ap.add_argument("--labels", action="append", default=[])
    a = ap.parse_args()
    main(a.reps_dir, a.out, a.patterns, a.labels)