File size: 4,682 Bytes
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
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"]