"""COCO loading and input normalization. Images are resized to a square `resolution` with bilinear interpolation and normalized with ImageNet statistics, matching the protocol every stage was measured under. """ from pathlib import Path from typing import Iterable, List, Sequence, Tuple, Union import numpy as np import torch from PIL import Image from .paths import COCO_ROOT MEAN = (0.485, 0.456, 0.406) STD = (0.229, 0.224, 0.225) PERSON_CATEGORY_ID = 1 def _stats(device: str) -> Tuple[torch.Tensor, torch.Tensor]: mean = torch.tensor(MEAN).view(1, 3, 1, 1).to(device) std = torch.tensor(STD).view(1, 3, 1, 1).to(device) return mean, std def normalize(img: Image.Image, resolution: int, device: str) -> torch.Tensor: """PIL image -> (1, 3, R, R) normalized float tensor.""" img = img.convert('RGB').resize((resolution, resolution), Image.BILINEAR) arr = np.asarray(img, dtype=np.uint8).copy() x = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).to(device).float() / 255.0 mean, std = _stats(device) return (x - mean) / std def load_image(image: Union[str, Path, Image.Image, np.ndarray, torch.Tensor], resolution: int, device: str) -> torch.Tensor: """Accept a path, PIL image, HWC array, or CHW tensor; return a batch of 1.""" if isinstance(image, (str, Path)): img = Image.open(image) elif isinstance(image, Image.Image): img = image elif isinstance(image, np.ndarray): img = Image.fromarray(image) elif isinstance(image, torch.Tensor): arr = image.cpu().numpy() if image.ndim == 3 else image[0].cpu().numpy() if arr.shape[0] == 3: arr = arr.transpose(1, 2, 0) img = Image.fromarray((arr * 255).astype('uint8')) else: raise TypeError(f'unsupported image type: {type(image)}') return normalize(img, resolution, device) def coco_split(split: str = 'val2017'): """Return (COCO handle, image-file lookup) for a COCO split.""" from pycocotools.coco import COCO coco = COCO(str(COCO_ROOT / 'annotations' / f'instances_{split}.json')) id_to_file = {i['id']: i['file_name'] for i in coco.loadImgs(coco.getImgIds())} return coco, id_to_file def person_labels(coco, img_ids: Sequence[int], device: str = 'cpu') -> torch.Tensor: """Image-level person presence for each id, as a bool tensor.""" labels = [ any(a['category_id'] == PERSON_CATEGORY_ID for a in coco.loadAnns(coco.getAnnIds(imgIds=i, iscrowd=False))) for i in img_ids ] return torch.tensor(labels, dtype=torch.bool, device=device) def image_paths(id_to_file: dict, img_ids: Iterable[int], split: str = 'val2017') -> List[Path]: """Absolute paths for a sequence of image ids within a split.""" root = COCO_ROOT / split return [root / id_to_file[i] for i in img_ids]