Add PyTorch SuperAnimal backend
Browse filesIntroduce a PyTorch inference path for SuperAnimal models, including backend selection in the UI, cached pose/detector runners, and JSON export of predictions. The legacy TensorFlow flow now maps model names separately from local directories, and MegaDetector device selection was corrected for yolov5 loading. The DeepLabCut dependency is also pinned to 3.0.2 to match the new backend integration.
- app.py +73 -15
- detection_utils.py +3 -11
- pytorch_utils.py +104 -0
- requirements.txt +1 -1
- ui_utils.py +13 -5
- viz_utils.py +33 -0
app.py
CHANGED
|
@@ -16,9 +16,10 @@ import dlclive
|
|
| 16 |
from PIL import Image, ImageColor, ImageFont, ImageDraw
|
| 17 |
import requests
|
| 18 |
|
| 19 |
-
from viz_utils import save_results_as_json, draw_keypoints_on_image, draw_bbox_w_text, save_results_only_dlc
|
| 20 |
from detection_utils import predict_md, crop_animal_detections
|
| 21 |
from dlc_utils import predict_dlc
|
|
|
|
| 22 |
from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
|
| 23 |
|
| 24 |
from deeplabcut.utils import auxiliaryfunctions
|
|
@@ -42,14 +43,58 @@ image = Image.open(requests.get(url, stream=True).raw)
|
|
| 42 |
MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
|
| 43 |
'md_v5b': "MD_models/md_v5b.0.0.pt"}
|
| 44 |
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
|
| 50 |
|
| 51 |
#####################################################
|
| 52 |
def predict_pipeline(img_input,
|
|
|
|
| 53 |
mega_model_input,
|
| 54 |
dlc_model_input_str,
|
| 55 |
flag_dlc_only,
|
|
@@ -62,6 +107,21 @@ def predict_pipeline(img_input,
|
|
| 62 |
marker_size,
|
| 63 |
):
|
| 64 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 65 |
if not flag_dlc_only:
|
| 66 |
############################################################
|
| 67 |
# ### Run Megadetector
|
|
@@ -81,15 +141,12 @@ def predict_pipeline(img_input,
|
|
| 81 |
|
| 82 |
# If model is found: do not download (previous execution is likely within same day)
|
| 83 |
# TODO: can we ask the user whether to reload dlc model if a directory is found?
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
else:
|
| 88 |
-
path_to_DLCmodel = DLC_models_dict[dlc_model_input_str]
|
| 89 |
-
download_huggingface_model(dlc_model_input_str, path_to_DLCmodel)
|
| 90 |
|
| 91 |
# extract map label ids to strings
|
| 92 |
-
pose_cfg_path = os.path.join(
|
| 93 |
'pose_cfg.yaml')
|
| 94 |
with open(pose_cfg_path, "r") as stream:
|
| 95 |
pose_cfg_dict = yaml.safe_load(stream)
|
|
@@ -119,7 +176,7 @@ def predict_pipeline(img_input,
|
|
| 119 |
keypt_color=keypt_color,
|
| 120 |
marker_size=marker_size)
|
| 121 |
|
| 122 |
-
donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,
|
| 123 |
|
| 124 |
return img_input, donw_file
|
| 125 |
|
|
@@ -163,7 +220,7 @@ def predict_pipeline(img_input,
|
|
| 163 |
|
| 164 |
|
| 165 |
# Save detection results as json
|
| 166 |
-
download_file = save_results_as_json(md_results,list_kpts_per_crop,list_bboxes,map_label_id_to_str,
|
| 167 |
|
| 168 |
return img_background, download_file
|
| 169 |
|
|
@@ -171,7 +228,8 @@ def predict_pipeline(img_input,
|
|
| 171 |
|
| 172 |
#########################################################
|
| 173 |
# Define user interface and launch
|
| 174 |
-
inputs = gradio_inputs_for_MD_DLC(
|
|
|
|
| 175 |
list(DLC_models_dict.keys()))
|
| 176 |
outputs = gradio_outputs_for_MD_DLC()
|
| 177 |
[gr_title,
|
|
|
|
| 16 |
from PIL import Image, ImageColor, ImageFont, ImageDraw
|
| 17 |
import requests
|
| 18 |
|
| 19 |
+
from viz_utils import save_results_as_json, draw_keypoints_on_image, draw_bbox_w_text, save_results_only_dlc, save_results_pytorch
|
| 20 |
from detection_utils import predict_md, crop_animal_detections
|
| 21 |
from dlc_utils import predict_dlc
|
| 22 |
+
from pytorch_utils import predict_superanimal, PYTORCH_MODELS
|
| 23 |
from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
|
| 24 |
|
| 25 |
from deeplabcut.utils import auxiliaryfunctions
|
|
|
|
| 43 |
MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
|
| 44 |
'md_v5b': "MD_models/md_v5b.0.0.pt"}
|
| 45 |
|
| 46 |
+
BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
|
| 47 |
+
|
| 48 |
+
# TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
|
| 49 |
+
DLC_models_dict = {'superanimal_topviewmouse': ('superanimal_topviewmouse_dlcrnet', "DLC_models/sa-tvm"),
|
| 50 |
+
'superanimal_quadruped': ('superanimal_quadruped_dlcrnet', "DLC_models/sa-q")}
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
#####################################################
|
| 55 |
+
def predict_pipeline_pytorch(img_input,
|
| 56 |
+
superanimal,
|
| 57 |
+
flag_dlc_only,
|
| 58 |
+
flag_show_str_labels,
|
| 59 |
+
bbox_likelihood_th,
|
| 60 |
+
kpts_likelihood_th,
|
| 61 |
+
font_style,
|
| 62 |
+
font_size,
|
| 63 |
+
keypt_color,
|
| 64 |
+
marker_size,
|
| 65 |
+
):
|
| 66 |
+
# detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
|
| 67 |
+
img_output, animals, bodyparts = predict_superanimal(img_input,
|
| 68 |
+
superanimal,
|
| 69 |
+
bbox_likelihood_th,
|
| 70 |
+
kpts_likelihood_th,
|
| 71 |
+
full_image=flag_dlc_only)
|
| 72 |
+
map_label_id_to_str = dict(enumerate(bodyparts))
|
| 73 |
+
|
| 74 |
+
for animal in animals:
|
| 75 |
+
draw_keypoints_on_image(img_output,
|
| 76 |
+
animal['kpts'],
|
| 77 |
+
map_label_id_to_str,
|
| 78 |
+
flag_show_str_labels,
|
| 79 |
+
use_normalized_coordinates=False,
|
| 80 |
+
font_style=font_style,
|
| 81 |
+
font_size=font_size,
|
| 82 |
+
keypt_color=keypt_color,
|
| 83 |
+
marker_size=marker_size)
|
| 84 |
+
if not flag_dlc_only:
|
| 85 |
+
draw_bbox_w_text(img_output,
|
| 86 |
+
animal['bbox'],
|
| 87 |
+
font_size=font_size)
|
| 88 |
+
|
| 89 |
+
pose_model, detector = PYTORCH_MODELS[superanimal]
|
| 90 |
+
download_file = save_results_pytorch(animals, map_label_id_to_str, superanimal,
|
| 91 |
+
pose_model, None if flag_dlc_only else detector)
|
| 92 |
+
return img_output, download_file
|
| 93 |
|
| 94 |
|
| 95 |
#####################################################
|
| 96 |
def predict_pipeline(img_input,
|
| 97 |
+
backend,
|
| 98 |
mega_model_input,
|
| 99 |
dlc_model_input_str,
|
| 100 |
flag_dlc_only,
|
|
|
|
| 107 |
marker_size,
|
| 108 |
):
|
| 109 |
|
| 110 |
+
if backend == "PyTorch":
|
| 111 |
+
return predict_pipeline_pytorch(img_input,
|
| 112 |
+
dlc_model_input_str,
|
| 113 |
+
flag_dlc_only,
|
| 114 |
+
flag_show_str_labels,
|
| 115 |
+
bbox_likelihood_th,
|
| 116 |
+
kpts_likelihood_th,
|
| 117 |
+
font_style,
|
| 118 |
+
font_size,
|
| 119 |
+
keypt_color,
|
| 120 |
+
marker_size)
|
| 121 |
+
|
| 122 |
+
# TensorFlow (legacy): MegaDetector crops + DLCLive
|
| 123 |
+
dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
|
| 124 |
+
|
| 125 |
if not flag_dlc_only:
|
| 126 |
############################################################
|
| 127 |
# ### Run Megadetector
|
|
|
|
| 141 |
|
| 142 |
# If model is found: do not download (previous execution is likely within same day)
|
| 143 |
# TODO: can we ask the user whether to reload dlc model if a directory is found?
|
| 144 |
+
path_to_DLCmodel = dlc_model_dir
|
| 145 |
+
if not (os.path.isdir(dlc_model_dir) and len(os.listdir(dlc_model_dir)) > 0):
|
| 146 |
+
download_huggingface_model(dlc_model_name, path_to_DLCmodel)
|
|
|
|
|
|
|
|
|
|
| 147 |
|
| 148 |
# extract map label ids to strings
|
| 149 |
+
pose_cfg_path = os.path.join(dlc_model_dir,
|
| 150 |
'pose_cfg.yaml')
|
| 151 |
with open(pose_cfg_path, "r") as stream:
|
| 152 |
pose_cfg_dict = yaml.safe_load(stream)
|
|
|
|
| 176 |
keypt_color=keypt_color,
|
| 177 |
marker_size=marker_size)
|
| 178 |
|
| 179 |
+
donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_name)
|
| 180 |
|
| 181 |
return img_input, donw_file
|
| 182 |
|
|
|
|
| 220 |
|
| 221 |
|
| 222 |
# Save detection results as json
|
| 223 |
+
download_file = save_results_as_json(md_results,list_kpts_per_crop,list_bboxes,map_label_id_to_str,dlc_model_name,mega_model_input)
|
| 224 |
|
| 225 |
return img_background, download_file
|
| 226 |
|
|
|
|
| 228 |
|
| 229 |
#########################################################
|
| 230 |
# Define user interface and launch
|
| 231 |
+
inputs = gradio_inputs_for_MD_DLC(BACKENDS,
|
| 232 |
+
list(MD_models_dict.keys()),
|
| 233 |
list(DLC_models_dict.keys()))
|
| 234 |
outputs = gradio_outputs_for_MD_DLC()
|
| 235 |
[gr_title,
|
detection_utils.py
CHANGED
|
@@ -24,11 +24,8 @@ def predict_md(im,
|
|
| 24 |
g = (size / max(im.size)) # multipl factor to make max size of the image equal to input size
|
| 25 |
im = im.resize((int(x * g) for x in im.size),
|
| 26 |
PIL.Image.Resampling.LANCZOS) # resize
|
| 27 |
-
# device
|
| 28 |
-
if torch.cuda.is_available()
|
| 29 |
-
md_device = torch.device('cuda')
|
| 30 |
-
else:
|
| 31 |
-
md_device = torch.device('cpu')
|
| 32 |
|
| 33 |
# megadetector
|
| 34 |
MD_model = torch.hub.load('ultralytics/yolov5', # repo_or_dir
|
|
@@ -37,12 +34,7 @@ def predict_md(im,
|
|
| 37 |
skip_validation=True, # avoid GitHub API rate limit (403)
|
| 38 |
device=md_device,
|
| 39 |
trust_repo=True
|
| 40 |
-
)
|
| 41 |
-
|
| 42 |
-
# send model to gpu if possible
|
| 43 |
-
if (md_device == torch.device('cuda')):
|
| 44 |
-
print('Sending model to GPU')
|
| 45 |
-
MD_model.to(md_device)
|
| 46 |
|
| 47 |
## detect objects
|
| 48 |
results = MD_model(im) # inference # vars(results).keys()= dict_keys(['imgs', 'pred', 'names', 'files', 'times', 'xyxy', 'xywh', 'xyxyn', 'xywhn', 'n', 't', 's'])
|
|
|
|
| 24 |
g = (size / max(im.size)) # multipl factor to make max size of the image equal to input size
|
| 25 |
im = im.resize((int(x * g) for x in im.size),
|
| 26 |
PIL.Image.Resampling.LANCZOS) # resize
|
| 27 |
+
# device: yolov5's select_device expects a CUDA index ('0') or 'cpu', not 'cuda'
|
| 28 |
+
md_device = '0' if torch.cuda.is_available() else 'cpu'
|
|
|
|
|
|
|
|
|
|
| 29 |
|
| 30 |
# megadetector
|
| 31 |
MD_model = torch.hub.load('ultralytics/yolov5', # repo_or_dir
|
|
|
|
| 34 |
skip_validation=True, # avoid GitHub API rate limit (403)
|
| 35 |
device=md_device,
|
| 36 |
trust_repo=True
|
| 37 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 38 |
|
| 39 |
## detect objects
|
| 40 |
results = MD_model(im) # inference # vars(results).keys()= dict_keys(['imgs', 'pred', 'names', 'files', 'times', 'xyxy', 'xywh', 'xyxyn', 'xywhn', 'n', 't', 's'])
|
pytorch_utils.py
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import threading
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
import PIL
|
| 5 |
+
from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners
|
| 6 |
+
from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
|
| 7 |
+
from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
# SuperAnimal (pose model, detector) used by the PyTorch backend
|
| 11 |
+
PYTORCH_MODELS = {'superanimal_quadruped': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2'),
|
| 12 |
+
'superanimal_topviewmouse': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2')}
|
| 13 |
+
|
| 14 |
+
MAX_INDIVIDUALS = 10
|
| 15 |
+
MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
|
| 16 |
+
|
| 17 |
+
_runners = {}
|
| 18 |
+
_build_lock = threading.Lock()
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
##########################################
|
| 22 |
+
def load_superanimal(superanimal, device="auto"):
|
| 23 |
+
"""Build (once) the detector and pose runners for a SuperAnimal model.
|
| 24 |
+
|
| 25 |
+
Weights are downloaded on first use to deeplabcut/modelzoo/checkpoints.
|
| 26 |
+
"""
|
| 27 |
+
with _build_lock:
|
| 28 |
+
if superanimal not in _runners:
|
| 29 |
+
pose_model, detector = PYTORCH_MODELS[superanimal]
|
| 30 |
+
cfg = PoseConfig.build_for_superanimal_inference(superanimal,
|
| 31 |
+
model_name=pose_model,
|
| 32 |
+
detector_name=detector,
|
| 33 |
+
max_individuals=MAX_INDIVIDUALS,
|
| 34 |
+
device=device)
|
| 35 |
+
# keep low-score boxes: the UI threshold filters them afterwards
|
| 36 |
+
cfg["detector"]["model"]["box_score_thresh"] = 0.05
|
| 37 |
+
pose_runner, detector_runner = get_inference_runners(
|
| 38 |
+
cfg,
|
| 39 |
+
snapshot_path=get_super_animal_snapshot_path(superanimal, pose_model),
|
| 40 |
+
detector_path=get_super_animal_snapshot_path(superanimal, detector),
|
| 41 |
+
max_individuals=MAX_INDIVIDUALS,
|
| 42 |
+
inference_cfg={"multithreading": {"enabled": False}},
|
| 43 |
+
)
|
| 44 |
+
_runners[superanimal] = {"pose": pose_runner,
|
| 45 |
+
"detector": detector_runner,
|
| 46 |
+
"bodyparts": list(cfg["metadata"]["bodyparts"]),
|
| 47 |
+
"lock": threading.Lock()} # runners are not thread-safe
|
| 48 |
+
return _runners[superanimal]
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
##########################################
|
| 52 |
+
def resize_max_side(img, max_size=MAX_IMAGE_SIZE):
|
| 53 |
+
scale = max_size / max(img.size)
|
| 54 |
+
if scale >= 1:
|
| 55 |
+
return img
|
| 56 |
+
return img.resize([int(x * scale) for x in img.size], PIL.Image.Resampling.LANCZOS)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
##########################################
|
| 60 |
+
def predict_superanimal(img_input,
|
| 61 |
+
superanimal,
|
| 62 |
+
bbox_likelihood_th,
|
| 63 |
+
kpts_likelihood_th,
|
| 64 |
+
full_image=False):
|
| 65 |
+
"""Detect animals and estimate their pose with a PyTorch SuperAnimal model.
|
| 66 |
+
|
| 67 |
+
Returns the (resized) RGB image the predictions refer to, the list of animals
|
| 68 |
+
as dicts {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3) array of x,y,llk
|
| 69 |
+
in image pixels, NaN below kpts_likelihood_th}, and the bodypart names.
|
| 70 |
+
"""
|
| 71 |
+
img = resize_max_side(img_input.convert("RGB"))
|
| 72 |
+
img_np = np.asarray(img)
|
| 73 |
+
runners = load_superanimal(superanimal)
|
| 74 |
+
|
| 75 |
+
with runners["lock"]:
|
| 76 |
+
if full_image:
|
| 77 |
+
# skip the detector and treat the whole image as one animal
|
| 78 |
+
h, w = img_np.shape[:2]
|
| 79 |
+
detections = {"bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
|
| 80 |
+
"bbox_scores": np.array([1.0], dtype=np.float32)}
|
| 81 |
+
else:
|
| 82 |
+
detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
|
| 83 |
+
keep = detections["bbox_scores"] >= bbox_likelihood_th
|
| 84 |
+
detections = {"bboxes": detections["bboxes"][keep],
|
| 85 |
+
"bbox_scores": detections["bbox_scores"][keep]}
|
| 86 |
+
|
| 87 |
+
if len(detections["bboxes"]) == 0:
|
| 88 |
+
return img, [], runners["bodyparts"]
|
| 89 |
+
|
| 90 |
+
predictions = runners["pose"].inference([(img_np, detections)])[0]
|
| 91 |
+
|
| 92 |
+
animals = []
|
| 93 |
+
# outputs are padded to MAX_INDIVIDUALS with -1
|
| 94 |
+
for kpts, (x, y, w, h), score in zip(predictions["bodyparts"],
|
| 95 |
+
predictions["bboxes"],
|
| 96 |
+
predictions["bbox_scores"]):
|
| 97 |
+
if score < 0:
|
| 98 |
+
continue
|
| 99 |
+
kpts = kpts.astype(float)
|
| 100 |
+
kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
|
| 101 |
+
animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)],
|
| 102 |
+
"kpts": kpts})
|
| 103 |
+
|
| 104 |
+
return img, animals, runners["bodyparts"]
|
requirements.txt
CHANGED
|
@@ -1,7 +1,7 @@
|
|
| 1 |
gradio
|
| 2 |
gitpython>=3.1.30
|
| 3 |
seaborn
|
| 4 |
-
deeplabcut[modelzoo,tf]
|
| 5 |
deeplabcut-live
|
| 6 |
ruamel.yaml==0.17.21
|
| 7 |
dlclibrary
|
|
|
|
| 1 |
gradio
|
| 2 |
gitpython>=3.1.30
|
| 3 |
seaborn
|
| 4 |
+
deeplabcut[modelzoo,tf]==3.0.2
|
| 5 |
deeplabcut-live
|
| 6 |
ruamel.yaml==0.17.21
|
| 7 |
dlclibrary
|
ui_utils.py
CHANGED
|
@@ -1,21 +1,27 @@
|
|
| 1 |
import gradio as gr
|
| 2 |
|
| 3 |
|
| 4 |
-
def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
|
| 5 |
# Input image
|
| 6 |
gr_image_input = gr.Image(type="pil", label="Input Image")
|
| 7 |
|
| 8 |
# Models
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
gr_mega_model_input = gr.Dropdown(
|
| 10 |
choices=md_models_list,
|
| 11 |
value="md_v5a",
|
| 12 |
type="value",
|
| 13 |
-
label="Select Detector model",
|
| 14 |
)
|
| 15 |
|
| 16 |
gr_dlc_model_input = gr.Dropdown(
|
| 17 |
choices=dlc_models_list,
|
| 18 |
-
value="
|
| 19 |
type="value",
|
| 20 |
label="Select DeepLabCut model",
|
| 21 |
)
|
|
@@ -23,7 +29,7 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
|
|
| 23 |
# Other inputs
|
| 24 |
gr_dlc_only_checkbox = gr.Checkbox(
|
| 25 |
value=False,
|
| 26 |
-
label="Run
|
| 27 |
)
|
| 28 |
|
| 29 |
gr_str_labels_checkbox = gr.Checkbox(
|
|
@@ -79,6 +85,7 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
|
|
| 79 |
|
| 80 |
return [
|
| 81 |
gr_image_input,
|
|
|
|
| 82 |
gr_mega_model_input,
|
| 83 |
gr_dlc_model_input,
|
| 84 |
gr_dlc_only_checkbox,
|
|
@@ -111,8 +118,9 @@ def gradio_description_and_examples():
|
|
| 111 |
|
| 112 |
examples = [[
|
| 113 |
"examples/dog.jpeg",
|
|
|
|
| 114 |
"md_v5a",
|
| 115 |
-
"
|
| 116 |
False,
|
| 117 |
True,
|
| 118 |
0.5,
|
|
|
|
| 1 |
import gradio as gr
|
| 2 |
|
| 3 |
|
| 4 |
+
def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
|
| 5 |
# Input image
|
| 6 |
gr_image_input = gr.Image(type="pil", label="Input Image")
|
| 7 |
|
| 8 |
# Models
|
| 9 |
+
gr_backend_input = gr.Radio(
|
| 10 |
+
choices=backends_list,
|
| 11 |
+
value=backends_list[0],
|
| 12 |
+
label="Select backend",
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
gr_mega_model_input = gr.Dropdown(
|
| 16 |
choices=md_models_list,
|
| 17 |
value="md_v5a",
|
| 18 |
type="value",
|
| 19 |
+
label="Select Detector model (TensorFlow legacy only)",
|
| 20 |
)
|
| 21 |
|
| 22 |
gr_dlc_model_input = gr.Dropdown(
|
| 23 |
choices=dlc_models_list,
|
| 24 |
+
value="superanimal_quadruped",
|
| 25 |
type="value",
|
| 26 |
label="Select DeepLabCut model",
|
| 27 |
)
|
|
|
|
| 29 |
# Other inputs
|
| 30 |
gr_dlc_only_checkbox = gr.Checkbox(
|
| 31 |
value=False,
|
| 32 |
+
label="Run DeepLabCut only, directly on input image?",
|
| 33 |
)
|
| 34 |
|
| 35 |
gr_str_labels_checkbox = gr.Checkbox(
|
|
|
|
| 85 |
|
| 86 |
return [
|
| 87 |
gr_image_input,
|
| 88 |
+
gr_backend_input,
|
| 89 |
gr_mega_model_input,
|
| 90 |
gr_dlc_model_input,
|
| 91 |
gr_dlc_only_checkbox,
|
|
|
|
| 118 |
|
| 119 |
examples = [[
|
| 120 |
"examples/dog.jpeg",
|
| 121 |
+
"PyTorch",
|
| 122 |
"md_v5a",
|
| 123 |
+
"superanimal_quadruped",
|
| 124 |
False,
|
| 125 |
True,
|
| 126 |
0.5,
|
viz_utils.py
CHANGED
|
@@ -185,4 +185,37 @@ def save_results_only_dlc(dlc_outputs,map_label_id_to_str,model,output_file = 'd
|
|
| 185 |
return output_file
|
| 186 |
|
| 187 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 188 |
###########################################
|
|
|
|
| 185 |
return output_file
|
| 186 |
|
| 187 |
|
| 188 |
+
def save_results_pytorch(animals, map_label_id_to_str, model, pose_model, detector, path_to_output_file = 'download_predictions.json'):
|
| 189 |
+
|
| 190 |
+
"""
|
| 191 |
+
Output PyTorch SuperAnimal predictions as json file (same layout as save_results_as_json)
|
| 192 |
+
|
| 193 |
+
animals: list of {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3)}, in image coords
|
| 194 |
+
detector: None if the detector was skipped (whole image used as one animal)
|
| 195 |
+
"""
|
| 196 |
+
info = {}
|
| 197 |
+
info['date'] = str(today)
|
| 198 |
+
info['backend'] = 'pytorch'
|
| 199 |
+
info['dlc_model'] = model
|
| 200 |
+
info['pose_model'] = pose_model
|
| 201 |
+
info['detector'] = detector
|
| 202 |
+
info['number_of_bb'] = len(animals)
|
| 203 |
+
labels = [n for n in map_label_id_to_str.values()]
|
| 204 |
+
|
| 205 |
+
for i, animal in enumerate(animals):
|
| 206 |
+
corner_x1, corner_y1, corner_x2, corner_y2, confidence = animal['bbox']
|
| 207 |
+
aux = {}
|
| 208 |
+
aux['corner_1'] = (corner_x1, corner_y1)
|
| 209 |
+
aux['corner_2'] = (corner_x2, corner_y2)
|
| 210 |
+
aux['confidence'] = confidence
|
| 211 |
+
aux['dlc_pred'] = dict(zip(labels, [[float(v) for v in kpt] for kpt in animal['kpts']]))
|
| 212 |
+
info['bb_' + str(i)] = aux
|
| 213 |
+
|
| 214 |
+
with open(path_to_output_file, 'w') as f:
|
| 215 |
+
json.dump(info, f, indent=1)
|
| 216 |
+
print('Output file saved at {}'.format(path_to_output_file))
|
| 217 |
+
|
| 218 |
+
return path_to_output_file
|
| 219 |
+
|
| 220 |
+
|
| 221 |
###########################################
|