| """Collect a PushT dataset with stable_worldmodel's WeakPolicy. | |
| Writes shards under $STABLEWM_HOME/datasets/<name>/ in lance format. | |
| """ | |
| import argparse | |
| import os | |
| from pathlib import Path | |
| import numpy as np | |
| def main(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument('--name', default='pusht_weak_100') | |
| p.add_argument('--episodes', type=int, default=2000) | |
| p.add_argument('--shards', type=int, default=10) | |
| p.add_argument('--num-envs', type=int, default=10) | |
| p.add_argument('--max-episode-steps', type=int, default=100) | |
| p.add_argument('--dist-constraint', type=int, default=100) | |
| p.add_argument('--seed', type=int, default=3072) | |
| p.add_argument('--cache-dir', default=None) | |
| args = p.parse_args() | |
| import stable_worldmodel as swm | |
| from stable_worldmodel.envs.pusht import WeakPolicy | |
| world = swm.World( | |
| 'swm/PushT-v1', | |
| num_envs=args.num_envs, | |
| image_shape=(224, 224), | |
| max_episode_steps=args.max_episode_steps, | |
| render_mode='rgb_array', | |
| ) | |
| world.set_policy(WeakPolicy(dist_constraint=args.dist_constraint)) | |
| root = Path( | |
| args.cache_dir | |
| or os.getenv('STABLEWM_HOME') | |
| or swm.data.utils.get_cache_dir() | |
| ) | |
| out_dir = root / 'datasets' / args.name | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| per_shard = args.episodes // args.shards | |
| rng = np.random.default_rng(args.seed) | |
| for i in range(args.shards): | |
| shard = out_dir / f'shard_{i}.lance' | |
| if shard.exists(): | |
| print(f'[skip] {shard} exists') | |
| rng.integers(0, 1_000_000) | |
| continue | |
| print(f'[collect] shard {i + 1}/{args.shards} -> {shard}') | |
| world.collect( | |
| shard, | |
| episodes=per_shard, | |
| seed=rng.integers(0, 1_000_000).item(), | |
| ) | |
| print(f'done: {out_dir}') | |
| if __name__ == '__main__': | |
| main() | |