soft-decider-421m / scripts /preprocess_aug.py
winwinwinbb's picture
soft-decider-421m: RLCD fine-tune of laya on typed-decisions with holdout calibration
ad67f6b verified
Raw History Blame Contribute Delete
8.51 kB
"""Gold-preserving data augmentation (open-jev recipe) -> train_items_aug.pt + calib_items.
The five augmentations, each label-preserving by construction, applied at p=0.7 like
com-kotobalabs/open-jev-deberta-v3-large (their in-domain 0.854 / OOD 0.690 gap is what
this family of tricks attacks):
1. option shuffle - permute option render order; target permuted identically
2. instruction paraphrase - deterministic templates, semantics untouched
3. distractor drop - choice only: drop non-gold options with p<0.05, renormalise
4. level relabel - score only: shift the "level N" numbers by a constant
5. noul negation - "It is NOT the case that ..." with flipped target
The calibration slice uses the SAME seed/order as train_single.py's internal split and is
kept CLEAN (never augmented, never trained on), so temperature fitting stays honest.
Run: D:\\venvs\\voxcpm-cuda\\Scripts\\python.exe d:\\QoderCN\\jev_finetune\\preprocess_aug.py
"""
from __future__ import annotations
import json
import os
import random
import sys
import time
os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
os.environ.setdefault("USE_TF", "0")
sys.stdout.reconfigure(encoding="utf-8")
WORK = os.path.dirname(os.path.abspath(__file__))
TRAIN_PARQUET = os.path.join(WORK, "td_all_train.parquet")
OUT = os.path.join(WORK, "train_items_aug.pt")
CALIB_MAX = 400
CALIB_SEED = 20260922 # must match train_single.py so the calib slice is identical
P_AUG = 0.7
PARAPHRASE = [
"{ins}",
"Considering the state, {ins_lower}",
"Decide: {ins_lower}",
"Based only on the information given, {ins_lower}",
]
def variants_for(q, gold_q, rng):
"""Yield (question_dict, target_list) variants of one (question, gold) pair, always incl. the base."""
t = q["type"]
crit = q.get("criteria", {})
ins = q["instructions"]
# base distribution aligned to render order
if t == "choice":
keys = list(crit.keys())
base = [gold_q["probabilities"].get(k, 0.0) for k in keys]
opts = {"t": t, "crit": crit}
elif t == "noul":
base = [gold_q["probabilities"].get("false", 0.5), gold_q["probabilities"].get("true", 0.5)]
opts = {"t": t, "crit": crit}
else:
n = len(crit) if isinstance(crit, list) else 4
base = [gold_q["probabilities"].get(str(i), 0.0) for i in range(n)]
opts = {"t": t, "crit": crit}
s = sum(base)
base = [v / s for v in base] if s > 0 else [1.0 / len(base)] * len(base)
yield {"t": t, "ins": ins, "crit": opts["crit"]}, base
if rng.random() > P_AUG:
return
vs = [(ins, base, opts["crit"])]
# 1. option shuffle (choice only; score levels are ordinal and stay in level order)
if t == "choice" and len(base) >= 2 and rng.random() < 0.8:
perm = list(range(len(base)))
rng.shuffle(perm)
keys = list(crit.keys())
new_crit = {keys[perm[m]]: crit[keys[perm[m]]] for m in range(len(base))}
new_target = [base[perm[m]] for m in range(len(base))]
vs.append((ins, new_target, new_crit))
# 2. instruction paraphrase
if rng.random() < 0.6:
tpl = rng.choice(PARAPHRASE[1:])
new_ins = tpl.format(ins=ins, ins_lower=ins[0].lower() + ins[1:] if ins else ins)
vs.append((new_ins, vs[-1][1], vs[-1][2]))
# 3. distractor drop (choice, >=5 opts): drop non-gold with prob < 0.05
if t == "choice" and len(base) >= 5 and rng.random() < 0.5:
keys = list(crit.keys())
top = max(range(len(base)), key=lambda i: base[i])
drop = {i for i, v in enumerate(base) if i != top and v < 0.05}
drop -= {top}
if drop and len(base) - len(drop) >= 3:
keep = [i for i in range(len(base)) if i not in drop]
new_crit = {keys[i]: crit[keys[i]] for i in keep}
nt = [base[i] for i in keep]
vs.append((ins, [v / sum(nt) for v in nt], new_crit))
# 4. level relabel (score): shift the "level N" numbers; keys keep level order so the
# target is untouched; rendered through the choice path with pre-numbered option texts
if t == "score" and rng.random() < 0.5:
sh = rng.choice([1, 2])
new_crit = {str(i): "level %d: %s" % (i + sh, c) for i, c in enumerate(crit)}
vs.append((ins, base, new_crit))
# 5. noul negation with flipped gold
if t == "noul" and rng.random() < 0.5:
vs.append(("It is NOT the case that: " + ins[0].lower() + ins[1:], [base[1], base[0]], crit))
seen = set()
for ins2, tgt2, crit2 in vs[1:]:
key = (ins2, tuple(round(x, 6) for x in tgt2), json.dumps(crit2, sort_keys=True) if isinstance(crit2, dict) else str(crit2))
if key in seen:
continue
seen.add(key)
yield {"t": t, "ins": ins2, "crit": crit2}, tgt2
def to_item(tok, cfg, state, q, target):
from laya.common import build_sequence, render_options, QTYPES
t = q["t"]
# score relabelling arrives as a dict of pre-numbered strings -> render via the choice path
tv = "choice" if (t == "score" and isinstance(q["crit"], dict)) else t
kk = len(render_options({"t": tv, "crit": q["crit"]}))
seq, markers = build_sequence(tok, state, {"t": tv, "ins": q["ins"], "crit": q["crit"]},
cfg["max_len"], cfg["head_max_len"])
if len(markers) != kk:
return None
qt = QTYPES["score"] if t == "score" else QTYPES[t]
label = target.index(max(target))
return {"ids": seq, "markers": markers, "qtype": qt, "target": target, "label": label}
def main() -> None:
import torch
import pandas as pd
from transformers import AutoTokenizer
from huggingface_hub import snapshot_download
from laya.agent import _fix_tokenizer_config
model_dir = snapshot_download("convaiinnovations/laya")
_fix_tokenizer_config(model_dir)
tok = AutoTokenizer.from_pretrained(os.path.join(model_dir, "tokenizer"))
cfg = json.load(open(os.path.join(model_dir, "rl_agent_config.json")))
cfg["max_len"] = 1024
cfg["head_max_len"] = 256
df = pd.read_parquet(TRAIN_PARQUET)
rng = random.Random(20260924)
# Pass 1: base items in the exact same iteration order as preprocess.py (keeps the split identical)
print("[1/3] base items in canonical order ...")
base = []
row_ctx = []
for row in df.itertuples(index=False):
state = json.loads(row.state)
questions = json.loads(row.questions)
gold = json.loads(row.gold)
for qid, q in questions.items():
if qid in gold:
row_ctx.append((state, q, gold[qid]))
base.append(len(base))
# rebuild items canonically (identical call order to preprocess.py)
from preprocess import build_training_item # canonical builder, same output as preprocess.py run
items_base = [build_training_item(tok, cfg, *ctx) for ctx in row_ctx]
print(f" base: {len(items_base)} built, {sum(x is None for x in items_base)} skipped")
# NOTE: None entries can differ from preprocess.py only if builders diverge; they don't.
pairs = [(it, ctx) for it, ctx in zip(items_base, row_ctx) if it is not None]
n = len(pairs)
# Calibration slice: same shuffle & indices as train_single's internal split over canonical list
order = list(range(n))
random.Random(CALIB_SEED).shuffle(order)
n_calib = min(CALIB_MAX, n // 10)
calib_idx = set(order[:n_calib])
calib_items = [pairs[i][0] for i in sorted(calib_idx)]
train_idx = [i for i in range(n) if i not in calib_idx]
print(f"[2/3] split: train={len(train_idx)} calib(clean)={len(calib_items)}; augmenting ...")
t0 = time.time()
aug_items = []
for i in train_idx:
state, q, gold_q = pairs[i][1]
for v, tgt in variants_for(q, gold_q, rng):
it = to_item(tok, cfg, state, v, tgt)
if it:
aug_items.append(it)
rng.shuffle(aug_items)
print(f" augmented: {len(aug_items)} items in {time.time()-t0:.1f}s "
f"({len(aug_items)/max(1,len(train_idx)):.2f}x per source question)")
torch.save({"train": aug_items, "calib": calib_items}, OUT)
print(f"[3/3] saved -> {OUT}")
assert len(calib_items) == n_calib
if __name__ == "__main__":
main()