fsi-anomaly / train /train_sft_v4.py
FerrellSyntheticIntelligence's picture
backup all: 37 files (final)
97c39f2 verified
Raw
History Blame Contribute Delete
11.8 kB
"""Conservative SFT for TinyLiquid v4.
Designed for the 7.8M TinyLiquid model where full SFT caused catastrophic
forgetting/model-collapse into broken forensic jargon. Defaults freeze most of
the model, train only persona embeddings + final liquid block + output norm,
and add a base-model KL anchor.
"""
import argparse
import json
import random
import time
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from model.config import CONFIGS, TinyLiquidConfig
from model.tiny_liquid import TinyLiquid
from model.utils import latest_ckpt
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('--data', default='data/sft_mix_v4.jsonl')
ap.add_argument('--tok', default='data/tokenizer.json')
ap.add_argument('--ckpt', default='ckpt/v4')
ap.add_argument('--val-bin', default='data/valid.bin')
ap.add_argument('--epochs', type=int, default=2)
ap.add_argument('--batch', type=int, default=8)
ap.add_argument('--seq', type=int, default=256)
ap.add_argument('--lr', type=float, default=8e-6)
ap.add_argument('--kl', type=float, default=0.05)
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=45.0)
ap.add_argument('--val-batches', type=int, default=2)
ap.add_argument('--resume', default=None, help='continue from this checkpoint (model weights + step); teacher/KL anchor stays --base')
ap.add_argument('--resume-best-sft', type=float, default=None)
ap.add_argument('--resume-best-ppl', type=float, default=None)
ap.add_argument('--resume-iter', type=int, default=None, help='iterations already trained (needed when saved step is a shifted label)')
ap.add_argument('--train-scope', choices=['head', 'last', 'last2', 'full'], default='last')
ap.add_argument('--seed', type=int, default=11)
ap.add_argument('--threads', type=int, default=8)
return ap.parse_args()
def resolve_ckpt(path):
p = Path(path)
if p.is_file():
return p
ck = latest_ckpt(p)
assert ck, f'no checkpoints in {path}'
return ck
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_name = ex.get('persona', 'analyst')
persona = PERSONA_T.get(persona_name, PERSONA_T['analyst'])
p_id = P_IDS.get(persona_name, 1)
user_ids = tok.encode(ex['user']).ids
asst_ids = tok.encode(ex['assistant']).ids
p_ids = tok.encode(persona).ids if persona else []
ids = p_ids + [u_id] + user_ids + [a_id] + asst_ids + [eot_id]
asst_start = len(p_ids) + 1 + len(user_ids) + 1
if len(ids) > seq:
return None
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
if int(mask.sum().item()) < 16:
return None
return x, y, mask, 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)
def set_train_scope(model, scope):
for p in model.parameters():
p.requires_grad = False
for p in model.persona_emb.parameters():
p.requires_grad = True
if scope in {'last', 'last2', 'full'}:
for p in model.blocks[-1].parameters():
p.requires_grad = True
for p in model.norm_out.parameters():
p.requires_grad = True
if scope in {'last2', 'full'}:
for p in model.blocks[-2].parameters():
p.requires_grad = True
if scope == 'full':
for p in model.parameters():
p.requires_grad = True
return sum(p.numel() for p in model.parameters() if p.requires_grad)
@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 _ 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, model.cfg.vocab_size), y.reshape(-1))
total += loss.item() * y.numel()
cnt += y.numel()
return float(np.exp(total / cnt))
@torch.no_grad()
def sample_text(model, tok, prompt, persona_id):
ids = tok.encode(prompt).ids
out = model.generate(tok, ids, persona_id=persona_id, max_new=60, temperature=0.35, top_k=20, repetition_penalty=1.25, no_repeat_ngram_size=4)
return tok.decode(out[len(ids):]).replace('\n', ' ').strip()[:220]
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)
raw_exs = [json.loads(line) for line in open(args.data, encoding='utf-8') if line.strip()]
rng.shuffle(raw_exs)
base_path = resolve_ckpt(args.base)
base = torch.load(base_path, map_location='cpu')
resume_path = resolve_ckpt(args.resume) if args.resume else None
resume = torch.load(resume_path, map_location='cpu') if resume_path else None
config = (resume or base).get('config') or CONFIGS['tiny10m']
cfg = TinyLiquidConfig(vocab_size=tok.get_vocab_size(),
**{k: v for k, v in config.items() if k != 'vocab_size'})
model = TinyLiquid(cfg)
model.load_state_dict((resume or base)['model'])
teacher = None
if args.kl > 0:
teacher = TinyLiquid(cfg)
teacher.load_state_dict(base['model'])
teacher.eval()
for p in teacher.parameters():
p.requires_grad = False
trainable = set_train_scope(model, args.train_scope)
opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=args.lr, betas=(0.9, 0.95), weight_decay=0.02)
print(f'loaded base {base_path} trainable {trainable}/{sum(p.numel() for p in model.parameters())} scope={args.train_scope}', flush=True)
items_all = [tokenize_example(tok, e, args.seq, u_id, a_id, eot_id) for e in raw_exs]
items_all = [item for item in items_all if item is not None]
rng.shuffle(items_all)
n_eval = min(128, max(16, len(items_all) // 12))
eval_items, train_items = items_all[:n_eval], items_all[n_eval:]
print(f'train {len(train_items)} eval {len(eval_items)} filtered {len(raw_exs) - len(items_all)}', flush=True)
out = Path(args.ckpt)
out.mkdir(parents=True, exist_ok=True)
best_score = args.resume_best_sft if args.resume_best_sft is not None else float('inf')
best_ppl = args.resume_best_ppl if args.resume_best_ppl is not None else float('inf')
step = resume.get('step', 0) if resume is not None else 0
start_iter = args.resume_iter if args.resume_iter is not None else (resume.get('iter', step) if resume is not None else 0)
start_step = step
t0 = time.time()
total_steps = (len(train_items) // args.batch) * args.epochs
if start_step:
print(f'resuming from {resume_path} at step {start_step}/{total_steps} iter {start_iter}', flush=True)
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)
logits = model(x, persona_ids=p).reshape(-1, model.cfg.vocab_size)
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()
ppl = val_ppl(model, args.val_bin, n_batches=args.val_batches, seed=args.seed + step)
model.train()
return total / n, ppl
model.train()
iter_no = 0
for ep in range(args.epochs):
rng.shuffle(train_items)
usable = len(train_items) - len(train_items) % args.batch
for i in range(0, usable, args.batch):
iter_no += 1
if iter_no <= start_iter:
continue
step += 1
x, y, m, p = collate(train_items[i:i + args.batch], args.seq)
opt.zero_grad(set_to_none=True)
logits = model(x, persona_ids=p)
sft_loss = F.cross_entropy(logits.reshape(-1, model.cfg.vocab_size), y.reshape(-1), reduction='none')
sft_loss = (sft_loss * m.reshape(-1)).sum() / m.sum()
loss = sft_loss
if teacher is not None:
with torch.no_grad():
t_logits = teacher(x, persona_ids=p)
kl = F.kl_div(
F.log_softmax(logits.float(), dim=-1),
F.softmax(t_logits.float(), dim=-1),
reduction='none',
).sum(dim=-1)
kl = (kl * m).sum() / m.sum()
loss = loss + args.kl * kl
loss.backward()
torch.nn.utils.clip_grad_norm_([p for p in model.parameters() if p.requires_grad], 0.5)
opt.step()
if step % args.log_every == 0:
print(f'step {step}/{total_steps} loss {loss.item():.4f} sft {sft_loss.item():.4f} {args.batch * args.seq * args.log_every / max(1e-6, time.time() - t0):.0f} tok/s', flush=True)
t0 = time.time()
if step % args.eval_every == 0:
sft_vl, ppl = run_eval()
tag = ''
if ppl < best_ppl:
best_ppl = ppl
torch.save({'model': model.state_dict(), 'step': step, 'iter': iter_no, 'config': cfg.__dict__}, out / 'best_ppl.pt')
tag += ' [best ppl]'
if ppl <= args.ppl_guard and sft_vl < best_score:
best_score = sft_vl
torch.save({'model': model.state_dict(), 'step': step, 'iter': iter_no, 'config': cfg.__dict__}, out / 'best.pt')
tag += ' [new best]'
torch.save({'model': model.state_dict(), 'step': step, 'iter': iter_no, '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)
print(' sample:', sample_text(model, tok, '<|analyst|><|user|>Find discrepancies between: Account A: The meeting ended at 11am. Account B: The meeting ended at noon.<|assistant|>', 1), flush=True)
torch.save({'model': model.state_dict(), 'step': step, 'iter': iter_no, 'config': cfg.__dict__}, out / 'model_final.pt')
print(f'done -> {out} best_sft={best_score:.4f} best_ppl={best_ppl:.2f}', flush=True)
if __name__ == '__main__':
main()