DevQueen/deepfake-server
0
1"""Rebuild data/metadata.csv from all .npz files in data/processed/."""2import csv3import sys4from pathlib import Path5 6import numpy as np7 8ROOT = Path(__file__).parent.parent9processed = ROOT / "data" / "processed"10out_csv = ROOT / "data" / "metadata.csv"11 12rows = []13for npz in sorted(processed.glob("*.npz")):14 try:15 d = np.load(npz, allow_pickle=True)16 label = int(d["label"])17 video_id = str(d["video_id"])18 rows.append({"npz_path": str(npz.resolve()), "label": label, "video_id": video_id})19 except Exception as e:20 print(f"skipping {npz.name}: {e}", file=sys.stderr)21 22# Identity-disjoint split23rng = np.random.default_rng(42)24unique_ids = sorted({r["video_id"] for r in rows})25rng.shuffle(unique_ids)26n = len(unique_ids)27train_ids = set(unique_ids[: int(0.7 * n)])28val_ids = set(unique_ids[int(0.7 * n) : int(0.85 * n)])29 30for r in rows:31 if r["video_id"] in train_ids:32 r["split"] = "train"33 elif r["video_id"] in val_ids:34 r["split"] = "val"35 else:36 r["split"] = "test"37 38out_csv.parent.mkdir(parents=True, exist_ok=True)39with out_csv.open("w", newline="", encoding="utf-8") as f:40 writer = csv.DictWriter(f, fieldnames=["npz_path", "label", "video_id", "split"])41 writer.writeheader()42 writer.writerows(rows)43 44from collections import Counter45label_counts = Counter(r["label"] for r in rows)46split_counts = Counter(r["split"] for r in rows)47print(f"Rebuilt {len(rows)} sequences")48print(f" real (0): {label_counts[0]}, fake (1): {label_counts[1]}")49print(f" train: {split_counts['train']}, val: {split_counts['val']}, test: {split_counts['test']}")50print(f" written to: {out_csv}")51 