Spaces:
Running on Zero
Running on Zero
Download oev/distill.py from divyanshudhruv/oev-demo: direct link, hf CLI and curl.
- Browser
- Download file 9.9 kB
-
https://huggingface.co/spaces/divyanshudhruv/oev-demo/resolve/main/oev/distill.py
- Command line
-
hf download hf://spaces/divyanshudhruv/oev-demo/oev/distill.py
-
curl -L -o distill.py https://huggingface.co/spaces/divyanshudhruv/oev-demo/resolve/main/oev/distill.py
9.9 kB
| import argparse | |
| import json | |
| import os | |
| import random | |
| import time | |
| import torch | |
| import torch.nn.functional as F | |
| from oev.evaluate import load_model, pack_question | |
| from oev.tokenizer_hf import HFTokenPacker | |
| class Teacher: | |
| # one frozen teacher: forward a packed case, return per-question probs | |
| def __init__(self, path, device): | |
| self.model = load_model(path, device) | |
| self.packer = HFTokenPacker(self.model.cfg["backbone"]) | |
| self.device = device | |
| def question_probs(self, state, question, max_len, tau): | |
| # distill cases always carry answers; the default mirrors rotate()'s contract | |
| q = pack_question(dict(question, answer=question.get("answer", question["options"][0]))) | |
| ids, anchors, _ = self.packer.pack(state, q, min(max_len, self.model.cfg["max_len"])) | |
| tids = torch.tensor([ids], device=self.device) | |
| pmask = torch.zeros(1, len(ids), dtype=torch.bool, device=self.device) | |
| apos = torch.tensor([anchors], device=self.device) | |
| logits = self.model(tids, pmask, apos)[0] | |
| return F.softmax(logits.float() / tau, dim=-1) | |
| def rotate(q, rng, p): | |
| # order-invariance augmentation: rotate options + answer in lockstep | |
| n = len(q["options"]) | |
| if n < 3 or rng.random() > p: | |
| return q | |
| k = rng.randrange(1, n) | |
| ops = q["options"] | |
| rot = ops[k:] + ops[:k] | |
| out = dict(q) | |
| out["options"] = rot | |
| ans = q.get("answer") | |
| if ans in ops: | |
| out["answer"] = rot[(ops.index(ans) - k) % n] | |
| return out | |
| def load_cases(data_dir, domains, max_cases, seed, per_domain=0): | |
| # with per_domain>0, balance every domain to that count: big domains | |
| # sampled down, small ones repeated (rotation makes repeats non-identical) | |
| rng = random.Random(seed) | |
| pools = [] | |
| for d in domains: | |
| path = os.path.join(data_dir, d, "train.jsonl") | |
| if not os.path.exists(path): | |
| print(f"[distill] skipping {d} (no train split)", flush=True) | |
| continue | |
| with open(path, encoding="utf-8") as f: | |
| rows = [json.loads(line) for line in f] | |
| if per_domain: | |
| if len(rows) >= per_domain: | |
| rows = rng.sample(rows, per_domain) | |
| else: | |
| rows = (rows * (per_domain // len(rows) + 1))[:per_domain] | |
| print(f"[distill] {d}: {len(rows)} cases (balanced)", flush=True) | |
| pools.append(rows) | |
| cases = [c for rows in pools for c in rows] | |
| rng.shuffle(cases) | |
| if max_cases and len(cases) > max_cases: | |
| cases = cases[:max_cases] | |
| return cases | |
| def run_epoch(student, teachers, cases, packer, device, opt, scaler, args, epoch, f): | |
| student.train() | |
| rng = random.Random(epoch + 1) | |
| order = torch.randperm(len(cases)) | |
| running, t0, n_seen = 0.0, time.time(), 0 | |
| running_ce, running_kl, n_ce, n_kl = 0.0, 0.0, 0, 0 | |
| for bi, start in enumerate(range(0, len(order), args.batch_size)): | |
| batch = [cases[i] for i in order[start:start + args.batch_size].tolist()] | |
| opt.zero_grad(set_to_none=True) | |
| with torch.autocast("cuda", enabled=device == "cuda"): | |
| losses, ce_parts, kl_parts = [], [], [] | |
| for case in batch: | |
| q = rotate(case["questions"][0], rng, args.rotate) | |
| state = case["state"] | |
| sq = dict(q, answer=q.get("answer", q["options"][0])) | |
| ids, anchors, _ = packer.pack(state, sq, student.cfg["max_len"]) | |
| tids = torch.tensor([ids], device=device) | |
| pmask = torch.zeros(1, len(ids), dtype=torch.bool, device=device) | |
| apos = torch.tensor([anchors], device=device) | |
| slogits = student(tids, pmask, apos)[0] | |
| logq = F.log_softmax(slogits.float() / args.tau, dim=-1) | |
| tprobs, n_scored = None, 0 | |
| for t in teachers: | |
| tp = t.question_probs(state, q, student.cfg['max_len'], args.tau) | |
| if tp.shape[0] == len(q["options"]): | |
| tprobs = tp if tprobs is None else (tprobs + tp) | |
| n_scored += 1 | |
| if tprobs is None: | |
| continue | |
| tprobs = tprobs / n_scored # mean over teachers that actually scored this option set | |
| kl = -(tprobs * logq).sum() * (args.tau ** 2) | |
| kl_parts.append(kl) | |
| ans = q.get("answer") | |
| if args.alpha and ans in q["options"]: | |
| gold = q["options"].index(ans) | |
| ce = F.cross_entropy(slogits.float().unsqueeze(0), | |
| torch.tensor([gold], device=device)) | |
| ce_parts.append(ce) | |
| losses.append((1.0 - args.alpha) * kl + args.alpha * ce) | |
| else: | |
| losses.append(kl) | |
| if not losses: | |
| continue | |
| loss = torch.stack(losses).mean() | |
| scaler.scale(loss).backward() | |
| scaler.step(opt) | |
| scaler.update() | |
| running += loss.item() * len(losses) | |
| n_seen += len(losses) | |
| if kl_parts: | |
| running_kl += torch.stack(kl_parts).sum().item() | |
| n_kl += len(kl_parts) | |
| if ce_parts: | |
| running_ce += torch.stack(ce_parts).sum().item() | |
| n_ce += len(ce_parts) | |
| if bi % args.log_every == 0: | |
| msg = (f"epoch {epoch} step {bi} loss {running / max(n_seen, 1):.4f} " | |
| f"(ce {running_ce / max(n_ce, 1):.3f} kl {running_kl / max(n_kl, 1):.3f}) " | |
| f"({n_seen / max(time.time() - t0, 1):.2f} cases/s)") | |
| print(msg, flush=True) | |
| f.write(msg + "\n") | |
| f.flush() | |
| if args.save_every and bi and bi % args.save_every == 0: | |
| torch.save({"config": student.cfg, "state": student.state_dict()}, | |
| os.path.join(args.out, f"student-e{epoch}-s{bi}.pt")) | |
| return running / max(n_seen, 1) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--teachers", required=True, help="comma-separated teacher checkpoint paths") | |
| ap.add_argument("--student-init", default=None, help="optional warm-start checkpoint for the student") | |
| ap.add_argument("--data-dir", default="data") | |
| ap.add_argument("--domains", default="banking77,ag_news,emotion", | |
| help="comma-separated domain dirs under --data-dir, each with train.jsonl") | |
| ap.add_argument("--out", default="checkpoints_distill") | |
| ap.add_argument("--epochs", type=int, default=1) | |
| ap.add_argument("--batch-size", type=int, default=8) | |
| ap.add_argument("--lr", type=float, default=2e-5) | |
| ap.add_argument("--rotate", type=float, default=0.5, help="option-rotation augmentation probability") | |
| ap.add_argument("--tau", type=float, default=1.0, help="distillation temperature") | |
| ap.add_argument("--alpha", type=float, default=0.3, | |
| help="weight on gold cross-entropy; (1-alpha) goes to teacher KL. 0 = pure mimicry") | |
| ap.add_argument("--per-domain", type=int, default=0, | |
| help="if >0, balance each domain to this many train cases") | |
| ap.add_argument("--max-cases", type=int, default=27000) | |
| ap.add_argument("--seed", type=int, default=1234) | |
| ap.add_argument("--log-every", type=int, default=100) | |
| ap.add_argument("--save-every", type=int, default=2000) | |
| ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") | |
| args = ap.parse_args() | |
| os.makedirs(args.out, exist_ok=True) | |
| device = args.device | |
| log = open(os.path.join(args.out, "distill.log"), "a") # noqa: SIM115 - kept open for the whole run | |
| if args.student_init: | |
| student = load_model(args.student_init, device) | |
| print(f"student warm-started from {args.student_init}", flush=True) | |
| else: | |
| first = load_model(args.teachers.split(",")[0].strip(), device) | |
| from oev.model import HFBackboneOEV | |
| student = HFBackboneOEV(first.cfg["backbone"]) | |
| student.cfg = first.cfg | |
| student.eval().to(device) | |
| print("student initialized from scratch", flush=True) | |
| teacher_paths = [p.strip() for p in args.teachers.split(",") if p.strip()] | |
| teachers = [] | |
| for i, tpath in enumerate(teacher_paths, 1): | |
| # each Teacher is a full backbone load: the slowest silent stretch in | |
| # this script, so every load announces itself | |
| print(f"loading teacher {i}/{len(teacher_paths)}: {tpath} (model load, 1-2 min each)", flush=True) | |
| teachers.append(Teacher(tpath, device)) | |
| print(f"{len(teachers)} teachers loaded", flush=True) | |
| packer = HFTokenPacker(student.cfg["backbone"]) | |
| cases = load_cases(args.data_dir, [d.strip() for d in args.domains.split(",")], | |
| args.max_cases, args.seed, per_domain=args.per_domain) | |
| print(f"{len(cases)} training cases", flush=True) | |
| if not cases: | |
| raise SystemExit("[distill] 0 training cases - check --data-dir/--domains; " | |
| "each domain needs {data_dir}/{domain}/train.jsonl " | |
| "(build with: python -m oev.convert && python -m oev.convert_banking77 " | |
| "&& python -m oev.convert_typed)") | |
| opt = torch.optim.AdamW(student.parameters(), lr=args.lr) | |
| scaler = torch.amp.GradScaler(enabled=device == "cuda") | |
| for epoch in range(args.epochs): | |
| tl = run_epoch(student, teachers, cases, packer, device, opt, scaler, args, epoch, log) | |
| msg = f"epoch {epoch} done - train loss {tl:.4f}" | |
| print(msg, flush=True) | |
| log.write(msg + "\n") | |
| log.flush() | |
| torch.save({"config": student.cfg, "state": student.state_dict()}, | |
| os.path.join(args.out, "student.pt")) | |
| print(f"saved -> {args.out}/student.pt", flush=True) | |
| if __name__ == "__main__": | |
| main() | |