PixelModel-v3 / fetch_coco_subset.py
TobiasLogic's picture
PixelModel v3: SIREN+FiLM CPPN, 919K params, beats v1 FID
63a1291 verified
Raw
History Blame Contribute Delete
9.11 kB
"""Build the PixelModel v3 train/eval data from MS-COCO.
Same source as v1: `sayakpaul/coco-30-val-2014` (MS-COCO val2014 caption/image
pairs). The set is split *deterministically by image hash* into a train subset
and an eval subset so the two are provably disjoint and reproducible - exactly
the regression-safe protocol v1 used.
Outputs (into --out, default ../pm-work):
coco_train.npz images (uint8, N x S x S x 3) + token ids (int32, N x T)
coco_eval.npz eval images + raw caption strings (for FID/CLIP)
vocab.json word-level tokenizer vocabulary (built from TRAIN only)
split_manifest.json hashes + counts, so disjointness is auditable
Run on the training machine (needs `datasets` + network):
python fetch_coco_subset.py --out ../pm-work
"""
from __future__ import annotations
import argparse
import hashlib
import io
import json
import os
import time
import numpy as np
from PIL import Image
from model import build_vocab, encode_caption, PAD_ID
def center_square_resize(img: Image.Image, size: int) -> np.ndarray:
"""Center-crop to a square, resize to size x size, return uint8 HWC RGB."""
img = img.convert("RGB")
w, h = img.size
s = min(w, h)
left = (w - s) // 2
top = (h - s) // 2
img = img.crop((left, top, left + s, top + s)).resize((size, size), Image.BICUBIC)
return np.asarray(img, dtype=np.uint8)
def image_hash(img: Image.Image) -> str:
"""Stable content hash of an image, independent of dataset row order."""
buf = io.BytesIO()
img.convert("RGB").resize((64, 64), Image.BICUBIC).save(buf, format="PNG")
return hashlib.md5(buf.getvalue()).hexdigest()
def caption_of(example) -> str:
"""COCO rows expose captions under a few different keys across mirrors."""
for key in ("caption", "captions", "text", "sentences", "annotations_captions"):
if key in example and example[key]:
v = example[key]
if isinstance(v, (list, tuple)):
return str(v[0])
return str(v)
raise KeyError(f"no caption field found; row keys = {list(example.keys())}")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--out", default="../pm-work", help="output directory")
ap.add_argument("--dataset", default="sayakpaul/coco-30-val-2014",
help="HF dataset id (MS-COCO caption/image pairs)")
ap.add_argument("--split", default="train", help="HF split to pull from")
ap.add_argument("--n-train", type=int, default=20000)
ap.add_argument("--n-eval", type=int, default=5000)
ap.add_argument("--store-size", type=int, default=144,
help="square size images are stored at (must be >= train crop)")
ap.add_argument("--vocab-size", type=int, default=8192, help="base model, incl. PAD+UNK")
ap.add_argument("--vocab-size-xl", type=int, default=16384, help="XL variant vocab cap")
ap.add_argument("--min-freq", type=int, default=2)
ap.add_argument("--max-tokens", type=int, default=20)
ap.add_argument("--eval-hash-mod", type=int, default=6,
help="1/eval-hash-mod of images are reserved for eval")
ap.add_argument("--streaming", action="store_true",
help="stream the split instead of downloading it first. "
"Faster to start, but a long collection loop can die "
"to a dropped connection; off by default.")
ap.add_argument("--load-retries", type=int, default=6)
ap.add_argument("--max-reconnects", type=int, default=8,
help="resume attempts if iteration dies mid-collection")
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
os.makedirs(args.out, exist_ok=True)
from datasets import load_dataset
def open_dataset():
"""Load the split, retrying the network part. Default is NON-streaming:
the shards are downloaded once (HF retries/resumes those internally) and
iteration is then purely local, so a long collection loop cannot die
halfway to a dropped HTTP connection."""
last = None
for t in range(args.load_retries):
try:
return load_dataset(args.dataset, split=args.split,
streaming=args.streaming)
except Exception as e:
last = e
wait = 5 * (t + 1)
print(f"[data] load attempt {t + 1}/{args.load_retries} failed "
f"({type(e).__name__}: {e}); retrying in {wait}s", flush=True)
time.sleep(wait)
raise RuntimeError(f"could not load {args.dataset}: {last}")
train_imgs, train_caps_raw, train_hashes = [], [], []
eval_imgs, eval_caps, eval_hashes = [], [], []
seen = set()
def need_more():
return len(train_imgs) < args.n_train or len(eval_imgs) < args.n_eval
last_report = 0
reconnects = 0
print(f"[data] loading {args.dataset} split={args.split} "
f"(streaming={args.streaming})", flush=True)
while need_more():
ds = open_dataset()
try:
for ex in ds:
if not need_more():
break
try:
img = ex.get("image") or ex.get("img")
if img is None:
continue
cap = caption_of(ex)
h = image_hash(img)
if h in seen:
continue
seen.add(h)
is_eval = (int(h, 16) % args.eval_hash_mod == 0)
if is_eval:
if len(eval_imgs) >= args.n_eval:
continue
eval_imgs.append(center_square_resize(img, 256))
eval_caps.append(cap)
eval_hashes.append(h)
else:
if len(train_imgs) >= args.n_train:
continue
train_imgs.append(center_square_resize(img, args.store_size))
train_caps_raw.append(cap)
train_hashes.append(h)
except Exception:
continue
total = len(train_imgs) + len(eval_imgs)
if total - last_report >= 1000:
last_report = total
print(f"[data] train={len(train_imgs)} eval={len(eval_imgs)}", flush=True)
break
except Exception as e:
reconnects += 1
print(f"[data] iteration died ({type(e).__name__}: {e}); "
f"reconnect {reconnects}/{args.max_reconnects} with "
f"train={len(train_imgs)} eval={len(eval_imgs)}", flush=True)
if reconnects >= args.max_reconnects:
print("[data] giving up on reconnects; proceeding with what we have",
flush=True)
break
time.sleep(5)
assert set(train_hashes).isdisjoint(set(eval_hashes)), "train/eval overlap!"
print(f"[data] collected train={len(train_imgs)} eval={len(eval_imgs)} "
f"(disjoint: {set(train_hashes).isdisjoint(set(eval_hashes))})")
vocab = build_vocab(train_caps_raw, args.vocab_size, args.min_freq)
with open(os.path.join(args.out, "vocab.json"), "w") as f:
json.dump(vocab, f)
vocab_xl = build_vocab(train_caps_raw, args.vocab_size_xl, args.min_freq)
with open(os.path.join(args.out, "vocab_xl.json"), "w") as f:
json.dump(vocab_xl, f)
print(f"[data] vocab size = {len(vocab)} (cap {args.vocab_size}), "
f"vocab_xl = {len(vocab_xl)} (cap {args.vocab_size_xl})")
train_tokens = np.stack(
[encode_caption(c, vocab, args.max_tokens) for c in train_caps_raw])
np.savez_compressed(
os.path.join(args.out, "coco_train.npz"),
images=np.stack(train_imgs),
tokens=train_tokens,
captions=np.array(train_caps_raw, dtype=object),
store_size=args.store_size,
max_tokens=args.max_tokens,
)
np.savez_compressed(
os.path.join(args.out, "coco_eval.npz"),
images=np.stack(eval_imgs),
captions=np.array(eval_caps, dtype=object),
)
with open(os.path.join(args.out, "split_manifest.json"), "w") as f:
json.dump({
"dataset": args.dataset,
"n_train": len(train_imgs),
"n_eval": len(eval_imgs),
"eval_hash_mod": args.eval_hash_mod,
"vocab_size": len(vocab),
"store_size": args.store_size,
"max_tokens": args.max_tokens,
"disjoint": bool(set(train_hashes).isdisjoint(set(eval_hashes))),
"train_hash_sample": train_hashes[:8],
"eval_hash_sample": eval_hashes[:8],
}, f, indent=2)
mb = os.path.getsize(os.path.join(args.out, "coco_train.npz")) / 1e6
print(f"[data] wrote coco_train.npz ({mb:.1f} MB), coco_eval.npz, vocab.json")
print(f"[data] done -> {args.out}")
if __name__ == "__main__":
main()