File size: 1,369 Bytes
5b6e415
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
python scripts/dump_data.py \
    --dataset pi05_hang_cloth \
    --output hang_cloth_data.pkl

python scripts/dump_data.py \
    --dataset pi05_flat_fold_cloth \
    --output flat_fold_cloth_data.pkl
"""

import argparse
import pickle

from tqdm import tqdm
from openpi.training import config as _config
from openpi.training import data_loader as _data_loader
import openpi.training.sharding as sharding
import jax

def main():
    parser = argparse.ArgumentParser(description="dump test data for model mixture")
    parser.add_argument("--dataset", required=True, type=str, help="Config names used to dump test data from.")
    parser.add_argument("--output", required=True, help="Output directory for dumped data.")

    args = parser.parse_args()
    datasets = args.dataset
    config = _config.get_config(datasets)
    mesh = sharding.make_mesh(config.fsdp_devices)
    data_sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec(sharding.DATA_AXIS))

    data_loader = _data_loader.create_data_loader(
        config,
        sharding=data_sharding,
        shuffle=True,
    )
    samples_list = []
    for i, samples in enumerate(tqdm(data_loader)):
        if i >= 50:
            break
        samples_list.append(samples)
    with open(args.output, 'wb') as f:
        pickle.dump(samples_list, f)


if __name__ == "__main__":
    main()