CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483
Raw
History Blame Contribute Delete
2.68 kB
"""Evaluation-pool loading."""
from dataclasses import dataclass
from pathlib import Path
from typing import Iterator, List, Optional
import torch
from PIL import Image
from .data import coco_split, image_paths, normalize, person_labels
from .features import RES
from .pools import Pool
@dataclass
class LoadedPool:
"""Image ids, paths and labels for one named pool, images optionally resident."""
pool: Pool
img_ids: List[int]
paths: List[Path]
labels: torch.Tensor
device: str
images: Optional[List[torch.Tensor]] = None
def __len__(self) -> int:
return len(self.img_ids)
def __iter__(self) -> Iterator[torch.Tensor]:
"""Yield each image as a normalized (1, 3, RES, RES) tensor."""
if self.images is not None:
yield from self.images
return
for path in self.paths:
yield normalize(Image.open(path), RES, self.device)
@property
def positive_rate(self) -> float:
return round(self.labels.float().mean().item(), 4)
def provenance(self) -> dict:
"""Pool fields recorded in an artifact's provenance block."""
return {'pool': self.pool.name, 'split': self.pool.split,
'n_images': len(self), 'positive_rate': self.positive_rate,
'selection': self.pool.selection}
def balanced_indices(labels: torch.Tensor, seed: int = 0) -> torch.Tensor:
"""Indices subsampling `labels` to equal positive and negative counts, seeded."""
generator = torch.Generator(device='cpu').manual_seed(seed)
cpu = labels.cpu()
pos = cpu.nonzero(as_tuple=True)[0]
neg = (~cpu).nonzero(as_tuple=True)[0]
n = min(len(pos), len(neg))
sel = torch.cat([pos[torch.randperm(len(pos), generator=generator)[:n]],
neg[torch.randperm(len(neg), generator=generator)[:n]]])
return sel[torch.randperm(len(sel), generator=generator)]
def load_pool(pool: Pool, device: str, preload: bool = False, seed: int = 0) -> LoadedPool:
"""Resolve a named pool to ids, paths and labels; `preload` holds images in memory."""
coco, id_to_file = coco_split(pool.split)
img_ids = sorted(coco.getImgIds())
if pool.n is not None:
img_ids = img_ids[:pool.n]
labels = person_labels(coco, img_ids, device)
if pool.balanced:
sel = balanced_indices(labels, seed)
img_ids = [img_ids[i] for i in sel.tolist()]
labels = labels[sel.to(labels.device)]
paths = image_paths(id_to_file, img_ids, pool.split)
images = [normalize(Image.open(p), RES, device) for p in paths] if preload else None
return LoadedPool(pool, img_ids, paths, labels, device, images)