File size: 1,954 Bytes
dc9f917 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 | """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()
|