"""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 # noqa: E402 from lejepa_control.data import LatentGoalDataset, split_episodes # noqa: E402 from lejepa_control.losses import BehaviorDensity, refinement_loss, support_loss # noqa: E402 from lejepa_control.world_model import load_lewm # noqa: E402 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) # exp7: tokens live directly in the world model's own latent coordinates # (width must equal latent_dim) instead of an arbitrary learned width p.add_argument('--no-latent-proj', action='store_true') # exp7: one fused operator Phi (single projection + one transformer, # depth doubled to match total F+G layers) instead of the split # consequence/refine two-network pipeline 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) # Horizon-matched arrival: index the goal term by the offset the goal was # relabeled from, instead of always scoring at block H. Off by default so # the original objective's ablations stay reproducible. p.add_argument('--arrival-hold', action='store_true') p.add_argument('--hold-weight', type=float, default=0.5) # in-process: the latent cache is resident, so workers would each copy # ~1 GB on spawn to save no real work 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 # mean distance profile over the plan, split by goal offset 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) # current behavior: every run anchored at 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) # frozen, eval, requires_grad_(False) # tanh bounds: raw PushT actions live in [-1, 1], and the world model was # trained on z-scored actions, so the bound moves into normalized space. 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: # every refinement after the first, every horizon step 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()