CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483 | """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() | |