C-Achard's picture
Add PyTorch SuperAnimal backend
50ce3c9
Raw History Blame
4.68 kB
import threading
import numpy as np
import PIL
from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners
from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path
# SuperAnimal (pose model, detector) used by the PyTorch backend
PYTORCH_MODELS = {'superanimal_quadruped': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2'),
'superanimal_topviewmouse': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2')}
MAX_INDIVIDUALS = 10
MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
_runners = {}
_build_lock = threading.Lock()
##########################################
def load_superanimal(superanimal, device="auto"):
"""Build (once) the detector and pose runners for a SuperAnimal model.
Weights are downloaded on first use to deeplabcut/modelzoo/checkpoints.
"""
with _build_lock:
if superanimal not in _runners:
pose_model, detector = PYTORCH_MODELS[superanimal]
cfg = PoseConfig.build_for_superanimal_inference(superanimal,
model_name=pose_model,
detector_name=detector,
max_individuals=MAX_INDIVIDUALS,
device=device)
# keep low-score boxes: the UI threshold filters them afterwards
cfg["detector"]["model"]["box_score_thresh"] = 0.05
pose_runner, detector_runner = get_inference_runners(
cfg,
snapshot_path=get_super_animal_snapshot_path(superanimal, pose_model),
detector_path=get_super_animal_snapshot_path(superanimal, detector),
max_individuals=MAX_INDIVIDUALS,
inference_cfg={"multithreading": {"enabled": False}},
)
_runners[superanimal] = {"pose": pose_runner,
"detector": detector_runner,
"bodyparts": list(cfg["metadata"]["bodyparts"]),
"lock": threading.Lock()} # runners are not thread-safe
return _runners[superanimal]
##########################################
def resize_max_side(img, max_size=MAX_IMAGE_SIZE):
scale = max_size / max(img.size)
if scale >= 1:
return img
return img.resize([int(x * scale) for x in img.size], PIL.Image.Resampling.LANCZOS)
##########################################
def predict_superanimal(img_input,
superanimal,
bbox_likelihood_th,
kpts_likelihood_th,
full_image=False):
"""Detect animals and estimate their pose with a PyTorch SuperAnimal model.
Returns the (resized) RGB image the predictions refer to, the list of animals
as dicts {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3) array of x,y,llk
in image pixels, NaN below kpts_likelihood_th}, and the bodypart names.
"""
img = resize_max_side(img_input.convert("RGB"))
img_np = np.asarray(img)
runners = load_superanimal(superanimal)
with runners["lock"]:
if full_image:
# skip the detector and treat the whole image as one animal
h, w = img_np.shape[:2]
detections = {"bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
"bbox_scores": np.array([1.0], dtype=np.float32)}
else:
detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
keep = detections["bbox_scores"] >= bbox_likelihood_th
detections = {"bboxes": detections["bboxes"][keep],
"bbox_scores": detections["bbox_scores"][keep]}
if len(detections["bboxes"]) == 0:
return img, [], runners["bodyparts"]
predictions = runners["pose"].inference([(img_np, detections)])[0]
animals = []
# outputs are padded to MAX_INDIVIDUALS with -1
for kpts, (x, y, w, h), score in zip(predictions["bodyparts"],
predictions["bboxes"],
predictions["bbox_scores"]):
if score < 0:
continue
kpts = kpts.astype(float)
kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)],
"kpts": kpts})
return img, animals, runners["bodyparts"]