AstroAUmin/self-forcing
0
1"""2python create_lmdb_14b_shards.py \3--data_path /mnt/localssd/wanx_14b_data \4--lmdb_path /mnt/localssd/wanx_14B_shift-3.0_cfg-5.0_lmdb5"""6from tqdm import tqdm7import numpy as np8import argparse9import torch10import lmdb11import glob12import os13 14from utils.lmdb import store_arrays_to_lmdb, process_data_dict15 16 17def main():18 """19 Aggregate all ode pairs inside a folder into a lmdb dataset.20 Each pt file should contain a (key, value) pair representing a21 video's ODE trajectories.22 """23 parser = argparse.ArgumentParser()24 parser.add_argument("--data_path", type=str,25 required=True, help="path to ode pairs")26 parser.add_argument("--lmdb_path", type=str,27 required=True, help="path to lmdb")28 parser.add_argument("--num_shards", type=int,29 default=16, help="num_shards")30 31 args = parser.parse_args()32 33 all_dirs = sorted(os.listdir(args.data_path))34 35 # figure out the maximum map size needed36 map_size = int(1e12) # adapt to your need, set to 1TB by default37 os.makedirs(args.lmdb_path, exist_ok=True)38 # 1) Open one LMDB env per shard39 envs = []40 num_shards = args.num_shards41 for shard_id in range(num_shards):42 print("shard_id ", shard_id)43 path = os.path.join(args.lmdb_path, f"shard_{shard_id}")44 env = lmdb.open(path,45 map_size=map_size,46 subdir=True, # set to True if you want a directory per env47 readonly=False,48 metasync=True,49 sync=True,50 lock=True,51 readahead=False,52 meminit=False)53 envs.append(env)54 55 counters = [0] * num_shards56 seen_prompts = set() # for deduplication57 total_samples = 058 all_files = []59 60 for part_dir in all_dirs:61 all_files += sorted(glob.glob(os.path.join(args.data_path, part_dir, "*.pt")))62 63 # 2) Prepare a write transaction for each shard64 for idx, file in tqdm(enumerate(all_files)):65 try:66 data_dict = torch.load(file)67 data_dict = process_data_dict(data_dict, seen_prompts)68 except Exception as e:69 print(f"Error processing {file}: {e}")70 continue71 72 if data_dict["latents"].shape != (1, 21, 16, 60, 104):73 continue74 75 shard_id = idx % num_shards76 # write to lmdb file77 store_arrays_to_lmdb(envs[shard_id], data_dict, start_index=counters[shard_id])78 counters[shard_id] += len(data_dict['prompts'])79 data_shape = data_dict["latents"].shape80 81 total_samples += len(all_files)82 83 print(len(seen_prompts))84 85 # save each entry's shape to lmdb86 for shard_id, env in enumerate(envs):87 with env.begin(write=True) as txn:88 for key, val in (data_dict.items()):89 assert len(data_shape) == 590 array_shape = np.array(data_shape) # val.shape)91 array_shape[0] = counters[shard_id]92 shape_key = f"{key}_shape".encode()93 print(shape_key, array_shape)94 shape_str = " ".join(map(str, array_shape))95 txn.put(shape_key, shape_str.encode())96 97 print(f"Finished writing {total_samples} examples into {num_shards} shards under {args.lmdb_path}")98 99 100if __name__ == "__main__":101 main()102 