Leplanner / code /scripts /collect_pusht.py
nottygian's picture
Push scripts
dc9f917 verified
Raw
History Blame Contribute Delete
1.95 kB
"""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()