"""LAYER-LOCAL training of the Gemma hypernetwork student on one captured chunk. For each requested decoder layer, train its GemmaHyperExpert to reproduce the teacher's cached block output: minimize relMSE(student(X), Y) over the chunk's cached (X, Y) activations, with Adafactor (bitsandbytes SEGFAULTs on this Blackwell GPU). Training is CONTINUAL: per-layer (model+optimizer) checkpoints persist in --ckpt-dir and are reloaded each chunk, so the student improves chunk over chunk. Only ONE layer's expert + data is on the GPU at a time, so the full 30-layer, ~7.6B student trains comfortably on the 16 GB card (no teacher resident -- Y is cached). The chunk's per-layer X,Y are held in CPU RAM; minibatches stream to GPU. Held-out eval: if --eval-dir is given, after training each layer we measure its relMSE on the FIXED held-out chunk (captured once, never trained), and append a row to --progress (jsonl). This is the per-layer fidelity tracked over the whole run. Usage: python train_layers.py --chunk-dir /mnt/data/cache/gemma_cap/chunk00 \ --ckpt-dir ./student --layers 0-29 --passes 3 \ --eval-dir /mnt/data/cache/gemma_cap_eval --progress ./student/progress.jsonl \ --chunk-idx 0 --c 9856 --r 4 --b 640 """ import argparse, glob, json, os, time import torch, torch.nn.functional as F from gemma_hyper import GemmaHyperExpert, expert_params def parse_layers(spec, default_n): if not spec or spec == "all": return list(range(default_n)) out = [] for part in spec.split(","): if "-" in part: a, b = part.split("-"); out.extend(range(int(a), int(b) + 1)) else: out.append(int(part)) return out def load_layer_tensors(chunk_dir, layer, tag): """Concatenate all shards layerNN/_*.pt -> one [n, d] bf16 CPU tensor.""" d = os.path.join(chunk_dir, f"layer{layer:02d}") files = sorted(glob.glob(os.path.join(d, f"{tag}_*.pt"))) if not files: raise FileNotFoundError(f"no {tag} shards in {d}") parts = [torch.load(f, map_location="cpu") for f in files] return torch.cat(parts, dim=0) def make_optimizer(params, lr): from transformers.optimization import Adafactor return Adafactor(params, lr=lr, beta1=None, weight_decay=0.0, scale_parameter=False, relative_step=False, warmup_init=False) @torch.no_grad() def eval_relmse(expert, X, Y, dev, mb): """Full-set relMSE = sum((yhat-Y)^2) / sum(Y^2) over the eval chunk.""" was = expert.training; expert.eval() sse = 0.0; sy = 0.0 for i in range(0, X.shape[0], mb): xb = X[i:i + mb].to(dev, non_blocking=True) yb = Y[i:i + mb].to(dev, non_blocking=True) yhat = expert(xb) sse += (yhat.float() - yb.float()).pow(2).sum().item() sy += yb.float().pow(2).sum().item() if was: expert.train() return sse / max(sy, 1e-12) def train_one_layer(layer, args, dev): ckpt = os.path.join(args.ckpt_dir, f"layer{layer:02d}.pt") pdtype = torch.float32 if args.param_dtype == "fp32" else torch.bfloat16 expert = GemmaHyperExpert(args.hidden, args.c, args.r, args.b, dtype=pdtype).to(dev) params = list(expert.parameters()) opt = make_optimizer(params, args.lr) start_seen = 0 if os.path.exists(ckpt): st = torch.load(ckpt, map_location=dev) expert.load_state_dict(st["model"]) try: opt.load_state_dict(st["opt"]) except Exception as e: print(f" [L{layer}] opt state not restored ({str(e)[:60]}); fresh opt", flush=True) start_seen = st.get("tokens_seen", 0) X = load_layer_tensors(args.chunk_dir, layer, "input") Y = load_layer_tensors(args.chunk_dir, layer, "output") assert X.shape == Y.shape, f"L{layer} X{tuple(X.shape)} != Y{tuple(Y.shape)}" n = X.shape[0] if args.pin: X = X.pin_memory(); Y = Y.pin_memory() init_rel = eval_relmse(expert, X, Y, dev, args.mb) # train-chunk relMSE before this chunk expert.train() g = torch.Generator().manual_seed(1234 + layer) step = 0 total_steps = args.passes * ((n + args.mb - 1) // args.mb) t0 = time.time() last = init_rel for ep in range(args.passes): perm = torch.randperm(n, generator=g) for i in range(0, n, args.mb): idx = perm[i:i + args.mb] xb = X[idx].to(dev, non_blocking=True) yb = Y[idx].to(dev, non_blocking=True) yhat = expert(xb) num = (yhat.float() - yb.float()).pow(2).mean() den = yb.float().pow(2).mean().clamp_min(1e-12) loss = num / den # short warmup each chunk (Adafactor 2nd-moment resets across processes) lr = args.lr * min(1.0, (step + 1) / max(1, args.warmup)) for pg in opt.param_groups: pg["lr"] = lr opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(params, 1.0) opt.step() last = loss.item() step += 1 tok_seen = start_seen + args.passes * n torch.save({"model": expert.state_dict(), "opt": opt.state_dict(), "tokens_seen": tok_seen, "cfg": {"hidden": args.hidden, "c": args.c, "r": args.r, "b": args.b}}, ckpt) train_rel = eval_relmse(expert, X, Y, dev, args.mb) # train-chunk relMSE after eval_rel = None if args.eval_dir: EX = load_layer_tensors(args.eval_dir, layer, "input") EY = load_layer_tensors(args.eval_dir, layer, "output") eval_rel = eval_relmse(expert, EX, EY, dev, args.mb) dt = time.time() - t0 tps = (args.passes * n) / dt print(f"[L{layer:02d}] n={n} init_rel={init_rel:.4f} -> train_rel={train_rel:.4f}" + (f" | EVAL_rel={eval_rel:.4f}" if eval_rel is not None else "") + f" | seen={tok_seen/1e6:.1f}M {tps/1000:.0f}k tok/s {dt:.0f}s", flush=True) row = {"chunk_idx": args.chunk_idx, "layer": layer, "n_tokens": n, "init_train_rel": round(init_rel, 5), "train_rel": round(train_rel, 5), "eval_rel": (round(eval_rel, 5) if eval_rel is not None else None), "tokens_seen": tok_seen, "last_loss": round(last, 5)} if args.progress: with open(args.progress, "a") as f: f.write(json.dumps(row) + "\n") del X, Y, expert, opt torch.cuda.empty_cache() return row def main(): ap = argparse.ArgumentParser() ap.add_argument("--chunk-dir", required=True) ap.add_argument("--ckpt-dir", required=True) ap.add_argument("--eval-dir", default="") ap.add_argument("--progress", default="") ap.add_argument("--layers", default="all") ap.add_argument("--chunk-idx", type=int, default=0) ap.add_argument("--hidden", type=int, default=2816) ap.add_argument("--c", type=int, default=9856) ap.add_argument("--r", type=int, default=4) ap.add_argument("--b", type=int, default=640) ap.add_argument("--passes", type=int, default=3) ap.add_argument("--mb", type=int, default=8192) ap.add_argument("--lr", type=float, default=2e-4) ap.add_argument("--warmup", type=int, default=100) ap.add_argument("--param-dtype", default="fp32", choices=["fp32", "bf16"], help="fp32 (default) is stable on outlier-heavy layers and " "affordable here (one layer trained at a time).") ap.add_argument("--pin", action="store_true", help="pin chunk tensors (faster H2D)") args = ap.parse_args() # detect #layers present in the chunk present = sorted(int(os.path.basename(p)[5:]) for p in glob.glob(os.path.join(args.chunk_dir, "layer*"))) n_present = (present[-1] + 1) if present else 0 layers = [l for l in parse_layers(args.layers, n_present) if l in present] os.makedirs(args.ckpt_dir, exist_ok=True) dev = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"=== TRAIN chunk={args.chunk_idx} dir={args.chunk_dir} layers={layers} " f"c={args.c} r={args.r} b={args.b} ({expert_params(args.hidden,args.c,args.r,args.b)/1e6:.0f}M/layer) " f"passes={args.passes} mb={args.mb} lr={args.lr} dev={dev} ===", flush=True) t0 = time.time() for layer in layers: train_one_layer(layer, args, dev) print(f"=== TRAIN DONE {len(layers)} layers in {time.time()-t0:.0f}s ===", flush=True) if __name__ == "__main__": main()