diff options
Diffstat (limited to 'worldalign/prepare.py')
| -rw-r--r-- | worldalign/prepare.py | 80 |
1 files changed, 80 insertions, 0 deletions
diff --git a/worldalign/prepare.py b/worldalign/prepare.py new file mode 100644 index 0000000..dbdd101 --- /dev/null +++ b/worldalign/prepare.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import argparse +from collections import Counter + +import numpy as np +from datasets import load_dataset + +from .common import DATASET_NAME, DATASET_SPLIT, write_json + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--output", default="artifacts/manifest.json") + p.add_argument("--dataset", default=DATASET_NAME) + p.add_argument("--seed", type=int, default=20260728) + p.add_argument("--unpaired-per-modality", type=int, default=12_000) + p.add_argument("--paired-train", type=int, default=12_000) + p.add_argument("--max-eval", type=int, default=1_000) + return p.parse_args() + + +def main() -> None: + args = parse_args() + dataset = load_dataset(args.dataset, split=DATASET_SPLIT) + by_split: dict[str, list[int]] = {} + for i, split in enumerate(dataset["split"]): + by_split.setdefault(split, []).append(i) + + counts = Counter(dataset["split"]) + print(f"Internal Flickr splits: {dict(counts)}") + train = np.asarray(by_split["train"], dtype=np.int64) + rng = np.random.default_rng(args.seed) + rng.shuffle(train) + + n_unpaired = min(args.unpaired_per_modality, len(train) // 2) + vision_only = train[:n_unpaired] + text_only = train[n_unpaired : 2 * n_unpaired] + assert not set(vision_only.tolist()) & set(text_only.tolist()) + + paired_n = min(args.paired_train, len(train)) + paired_train = train[:paired_n] + + val_key = "val" if "val" in by_split else "validation" + val = np.asarray(by_split[val_key], dtype=np.int64)[: args.max_eval] + test = np.asarray(by_split["test"], dtype=np.int64)[: args.max_eval] + + all_rows = sorted( + set(vision_only.tolist()) + | set(text_only.tolist()) + | set(paired_train.tolist()) + | set(val.tolist()) + | set(test.tolist()) + ) + manifest = { + "dataset": args.dataset, + "dataset_split": DATASET_SPLIT, + "seed": args.seed, + "vision_only_train": vision_only.tolist(), + "text_only_train": text_only.tolist(), + "paired_train": paired_train.tolist(), + "val": val.tolist(), + "test": test.tolist(), + "all_rows": all_rows, + "protocol": ( + "vision_only_train and text_only_train contain disjoint image IDs; " + "val/test pairs are held out from all bridge training" + ), + } + write_json(args.output, manifest) + print( + f"Wrote {args.output}: unpaired={n_unpaired}+{n_unpaired}, " + f"paired_upper_bound={paired_n}, val={len(val)}, test={len(test)}, " + f"features={len(all_rows)} rows" + ) + + +if __name__ == "__main__": + main() + |
