soft-decider-421m / scripts /preprocess.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
4.64 kB
"""Tokenize LocalLLaMA/typed-decisions train parquet -> train_items.pt (CPU).
Bit-identical item construction to the official notebook (laya.common.build_sequence).
Input : td_all_train.parquet (1,200 cases; copied from D:/dsh-home3/laya-ft if missing)
Output: train_items.pt (~6,000 sequences: ids / markers / qtype / soft target / label)
Run: D:\\venvs\\voxcpm-cuda\\Scripts\\python.exe d:\\QoderCN\\jev_finetune\\preprocess.py
"""
from __future__ import annotations
import json
import os
import shutil
import sys
import time
os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
os.environ.setdefault("USE_TF", "0") # avoid the TF-probe deadlock noted in the Laya README
sys.stdout.reconfigure(encoding="utf-8")
WORK = os.path.dirname(os.path.abspath(__file__))
TRAIN_PARQUET = os.path.join(WORK, "td_all_train.parquet")
FALLBACK_TRAIN = r"D:\dsh-home3\laya-ft\td_all_train.parquet"
def build_training_item(tok, cfg, state, q, gold_q):
"""One (case, question) -> one model item with the gold soft distribution as target."""
from laya.common import build_sequence, render_options, QTYPES
t = q["type"]
crit = q.get("criteria", {})
if t == "choice":
keys = list(crit.keys())
target = [gold_q["probabilities"].get(k, 0.0) for k in keys]
elif t == "noul":
target = [gold_q["probabilities"].get("false", 0.5), gold_q["probabilities"].get("true", 0.5)]
elif t == "score":
n_levels = len(crit) if isinstance(crit, list) else 4
target = [gold_q["probabilities"].get(str(i), 0.0) for i in range(n_levels)]
else:
return None
s = sum(target)
target = [v / s for v in target] if s > 0 else [1.0 / len(target)] * len(target)
label = target.index(max(target))
k = len(render_options({"t": t, "crit": crit}))
seq, markers = build_sequence(tok, state, {"t": t, "ins": q["instructions"], "crit": crit},
cfg["max_len"], cfg["head_max_len"])
if len(markers) != k:
return None # options did not fit; skip rather than train on a truncated answer space
return {"ids": seq, "markers": markers, "qtype": QTYPES[t], "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
if not os.path.exists(TRAIN_PARQUET):
if os.path.exists(FALLBACK_TRAIN):
shutil.copy(FALLBACK_TRAIN, TRAIN_PARQUET)
print(f"[copy] {FALLBACK_TRAIN} -> {TRAIN_PARQUET}")
else:
print("[ERR] td_all_train.parquet not found; run via dataset download first")
sys.exit(1)
print("[1/3] fetch tokenizer + config from convaiinnovations/laya (already in HF cache) ...")
model_dir = snapshot_download("convaiinnovations/laya")
_fix_tokenizer_config(model_dir)
tok = AutoTokenizer.from_pretrained(os.path.join(model_dir, "tokenizer"))
with open(os.path.join(model_dir, "rl_agent_config.json")) as f:
cfg = json.load(f)
# Match the fine-tuning run: 1024 ctx / 256 head budget (same as the official notebook)
cfg["max_len"] = 1024
cfg["head_max_len"] = 256
print(f" max_len={cfg['max_len']} head_max_len={cfg['head_max_len']} mask={tok.mask_token!r}")
print("[2/3] preprocessing train split ...")
df = pd.read_parquet(TRAIN_PARQUET)
assert len(df) == 1200, f"expected 1200 train cases, got {len(df)}"
t0 = time.time()
items, skipped, qtype_counts = [], 0, {"choice": 0, "score": 0, "noul": 0}
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:
it = build_training_item(tok, cfg, state, q, gold[qid])
if it:
items.append(it)
qtype_counts[q["type"]] += 1
else:
skipped += 1
out = os.path.join(WORK, "train_items.pt")
torch.save(items, out)
lens = [len(it["ids"]) for it in items]
print(f"[OK] {len(items)} items in {time.time()-t0:.1f}s (skipped {skipped}) -> {out}")
print(f" by qtype: {qtype_counts} | seq len min/avg/max = {min(lens)}/{sum(lens)//len(lens)}/{max(lens)}")
print("[3/3] sanity: re-load ok ->", len(torch.load(out, weights_only=False)))
if __name__ == "__main__":
main()