Download generate_data.py from cccat6/FlowMo-WM: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/cccat6/FlowMo-WM/resolve/main/generate_data.py
- Command line
-
hf download hf://cccat6/FlowMo-WM/generate_data.py
-
curl -L -o generate_data.py https://huggingface.co/cccat6/FlowMo-WM/resolve/main/generate_data.py
2.27 kB
| """Generate the recorded simulation protocols in a separate output directory.""" | |
| from pathlib import Path | |
| import argparse | |
| import hashlib | |
| import json | |
| import os | |
| import subprocess | |
| import sys | |
| def main(): | |
| p = argparse.ArgumentParser(description=__doc__) | |
| p.add_argument('--package', type=Path, default=Path(__file__).resolve().parent) | |
| p.add_argument('--output', type=Path, required=True) | |
| p.add_argument('--groups', nargs='+', choices=['all', 'paper', 'rebuttal', 'frequency', 'switch'], default=['all']) | |
| p.add_argument('--dry-run', action='store_true') | |
| a = p.parse_args() | |
| root, out = a.package.resolve(), a.output.resolve() | |
| 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', 'protocols']): | |
| raise ValueError('Choose a separate output directory') | |
| jobs = json.loads((root / 'protocols/data_commands.json').read_text()) | |
| jobs = [j for j in jobs if 'all' in a.groups or j['group'] in a.groups] | |
| for job in jobs: | |
| cmd = [sys.executable, *job['command'][1:]] | |
| destination = out / job['path'] | |
| cmd[cmd.index('--out') + 1] = str(destination) | |
| job['command'] = cmd | |
| if a.dry_run: | |
| print(json.dumps(jobs, indent=2)); return | |
| if any((out / j['path']).exists() for j in jobs): | |
| raise FileExistsError('Generated inputs already exist; select a new output directory') | |
| out.mkdir(parents=True, exist_ok=True) | |
| (out / 'generation_plan.json').write_text(json.dumps(jobs, indent=2) + '\n') | |
| receipts = [] | |
| for job in jobs: | |
| env = os.environ.copy() | |
| env['PYTHONPATH'] = str(root / job['source_tree']) + os.pathsep + str(root) | |
| env['PYTHONDONTWRITEBYTECODE'] = '1' | |
| subprocess.run(job['command'], cwd=out, env=env, check=True) | |
| path = out / job['path'] | |
| receipts.append({'path': job['path'], 'sha256': hashlib.sha256(path.read_bytes()).hexdigest(), | |
| 'status': 'new_generation', 'group': job['group']}) | |
| (out / 'generated_inputs.json').write_text(json.dumps(receipts, indent=2) + '\n') | |
| print(json.dumps({'completed_datasets': len(receipts), 'output': str(out)})) | |
| if __name__ == '__main__': | |
| main() | |