| """Train the amortized iterative controller against the frozen LeWM.
|
|
|
| The controller never sees dataset actions as targets. It proposes a plan,
|
| the frozen predictor says what that plan would cause, and the controller is
|
| scored on how close the imagined outcome lands to the goal latent.
|
| """
|
|
|
| import argparse
|
| import json
|
| import sys
|
| import time
|
| 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.controller import IterativeController
|
| from lejepa_control.data import LatentGoalDataset, split_episodes
|
| from lejepa_control.losses import BehaviorDensity, refinement_loss, support_loss
|
| from lejepa_control.world_model import load_lewm
|
|
|
|
|
| def parse_args():
|
| p = argparse.ArgumentParser()
|
| p.add_argument('--latents', default='data/latents')
|
| p.add_argument('--density', default='data/runs/density/density.pt')
|
| p.add_argument('--out', default='data/runs/controller')
|
| p.add_argument('--steps', type=int, default=20000)
|
| p.add_argument('--batch-size', type=int, default=32)
|
| p.add_argument('--lr', type=float, default=1e-4)
|
| p.add_argument('--weight-decay', type=float, default=1e-4)
|
| p.add_argument('--horizon', type=int, default=5)
|
| p.add_argument('--refinements', type=int, default=3)
|
| p.add_argument('--width', type=int, default=256)
|
| p.add_argument('--depth', type=int, default=4)
|
| p.add_argument('--heads', type=int, default=8)
|
| p.add_argument('--dropout', type=float, default=0.1)
|
|
|
|
|
| p.add_argument('--no-latent-proj', action='store_true')
|
|
|
|
|
|
|
| p.add_argument('--fused', action='store_true')
|
| p.add_argument(
|
| '--train-seed', type=int, default=None,
|
| help='seed torch/numpy at startup; default None keeps unseeded '
|
| 'behavior (torch.manual_seed(0) below is unconditional and '
|
| 'unrelated — this only seeds data/init variation across reps)',
|
| )
|
| p.add_argument('--alpha', type=float, default=0.05)
|
| p.add_argument('--lambda-support', type=float, default=0.01)
|
|
|
|
|
|
|
| p.add_argument('--arrival-hold', action='store_true')
|
| p.add_argument('--hold-weight', type=float, default=0.5)
|
|
|
|
|
| p.add_argument('--workers', type=int, default=0)
|
| p.add_argument('--log-every', type=int, default=100)
|
| p.add_argument('--val-every', type=int, default=1000)
|
| p.add_argument(
|
| '--curriculum',
|
| default='0:2,0.25:3,0.5:5',
|
| help='fraction_of_training:max_goal_offset, comma separated',
|
| )
|
| p.add_argument('--wm-name', default='quentinll/lewm-pusht')
|
| return p.parse_args()
|
|
|
|
|
| def parse_curriculum(spec, total_steps):
|
| stages = []
|
| for part in spec.split(','):
|
| frac, offset = part.split(':')
|
| stages.append((int(float(frac) * total_steps), int(offset)))
|
| return sorted(stages)
|
|
|
|
|
| def current_offset(stages, step):
|
| offset = stages[0][1]
|
| for start, value in stages:
|
| if step >= start:
|
| offset = value
|
| return offset
|
|
|
|
|
| @torch.no_grad()
|
| def evaluate(controller, model, loader, device, max_batches=20):
|
| """Held-out terminal distance, per-refinement gain, and arrival profile.
|
|
|
| ``arrival`` is the distance at each sample's own goal offset q, which is
|
| what receding-horizon execution actually depends on; ``terminal`` is the
|
| distance at block H regardless of q. A controller that defers arrival
|
| scores well on terminal and badly on arrival.
|
| """
|
| controller.eval()
|
| terminal, first, arrival, batches = 0.0, 0.0, 0.0, 0
|
|
|
| profile = torch.zeros(controller.horizon + 1, controller.horizon, device=device)
|
| counts = torch.zeros(controller.horizon + 1, device=device)
|
| for batch in loader:
|
| q = batch['goal_offset'].to(device).clamp(1, controller.horizon)
|
| out = controller(
|
| model,
|
| batch['context'].to(device),
|
| batch['past_actions'].to(device),
|
| batch['goal'].to(device),
|
| )
|
| d = out['distances'][-1]
|
| terminal += d[:, -1].mean().item()
|
| first += out['distances'][0][:, -1].mean().item()
|
| arrival += d.gather(1, (q - 1).unsqueeze(1)).squeeze(1).mean().item()
|
| profile.index_add_(0, q, d)
|
| counts.index_add_(0, q, torch.ones_like(q, dtype=d.dtype))
|
| batches += 1
|
| if batches >= max_batches:
|
| break
|
| controller.train()
|
| profile = (profile / counts.clamp(min=1).unsqueeze(1)).cpu()
|
| return {
|
| 'terminal': terminal / batches,
|
| 'first': first / batches,
|
| 'arrival': arrival / batches,
|
| 'profile': {
|
| q: [round(v, 4) for v in profile[q].tolist()]
|
| for q in range(1, controller.horizon + 1)
|
| if counts[q] > 0
|
| },
|
| }
|
|
|
|
|
| def main():
|
| args = parse_args()
|
| device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
| if args.train_seed is None:
|
| torch.manual_seed(0)
|
| else:
|
| torch.manual_seed(args.train_seed)
|
| np.random.seed(args.train_seed)
|
|
|
| out_dir = Path(args.out)
|
| out_dir.mkdir(parents=True, exist_ok=True)
|
|
|
| stats = json.loads((Path(args.latents) / 'stats.json').read_text())
|
| latent_dim = stats['latent_dim']
|
|
|
| model = load_lewm(name=args.wm_name, device=device)
|
|
|
|
|
|
|
| a_std = torch.tensor(stats['action_std'])
|
| a_mean = torch.tensor(stats['action_mean'])
|
| controller = IterativeController(
|
| latent_dim=latent_dim,
|
| horizon=args.horizon,
|
| refinements=args.refinements,
|
| width=args.width,
|
| depth=args.depth,
|
| heads=args.heads,
|
| dropout=args.dropout,
|
| action_center=(-a_mean / a_std),
|
| action_scale=(1.0 / a_std),
|
| no_latent_proj=args.no_latent_proj,
|
| fused=args.fused,
|
| ).to(device)
|
| n_params = sum(p.numel() for p in controller.parameters())
|
| print(f'controller params: {n_params / 1e6:.2f}M')
|
|
|
| density, c95 = None, None
|
| if args.lambda_support > 0 and Path(args.density).exists():
|
| ckpt = torch.load(args.density, map_location=device)
|
| density = BehaviorDensity(
|
| latent_dim=latent_dim, components=ckpt['components']
|
| ).to(device)
|
| density.load_state_dict(ckpt['state_dict'])
|
| density.eval().requires_grad_(False)
|
| c95 = ckpt['c95']
|
| print(f'support model loaded, c95={c95:.4f}')
|
| else:
|
| print('no support model — training with goal loss only')
|
|
|
| train_eps, val_eps = split_episodes(stats['n_episodes'])
|
| train_set = LatentGoalDataset(
|
| args.latents, episodes=train_eps, horizon=args.horizon
|
| )
|
| val_set = LatentGoalDataset(
|
| args.latents, max_offset=5, episodes=val_eps, horizon=args.horizon
|
| )
|
| loader = DataLoader(
|
| train_set,
|
| batch_size=args.batch_size,
|
| shuffle=True,
|
| num_workers=args.workers,
|
| drop_last=True,
|
| persistent_workers=args.workers > 0,
|
| pin_memory=True,
|
| )
|
| val_loader = DataLoader(
|
| val_set, batch_size=args.batch_size, shuffle=True
|
| )
|
|
|
| opt = torch.optim.AdamW(
|
| controller.parameters(), lr=args.lr, weight_decay=args.weight_decay
|
| )
|
| sched = torch.optim.lr_scheduler.OneCycleLR(
|
| opt, max_lr=args.lr, total_steps=args.steps, pct_start=0.05
|
| )
|
| stages = parse_curriculum(args.curriculum, args.steps)
|
| print(f'curriculum: {stages}')
|
|
|
| step, t0 = 0, time.perf_counter()
|
| running = {}
|
| controller.train()
|
| while step < args.steps:
|
| for batch in loader:
|
| offset = current_offset(stages, step)
|
| if train_set.max_offset != offset:
|
| train_set.set_max_offset(offset)
|
|
|
| ctx = batch['context'].to(device, non_blocking=True)
|
| past = batch['past_actions'].to(device, non_blocking=True)
|
| goal = batch['goal'].to(device, non_blocking=True)
|
| q = batch['goal_offset'].to(device, non_blocking=True)
|
|
|
| out = controller(model, ctx, past, goal)
|
| if args.arrival_hold:
|
| loss_refine = refinement_loss(
|
| out['distances'],
|
| goal_offset=q,
|
| hold_weight=args.hold_weight,
|
| )
|
| else:
|
| loss_refine = refinement_loss(out['distances'], alpha=args.alpha)
|
|
|
| loss = loss_refine
|
| violation = torch.zeros((), device=device)
|
| if density is not None:
|
|
|
| contexts = torch.cat(out['contexts'][1:], dim=0).flatten(0, 1)
|
| blocks = torch.cat(out['plans'][1:], dim=0).flatten(0, 1)
|
| loss_support, violation = support_loss(
|
| density, contexts, blocks, c95
|
| )
|
| loss = loss + args.lambda_support * loss_support
|
| running['support'] = (
|
| running.get('support', 0.0) + loss_support.item()
|
| )
|
|
|
| opt.zero_grad(set_to_none=True)
|
| loss.backward()
|
| grad = torch.nn.utils.clip_grad_norm_(controller.parameters(), 1.0)
|
| opt.step()
|
| sched.step()
|
|
|
| d_first = out['distances'][0][:, -1].mean().item()
|
| d_last = out['distances'][-1][:, -1].mean().item()
|
| d_arrival = (
|
| out['distances'][-1]
|
| .gather(1, (q.clamp(1, args.horizon) - 1).unsqueeze(1))
|
| .squeeze(1)
|
| .mean()
|
| .item()
|
| )
|
| running['loss'] = running.get('loss', 0.0) + loss.item()
|
| running['terminal'] = running.get('terminal', 0.0) + d_last
|
| running['arrival'] = running.get('arrival', 0.0) + d_arrival
|
| running['gain'] = running.get('gain', 0.0) + (d_first - d_last)
|
| running['violation'] = (
|
| running.get('violation', 0.0) + violation.item()
|
| )
|
| running['grad'] = running.get('grad', 0.0) + grad.item()
|
|
|
| step += 1
|
| if step % args.log_every == 0:
|
| n = args.log_every
|
| rate = step / (time.perf_counter() - t0)
|
| msg = (
|
| f'step {step:6d} H_goal<={offset} '
|
| f'loss {running["loss"] / n:.4f} '
|
| f'terminal {running["terminal"] / n:.4f} '
|
| f'arrival {running["arrival"] / n:.4f} '
|
| f'gain {running["gain"] / n:+.4f} '
|
| f'viol {running["violation"] / n:.3f} '
|
| f'grad {running["grad"] / n:.2f} '
|
| f'{rate:.1f} it/s'
|
| )
|
| if 'support' in running:
|
| msg += f' support {running["support"] / n:.4f}'
|
| print(msg, flush=True)
|
| running = {}
|
|
|
| if step % args.val_every == 0 or step == args.steps:
|
| val = evaluate(controller, model, val_loader, device)
|
| print(
|
| f' [val] terminal K={args.refinements}: '
|
| f'{val["terminal"]:.4f} K=0: {val["first"]:.4f} '
|
| f'gain {val["first"] - val["terminal"]:+.4f} '
|
| f'arrival {val["arrival"]:.4f} '
|
| f'steps {controller.step_sizes_repr()}',
|
| flush=True,
|
| )
|
| for q_val, prof in val['profile'].items():
|
| print(f' q={q_val}: {prof}', flush=True)
|
| torch.save(
|
| {
|
| 'state_dict': controller.state_dict(),
|
| 'args': vars(args),
|
| 'step': step,
|
| 'val_terminal': val['terminal'],
|
| 'val_arrival': val['arrival'],
|
| 'val_profile': val['profile'],
|
| 'action_mean': stats['action_mean'],
|
| 'action_std': stats['action_std'],
|
| },
|
| out_dir / 'controller.pt',
|
| )
|
|
|
| if step >= args.steps:
|
| break
|
|
|
| print(f'done in {(time.perf_counter() - t0) / 60:.1f} min -> {out_dir}')
|
|
|
|
|
| if __name__ == '__main__':
|
| main()
|
|
|