"""Collect a PushT dataset with stable_worldmodel's WeakPolicy. Writes shards under $STABLEWM_HOME/datasets// 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()