fsi-anomaly / train /train_sft2.py
FerrellSyntheticIntelligence's picture
backup all: 37 files (final)
97c39f2 verified
Raw
History Blame Contribute Delete
9.07 kB
"""Guarded SFT for TinyLiquid: chat+SOP+forensic mix with TinyStories retention.
Differs from train_sft.py:
* supports raw full-loss retention examples ({"raw": text})
* evals BOTH masked SFT holdout loss AND TinyStories val PPL (coherence guard)
* keeps best.pt (min sft_val_loss while val_ppl < 90) and best_ppl.pt (min ppl)
Usage:
.venv/bin/python train/train_sft2.py --base ckpt/nlp --data data/sft_mix_v2.jsonl \
--ckpt ckpt/v2 --epochs 3 --lr 2e-5
"""
import argparse, json, math, random, time
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from model.config import TinyLiquidConfig, CONFIGS
from model.utils import latest_ckpt
from model.tiny_liquid import TinyLiquid
from data.tokenizer import load_tokenizer
USER_T, ASST_T, EOT_T = "<|user|>", "<|assistant|>", "<|endoftext|>"
PERSONA_T = {"analyst": "<|analyst|>", "skeptic": "<|skeptic|>", "spock": "<|analyst|>", "none": ""}
P_IDS = {"analyst": 1, "skeptic": 2, "spock": 1, "none": 0}
def parse_args():
ap = argparse.ArgumentParser()
ap.add_argument("--base", default="ckpt/nlp")
ap.add_argument("--resume", default="", help="resume from latest ckpt in this dir")
ap.add_argument("--data", default="data/sft_mix_v2.jsonl")
ap.add_argument("--tok", default="data/tokenizer.json")
ap.add_argument("--ckpt", default="ckpt/v2")
ap.add_argument("--val-bin", default="data/valid.bin")
ap.add_argument("--epochs", type=int, default=3)
ap.add_argument("--batch", type=int, default=8)
ap.add_argument("--seq", type=int, default=256)
ap.add_argument("--lr", type=float, default=2e-5)
ap.add_argument("--eval-every", type=int, default=25)
ap.add_argument("--log-every", type=int, default=25)
ap.add_argument("--ppl-guard", type=float, default=90.0)
ap.add_argument("--val-batches", type=int, default=2)
ap.add_argument("--seed", type=int, default=7)
ap.add_argument("--threads", type=int, default=8)
return ap.parse_args()
def tokenize_example(tok, ex, seq, u_id, a_id, eot_id):
if "raw" in ex:
ids = tok.encode(ex["raw"]).ids + [eot_id]
x = torch.tensor(ids[:-1], dtype=torch.long)
y = torch.tensor(ids[1:], dtype=torch.long)
mask = torch.ones_like(y, dtype=torch.bool)
return x[:seq], y[:seq], mask[:seq], 0
persona = PERSONA_T.get(ex.get("persona", "analyst"), PERSONA_T["analyst"])
p_id = P_IDS.get(ex.get("persona"), 1)
user_ids = tok.encode(ex["user"]).ids
asst_ids = tok.encode(ex["assistant"]).ids
if persona:
p_ids_ = tok.encode(persona).ids
ids = p_ids_ + [u_id] + user_ids + [a_id] + asst_ids + [eot_id]
asst_start = len(p_ids_) + 1 + len(user_ids) + 1
else:
ids = [u_id] + user_ids + [a_id] + asst_ids + [eot_id]
asst_start = 1 + len(user_ids) + 1
if len(ids) > seq:
ids = ids[:seq - 1] + [eot_id]
x = torch.tensor(ids[:-1], dtype=torch.long)
y = torch.tensor(ids[1:], dtype=torch.long)
mask = torch.zeros_like(y, dtype=torch.bool)
mask[asst_start - 1:] = True
return x[:seq], y[:seq], mask[:seq], p_id
def collate(items, seq):
xs, ys, ms, ps = [], [], [], []
for x, y, m, p in items:
xs.append(F.pad(x, (0, seq - x.shape[0]), value=0))
ys.append(F.pad(y, (0, seq - y.shape[0]), value=0))
ms.append(F.pad(m, (0, seq - m.shape[0]), value=False))
ps.append(p)
return torch.stack(xs), torch.stack(ys), torch.stack(ms), torch.tensor(ps, dtype=torch.long)
@torch.no_grad()
def val_ppl(model, val_bin, batch=4, seq=64, n_batches=2, seed=0):
mm = np.memmap(val_bin, dtype=np.uint16, mode="r")
total, cnt = 0.0, 0
rng = np.random.RandomState(seed)
n = (len(mm) - 1) // seq
for b in range(n_batches):
s = int(rng.randint(0, n - batch))
buf = torch.stack([torch.from_numpy(mm[s * seq + i * seq: s * seq + i * seq + seq].astype(np.int64))
for i in range(batch)])
x, y = buf[:, :-1], buf[:, 1:]
loss = F.cross_entropy(model(x).reshape(-1, 8192), y.reshape(-1))
total += loss.item() * y.numel(); cnt += y.numel()
return float(np.exp(total / cnt))
def main():
args = parse_args()
torch.set_num_threads(args.threads)
torch.manual_seed(args.seed); random.seed(args.seed)
rng = random.Random(args.seed)
tok = load_tokenizer(args.tok)
u_id, a_id, eot_id = tok.token_to_id(USER_T), tok.token_to_id(ASST_T), tok.token_to_id(EOT_T)
assert None not in (u_id, a_id, eot_id)
exs = [json.loads(l) for l in open(args.data, encoding="utf-8") if l.strip()]
rng.shuffle(exs)
n_eval = min(128, max(8, len(exs) // 12))
eval_ex, train_ex = exs[:n_eval], exs[n_eval:]
print(f"train {len(train_ex)} eval {len(eval_ex)}", flush=True)
base_path = latest_ckpt(args.resume or args.base)
base = torch.load(base_path, map_location="cpu")
base_cfg = base.get("config") or CONFIGS["tiny10m"]
cfg = TinyLiquidConfig(vocab_size=tok.get_vocab_size(),
**{k: v for k, v in base_cfg.items() if k != "vocab_size"})
model = TinyLiquid(cfg)
model.load_state_dict(base["model"])
opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.05)
print(f"loaded base {base_path.name}", flush=True)
out = Path(args.ckpt); out.mkdir(parents=True, exist_ok=True)
steps_per_epoch = max(1, len(train_ex) // args.batch)
total_steps = steps_per_epoch * args.epochs
def make_items(exs_):
return [tokenize_example(tok, e, args.seq, u_id, a_id, eot_id) for e in exs_]
eval_items = make_items(eval_ex)
def run_eval():
model.eval()
total, n = 0.0, 0
for i in range(0, len(eval_items), args.batch):
x, y, m, p = collate(eval_items[i:i + args.batch], args.seq)
with torch.no_grad():
logits = model(x, persona_ids=p).reshape(-1, 8192)
loss = F.cross_entropy(logits, y.reshape(-1), reduction="none")
loss = (loss * m.reshape(-1)).sum() / m.sum()
total += loss.item() * m.sum().item(); n += m.sum().item()
sft_vl = total / n
ppl = val_ppl(model, args.val_bin, n_batches=args.val_batches)
model.train()
return sft_vl, ppl
best_guard, best_ppl = float("inf"), float("inf")
t0 = time.time(); step = 0
for ep in range(args.epochs):
rng.shuffle(train_ex)
items = make_items(train_ex)
for i in range(0, len(items) - len(items) % args.batch, args.batch):
step += 1
x, y, m, p = collate(items[i:i + args.batch], args.seq)
opt.zero_grad(set_to_none=True)
logits = model(x, persona_ids=p).reshape(-1, 8192)
loss = F.cross_entropy(logits, y.reshape(-1), reduction="none")
loss = (loss * m.reshape(-1)).sum() / m.sum()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
if step % args.log_every == 0:
print(f"step {step}/{total_steps} loss {loss.item():.4f} "
f"{args.batch*args.seq*args.log_every/(time.time()-t0):.0f} tok/s", flush=True)
t0 = time.time()
if step % args.eval_every == 0:
sft_vl, ppl = run_eval()
try:
sp = tok.encode("<|analyst|><|user|>Find discrepancies between: Account A: The meeting ended at 11am. Account B: The meeting ended at noon.<|assistant|>").ids
with torch.no_grad():
sout = tok.decode(model.generate(tok, sp, persona_id=1, max_new=50, temperature=0.35,
top_k=20, repetition_penalty=1.25,
no_repeat_ngram_size=4)[len(sp):]).replace("\n", " ").strip()[:180]
print(f" sample: {sout}", flush=True)
except Exception:
pass
tag = ""
if ppl < args.ppl_guard and sft_vl < best_guard:
best_guard = sft_vl
torch.save({"model": model.state_dict(), "step": step, "config": cfg.__dict__}, out / "best.pt")
tag += " [new best]"
if ppl < best_ppl:
best_ppl = ppl
torch.save({"model": model.state_dict(), "step": step, "config": cfg.__dict__}, out / "best_ppl.pt")
tag += " [best ppl]"
torch.save({"model": model.state_dict(), "step": step, "config": cfg.__dict__}, out / f"model_{step}.pt")
print(f" [eval {step}] sft_val_loss {sft_vl:.4f} val_ppl {ppl:.2f}{tag}", flush=True)
torch.save({"model": model.state_dict(), "step": step, "config": cfg.__dict__}, out / "model_final.pt")
print(f"done -> {out} best_guard_sft_loss={best_guard:.4f} best_ppl={best_ppl:.2f}", flush=True)
if __name__ == "__main__":
main()