CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483
Raw
History Blame Contribute Delete
4.93 kB
"""Cache the pooled 768-D vector for every image in a COCO split.
python cache.py --split train2017
python cache.py --split val2017
Selection and evaluation both read these, so the backbone is forwarded once per
split rather than once per experiment. Values are written straight through to a
memmap alongside a per-row status byte, and meta.json is written last, so an
interrupted run resumes with --resume instead of restarting.
bfloat16 kernels select reduction orders by batch size, so cached values depend
on --batch. A cache must be built at one batch size throughout.
"""
import argparse
import json
import time
from pathlib import Path
import numpy as np
import torch
from PIL import Image
from common import COCO_ROOT, D, RES, coco_split, device, image_paths, normalize, pool
from common.artifacts import write_artifact
from common.models import load_backbone
from common.paths import BACKBONE
TODO, DONE, UNREADABLE = 0, 1, 2
class Images(torch.utils.data.Dataset):
"""Normalized images for a list of (row, path); unreadable files yield None."""
def __init__(self, work):
self.work = work
def __len__(self):
return len(self.work)
def __getitem__(self, i):
row, path = self.work[i]
try:
return row, normalize(Image.open(path), RES, 'cpu')[0]
except Exception: # noqa: BLE001
# One undecodable image is not worth losing an hour of forwards.
return row, None
def collate(batch):
good = [(r, x) for r, x in batch if x is not None]
bad = [r for r, x in batch if x is None]
if not good:
return None, bad
rows, xs = zip(*good)
return (torch.tensor(rows), torch.stack(xs)), bad
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument('--split', default='val2017', choices=('train2017', 'val2017'))
ap.add_argument('--backbone', default=BACKBONE)
ap.add_argument('--batch', type=int, default=8)
ap.add_argument('--workers', type=int, default=6)
ap.add_argument('--resume', action='store_true')
ap.add_argument('--out', type=Path, default=None)
args = ap.parse_args()
out = args.out or COCO_ROOT / f'pooled_{args.split}'
out.mkdir(parents=True, exist_ok=True)
dev = device()
coco, id_to_file = coco_split(args.split)
img_ids = sorted(coco.getImgIds())
paths = image_paths(id_to_file, img_ids, args.split)
n = len(img_ids)
present = all((out / f).exists() for f in ('pooled.dat', 'status.dat'))
if args.resume and not present:
raise SystemExit(f'--resume asked for but {out} has no partial run')
mode = 'r+' if (args.resume and present) else 'w+'
pooled = np.memmap(out / 'pooled.dat', np.float16, mode, shape=(n, D))
status = np.memmap(out / 'status.dat', np.uint8, mode, shape=(n,))
work = [(r, p) for r, p in enumerate(paths) if status[r] == TODO]
print(f'[init] {args.split}: {n} images, {len(work)} to compute', flush=True)
backbone = load_backbone(args.backbone).to(dev).eval()
loader = torch.utils.data.DataLoader(
Images(work), batch_size=args.batch, shuffle=False,
num_workers=args.workers, collate_fn=collate)
t0, seen = time.time(), 0
with torch.inference_mode():
for good, bad in loader:
for r in bad:
status[r] = UNREADABLE
if good is None:
continue
idx, x = good
with torch.autocast(dev, dtype=torch.bfloat16):
tok = backbone.forward_features(x.to(dev))['x_norm_patchtokens']
rows = idx.numpy()
pooled[rows] = pool(tok.float()).half().cpu().numpy()
status[rows] = DONE
seen += len(rows)
if seen % (args.batch * 50) < args.batch:
pooled.flush()
status.flush()
rate = seen / max(time.time() - t0, 1e-6)
print(f' {seen}/{len(work)} {rate:.1f} img/s '
f'ETA {(len(work) - seen) / max(rate, 1e-6) / 60:.1f} min',
flush=True)
pooled.flush()
status.flush()
done = int((np.asarray(status) == DONE).sum())
np.save(out / 'pooled.npy', np.asarray(pooled))
(out / 'img_ids.json').write_text(json.dumps([int(i) for i in img_ids]))
write_artifact(out / 'meta.json', {
'split': args.split,
'n_images': done,
'n_unreadable': int((np.asarray(status) == UNREADABLE).sum()),
'resolution': RES,
'files': {'pooled.npy': {'shape': [n, D], 'dtype': 'float16'}},
}, generator='cache.py', backbone=args.backbone, batch=args.batch,
protocol='live backbone forward at 768 px, layernorm over 768 channels '
'per patch, max over patches')
print(f'[done] {done}/{n} -> {out} ({time.time() - t0:.0f}s)', flush=True)
if __name__ == '__main__':
main()