mwmathis's picture C-Achard's picture
Add PyTorch SuperAnimal backend and refresh the Space (#14)
2e1b62d
Raw History Blame Contribute Delete
5.02 kB
import threading
from pathlib import Path
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 MODEL_FILENAME_MAPPING
from dlclibrary import download_huggingface_model
# 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)
# next to the TF models, not deeplabcut/modelzoo/checkpoints: site-packages is read-only for the Space's non-root user
WEIGHTS_DIR = Path(__file__).parent / "DLC_models" / "pytorch"
_runners = {}
_build_lock = threading.Lock()
##########################################
def snapshot_path(superanimal, model_name):
"""Path to a SuperAnimal snapshot in WEIGHTS_DIR, downloaded on first use (as deeplabcut does)."""
name = f"{superanimal}_{model_name}"
path = WEIGHTS_DIR / f"{name}.pt"
if not path.exists():
source = MODEL_FILENAME_MAPPING.get(name, path.name)
rename = None if source == path.name else {source: path.name}
download_huggingface_model(name, target_dir=str(WEIGHTS_DIR), rename_mapping=rename)
return path
##########################################
def load_superanimal(superanimal, device="auto"):
"""Build (once) the detector and pose runners for a SuperAnimal model; weights are downloaded on first use."""
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=snapshot_path(superanimal, pose_model),
detector_path=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"], strict=True
):
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"]