Ines-1 / training /code /prepare_mix.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
16.3 kB
"""Decision SFT v3: a balanced mix of public typed-decision data -> common case format (decisions.py).
python scripts/decisions/prepare_mix.py --publico $DATA_ROOT/decisions/public \
--out $DATA_ROOT/decisions/mix
Families and caps (questions), so no single source dominates:
typed typed-decisions train EN + our ES translation (publico/, prepare_public.py) + the
independent ES translation of telepatia-ai/typed-decisions-pt-es (Portuguese dropped)
jev samatv256/jev-decisions-v1 general-clean-50k, tool selection (validation ids unchanged)
tasksource tasksource/tasksource-jev-typed-decisions: only rows with license_use == commercial,
at most --per-source questions from each of its sources, no source that is itself a
typed-decisions or jev-decisions derivative (would leak the tests)
n4ze3m n4ze3m/typed-decisions-synth (MIT): soft targets from its teacher (3 samples, mean)
helmo helmo/synthetic-typed-decisions (MIT): score gold given as mean/variance -> discretised
normal over the levels
Cleaning: options 2..26, no empty instructions, exact duplicate (state, question) pairs dropped,
telepatia train cases whose structure matches a typed-decisions TEST or VALIDATION case dropped (fingerprint of the
non-text leaves, which survive translation).
Outputs: train_<family>.jsonl, val_<family>.jsonl, test_telepatia_es.jsonl, test_tasksource.jsonl,
STATS.json, SOURCES.md
"""
from __future__ import annotations
import argparse
import ast
import collections
import hashlib
import json
import math
import random
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from prepare_public import jev_cases, val_side, write # noqa: E402
TELEPATIA = "telepatia-ai/typed-decisions-pt-es"
TASKSOURCE = "tasksource/tasksource-jev-typed-decisions"
N4ZE3M = "n4ze3m/typed-decisions-synth"
HELMO = "helmo/synthetic-typed-decisions"
JEV = "samatv256/jev-decisions-v1"
def js(x):
return json.loads(x) if isinstance(x, str) else x
def dl(repo, f):
from huggingface_hub import hf_hub_download
return hf_hub_download(repo, f, repo_type="dataset")
def jsonl(path):
return [json.loads(x) for x in open(path, encoding="utf-8") if x.strip()]
def norm(d):
s = sum(d.values())
return {k: v / s for k, v in d.items()} if s > 0 else None
def argmax(d):
return max(d, key=d.get)
def n_questions(cases):
return sum(len(c["questions"]) for c in cases)
def fingerprint(state):
"""Non-text leaves and key paths of a state: identical across translations of the same case."""
out = []
def walk(x, path):
if isinstance(x, dict):
for k in sorted(x):
walk(x[k], path + "/" + str(k))
elif isinstance(x, list):
out.append("%s#%d" % (path, len(x)))
for i, v in enumerate(x):
walk(v, "%s[%d]" % (path, i))
elif not isinstance(x, str):
out.append("%s=%r" % (path, x))
walk(state, "")
return hashlib.sha1("|".join(out).encode()).hexdigest()
def parse_state(s):
if not isinstance(s, str):
return s
for f in (json.loads, ast.literal_eval):
try:
v = f(s)
if isinstance(v, (dict, list)):
return v
except Exception: # noqa: BLE001
pass
return s
def valid(q):
if not (q.get("instructions") or "").strip():
return False
if q["type"] == "choice":
return 2 <= len(q.get("criteria") or {}) <= 26
if q["type"] == "score":
return 2 <= len(q.get("criteria") or []) <= 26
return q["type"] == "noul"
# ------------------------------------------------------------------ telepatia (typed ES, another translation)
def telepatia(split, test_fps):
out, dropped = [], 0
for r in jsonl(dl(TELEPATIA, "data/%s.es.jsonl" % split)):
state = parse_state(r["state"])
if split == "train" and isinstance(state, dict) and fingerprint(state) in test_fps:
dropped += 1
continue
qs, gold, soft = {}, {}, {}
for qid, q in r["questions"].items():
g = r["gold"][qid]
if q["type"] == "noul":
lab = str(g["label"]).lower()
else:
lab = str(g["label"])
if not valid(q):
continue
qs[qid], gold[qid] = q, [lab]
p = g.get("probabilities")
if p:
soft[qid] = {(str(k).lower() if q["type"] == "noul" else str(k)): float(v) for k, v in p.items()}
if qs:
out.append({"id": "telepatia-" + r["id"], "grupo": "telepatia/" + r["workflow"], "lang": "es",
"state": state, "questions": qs, "gold": gold, "suave": soft})
return out, dropped
# ------------------------------------------------------------------ tasksource
def tasksource_rows(files, per_source, rng, leak_words=("typed-decisions", "typed_decisions", "jev")):
"""Reservoir of at most `per_source` commercial-licence questions per source."""
import pyarrow.parquet as pq
cols = ["state", "kind", "id", "options", "target", "question", "source", "license", "license_use", "variant"]
seen_src = collections.Counter()
keep = collections.defaultdict(list)
skipped = collections.Counter()
for f in files:
pf = pq.ParquetFile(f)
for b in pf.iter_batches(columns=cols, batch_size=50000):
for r in b.to_pylist():
src = r["source"] or "?"
if r["license_use"] != "commercial":
skipped["licence"] += 1
continue
if any(w in src.lower() for w in leak_words):
skipped["leak-source"] += 1
continue
k = r["kind"]
opts, tgt = r["options"] or [], r["target"] or []
if k == "noul":
ok = len(tgt) == 1
else:
ok = 2 <= len(opts) <= 26 and len(tgt) == len(opts) and sum(tgt) > 0
if not ok or not (r["question"] or "").strip():
skipped["shape"] += 1
continue
seen_src[src] += 1
n = seen_src[src]
if len(keep[src]) < per_source:
keep[src].append(r)
else:
j = rng.randrange(n)
if j < per_source:
keep[src][j] = r
return keep, skipped
def tasksource_case(r):
k, opts, tgt = r["kind"], r["options"] or [], [float(x) for x in r["target"] or []]
if k == "noul":
p = min(max(tgt[0], 0.0), 1.0)
q = {"type": "noul", "instructions": r["question"]}
soft = {"true": p, "false": 1 - p}
elif k == "choice":
q = {"type": "choice", "instructions": r["question"], "criteria": {"o%d" % i: str(o) for i, o in enumerate(opts)}}
soft = norm({"o%d" % i: v for i, v in enumerate(tgt)})
else:
q = {"type": "score", "instructions": r["question"], "criteria": [str(o) for o in opts]}
soft = norm({str(i): v for i, v in enumerate(tgt)})
if soft is None:
return None
return {"id": "ts-" + r["id"], "grupo": "tasksource/" + r["source"].split("/")[0], "lang": "en",
"state": r["state"], "questions": {"q": q}, "gold": {"q": [argmax(soft)]}, "suave": {"q": soft},
"license": r["license"]}
def tasksource(split, per_source, total, rng):
from huggingface_hub import HfApi
files = sorted(s.rfilename for s in HfApi().dataset_info(TASKSOURCE).siblings
if s.rfilename.startswith("data/%s-" % split) and s.rfilename.endswith(".parquet"))
keep, skipped = tasksource_rows([dl(TASKSOURCE, f) for f in files], per_source, rng)
cases = [c for rows in keep.values() for c in map(tasksource_case, rows) if c is not None]
rng.shuffle(cases)
return cases[:total], {"sources": len(keep), "skipped": dict(skipped), "available": len(cases)}
# ------------------------------------------------------------------ n4ze3m (teacher soft targets)
def n4ze3m(split):
out = []
for r in jsonl(dl(N4ZE3M, "data/%s.jsonl" % split)):
qs_all, gold_all, teacher = js(r["questions"]), js(r["gold"]), js(r.get("teacher") or "{}") or {}
qs, gold, soft = {}, {}, {}
for qid, q in qs_all.items():
if not valid(q) or qid not in gold_all:
continue
g, t = gold_all[qid], teacher.get(qid) or {}
if q["type"] == "noul":
lab = "true" if g in (True, "true", 1) else "false"
s = {"true": float(t["noul"]), "false": 1 - float(t["noul"])} if "noul" in t else None
else:
lab = str(g)
s = {str(k): float(v) for k, v in (t.get("probabilities") or {}).items()} or None
qs[qid], gold[qid] = q, [lab]
if s and norm(s):
soft[qid] = norm(s)
if qs:
out.append({"id": "n4-" + r["state_id"], "grupo": "n4ze3m/" + str(r.get("domain", "?"))[:40], "lang": "en",
"state": parse_state(r["state"]) if r.get("state_is_json") else r["state"],
"questions": qs, "gold": gold, "suave": soft})
return out
# ------------------------------------------------------------------ helmo
def helmo():
out = []
for i, r in enumerate(jsonl(dl(HELMO, "synthetic_train.jsonl"))):
qs_all, gold_all = js(r["questions"]), js(r["gold"])
qs, gold, soft = {}, {}, {}
for qid, q in qs_all.items():
if q["type"] == "noul" and not q.get("criteria"):
q = {k: v for k, v in q.items() if k != "criteria"}
if not valid(q) or qid not in gold_all:
continue
g = gold_all[qid]
if q["type"] == "noul":
p = float(g["noul"])
s = {"true": p, "false": 1 - p}
elif q["type"] == "choice":
s = norm({str(k): float(v) for k, v in g["probabilities"].items() if str(k) in q["criteria"]})
else:
n = len(q["criteria"])
mu, var = float(g["mean"]), max(float(g.get("variance") or 0.0), 0.05)
s = norm({str(j): math.exp(-((j - mu) ** 2) / (2 * var)) for j in range(n)})
if not s:
continue
qs[qid], gold[qid], soft[qid] = q, [argmax(s)], s
if qs:
out.append({"id": "helmo-%d" % i, "grupo": "helmo/" + str(r.get("topic", "?"))[:40], "lang": "en",
"state": r["state"], "questions": qs, "gold": gold, "suave": soft})
return out
def dedupe(cases, seen):
out = []
for c in cases:
st = c["state"] if isinstance(c["state"], str) else json.dumps(c["state"], ensure_ascii=False, sort_keys=True)
qs = {}
for qid, q in c["questions"].items():
h = hashlib.sha1((st + "\x00" + json.dumps(q, ensure_ascii=False, sort_keys=True)).encode()).hexdigest()
if h not in seen:
seen.add(h)
qs[qid] = q
if qs:
out.append(dict(c, questions=qs, gold={k: c["gold"][k] for k in qs},
suave={k: v for k, v in (c.get("suave") or {}).items() if k in qs}))
return out
def cap(cases, n_q, rng):
cases = list(cases)
rng.shuffle(cases)
out, k = [], 0
for c in cases:
if k >= n_q:
break
out.append(c)
k += len(c["questions"])
return out
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--publico", type=Path, required=True, help="output of prepare_public.py")
ap.add_argument("--out", type=Path, required=True)
ap.add_argument("--per-source", type=int, default=80)
ap.add_argument("--cap", type=int, default=18000, help="max train questions per non-typed family")
ap.add_argument("--cap-tasksource", type=int, default=25000)
ap.add_argument("--seed", type=int, default=17)
a = ap.parse_args()
a.out.mkdir(parents=True, exist_ok=True)
rng = random.Random(a.seed)
stats = {}
seen = set()
# tests first: their (state, question) pairs are reserved before any train data is deduped
typed_test = jsonl(a.publico / "typed_test_en.jsonl") + jsonl(a.publico / "typed_test_es.jsonl")
# telepatia's train is typed-decisions train translated, and that train holds our validation cases
# (115 of 1,200): both test and validation states are excluded from it
typed_val = jsonl(a.publico / "typed_val_en.jsonl")
test_fps = {fingerprint(c["state"]) for c in typed_test + typed_val if isinstance(c["state"], dict)}
tele_test, _ = telepatia("test", set())
ts_test, ts_test_info = tasksource("test", 10, 3000, rng)
for name, cases in (("test_telepatia_es", tele_test), ("test_tasksource", ts_test)):
cases = dedupe(cases, seen)
write(a.out / (name + ".jsonl"), cases)
stats[name] = n_questions(cases)
dedupe(typed_test, seen)
# validation: one file per family (train_decisions weighs files equally)
ts_val, _ = tasksource("validation", 6, 1500, rng)
n4_val = n4ze3m("validation")
he_all = helmo()
he_val = [c for c in he_all if val_side(c["id"], 0.05)]
vals = {"typed_en": jsonl(a.publico / "typed_val_en.jsonl"), "typed_es": jsonl(a.publico / "typed_val_es.jsonl"),
"jev": jsonl(a.publico / "jev_val.jsonl"), "tasksource": ts_val, "n4ze3m": cap(n4_val, 1500, rng),
"helmo": cap(he_val, 1000, rng)}
for k, v in vals.items():
v = dedupe(v, seen)
write(a.out / ("val_%s.jsonl" % k), v)
stats["val_" + k] = n_questions(v)
# train
tele_train, dropped = telepatia("train", test_fps)
stats["telepatia_train_dropped_test_val_overlap"] = dropped
from huggingface_hub import list_repo_files
jev_files = sorted(f for f in list_repo_files(JEV, repo_type="dataset") if f.startswith("clean50k/data/") and f.endswith(".parquet"))
jev_all = jev_cases([dl(JEV, f) for f in jev_files], 60000, random.Random(a.seed))
jev_ids_val = {c["id"] for c in vals["jev"]}
jev_train = [c for c in jev_all if not val_side(c["id"], 0.05) and c["id"] not in jev_ids_val]
ts_train, ts_info = tasksource("train", a.per_source, 10 ** 9, rng)
stats["tasksource_info"] = ts_info
trains = {
"typed_en": jsonl(a.publico / "typed_train_en.jsonl"),
"typed_es": jsonl(a.publico / "typed_train_es.jsonl"),
"typed_es_telepatia": tele_train,
"jev": cap(jev_train, a.cap, rng),
"tasksource": cap(ts_train, a.cap_tasksource, rng),
"n4ze3m": cap(n4ze3m("train"), a.cap, rng),
"helmo": cap([c for c in he_all if not val_side(c["id"], 0.05)], a.cap, rng),
}
types = collections.Counter()
for k, v in trains.items():
v = dedupe(v, seen)
write(a.out / ("train_%s.jsonl" % k), v)
stats["train_" + k] = n_questions(v)
types.update(q["type"] for c in v for q in c["questions"].values())
stats["train_types"] = dict(types)
stats["train_total"] = sum(v for k, v in stats.items() if k.startswith("train_") and isinstance(v, int))
lic = collections.Counter(c.get("license") for c in jsonl(a.out / "train_tasksource.jsonl"))
stats["tasksource_train_licenses"] = dict(lic.most_common())
(a.out / "STATS.json").write_text(json.dumps(stats, indent=2) + "\n")
print(json.dumps(stats, indent=2))
(a.out / "SOURCES.md").write_text(
"# Sources (decision SFT v3 mix)\n\n"
"- LocalLLaMA/typed-decisions @c76749ec (Apache-2.0) + our Spanish machine translation (apodex-mini).\n"
"- telepatia-ai/typed-decisions-pt-es (Apache-2.0), Spanish files only.\n"
"- samatv256/jev-decisions-v1, general-clean-50k (CC-BY-4.0, see its SOURCE_LICENSES.md).\n"
"- tasksource/tasksource-jev-typed-decisions: rows with license_use == commercial only; per-row licences\n"
" kept in the `license` field (counts in STATS.json).\n"
"- n4ze3m/typed-decisions-synth (MIT); helmo/synthetic-typed-decisions (MIT).\n")
if __name__ == "__main__":
sys.exit(main())