HyperPEER / gemma /train_layers.py
MikeyBeez's picture
Add HyperPEER pipeline, testbed code, results, docs, Gradio landing
e41a3a4 verified
Raw
History Blame Contribute Delete
8.48 kB
"""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)
@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()