| """Train the recursive latent planner against the frozen LeWM. |
| |
| Staged bring-up, per ``recursive_planner_design.pdf`` and the task brief's |
| section 5. Each stage turns on one more mechanism and has a gate that must |
| clear before the next one starts:: |
| |
| A arrival+hold + late path, H=3, T=1, no feedback |
| gate: loss falls, real-env success beats random |
| B + consequence feedback (2) + deep supervision (3), T=3 |
| gate: per-cycle distances strictly decreasing <-- the important one |
| C + manifold anchor (7) + saturation barrier (6) |
| gate: anchor ratio < 0.05 stable, |r| > 2 fraction < 10% |
| D + support hinge (8), H -> 5 |
| gate: violation fraction < 5% and flat, real-env success still rising |
| E warm-start verification (cold-start control run) |
| gate: equal-or-better success, per-cycle curve flattens earlier |
| |
| Stage presets only fill arguments left unset, so any flag given explicitly on |
| the command line wins. |
| |
| Every run writes ``config.json``, ``metrics.jsonl`` (one record per log |
| interval, with the full diagnostic set from section 11) and a resumable |
| ``planner.pt``. |
| """ |
|
|
| import argparse |
| import json |
| import random |
| 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[2])) |
|
|
| from lejepa_control.data import LatentGoalDataset, split_episodes |
| from lejepa_control.losses import BehaviorDensity |
| from lejepa_control.world_model import load_lewm |
| from lejepa_control_2.losses import DistanceScale, planner_loss |
| from lejepa_control_2.planner import RecursivePlanner |
|
|
| |
| |
| STAGES = { |
| 'A': dict( |
| cycles=1, |
| horizon_curriculum='0:3', |
| use_feedback=False, |
| lambda_cycle=0.0, |
| lambda_anchor=0.0, |
| lambda_sat=0.0, |
| lambda_support=0.0, |
| ), |
| 'B': dict( |
| cycles=3, |
| horizon_curriculum='0:3', |
| use_feedback=True, |
| lambda_cycle=0.3, |
| lambda_anchor=0.0, |
| lambda_sat=0.0, |
| lambda_support=0.0, |
| ), |
| 'C': dict( |
| cycles=3, |
| horizon_curriculum='0:3', |
| use_feedback=True, |
| lambda_cycle=0.3, |
| lambda_anchor=0.05, |
| lambda_sat=1e-3, |
| lambda_support=0.0, |
| ), |
| 'D': dict( |
| cycles=3, |
| horizon_curriculum='0:3,0.5:5', |
| use_feedback=True, |
| lambda_cycle=0.3, |
| lambda_anchor=0.05, |
| lambda_sat=1e-3, |
| lambda_support=0.01, |
| ), |
| |
| |
| 'E': dict( |
| cycles=3, |
| horizon_curriculum='0:3,0.5:5', |
| use_feedback=True, |
| lambda_cycle=0.3, |
| lambda_anchor=0.05, |
| lambda_sat=1e-3, |
| lambda_support=0.01, |
| warm_start=False, |
| ), |
| } |
|
|
|
|
| def parse_args(argv=None): |
| 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/planner') |
| p.add_argument('--stage', choices=sorted(STAGES), default=None) |
| p.add_argument('--resume', action='store_true') |
| p.add_argument('--seed', type=int, default=0) |
|
|
| |
| p.add_argument('--width', type=int, default=256) |
| p.add_argument('--hidden', type=int, default=512) |
|
|
| |
| p.add_argument('--inner', type=int, default=6) |
| p.add_argument('--cycles', type=int, default=None) |
| p.add_argument( |
| '--horizon-curriculum', |
| default=None, |
| help='fraction_of_training:horizon, comma separated (base: 0:3,0.5:5)', |
| ) |
| 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('--hold-weight', type=float, default=0.5) |
| p.add_argument('--alpha', type=float, default=0.05) |
| p.add_argument('--lambda-cycle', type=float, default=None) |
| p.add_argument('--lambda-support', type=float, default=None) |
| p.add_argument('--lambda-anchor', type=float, default=None) |
| p.add_argument('--lambda-sat', type=float, default=None) |
| p.add_argument('--sat-limit', type=float, default=2.0) |
|
|
| |
| 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('--grad-clip', type=float, default=1.0) |
| p.add_argument('--workers', type=int, default=0) |
|
|
| |
| p.add_argument('--pretrain-steps', type=int, default=400) |
| p.add_argument('--pretrain-lr', type=float, default=1e-3) |
|
|
| |
| p.add_argument('--no-warm-start', dest='warm_start', action='store_false') |
| p.add_argument('--lambda-z', type=float, default=0.0) |
| p.add_argument('--learn-lambda-z', action='store_true') |
| p.set_defaults(warm_start=None) |
|
|
| |
| |
| p.add_argument('--no-feedback', dest='use_feedback', action='store_false') |
| p.set_defaults(use_feedback=None) |
| |
| |
| p.add_argument( |
| '--detach-schedule', |
| choices=['last-cycle', 'one-step', 'full'], |
| default='last-cycle', |
| ) |
| |
| p.add_argument('--terminal-only', action='store_true') |
| |
| p.add_argument( |
| '--path-weighting', choices=['late', 'discount'], default='late' |
| ) |
| p.add_argument('--gamma', type=float, default=0.9) |
| |
|
|
| |
| p.add_argument('--log-every', type=int, default=100) |
| p.add_argument('--val-every', type=int, default=1000) |
| p.add_argument('--val-batches', type=int, default=20) |
| p.add_argument('--in-memory', action='store_true', default=True) |
| p.add_argument('--mmap', dest='in_memory', action='store_false') |
|
|
| args = p.parse_args(argv) |
| return apply_stage(args) |
|
|
|
|
| def apply_stage(args): |
| """Fill unset arguments from the stage preset; explicit flags win.""" |
| preset = STAGES.get(args.stage, {}) |
| fallback = dict( |
| cycles=3, |
| horizon_curriculum='0:3,0.5:5', |
| use_feedback=True, |
| warm_start=True, |
| lambda_cycle=0.3, |
| lambda_support=0.01, |
| lambda_anchor=0.05, |
| lambda_sat=1e-3, |
| ) |
| for key, default in fallback.items(): |
| if getattr(args, key) is None: |
| setattr(args, key, preset.get(key, default)) |
| return args |
|
|
|
|
| def parse_curriculum(spec, total_steps): |
| stages = [] |
| for part in spec.split(','): |
| frac, value = part.split(':') |
| stages.append((int(float(frac) * total_steps), int(value))) |
| return sorted(stages) |
|
|
|
|
| def current_value(stages, step): |
| value = stages[0][1] |
| for start, v in stages: |
| if step >= start: |
| value = v |
| return value |
|
|
|
|
| def seed_everything(seed): |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| torch.cuda.manual_seed_all(seed) |
|
|
|
|
| def build_planner(args, stats, latent_dim, device): |
| a_mean = torch.tensor(stats['action_mean']) |
| a_std = torch.tensor(stats['action_std']) |
| return RecursivePlanner( |
| latent_dim=latent_dim, |
| width=args.width, |
| hidden=args.hidden, |
| inner=args.inner, |
| cycles=args.cycles, |
| horizon=max(v for _, v in parse_curriculum( |
| args.horizon_curriculum, args.steps |
| )), |
| use_feedback=args.use_feedback, |
| warm_start=args.warm_start, |
| lambda_z=args.lambda_z, |
| learn_lambda_z=args.learn_lambda_z, |
| detach_schedule=args.detach_schedule, |
| |
| |
| |
| action_center=(-a_mean / a_std), |
| action_scale=(1.0 / a_std), |
| ).to(device) |
|
|
|
|
| def pretrain_action_embedding(planner, loader, args, device, log=print): |
| """Change 7's prerequisite: fit ``phi``/``psi`` as a plain autoencoder. |
| |
| The anchor pulls ``y`` toward whatever ``phi(psi(y))`` happens to be. If |
| the pair starts random that is a meaningless target and the anchor spends |
| the early run fighting the goal loss over where the manifold even is. A |
| few hundred steps on real dataset blocks settles it first. |
| """ |
| if args.pretrain_steps <= 0: |
| return None |
| embed = planner.action_embed |
| opt = torch.optim.AdamW(embed.parameters(), lr=args.pretrain_lr) |
| losses, step = [], 0 |
| while step < args.pretrain_steps: |
| for batch in loader: |
| |
| |
| blocks = torch.cat( |
| [ |
| batch['past_actions'].flatten(0, 1), |
| batch['real_action'], |
| ], |
| dim=0, |
| ).to(device) |
| recon = embed.decode(embed.encode(blocks)) |
| loss = (recon - blocks).pow(2).mean() |
| opt.zero_grad(set_to_none=True) |
| loss.backward() |
| opt.step() |
| losses.append(loss.item()) |
| step += 1 |
| if step >= args.pretrain_steps: |
| break |
| log( |
| f'phi/psi pretrain: {losses[0]:.4f} -> ' |
| f'{sum(losses[-20:]) / min(20, len(losses)):.4f} ' |
| f'over {step} steps' |
| ) |
| return {'first': losses[0], 'last': sum(losses[-20:]) / min(20, len(losses))} |
|
|
|
|
| @torch.no_grad() |
| def evaluate(planner, model, loader, scale, args, device, horizon): |
| """Held-out imagined metrics. Real-env success is a separate script. |
| |
| ``arrival`` is the distance at each sample's own relabel offset ``q``, |
| which is what receding-horizon execution depends on. ``terminal`` is the |
| distance at step ``H`` regardless of ``q`` — a planner that defers arrival |
| scores well on terminal and badly on arrival, and the gap between them is |
| the Change-4 diagnostic. |
| """ |
| planner.eval() |
| total = {'terminal': 0.0, 'arrival': 0.0} |
| per_step = torch.zeros(horizon, device=device) |
| per_cycle = None |
| batches = 0 |
| for batch in loader: |
| q = batch['goal_offset'].to(device).clamp(1, horizon) |
| out = planner( |
| model, |
| batch['context'].to(device), |
| batch['past_actions'].to(device), |
| batch['goal'].to(device), |
| horizon=horizon, |
| ) |
| d = scale.normalize(out['distances']) |
| total['terminal'] += d[:, -1].mean().item() |
| total['arrival'] += ( |
| d.gather(1, (q - 1).unsqueeze(1)).squeeze(1).mean().item() |
| ) |
| per_step += d.mean(dim=0) |
| if out['cycle_distances'] is not None: |
| c = scale.normalize(out['cycle_distances']).mean(dim=(0, 1)) |
| per_cycle = c if per_cycle is None else per_cycle + c |
| batches += 1 |
| if batches >= args.val_batches: |
| break |
| planner.train() |
| result = {k: v / batches for k, v in total.items()} |
| result['per_step'] = (per_step / batches).tolist() |
| if per_cycle is not None: |
| result['per_cycle'] = (per_cycle / batches).tolist() |
| return result |
|
|
|
|
| def main(argv=None): |
| args = parse_args(argv) |
| device = 'cuda' if torch.cuda.is_available() else 'cpu' |
| seed_everything(args.seed) |
|
|
| out_dir = Path(args.out) |
| out_dir.mkdir(parents=True, exist_ok=True) |
| ckpt_path = out_dir / 'planner.pt' |
| log_path = out_dir / 'metrics.jsonl' |
|
|
| stats = json.loads((Path(args.latents) / 'stats.json').read_text()) |
| latent_dim = stats['latent_dim'] |
|
|
| model = load_lewm(device=device) |
| planner = build_planner(args, stats, latent_dim, device) |
| scale = DistanceScale().to(device) |
|
|
| horizon_stages = parse_curriculum(args.horizon_curriculum, args.steps) |
| offset_stages = parse_curriculum(args.curriculum, args.steps) |
| max_horizon = max(v for _, v in horizon_stages) |
|
|
| n_params = sum(p.numel() for p in planner.parameters()) |
| print( |
| f'stage {args.stage or "custom"} planner {n_params / 1e6:.2f}M params ' |
| f'n={args.inner} T={args.cycles} H={horizon_stages} ' |
| f'feedback={args.use_feedback} warm_start={args.warm_start}' |
| ) |
|
|
| |
| |
| |
| density, c95 = None, None |
| if Path(args.density).exists(): |
| d_ckpt = torch.load(args.density, map_location=device) |
| density = BehaviorDensity( |
| latent_dim=latent_dim, components=d_ckpt['components'] |
| ).to(device) |
| density.load_state_dict(d_ckpt['state_dict']) |
| density.eval().requires_grad_(False) |
| c95 = d_ckpt['c95'] |
| print( |
| f'support model loaded, c95={c95:.4f} ' |
| f'(lambda_support={args.lambda_support})' |
| ) |
| else: |
| print(f'no support model at {args.density} — support term disabled') |
|
|
| train_eps, val_eps = split_episodes(stats['n_episodes']) |
| train_set = LatentGoalDataset( |
| args.latents, |
| episodes=train_eps, |
| horizon=max_horizon, |
| in_memory=args.in_memory, |
| ) |
| val_set = LatentGoalDataset( |
| args.latents, |
| max_offset=5, |
| episodes=val_eps, |
| horizon=max_horizon, |
| in_memory=args.in_memory, |
| ) |
| gen = torch.Generator().manual_seed(args.seed) |
| 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, |
| generator=gen, |
| ) |
| val_loader = DataLoader( |
| val_set, |
| batch_size=args.batch_size, |
| shuffle=True, |
| generator=torch.Generator().manual_seed(args.seed + 1), |
| ) |
|
|
| opt = torch.optim.AdamW( |
| planner.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 |
| ) |
|
|
| step = 0 |
| pretrain = None |
| if args.resume and ckpt_path.exists(): |
| ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) |
| saved = ckpt['args'] |
| |
| |
| |
| for key in ('steps', 'curriculum', 'horizon_curriculum', 'lr'): |
| if saved[key] != getattr(args, key): |
| raise SystemExit( |
| f'cannot resume: --{key.replace("_", "-")} was ' |
| f'{saved[key]!r} in the checkpoint, now ' |
| f'{getattr(args, key)!r}. Start a new run instead.' |
| ) |
| planner.load_state_dict(ckpt['state_dict']) |
| scale.load_state_dict(ckpt['scale']) |
| opt.load_state_dict(ckpt['optimizer']) |
| sched.load_state_dict(ckpt['scheduler']) |
| step = ckpt['step'] |
| torch.set_rng_state(ckpt['rng_torch'].cpu().to(torch.uint8)) |
| np.random.set_state(ckpt['rng_numpy']) |
| random.setstate(ckpt['rng_python']) |
| if ckpt.get('rng_cuda') is not None and torch.cuda.is_available(): |
| torch.cuda.set_rng_state_all( |
| [s.cpu().to(torch.uint8) for s in ckpt['rng_cuda']] |
| ) |
| print(f'resumed from {ckpt_path} at step {step}') |
| else: |
| pretrain = pretrain_action_embedding(planner, loader, args, device) |
| (out_dir / 'config.json').write_text(json.dumps(vars(args), indent=2)) |
| log_path.write_text('') |
|
|
| def write_record(record): |
| with log_path.open('a') as fh: |
| fh.write(json.dumps(record) + '\n') |
|
|
| def save(extra=None): |
| torch.save( |
| { |
| 'state_dict': planner.state_dict(), |
| 'scale': scale.state_dict(), |
| 'optimizer': opt.state_dict(), |
| 'scheduler': sched.state_dict(), |
| 'step': step, |
| 'args': vars(args), |
| 'action_mean': stats['action_mean'], |
| 'action_std': stats['action_std'], |
| 'latent_dim': latent_dim, |
| 'horizon': planner.horizon, |
| 'rng_torch': torch.get_rng_state(), |
| 'rng_numpy': np.random.get_state(), |
| 'rng_python': random.getstate(), |
| 'rng_cuda': ( |
| torch.cuda.get_rng_state_all() |
| if torch.cuda.is_available() |
| else None |
| ), |
| **(extra or {}), |
| }, |
| ckpt_path, |
| ) |
|
|
| running, t0 = {}, time.perf_counter() |
| planner.train() |
| while step < args.steps: |
| for batch in loader: |
| horizon = current_value(horizon_stages, step) |
| |
| |
| |
| offset = min(current_value(offset_stages, step), horizon) |
| 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) |
|
|
| capture = (step + 1) % args.log_every == 0 |
| if capture: |
| planner.start_capture() |
|
|
| scale.update(ctx, goal) |
| out = planner(model, ctx, past, goal, horizon=horizon) |
| loss, metrics = planner_loss( |
| out, |
| q, |
| scale, |
| planner.action_embed, |
| density=density, |
| c95=c95, |
| hold_weight=args.hold_weight, |
| alpha=args.alpha, |
| lambda_cycle=args.lambda_cycle, |
| lambda_support=args.lambda_support, |
| lambda_anchor=args.lambda_anchor, |
| lambda_sat=args.lambda_sat, |
| sat_limit=args.sat_limit, |
| terminal_only=args.terminal_only, |
| path_weighting=args.path_weighting, |
| gamma=args.gamma, |
| ) |
|
|
| opt.zero_grad(set_to_none=True) |
| loss.backward() |
| grad = torch.nn.utils.clip_grad_norm_( |
| planner.parameters(), args.grad_clip |
| ) |
| opt.step() |
| sched.step() |
| step += 1 |
|
|
| for key in ('loss', 'arrival', 'path', 'cycle', 'support', |
| 'violation', 'anchor', 'anchor_ratio', 'sat', |
| 'sat_fraction'): |
| if key in metrics: |
| running[key] = running.get(key, 0.0) + metrics[key].item() |
| running['grad'] = running.get('grad', 0.0) + grad.item() |
| running['n'] = running.get('n', 0) + 1 |
|
|
| if capture: |
| planner.stop_capture() |
| taps = planner.taps |
| n = running.pop('n') |
| rate = step / (time.perf_counter() - t0) |
| record = { |
| 'step': step, |
| 'horizon': horizon, |
| 'max_offset': offset, |
| 'lr': sched.get_last_lr()[0], |
| 'it_per_s': round(rate, 3), |
| 'scale': scale.scale.item(), |
| **{k: v / n for k, v in running.items()}, |
| |
| 'per_step': [round(v, 5) for v in |
| metrics['per_step'].tolist()], |
| 'gates': planner.gate_values(), |
| |
| 'z_norms': [round(v, 4) for v in taps['z_norms'][:32]], |
| |
| |
| 'grad_first': taps['grad_first'][:1], |
| 'grad_last': taps['grad_last'][:1], |
| } |
| if 'per_cycle' in metrics: |
| record['per_cycle'] = [ |
| round(v, 5) for v in metrics['per_cycle'].tolist() |
| ] |
| write_record(record) |
|
|
| msg = ( |
| f'step {step:6d} H={horizon} q<={offset} ' |
| f'loss {record["loss"]:.4f} ' |
| f'arrival {record["arrival"]:.4f} ' |
| f'grad {record["grad"]:.2f} {rate:.2f} it/s' |
| ) |
| if 'per_cycle' in record: |
| msg += f' cycles {record["per_cycle"]}' |
| if 'violation' in record: |
| msg += f' viol {record["violation"]:.3f}' |
| if 'anchor_ratio' in record: |
| msg += f' anchor {record["anchor_ratio"]:.4f}' |
| if 'sat_fraction' in record: |
| msg += f' sat {record["sat_fraction"]:.3f}' |
| print(msg, flush=True) |
| running = {} |
|
|
| if step % args.val_every == 0 or step >= args.steps: |
| val = evaluate( |
| planner, model, val_loader, scale, args, device, horizon |
| ) |
| print( |
| f' [val] arrival {val["arrival"]:.4f} ' |
| f'terminal {val["terminal"]:.4f} ' |
| f'per_step {[round(v, 4) for v in val["per_step"]]}' |
| + ( |
| f' per_cycle ' |
| f'{[round(v, 4) for v in val["per_cycle"]]}' |
| if 'per_cycle' in val |
| else '' |
| ), |
| flush=True, |
| ) |
| write_record({'step': step, 'split': 'val', **val}) |
| save({'val': val, 'pretrain': pretrain}) |
|
|
| if step >= args.steps: |
| break |
|
|
| save({'pretrain': pretrain}) |
| print( |
| f'done in {(time.perf_counter() - t0) / 60:.1f} min -> {out_dir}', |
| flush=True, |
| ) |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|