| """Train PixelModel v3. The canonical output is model.png (weights-as-pixels). |
| |
| Designed for a single 4 GB RTX 3050: mixed precision, a configurable batch and |
| crop size, a cosine LR schedule, and a peak-VRAM budget it actively watches. |
| |
| The default run is a multi-hour job that meaningfully uses the GPU - not v1's |
| 30-minute CPU toy. Shorten it with --epochs for a quick sanity pass. |
| |
| python train.py --data ../pm-work/coco_train.npz --vocab ../pm-work/vocab.json |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import math |
| import os |
| import time |
|
|
| import numpy as np |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
| from model import ( |
| ModelConfig, PixelModelV3, make_coord_grid, save_model_png, load_vocab, |
| ) |
|
|
|
|
| def get_args(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--data", default="../pm-work/coco_train.npz") |
| ap.add_argument("--vocab", default="../pm-work/vocab.json") |
| ap.add_argument("--out-png", default="model.png") |
| ap.add_argument("--out-config", default="config.json") |
|
|
| ap.add_argument("--epochs", type=int, default=80, |
| help="default is a multi-hour run; lower it for a quick pass") |
| ap.add_argument("--batch-size", type=int, default=16) |
| ap.add_argument("--crop", type=int, default=128, help="training crop size") |
| ap.add_argument("--pixels-per-step", type=int, default=0, |
| help="decode only this many random pixels per image each step " |
| "(0 = the full crop). Decouples VRAM from crop size, so " |
| "high-res crops train with a usable batch (NeRF-style).") |
| ap.add_argument("--lr", type=float, default=2e-4) |
| ap.add_argument("--min-lr", type=float, default=1e-6, help="cosine floor") |
| ap.add_argument("--warmup-steps", type=int, default=500) |
| ap.add_argument("--grad-clip", type=float, default=1.0) |
| ap.add_argument("--steps-per-epoch", type=int, default=0, |
| help="0 = one pass over the data per epoch") |
|
|
| ap.add_argument("--embed-dim", type=int, default=None) |
| ap.add_argument("--text-hidden", type=int, default=None) |
| ap.add_argument("--z-dim", type=int, default=None) |
| ap.add_argument("--decoder-width", type=int, default=None) |
| ap.add_argument("--num-sine-layers", type=int, default=None) |
| ap.add_argument("--num-freq", type=int, default=None) |
| ap.add_argument("--w0-first", type=float, default=None) |
| ap.add_argument("--w0-hidden", type=float, default=None) |
|
|
| ap.add_argument("--device", default="cuda") |
| ap.add_argument("--data-device", choices=["auto", "gpu", "cpu"], default="auto", |
| help="where images live for cropping. 'gpu' keeps the whole " |
| "image tensor on-device and crops there (removes the CPU " |
| "bottleneck on fast GPUs); 'auto' = gpu when CUDA is used") |
| ap.add_argument("--amp", dest="amp", action="store_true", default=True) |
| ap.add_argument("--no-amp", dest="amp", action="store_false") |
| ap.add_argument("--grad-checkpoint", action="store_true", |
| help="checkpoint sine layers to cut activation memory " |
| "(lets XL use a bigger batch on a small/shared GPU)") |
| ap.add_argument("--vram-budget-gb", type=float, default=3.5) |
| ap.add_argument("--log-interval", type=int, default=50) |
| ap.add_argument("--save-every-epochs", type=int, default=1) |
| ap.add_argument("--limit", type=int, default=0, help="use only N samples (debug)") |
| ap.add_argument("--seed", type=int, default=0) |
| return ap.parse_args() |
|
|
|
|
| def build_config(args, vocab_size, max_tokens) -> ModelConfig: |
| cfg = ModelConfig(vocab_size=vocab_size, max_tokens=max_tokens) |
| for name in ("embed_dim", "text_hidden", "z_dim", "decoder_width", |
| "num_sine_layers", "num_freq", "w0_first", "w0_hidden"): |
| v = getattr(args, name) |
| if v is not None: |
| setattr(cfg, name, v) |
| return cfg |
|
|
|
|
| def cosine_lr(step, total_steps, base_lr, min_lr, warmup): |
| if step < warmup: |
| return base_lr * (step + 1) / max(1, warmup) |
| t = (step - warmup) / max(1, total_steps - warmup) |
| t = min(1.0, t) |
| return min_lr + 0.5 * (base_lr - min_lr) * (1 + math.cos(math.pi * t)) |
|
|
|
|
| def random_crops(images_u8, idx, crop, rng): |
| """images_u8: (N,S,S,3) uint8 -> batch (B, crop, crop, 3) float32 in [0,1].""" |
| B = len(idx) |
| S = images_u8.shape[1] |
| out = np.empty((B, crop, crop, 3), dtype=np.float32) |
| for i, j in enumerate(idx): |
| top = rng.integers(0, S - crop + 1) |
| left = rng.integers(0, S - crop + 1) |
| out[i] = images_u8[j, top:top + crop, left:left + crop, :].astype(np.float32) / 255.0 |
| return out |
|
|
|
|
| def gpu_random_crops(images_u8, idx, crop, gen): |
| """Fully vectorised random crop on-device. images_u8: (N,S,S,3) uint8 on GPU, |
| idx: (B,) long on GPU -> (B, crop, crop, 3) float32 in [0,1]. No python loop, |
| no host<->device copy per step, so a fast GPU is not left waiting on the CPU.""" |
| B = idx.shape[0] |
| S = images_u8.shape[1] |
| sel = images_u8[idx] |
| span = S - crop + 1 |
| top = torch.randint(0, span, (B,), generator=gen, device=sel.device) |
| left = torch.randint(0, span, (B,), generator=gen, device=sel.device) |
| ar = torch.arange(crop, device=sel.device) |
| rows = (top[:, None] + ar)[:, :, None] |
| cols = (left[:, None] + ar)[:, None, :] |
| b = torch.arange(B, device=sel.device)[:, None, None] |
| out = sel[b, rows, cols] |
| return out.float() / 255.0 |
|
|
|
|
| def print_activation_stats(model, tokens, coords, tag): |
| """SIREN sanity check: per-layer sine-output mean/std. Healthy SIREN layers |
| sit near mean~0, std~0.5-0.7. Dead (~0 std) or exploding (>>1) means the |
| init is wrong - catch it here, not three hours into a bad loss curve.""" |
| model.eval() |
| with torch.no_grad(): |
| _, stats = model(tokens, coords, return_stats=True) |
| model.train() |
| line = " | ".join(f"L{i}: mean={m:+.3f} std={s:.3f}" for i, (m, s) in enumerate(stats)) |
| print(f"[siren-stats {tag}] {line}") |
|
|
|
|
| def main(): |
| args = get_args() |
| torch.manual_seed(args.seed) |
| rng = np.random.default_rng(args.seed) |
| device = args.device if torch.cuda.is_available() or args.device == "cpu" else "cpu" |
| if device == "cpu": |
| print("[warn] CUDA not available -> running on CPU (AMP disabled)") |
| args.amp = False |
|
|
| from model import encode_caption |
| data = np.load(args.data, allow_pickle=True) |
| images = data["images"] |
| max_tokens = int(data["max_tokens"]) if "max_tokens" in data else 20 |
| if args.limit: |
| images = images[:args.limit] |
| N, S = images.shape[0], images.shape[1] |
| assert S >= args.crop, f"stored size {S} < crop {args.crop}" |
|
|
| vocab = load_vocab(args.vocab) |
| cfg = build_config(args, vocab_size=len(vocab), max_tokens=max_tokens) |
|
|
| if "captions" in data: |
| caps = [str(c) for c in data["captions"]] |
| if args.limit: |
| caps = caps[:args.limit] |
| tokens_all = np.stack([encode_caption(c, vocab, max_tokens) for c in caps]).astype(np.int64) |
| else: |
| tokens_all = data["tokens"].astype(np.int64) |
| if args.limit: |
| tokens_all = tokens_all[:args.limit] |
|
|
| data_on_gpu = args.data_device == "gpu" or (args.data_device == "auto" and device == "cuda") |
| gen = torch.Generator(device=device).manual_seed(args.seed) |
| if data_on_gpu: |
| images_dev = torch.from_numpy(np.ascontiguousarray(images)).to(device) |
| tokens_dev = torch.from_numpy(tokens_all).to(device) |
| gb = images_dev.numel() / 1e9 |
| print(f"[train] images resident on {device}: {gb:.2f} GB uint8") |
| else: |
| images_dev = tokens_dev = None |
|
|
| print(f"[train] {N} images @ {S}px, crop={args.crop}, vocab={len(vocab)}, " |
| f"device={device}, amp={args.amp}, data_on_gpu={data_on_gpu}") |
|
|
| model = PixelModelV3(cfg).to(device) |
| model.grad_checkpoint = args.grad_checkpoint |
| n_params = sum(p.numel() for p in model.parameters()) |
| print(f"[train] model params = {n_params:,}, grad_checkpoint={args.grad_checkpoint}") |
|
|
| opt = torch.optim.Adam(model.parameters(), lr=args.lr, betas=(0.9, 0.99)) |
| scaler = torch.amp.GradScaler("cuda", enabled=args.amp) |
|
|
| steps_per_epoch = args.steps_per_epoch or max(1, N // args.batch_size) |
| total_steps = steps_per_epoch * args.epochs |
|
|
| coords = make_coord_grid(args.crop, args.crop, device=device, dtype=torch.float32) |
| coords = coords.unsqueeze(0) |
| HW = args.crop * args.crop |
| if args.pixels_per_step: |
| print(f"[train] pixel subsampling: {args.pixels_per_step}/{HW} pixels per image per step") |
|
|
| print(f"[train] {steps_per_epoch} steps/epoch x {args.epochs} epochs " |
| f"= {total_steps} steps") |
|
|
| step = 0 |
| t_start = time.time() |
| for epoch in range(args.epochs): |
| perm = rng.permutation(N) |
| model.train() |
| running = 0.0 |
| for it in range(steps_per_epoch): |
| batch_idx = perm[(it * args.batch_size) % N: |
| (it * args.batch_size) % N + args.batch_size] |
| if len(batch_idx) < args.batch_size: |
| batch_idx = rng.integers(0, N, size=args.batch_size) |
|
|
| B = len(batch_idx) |
| if data_on_gpu: |
| idx_t = torch.as_tensor(batch_idx, device=device, dtype=torch.long) |
| crops = gpu_random_crops(images_dev, idx_t, args.crop, gen) |
| target_full = crops.reshape(B, -1, 3) |
| tokens = tokens_dev[idx_t] |
| else: |
| crops = random_crops(images, batch_idx, args.crop, rng) |
| target_full = torch.from_numpy(crops.reshape(B, -1, 3)).to(device) |
| tokens = torch.from_numpy(tokens_all[batch_idx]).to(device) |
|
|
| if 0 < args.pixels_per_step < HW: |
| pix = torch.randint(0, HW, (B, args.pixels_per_step), device=device, generator=gen) |
| bar = torch.arange(B, device=device)[:, None] |
| coords_b = coords[0][pix] |
| target = target_full[bar, pix] |
| else: |
| coords_b = coords.expand(B, -1, -1) |
| target = target_full |
|
|
| lr = cosine_lr(step, total_steps, args.lr, args.min_lr, args.warmup_steps) |
| for g in opt.param_groups: |
| g["lr"] = lr |
|
|
| opt.zero_grad(set_to_none=True) |
| with torch.amp.autocast("cuda", enabled=args.amp): |
| pred = model(tokens, coords_b) |
| loss = F.mse_loss(pred, target) |
| scaler.scale(loss).backward() |
| if args.grad_clip: |
| scaler.unscale_(opt) |
| torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) |
| scaler.step(opt) |
| scaler.update() |
|
|
| running += loss.item() |
|
|
| if step in (0, 1, 5, 20, 100) or (epoch == 0 and it % 200 == 0): |
| print_activation_stats(model, tokens, coords_b, tag=f"step{step}") |
|
|
| if step % args.log_interval == 0: |
| msg = f"[e{epoch:03d} s{step:06d}] loss={loss.item():.5f} lr={lr:.2e}" |
| if device == "cuda": |
| peak = torch.cuda.max_memory_allocated() / 1e9 |
| msg += f" vram_peak={peak:.2f}GB" |
| if peak > args.vram_budget_gb: |
| msg += f" !! over {args.vram_budget_gb}GB budget" |
| print(msg) |
| step += 1 |
|
|
| if (epoch + 1) % args.save_every_epochs == 0 or epoch == args.epochs - 1: |
| info = save_model_png(model, args.out_png, args.out_config, verbose=(epoch == 0)) |
| elapsed = (time.time() - t_start) / 60.0 |
| print(f"[e{epoch:03d}] avg_loss={running/steps_per_epoch:.5f} " |
| f"-> wrote {args.out_png} ({info['png_bytes']/1e6:.3f} MB) " |
| f"| {elapsed:.1f} min elapsed") |
|
|
| print(f"[train] done in {(time.time()-t_start)/60.0:.1f} min. " |
| f"Canonical model -> {args.out_png}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|