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 <noreply@anthropic.com>
This commit is contained in:
@@ -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,
|
||||
|
||||
+1
-1
@@ -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,
|
||||
|
||||
@@ -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"],
|
||||
)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
+34
-4
@@ -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}
|
||||
|
||||
+3
-1
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user