Download code/scripts/train_density.py from nottygian/Leplanner: direct link, hf CLI and curl.
- Browser
- Download file 4.11 kB
-
https://huggingface.co/nottygian/Leplanner/resolve/main/code/scripts/train_density.py
- Command line
-
hf download hf://nottygian/Leplanner/code/scripts/train_density.py
-
curl -L -o train_density.py https://huggingface.co/nottygian/Leplanner/resolve/main/code/scripts/train_density.py
4.11 kB
| """Fit the conditional behavior-density model beta(b | C) on real transitions. | |
| This is a support model for the controller's plans, not a policy. Training it | |
| separately keeps the controller objective free of behavior cloning: the | |
| controller is only penalized when a plan leaves the region the dataset covers, | |
| measured against the 95th percentile of held-out real transitions. | |
| """ | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import DataLoader | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) | |
| from lejepa_control.data import LatentGoalDataset, split_episodes # noqa: E402 | |
| from lejepa_control.losses import BehaviorDensity # noqa: E402 | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--latents', default='data/latents') | |
| parser.add_argument('--out', default='data/runs/density') | |
| parser.add_argument('--steps', type=int, default=4000) | |
| parser.add_argument('--batch-size', type=int, default=256) | |
| parser.add_argument('--lr', type=float, default=1e-3) | |
| parser.add_argument('--components', type=int, default=16) | |
| # in-process is fastest here: the latent cache is resident, so workers | |
| # would each copy ~1 GB on spawn to save no real work | |
| parser.add_argument('--workers', type=int, default=0) | |
| args = parser.parse_args() | |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' | |
| torch.manual_seed(0) | |
| stats = json.loads((Path(args.latents) / 'stats.json').read_text()) | |
| train_eps, val_eps = split_episodes(stats['n_episodes']) | |
| train_set = LatentGoalDataset(args.latents, episodes=train_eps) | |
| val_set = LatentGoalDataset(args.latents, episodes=val_eps) | |
| print(f'train clips {len(train_set)} val clips {len(val_set)}') | |
| loader = DataLoader( | |
| train_set, | |
| batch_size=args.batch_size, | |
| shuffle=True, | |
| num_workers=args.workers, | |
| drop_last=True, | |
| persistent_workers=args.workers > 0, | |
| ) | |
| density = BehaviorDensity( | |
| latent_dim=stats['latent_dim'], components=args.components | |
| ).to(device) | |
| opt = torch.optim.AdamW(density.parameters(), lr=args.lr, weight_decay=1e-4) | |
| step = 0 | |
| density.train() | |
| while step < args.steps: | |
| for batch in loader: | |
| ctx = batch['context'].to(device) | |
| block = batch['real_action'].to(device) | |
| loss = density.nll_per_dim(ctx, block).mean() | |
| opt.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(density.parameters(), 1.0) | |
| opt.step() | |
| step += 1 | |
| if step % 500 == 0: | |
| print(f'step {step:5d} nll/dim {loss.item():.4f}', flush=True) | |
| if step >= args.steps: | |
| break | |
| # --- calibrate c_95 on held-out real transitions ---------------------- | |
| density.eval() | |
| val_loader = DataLoader( | |
| val_set, batch_size=512, shuffle=True, num_workers=args.workers | |
| ) | |
| scores = [] | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| s = density.nll_per_dim( | |
| batch['context'].to(device), batch['real_action'].to(device) | |
| ) | |
| scores.append(s.cpu().numpy()) | |
| if sum(len(x) for x in scores) >= 200_000: | |
| break | |
| scores = np.concatenate(scores) | |
| c95 = float(np.percentile(scores, 95)) | |
| out_dir = Path(args.out) | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| torch.save( | |
| { | |
| 'state_dict': density.state_dict(), | |
| 'c95': c95, | |
| 'components': args.components, | |
| 'latent_dim': stats['latent_dim'], | |
| }, | |
| out_dir / 'density.pt', | |
| ) | |
| print( | |
| f'held-out nll/dim: mean {scores.mean():.4f} ' | |
| f'p50 {np.percentile(scores, 50):.4f} ' | |
| f'p95 {c95:.4f} p99 {np.percentile(scores, 99):.4f}' | |
| ) | |
| print(f'saved -> {out_dir / "density.pt"}') | |
| if __name__ == '__main__': | |
| main() | |