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()