File size: 5,018 Bytes
50ce3c9
26bd0a2
50ce3c9
 
 
 
 
26bd0a2
 
50ce3c9
 
752f8f0
 
 
 
50ce3c9
 
 
26bd0a2
 
50ce3c9
 
 
 
 
 
26bd0a2
 
 
 
 
 
 
 
 
 
50ce3c9
26bd0a2
 
 
50ce3c9
 
 
752f8f0
 
 
 
 
 
 
50ce3c9
 
 
 
26bd0a2
 
50ce3c9
 
 
752f8f0
 
 
 
 
 
50ce3c9
 
 
 
 
 
 
 
 
 
 
 
752f8f0
50ce3c9
 
 
 
 
 
 
 
 
 
 
 
 
 
752f8f0
 
 
 
50ce3c9
 
 
752f8f0
50ce3c9
 
 
 
 
 
 
 
752f8f0
 
 
50ce3c9
 
 
 
752f8f0
50ce3c9
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
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"]