From 893d91e7494d35ff31a4a69c340cf92fb0ac5789 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Thu, 18 Jun 2026 13:45:12 +0200 Subject: [PATCH] Add KL divergence to marginal validation and hook it into the training loop validate_marginals now estimates a per-dimension KL(real || generated) via a shared histogram, alongside the existing mean/std comparison, so distribution-shape drift shows up even when the first two moments match. Wire it into giant/train.py: every validate_every epochs (default 10, 0 disables), the training loop runs validate_marginals against val_loader and prints the table. validate_every flows through DEFAULT_CONFIG/config.toml and is exposed as --validate-every on both giant train and scripts/train.py. Co-Authored-By: Claude Sonnet 4.6 --- giant/cli.py | 2 ++ giant/config.py | 2 +- giant/pipeline.py | 1 + giant/train.py | 6 ++++++ giant/validate.py | 38 ++++++++++++++++++++++++++++++++++---- scripts/train.py | 4 +++- 6 files changed, 47 insertions(+), 6 deletions(-) diff --git a/giant/cli.py b/giant/cli.py index ef0a3cc..4eaf6bd 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -60,6 +60,7 @@ def train( emb_dim: Annotated[Optional[int], typer.Option()] = None, val_fraction: Annotated[Optional[float], typer.Option()] = None, seed: Annotated[Optional[int], typer.Option(help="Random seed for reproducibility")] = None, + validate_every: Annotated[Optional[int], typer.Option(help="Run marginal+KL validation every N epochs (0 disables)")] = None, shuffle_buffer: Annotated[int, typer.Option(help="Rows held in RAM per worker for shuffling")] = 65536, out: Annotated[Optional[Path], typer.Option(help="Checkpoint dir (default: auto from hyperparams)")] = None, device: Annotated[Optional[str], typer.Option(help="cpu | cuda | mps (default: auto)")] = None, @@ -71,6 +72,7 @@ def train( "mode": mode.value if mode is not None else None, "epochs": epochs, "batch_size": batch_size, "lr": lr, "val_fraction": val_fraction, "num_workers": num_workers, "seed": seed, + "validate_every": validate_every, }.items() if v is not None} cli_model = {k: v for k, v in { "hidden_dim": hidden_dim, "n_blocks": n_blocks, "emb_dim": emb_dim, diff --git a/giant/config.py b/giant/config.py index 87cd3e9..2b413a3 100644 --- a/giant/config.py +++ b/giant/config.py @@ -11,7 +11,7 @@ import torch DEFAULT_CONFIG: dict = { "train": { "mode": "flow", "epochs": 100, "batch_size": 4096, "lr": 3e-4, - "val_fraction": 0.1, "num_workers": 4, "seed": 0, + "val_fraction": 0.1, "num_workers": 4, "seed": 0, "validate_every": 10, }, "model": { "hidden_dim": 256, "n_blocks": 6, "emb_dim": 16, diff --git a/giant/pipeline.py b/giant/pipeline.py index db8d5b9..1e5c8d2 100644 --- a/giant/pipeline.py +++ b/giant/pipeline.py @@ -126,4 +126,5 @@ def run_train_job( mat_map={str(k): v for k, v in mat_map.items()}, model_config=model_config, resume_path=resume, + validate_every=t["validate_every"], ) diff --git a/giant/train.py b/giant/train.py index 12d0b5b..3266cf8 100644 --- a/giant/train.py +++ b/giant/train.py @@ -7,6 +7,7 @@ import torch.optim as optim from torch.utils.data import DataLoader from giant.model.schedule import CosineSchedule, flow_matching_loss +from giant.validate import validate_marginals _METRICS_FIELDS = ["epoch", "train_loss", "val_loss", "lr", "epoch_time_s"] @@ -25,6 +26,7 @@ def train( mat_map: dict | None = None, model_config: dict | None = None, resume_path: str | Path | None = None, + validate_every: int = 0, ) -> None: out_dir = Path(out_dir) out_dir.mkdir(parents=True, exist_ok=True) @@ -105,6 +107,10 @@ def train( }) metrics_file.flush() + if validate_every > 0 and epoch % validate_every == 0: + print(f"[epoch {epoch}] marginal validation:") + validate_marginals(model, val_loader, mode=mode, schedule=ddpm_schedule, device=device) + ckpt: dict = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), diff --git a/giant/validate.py b/giant/validate.py index c9b64a3..d2ed619 100644 --- a/giant/validate.py +++ b/giant/validate.py @@ -6,6 +6,22 @@ from giant.constants import LOCAL_TARGET_NAMES from giant.sample import sample_flow, sample_ddpm, sample_ddim +def _histogram_kl(p_samples: np.ndarray, q_samples: np.ndarray, bins: int = 50, eps: float = 1e-8) -> float: + """KL(P || Q) between two 1D samples, estimated via a shared histogram.""" + lo = min(p_samples.min(), q_samples.min()) + hi = max(p_samples.max(), q_samples.max()) + if hi <= lo: + return 0.0 + edges = np.linspace(lo, hi, bins + 1) + p_hist, _ = np.histogram(p_samples, bins=edges) + q_hist, _ = np.histogram(q_samples, bins=edges) + p = p_hist.astype(np.float64) + eps + q = q_hist.astype(np.float64) + eps + p /= p.sum() + q /= q.sum() + return float(np.sum(p * np.log(p / q))) + + def validate_marginals( model: torch.nn.Module, val_loader: DataLoader, @@ -13,10 +29,13 @@ def validate_marginals( schedule=None, device: torch.device | None = None, n_batches: int | None = None, + kl_bins: int = 50, ) -> dict[str, np.ndarray]: """Compare per-dimension marginals of generated vs. real steps. - Returns {"real": (N,9), "generated": (N,9)} in normalised space. + Returns {"real": (N,9), "generated": (N,9), "kl_divergence": (9,)} in + normalised space. `kl_divergence[j]` is KL(real || generated) for + dimension j, estimated from a shared histogram over both samples. """ if device is None: device = next(model.parameters()).device @@ -42,11 +61,22 @@ def validate_marginals( real = np.concatenate(all_real, axis=0) generated = np.concatenate(all_gen, axis=0) - header = f"{'Dim':<20} {'real_mean':>10} {'gen_mean':>10} {'real_std':>10} {'gen_std':>10}" + kl_divergence = np.array([ + _histogram_kl(real[:, j], generated[:, j], bins=kl_bins) + for j in range(real.shape[1]) + ]) + + header = ( + f"{'Dim':<20} {'real_mean':>10} {'gen_mean':>10} " + f"{'real_std':>10} {'gen_std':>10} {'KL(real||gen)':>14}" + ) print(f"\n{header}") print("-" * len(header)) for j, name in enumerate(LOCAL_TARGET_NAMES): r, g = real[:, j], generated[:, j] - print(f"{name:<20} {r.mean():>10.4f} {g.mean():>10.4f} {r.std():>10.4f} {g.std():>10.4f}") + print( + f"{name:<20} {r.mean():>10.4f} {g.mean():>10.4f} " + f"{r.std():>10.4f} {g.std():>10.4f} {kl_divergence[j]:>14.4f}" + ) - return {"real": real, "generated": generated} + return {"real": real, "generated": generated, "kl_divergence": kl_divergence} diff --git a/scripts/train.py b/scripts/train.py index 87fc88a..9a5bae7 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -20,6 +20,8 @@ def main() -> None: parser.add_argument("--emb-dim", type=int) parser.add_argument("--val-fraction", type=float) parser.add_argument("--seed", type=int, help="Random seed for reproducibility") + parser.add_argument("--validate-every", type=int, + help="Run marginal+KL validation every N epochs (0 disables)") parser.add_argument("--shuffle-buffer", type=int, default=65536, help="Rows held in RAM for shuffling per worker (default: 65536)") parser.add_argument("--out", default=None, help="Checkpoint output directory (default: auto from hyperparams)") @@ -31,7 +33,7 @@ def main() -> None: cli_train = {k: v for k, v in { "mode": args.mode, "epochs": args.epochs, "batch_size": args.batch_size, "lr": args.lr, "val_fraction": args.val_fraction, "num_workers": args.num_workers, - "seed": args.seed, + "seed": args.seed, "validate_every": args.validate_every, }.items() if v is not None} cli_model = {k: v for k, v in { "hidden_dim": args.hidden_dim, "n_blocks": args.n_blocks, "emb_dim": args.emb_dim,