FlowMo-WM / control.py
cccat6's picture
Complete FlowMo reproducibility materials and record input provenance
eb4fc81 verified
Raw History Blame Contribute Delete
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()