Spaces:
Sleeping
Sleeping
| """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/<tag>_*.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) | |
| 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() | |