CharlesCNorton
One coordinate comparison per COCO category on EUPE-ViT-B features
8f61e02
Raw
History Blame Contribute Delete
2.72 kB
"""COCO loading, input normalization and per-category image-level labels.
Images are resized to a square `resolution` with bilinear interpolation and
normalized with ImageNet statistics.
"""
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)
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 category_labels(coco, img_ids: Sequence[int], cat_id: int) -> torch.Tensor:
"""Image-level presence of one category for each id, as a bool tensor."""
have = set(coco.getImgIds(catIds=[cat_id]))
return torch.tensor([i in have for i in img_ids], dtype=torch.bool)
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]