Zap / training /label /select.py
ProCreations's picture
Zap v1: 20.7M-param web form & field classifier for autofill (bf16, MIT)
e4805d0 verified
Raw History Blame Contribute Delete
4.72 kB
"""Dedup harvested forms, assign host-level splits, and build prioritized teacher-labeling chunks.
Usage: python select.py OUT_DIR PATTERN [PATTERN ...] [--labeled DIR ...]
Signatures already labeled in a previous round (DIR/queue-*.jsonl + DIR/labels-*.jsonl) are not queued again.
Writes OUT_DIR/queue-NNN.jsonl (teacher input, 10k rows each, interleaved priorities),
OUT_DIR/reps.jsonl.gz (one representative per signature, with split), OUT_DIR/select_stats.json.
"""
import gzip
import hashlib
import json
import os
import random
import sys
from collections import Counter
sys.path.insert(0, "/job/work")
from label.common import load_forms # noqa: E402
SEARCHY = {"q", "s", "search", "query", "keyword", "keywords", "search query", "k", "term", "search term"}
MIX = [(0, 0.40), (1, 0.25), (2, 0.17), (3, 0.12), (4, 0.06)]
CAPS = {3: 60000, 4: 20000}
def split_of(host):
h = int(hashlib.md5(host.encode()).hexdigest()[:8], 16) % 100
return "test" if h < 4 else ("val" if h < 6 else "train")
def prio(r):
fs = r["fields"]
types = [f["type"] for f in fs]
if "input password" in types:
return 0
if len(fs) >= 3:
return 1
if len(fs) == 2 or any(t in ("input email", "input tel") for t in types):
return 2
f = fs[0]
if f["type"] == "input search" or f["name"] in SEARCHY or f["id"] in SEARCHY:
return 4
return 3
def labeled_sigs(dirs):
import glob
out = {}
for d in dirs:
uid2sig = {}
for q in glob.glob(os.path.join(d, "queue-*.jsonl")):
for line in open(q):
j = json.loads(line)
uid2sig[j["uid"]] = j["sig"]
for lf in glob.glob(os.path.join(d, "labels-*.jsonl")):
for line in open(lf):
try:
j = json.loads(line)
except json.JSONDecodeError: # partial last line of an interrupted run
continue
if j.get("form") and j.get("fields") and j["uid"] in uid2sig:
out[uid2sig[j["uid"]]] = j
return out
def main(out_dir, patterns, labeled_dirs=()):
os.makedirs(out_dir, exist_ok=True)
done = labeled_sigs(labeled_dirs)
print(f"{len(done)} signatures already labeled", flush=True)
reps, copies, seen = {}, Counter(), set()
n = 0
for pat in patterns:
for r in load_forms(pat):
n += 1
if r["uid"] in seen:
continue
seen.add(r["uid"])
copies[r["sig"]] += 1
cur = reps.get(r["sig"])
if cur is None or (not cur.get("html") and r.get("html")):
reps[r["sig"]] = r
print(f"{n} rows, {len(seen)} unique uids, {len(reps)} unique signatures", flush=True)
rng = random.Random(13)
pools = {p: [] for p, _ in MIX}
with gzip.open(os.path.join(out_dir, "reps.jsonl.gz"), "wt") as g:
for sig, r in reps.items():
r["split"] = split_of(r["host"])
r["copies"] = copies[sig]
r["prio"] = prio(r)
g.write(json.dumps(r, ensure_ascii=False) + "\n")
if r.get("html") and sig not in done:
pools[r["prio"]].append(r)
for p in pools:
rng.shuffle(pools[p])
if p in CAPS:
pools[p] = pools[p][:CAPS[p]]
stats = {"rows": n, "uids": len(seen), "sigs": len(reps), "pools": {p: len(v) for p, v in pools.items()},
"splits": Counter(split_of(r["host"]) for r in reps.values())}
print(stats, flush=True)
order, idx = [], {p: 0 for p in pools}
while any(idx[p] < len(pools[p]) for p in pools):
x, acc = rng.random(), 0.0
live = [(p, w) for p, w in MIX if idx[p] < len(pools[p])]
tot = sum(w for _, w in live)
for p, w in live:
acc += w / tot
if x <= acc:
order.append(pools[p][idx[p]])
idx[p] += 1
break
keep = ("uid", "url", "host", "ctx", "fields", "html", "sig", "split", "prio")
for c in range(0, len(order), 10000):
with open(os.path.join(out_dir, f"queue-{c // 10000:03d}.jsonl"), "w") as f:
for r in order[c:c + 10000]:
f.write(json.dumps({k: r[k] for k in keep}, ensure_ascii=False) + "\n")
stats["queued"] = len(order)
json.dump(stats, open(os.path.join(out_dir, "select_stats.json"), "w"), indent=1)
print("queued", len(order), flush=True)
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser()
ap.add_argument("out")
ap.add_argument("patterns", nargs="+")
ap.add_argument("--labeled", action="append", default=[])
a = ap.parse_args()
main(a.out, a.patterns, a.labeled)