Download control.py from cccat6/FlowMo-WM: direct link, hf CLI and curl.
- Browser
- Download file 5.87 kB
-
https://huggingface.co/cccat6/FlowMo-WM/resolve/main/control.py
- Command line
-
hf download hf://cccat6/FlowMo-WM/control.py
-
curl -L -o control.py https://huggingface.co/cccat6/FlowMo-WM/resolve/main/control.py
5.87 kB
| """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() | |