QLD-WOB-v1 / inference.py
Addax-Data-Science's picture
Upload 4 files
417a24b verified
Raw
History Blame Contribute Delete
5.44 kB
"""
Inference script for QLD-WOB-v1 (Queensland Wet Tropics, WildObs)
Camera-trap classifier for the Wet Tropics of Queensland, Australia.
Model: WildObs QLD WetTropics
Input: 480x480 RGB, NHWC, unnormalised
Framework: PyTorch (whole pickled model object, not a state_dict)
Classes: 15 Wet Tropics species
Developer: Prakash Palanivelu Rajmohan and Renuka Sharma (WildObs)
Info: https://huggingface.co/WildObs/WildObs_QLD_WetTropics
Ported from AddaxAI's legacy classify_detections.py (wildobs-qld-wettropics),
whose crop function came from the authors' own evaluation notebook:
https://huggingface.co/WildObs/WildObs_QLD_WetTropics/blob/main/Evaluate_WetTropics_hf.ipynb
Author: Peter van Lunteren
"""
from __future__ import annotations
from pathlib import Path
import numpy as np
import torch
from PIL import Image, ImageFile
from torchvision import transforms
# Don't freak out over truncated images
ImageFile.LOAD_TRUNCATED_IMAGES = True
# The size the crop is squared to before preprocessing. The model then
# resizes again to 480. Both steps are kept because that is what the
# authors' notebook does, and the model was evaluated through it.
CROP_SIZE = 600
INPUT_SIZE = 480
# Class order is the model's output order and must not be reordered.
CLASS_NAMES = [
'Alectura_lathami', 'Bos_taurus', 'Canis_familiaris', 'Casuarius_casuarius',
'Felis_catus', 'Heteromyias_cinereifrons', 'Homo_sapiens',
'Hypsiprymnodon_moschatus', 'Megapodius_reinwardt', 'Orthonyx_spaldingii',
'Perameles_nasuta', 'Sus_scrofa', 'Thylogale_stigmatica',
'Uromys_caudimaculatus', 'Wallabia_bicolor'
]
class ModelInference:
"""WildObs Wet Tropics classifier."""
def __init__(self, model_dir: Path, model_path: Path) -> None:
self.model_dir = Path(model_dir)
self.model_path = Path(model_path)
self.model = None
self.device: torch.device | None = None
# Resize plus ToTensor, and deliberately no Normalize: the model
# was trained on raw 0-1 pixel values.
self.preprocess = transforms.Compose([
transforms.Resize((INPUT_SIZE, INPUT_SIZE)),
transforms.ToTensor(),
])
# ------------------------------------------------------------------
# Required interface
# ------------------------------------------------------------------
def check_gpu(self) -> bool:
try:
if torch.backends.mps.is_built() and torch.backends.mps.is_available():
return True
except Exception:
pass
return torch.cuda.is_available()
def load_model(self) -> None:
if self.check_gpu():
self.device = torch.device(
"mps" if torch.backends.mps.is_available() else "cuda"
)
else:
self.device = torch.device("cpu")
# The checkpoint is a pickled nn.Module, not a state_dict, so it
# has to be loaded with weights_only=False.
self.model = torch.load(
self.model_path, map_location=self.device, weights_only=False
)
self.model.eval()
self.model.to(self.device)
def get_crop(
self, image: Image.Image, bbox: tuple[float, float, float, float]
) -> Image.Image:
"""Crop the bbox, then square it to 600x600 as the authors do."""
width, height = image.size
x, y, w, h = bbox
left = int(x * width)
top = int(y * height)
right = int((x + w) * width)
bottom = int((y + h) * height)
crop = image.crop((left, top, right, bottom))
return crop.resize((CROP_SIZE, CROP_SIZE), Image.BILINEAR)
def get_classification(self, crop: Image.Image) -> list[list]:
"""Per-crop inference. Returns [[name, prob], ...] for all classes."""
assert self.model is not None
batch = self._to_nhwc(np.stack([self.get_tensor(crop)]))
probs = self._forward(batch)[0]
return [[CLASS_NAMES[i], float(probs[i])] for i in range(len(probs))]
def get_class_names(self) -> dict[str, str]:
"""1-indexed mapping {id: class_name} for the output JSON."""
return {str(i + 1): name for i, name in enumerate(CLASS_NAMES)}
# ------------------------------------------------------------------
# Optional batch interface
# ------------------------------------------------------------------
def get_tensor(self, crop: Image.Image) -> np.ndarray:
if crop.mode != "RGB":
crop = crop.convert("RGB")
return self.preprocess(crop).numpy()
def classify_batch(self, batch: np.ndarray) -> list[list[list]]:
assert self.model is not None
probs = self._forward(self._to_nhwc(batch))
return [
[[CLASS_NAMES[j], float(p[j])] for j in range(len(p))]
for p in probs
]
# ------------------------------------------------------------------
# Internals
# ------------------------------------------------------------------
def _to_nhwc(self, batch: np.ndarray) -> torch.Tensor:
"""B,C,H,W -> B,H,W,C. This model takes channels last."""
return torch.from_numpy(batch).permute(0, 2, 3, 1)
def _forward(self, batch: torch.Tensor) -> np.ndarray:
assert self.model is not None
with torch.no_grad():
logits = self.model(batch.to(self.device))
return torch.softmax(logits, dim=1).cpu().numpy()