#!/usr/bin/env python3 """Sample a stratified subset from data/in1k_hybrid_1k/ for ablation sweep. Picks N images from the hybrid dataset, prioritizing synsets with Guillaumin GT mask (so foreground analysis remains valid for as many subset rows as possible). Within each preference tier, samples are deterministic given --seed. Output: a new dir with N images + metadata.json (same schema as parent dataset). Usage: python scripts/sample_hybrid_subset.py --n 200 \ --in data/in1k_hybrid_1k \ --out data/in1k_hybrid_1k_subset200 \ --seed 42 """ from __future__ import annotations import argparse import json import random import shutil import sys from pathlib import Path def _find_project_root() -> Path: cur = Path(__file__).resolve().parent for parent in [cur, *cur.parents]: if (parent / "requirements.txt").exists(): return parent raise RuntimeError(f"project root not found from {__file__}") PROJECT_ROOT = _find_project_root() def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--in", dest="in_dir", type=Path, default=PROJECT_ROOT / "data" / "in1k_hybrid_1k") parser.add_argument("--out", type=Path, default=PROJECT_ROOT / "data" / "in1k_hybrid_1k_subset200") parser.add_argument("--n", type=int, default=200) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--use-symlinks", action="store_true") args = parser.parse_args() rng = random.Random(args.seed) parent_meta = args.in_dir / "metadata.json" if not parent_meta.exists(): print(f"ERROR: parent metadata not found: {parent_meta}") print(f"Run `python scripts/build_in1k_hybrid_1k.py` first.") return 1 payload = json.loads(parent_meta.read_text()) samples = payload.get("samples") or [] if len(samples) != 1000: print(f"ERROR: parent has {len(samples)} samples, expected 1000") return 1 # Stratified sampling: pegar TODOS os has_mask=True primeiro (~95), # depois preencher com aleatórios sem mask até atingir n. with_mask = [s for s in samples if s.get("has_mask")] without_mask = [s for s in samples if not s.get("has_mask")] print(f"Parent dataset: {len(with_mask)} com mask, {len(without_mask)} sem mask") if args.n <= len(with_mask): # Subset cabe inteirinho dentro dos com mask: amostra deterministicamente chosen = sorted(rng.sample(with_mask, args.n), key=lambda s: s["imagenet_id"]) else: n_extra = args.n - len(with_mask) extras = sorted(rng.sample(without_mask, n_extra), key=lambda s: s["imagenet_id"]) chosen = sorted(with_mask + extras, key=lambda s: s["imagenet_id"]) n_chosen_mask = sum(1 for s in chosen if s.get("has_mask")) print(f"Selected: {len(chosen)} ({n_chosen_mask} com mask)") args.out.mkdir(parents=True, exist_ok=True) for s in chosen: src = args.in_dir / s["filename"] dst = args.out / s["filename"] if not src.exists(): print(f"ERROR: source missing: {src}") return 1 if args.use_symlinks: if dst.exists() or dst.is_symlink(): dst.unlink() dst.symlink_to(src.resolve()) else: shutil.copy2(src, dst) # Metadata schema: same as parent for compat with run_attack_sweep loader out_meta = { "description": ( f"Stratified subset of in1k_hybrid_1k (N={len(chosen)}). " f"Prioritizes Guillaumin∩IN-1k synsets (has_mask=True) for " f"foreground analysis coverage. Same schema as parent dataset." ), "version": "1.0", "seed": args.seed, "parent": str(args.in_dir.relative_to(PROJECT_ROOT)), "total": len(chosen), "n_with_mask": n_chosen_mask, "n_without_mask": len(chosen) - n_chosen_mask, "downloaded_files": [s["filename"] for s in chosen], "suggested_classes": [ { "synset": s["synset"], "imagenet_id": s["imagenet_id"], "in_imagenet_1k": True, "source": s["source"], "has_mask": s["has_mask"], } for s in chosen ], "samples": chosen, } (args.out / "metadata.json").write_text(json.dumps(out_meta, indent=2)) print(f"\n✓ Subset OK") print(f" {len(chosen)} imgs em {args.out} ({n_chosen_mask} com mask)") return 0 if __name__ == "__main__": sys.exit(main())