"""Run the original planning protocols with their fixed CEM settings.""" from pathlib import Path import argparse import json import os import subprocess import sys import time INITIAL = ['flowmo', 'leworldmodel', 'planet', 'tdmpc2', 'pid_los_controller', 'no_flow_los_controller', 'current_estimator_los_controller', 'oracle_flow_los_controller'] FORMAL = ['flowmo', 'history32', 'flowmo_additive', 'flowmo_film'] def main(): parser = argparse.ArgumentParser() parser.add_argument('--package', type=Path, default=Path(__file__).resolve().parent) parser.add_argument('--output', type=Path, required=True) parser.add_argument('--suite', choices=['initial', 'formal'], default='formal') parser.add_argument('--seeds', nargs='+', type=int, default=list(range(1, 9))) parser.add_argument('--methods', nargs='+', choices=sorted(set(INITIAL + FORMAL))) parser.add_argument('--device', default='cuda') parser.add_argument('--training-output', type=Path, help='Output directory produced by train.py') parser.add_argument('--smoke', action='store_true', help='One episode per selected setting, with the original planner budget') parser.add_argument('--dry-run', action='store_true') args = parser.parse_args() root, out = args.package.resolve(), args.output.resolve() work = args.training_output.resolve() if args.training_output else root if out == root or out in root.parents or any(out == root / d or root / d in out.parents for d in ['data', 'experiments', 'paper', 'assets', 'provenance']): raise ValueError('Choose a separate control output directory') if any(s not in range(1, 9) for s in args.seeds): raise ValueError('Formal seeds are 1 through 8') cfg = json.loads((root / 'experiments/shared/config' / ('paper_image.json' if args.suite == 'initial' else 'rebuttal_v1.json')).read_text()) plan = cfg['planning_eval'] conditions = [] if args.suite == 'initial': representative = {'reach_target': 'uniform', 'station_keeping': 'vortex_center', 'waypoint_square': 'gradient', 'waypoint_zigzag': 'random_fourier'} for task in plan['tasks']: for boat in plan['boats']: for flow in ([representative[task]] if args.smoke else cfg['flow_families']): conditions.append(dict(name=f'{task}_{boat}_{flow}', task=task, boat=boat, flow_type=flow, magnitude_scale=1.0, dynamic=False, max_steps=plan['task_max_steps'][task])) seeds = [None] methods = args.methods or INITIAL else: conditions = plan['conditions'] seeds = args.seeds methods = args.methods or FORMAL jobs = [] for seed in seeds: cp = 'paper.pt' if seed is None else f'rebuttal_core_seed_{seed:04d}.pt' for i, cond in enumerate(conditions): name = f'initial/{cond["name"]}' if seed is None else f'formal/seed_{seed:04d}/{cond["name"]}' destination = out / name cmd = [sys.executable, '-m', 'experiments.evaluate_image_planning', '--methods', *methods, '--task', cond['task'], '--boat', cond['boat'], '--flow-type', cond['flow_type'], '--flow-magnitude-scale', str(cond['magnitude_scale']), '--episodes', '1' if args.smoke else str(plan['episodes']), '--max-steps', str(cond['max_steps']), '--history-len', '32', '--image-size', '160', '--visual-scale', '2.5', '--checkpoint-name', cp, '--context-modes', 'inferred', '--success-radius', str(plan.get('success_radius', .65)), '--make-gifs', '0', '--seed', '33' if seed is None else str(7300 + i*100), '--device', args.device, '--precision', 'fp32', '--out', str(destination)] for key, value in plan.items(): if key.startswith('cem_'): cmd += ['--' + key.replace('_', '-'), str(value)] if args.suite == 'initial': for flag in ['--flow-magnitude-scale', '--context-modes']: index = cmd.index(flag) del cmd[index:index + 2] if cond['dynamic']: cmd += ['--flow-temporal-amplitude', '0.45', '--flow-temporal-frequency', '0.08', '--flow-rotation-amplitude', '0.65'] jobs.append(dict(name=name, methods=methods, checkpoint=cp, command=cmd, source_tree='protocols/initial/src' if args.suite == 'initial' else '.')) if args.dry_run: print(json.dumps({'smoke': args.smoke, 'jobs': jobs}, indent=2)) return out.mkdir(parents=True, exist_ok=True) env = os.environ.copy() env['PYTHONPATH'] = str(root) + os.pathsep + env.get('PYTHONPATH', '') env['OMP_NUM_THREADS'] = '1' env['PYTHONPYCACHEPREFIX'] = str(out / 'pycache') completed = [] for job in jobs: env['PYTHONPATH'] = str(root / job['source_tree']) + os.pathsep + str(root) dest = out / job['name'] if dest.is_dir() and any(dest.glob('*.json')): raise FileExistsError(f'Control outputs already exist: {dest}') dest.mkdir(parents=True, exist_ok=True) started = time.monotonic() with (dest / 'execution.log').open('w') as log: subprocess.run(job['command'], cwd=work, env=env, stdout=log, stderr=subprocess.STDOUT, check=True) assert list(dest.glob('*.json')) completed.append(dict(name=job['name'], elapsed_seconds=time.monotonic() - started)) (out / 'control_receipt.json').write_text(json.dumps({'smoke': args.smoke, 'suite': args.suite, 'completed': len(completed), 'planned': len(jobs), 'jobs': completed}, indent=2)) print(json.dumps(completed[-1]), flush=True) if __name__ == '__main__': main()