#!/usr/bin/env python3 """Qyvos training: head-only fine-tune of Julia-1 on Open-Jev (release-v2-redistributable). RAM strategy (3.9 GB hard limit): - Encoder + act_head frozen (no grads, no optimizer state for 140.5M params) - Trainable: head (2x TransformerEncoderLayer), type_emb, scorer (~3.7M params, 2.6%) - Encoder forward runs under torch.no_grad() (frozen -> no activations retained) - bs=1 micro-batches, grad accumulation (default 8), num_workers=0 - Parquet read via pyarrow mmap, one row-group (~1000 rows, few MB) at a time - Lazy per-row tokenization, no dataset-wide RAM cache - psutil RAM guard: abort+checkpoint if system available < 350 MB - Head-only checkpoints (~45 MB fp32 incl. optimizer state) Deterministic resume: data order is a pure function of (seed, shard layout): row-group order = rng(seed).permutation(n_row_groups) within-group order = rng((seed, rg_index)).sample(rows) Checkpoints are taken at accumulation boundaries; resume fast-forwards by replaying the deterministic order (counting only, no forward) to rows_seen. """ import argparse import hashlib import json import os import random import resource import sys import time from pathlib import Path import numpy as np import psutil import pyarrow.parquet as pq import torch from torch import nn BASE = Path(os.environ.get("QYVOS_HOME", "/home/z/my-project/download/qyvos")) JULIA_DIR = BASE / "Julia-1" DATA = BASE / "data" / "data" / "release-v2-redistributable" CKPT_DIR = BASE / "models" / "qyvos_ckpt" LOG_DIR = BASE / "logs" MIN_AVAILABLE_BYTES = 350_000_000 # abort below this much system RAM def log(msg: str) -> None: print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True) def row_group_order(n_groups: int, seed: int): rng = random.Random(seed) return rng.sample(range(n_groups), n_groups) def group_row_order(n_rows: int, seed: int, rg_index: int): rng = random.Random(f"{seed}:{rg_index}") idx = list(range(n_rows)) rng.shuffle(idx) return idx def to_engine_row(row: dict) -> dict: """Open-Jev parquet row -> official Julia engine request row.""" state = row["state_json"] if isinstance(state, str): state = json.loads(state) # state_json is a JSON-encoded string return { "state": state, "question": row["question"], "options": list(row["options"]), "type": row["kind"], } def stream_rows(parquet_path: Path, seed: int): """Yield (rg_index, pos, row_dict) in deterministic shuffled order, mmap one rg at a time.""" pf = pq.ParquetFile(parquet_path) # mmap n_groups = pf.metadata.num_row_groups cols = ["kind", "question", "options", "target", "state_json"] for rg in row_group_order(n_groups, seed): tbl = pf.read_row_group(rg, columns=cols) # ~1000 rows, few MB rows = tbl.to_pylist() del tbl for pos in group_row_order(len(rows), seed, rg): yield rg, pos, rows[pos] del rows def take_eval_rows(parquet_path: Path, n_rows: int, seed: int): """Shuffled-mixture eval sample: round-robin across ALL row-groups so every region contributes equally (avoids the single-region bias that once faked 97%).""" pf = pq.ParquetFile(parquet_path) n_groups = pf.metadata.num_row_groups cols = ["kind", "question", "options", "target", "state_json"] buffers = [] for rg in row_group_order(n_groups, seed): tbl = pf.read_row_group(rg, columns=cols) rows = tbl.to_pylist() del tbl idx = group_row_order(len(rows), seed, rg) buffers.append([rows[i] for i in idx]) out = [] cursor = 0 while len(out) < n_rows: added = False for buf in buffers: if cursor < len(buf): out.append(buf[cursor]) added = True if len(out) >= n_rows: break if not added: break cursor += 1 return out def build_trainable(model: nn.Module): """Freeze encoder + act_head; return trainable params (head, type_emb, scorer).""" for p in model.encoder.parameters(): p.requires_grad_(False) for p in model.act_head.parameters(): p.requires_grad_(False) for p in model.head.parameters(): p.requires_grad_(True) for p in model.type_emb.parameters(): p.requires_grad_(True) for p in model.scorer.parameters(): p.requires_grad_(True) trainable = [p for p in model.parameters() if p.requires_grad] return trainable def trainable_state_dict(model: nn.Module): keys = ("head", "type_emb", "scorer") sd = model.state_dict() return {k: v for k, v in sd.items() if k.split(".")[0] in keys} def soft_ce(scores: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: return -(targets * torch.log_softmax(scores, dim=-1)).sum(-1).mean() @torch.no_grad() def evaluate(model, tok, args, kind_names=("choice", "score", "noul"), split="validation"): """Shuffled-mixture evaluation (honest, non-region-biased). Length-bucketed batching: rows are sorted by encoded length before batching so padding waste stays minimal on CPU (3-4x faster than naive batching). Metrics are order-independent, so sorting is safe. """ model.eval() rows = take_eval_rows(DATA / f"{split}-00000-of-00001.parquet", args.eval_rows, seed=999) from julia.data import Collator, sequence collate = Collator(tok, args.max_length, args.head_length) pairs = [] for r in rows: try: req = to_engine_row(r) enc = sequence(tok, req, args.max_length, args.head_length) pairs.append((len(enc["ids"]), req, [float(x) for x in r["target"]])) except Exception: continue pairs.sort(key=lambda p: p[0]) n_ok = 0 n_tot = 0 ce_sum = 0.0 per_kind = {k: [0, 0] for k in kind_names} for i in range(0, len(pairs), args.eval_batch): chunk = pairs[i : i + args.eval_batch] reqs = [c[1] for c in chunk] targets = [c[2] for c in chunk] batch = collate(reqs, include_targets=False) with torch.no_grad(): scores = model(**batch) # (B, Kmax) for b, (req, tgt) in enumerate(zip(reqs, targets)): k = len(req["options"]) s = scores[b, :k] t = torch.tensor(tgt[:k], dtype=torch.float32) ce = -(t * torch.log_softmax(s, dim=-1)).sum().item() pred = int(s.argmax().item()) gold = int(t.argmax().item()) n_tot += 1 n_ok += int(pred == gold) ce_sum += ce kt = req["type"] if kt in per_kind: per_kind[kt][0] += int(pred == gold) per_kind[kt][1] += 1 model.train() model.encoder.eval() # encoder stays in eval mode (frozen, no dropout) acc = n_ok / max(1, n_tot) ce = ce_sum / max(1, n_tot) detail = {k: (v[0] / v[1] if v[1] else float("nan"), v[1]) for k, v in per_kind.items()} return acc, ce, detail def save_ckpt(path: Path, model, optimizer, step, rows_seen, args, extra=None): CKPT_DIR.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(".tmp") payload = { "step": step, "rows_seen": rows_seen, "model": trainable_state_dict(model), "optimizer": optimizer.state_dict(), "rng_torch": torch.get_rng_state(), "rng_python": random.getstate(), "config": vars(args), "extra": extra or {}, } torch.save(payload, tmp) tmp.replace(path) # atomic log(f"[ckpt] saved {path.name} step={step} rows_seen={rows_seen}") def load_ckpt(path: Path, model, optimizer, args): payload = torch.load(path, map_location="cpu", weights_only=False) model.load_state_dict(payload["model"], strict=False) optimizer.load_state_dict(payload["optimizer"]) torch.set_rng_state(payload["rng_torch"]) random.setstate(payload["rng_python"]) log(f"[ckpt] resumed {path.name} step={payload['step']} rows_seen={payload['rows_seen']}") return payload["step"], payload["rows_seen"] def ram_guard(tag: str): avail = psutil.virtual_memory().available if avail < MIN_AVAILABLE_BYTES: log(f"[RAM] available={avail/1e9:.2f}GB < threshold — aborting ({tag})") return False return True def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--max-rows", type=int, default=30_000) ap.add_argument("--accum", type=int, default=8) ap.add_argument("--lr", type=float, default=1e-4) ap.add_argument("--weight-decay", type=float, default=0.01) ap.add_argument("--warmup", type=int, default=100) ap.add_argument("--eval-every", type=int, default=400, help="optimizer steps between evals") ap.add_argument("--eval-rows", type=int, default=300) ap.add_argument("--eval-batch", type=int, default=8) ap.add_argument("--ckpt-every", type=int, default=200, help="optimizer steps between ckpts") ap.add_argument("--max-length", type=int, default=1024) ap.add_argument("--head-length", type=int, default=512) ap.add_argument("--seed", type=int, default=17) ap.add_argument("--time-budget", type=int, default=440, help="seconds; 0 = unlimited") ap.add_argument("--resume", default="auto", help="auto|off") ap.add_argument("--smoke", action="store_true", help="tiny run: 24 rows, eval 40, no resume") args = ap.parse_args() if args.smoke: args.max_rows = 24 args.eval_rows = 40 args.eval_every = 2 args.ckpt_every = 2 args.resume = "off" args.time_budget = 0 torch.manual_seed(args.seed) random.seed(args.seed) np.random.seed(args.seed % (2**31)) CKPT_DIR.mkdir(parents=True, exist_ok=True) LOG_DIR.mkdir(parents=True, exist_ok=True) sys.path.insert(0, str(JULIA_DIR)) from transformers import AutoTokenizer from julia.model import JuliaDecisionModel t0 = time.time() log("loading base Julia-1 (fp32, frozen backbone) ...") model = JuliaDecisionModel.from_pretrained(JULIA_DIR) model.eval() # encoder always eval (no dropout); head toggled below tok = AutoTokenizer.from_pretrained(JULIA_DIR / "tokenizer") trainable = build_trainable(model) n_train = sum(p.numel() for p in trainable) n_total = sum(p.numel() for p in model.parameters()) log(f"trainable {n_train:,} / {n_total:,} params ({100*n_train/n_total:.2f}%)") decay, no_decay = [], [] for n, p in model.named_parameters(): if not p.requires_grad: continue (no_decay if (p.ndim <= 1 or "type_emb" in n) else decay).append(p) optimizer = torch.optim.AdamW( [{"params": decay, "weight_decay": args.weight_decay}, {"params": no_decay, "weight_decay": 0.0}], lr=args.lr, betas=(0.9, 0.999), eps=1e-8, ) ckpt_path = CKPT_DIR / "head_ckpt.pt" step, rows_seen = 0, 0 if args.resume == "auto" and ckpt_path.exists(): step, rows_seen = load_ckpt(ckpt_path, model, optimizer, args) if rows_seen >= args.max_rows: log(f"[DONE] checkpoint already at rows_seen={rows_seen} >= max_rows={args.max_rows}") return 0 model.train() model.encoder.eval() # frozen encoder: keep eval mode (dropout-free) always # deterministic LR schedule: linear warmup -> linear decay to 10% max_steps = (args.max_rows + args.accum - 1) // args.accum def lr_at(s: int) -> float: if s < args.warmup: return args.lr * (s + 1) / args.warmup frac = (s - args.warmup) / max(1, max_steps - args.warmup) return args.lr * (1.0 - 0.9 * min(1.0, frac)) if step == 0: log("eval baseline (step 0) ...") acc, ce, detail = evaluate(model, tok, args, split="validation") log(f"[eval] step={step} rows={rows_seen} acc={acc:.4f} softCE={ce:.4f} " + " ".join(f"{k}={v[0]:.3f}({v[1]})" for k, v in detail.items())) else: log(f"resuming at step={step} (baseline eval skipped; periodic eval continues)") from julia.data import Collator collate = Collator(tok, args.max_length, args.head_length) train_path = DATA / "train-00000-of-00001.parquet" t_start = time.time() loss_window = [] eval_hist = [] skipped = 0 stopped = "done" stream = stream_rows(train_path, args.seed) log(f"training from rows_seen={rows_seen} (skip fast-forward) ...") pending = rows_seen # rows to skip before resuming gradient work for rg_i, pos, row in stream: if pending > 0: pending -= 1 continue if rows_seen >= args.max_rows: stopped = "done" break if args.time_budget and time.time() - t_start > args.time_budget: stopped = "budget" break if rows_seen % 100 == 0 and not ram_guard(f"row {rows_seen}"): stopped = "ram" break try: req = to_engine_row(row) k = len(req["options"]) tgt = torch.tensor([float(x) for x in row["target"][:k]], dtype=torch.float32) if k < 2 or abs(sum(tgt.tolist()) - 1.0) > 0.05: raise ValueError(f"bad row shape k={k} target_sum={tgt.sum():.3f}") except Exception as e: skipped += 1 rows_seen += 1 continue batch = collate([req], include_targets=False) scores = model(**batch)[0, :k] loss = soft_ce(scores[None], tgt[None]) / args.accum loss.backward() loss_window.append(loss.item() * args.accum) rows_seen += 1 if rows_seen % args.accum == 0: lr = lr_at(step) for g in optimizer.param_groups: g["lr"] = lr torch.nn.utils.clip_grad_norm_(trainable, 1.0) optimizer.step() optimizer.zero_grad(set_to_none=True) step += 1 if step % args.eval_every == 0: acc, ce, detail = evaluate(model, tok, args, split="validation") eval_hist.append({"step": step, "rows": rows_seen, "acc": round(acc, 4), "ce": round(ce, 4), "kinds": {k: round(v[0], 4) for k, v in detail.items()}}) log(f"[eval] step={step} rows={rows_seen} acc={acc:.4f} softCE={ce:.4f} " + " ".join(f"{k}={v[0]:.3f}({v[1]})" for k, v in detail.items())) if step % args.ckpt_every == 0: # roll back to accumulation boundary for exact resume save_ckpt(ckpt_path, model, optimizer, step, rows_seen, args, extra={"eval_hist": eval_hist, "skipped": skipped, "loss_tail": loss_window[-50:]}) boundary_rows = (rows_seen // args.accum) * args.accum save_ckpt(ckpt_path, model, optimizer, step, boundary_rows, args, extra={"eval_hist": eval_hist, "skipped": skipped, "loss_tail": loss_window[-50:], "stopped": stopped}) elapsed = time.time() - t_start peak_mb = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024 log(f"[{stopped.upper()}] rows_seen={rows_seen} step={step} skipped={skipped} " f"elapsed={elapsed:.0f}s peakRSS={peak_mb:.0f}MB " f"mean_loss(last50)={sum(loss_window[-50:])/max(1,len(loss_window[-50:])):.4f}") if stopped == "done" and rows_seen >= args.max_rows: log("[DONE] training target reached") return 0 return 0 if stopped in ("budget",) else (3 if stopped == "ram" else 0) if __name__ == "__main__": sys.exit(main())