Add PyTorch SuperAnimal backend and refresh the Space
#14
by C-Achard - opened
- .gitignore +14 -1
- .pre-commit-config.yaml +28 -0
- README.md +2 -2
- app.py +261 -140
- detection_utils.py +43 -61
- dlc_utils.py +16 -20
- pre-requirements.txt +1 -0
- pyproject.toml +33 -0
- pytorch_utils.py +118 -0
- requirements.txt +6 -3
- save_results.py +0 -56
- ui_utils.py +162 -60
- viz_utils.py +222 -142
.gitignore
CHANGED
|
@@ -1,3 +1,16 @@
|
|
| 1 |
# Byte-compiled / optimized / DLL files
|
| 2 |
__pycache__/
|
| 3 |
-
model/__pycache__/
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
# Byte-compiled / optimized / DLL files
|
| 2 |
__pycache__/
|
| 3 |
+
model/__pycache__/
|
| 4 |
+
.vscode/settings.json
|
| 5 |
+
|
| 6 |
+
# Local environment and tool caches
|
| 7 |
+
.venv/
|
| 8 |
+
.ruff_cache/
|
| 9 |
+
|
| 10 |
+
# Written by the app at runtime
|
| 11 |
+
.gradio/
|
| 12 |
+
download_predictions.json
|
| 13 |
+
dowload_predictions_dlc.json
|
| 14 |
+
download_annotated.png
|
| 15 |
+
# TF (legacy) models downloaded on first use; DLC_models/readme.md stays tracked
|
| 16 |
+
DLC_models/*/
|
.pre-commit-config.yaml
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Adapted from DeepLabCut's .pre-commit-config.yaml.
|
| 2 |
+
repos:
|
| 3 |
+
- repo: https://github.com/pre-commit/pre-commit-hooks
|
| 4 |
+
rev: v6.0.0
|
| 5 |
+
hooks:
|
| 6 |
+
- id: check-added-large-files
|
| 7 |
+
- id: check-yaml
|
| 8 |
+
- id: check-toml
|
| 9 |
+
- id: check-merge-conflict
|
| 10 |
+
- id: end-of-file-fixer
|
| 11 |
+
- id: trailing-whitespace
|
| 12 |
+
|
| 13 |
+
- repo: https://github.com/tox-dev/pyproject-fmt
|
| 14 |
+
rev: v2.19.0
|
| 15 |
+
hooks:
|
| 16 |
+
- id: pyproject-fmt
|
| 17 |
+
|
| 18 |
+
- repo: https://github.com/abravalheri/validate-pyproject
|
| 19 |
+
rev: v0.25
|
| 20 |
+
hooks:
|
| 21 |
+
- id: validate-pyproject
|
| 22 |
+
|
| 23 |
+
- repo: https://github.com/astral-sh/ruff-pre-commit
|
| 24 |
+
rev: v0.15.6
|
| 25 |
+
hooks:
|
| 26 |
+
- id: ruff-check
|
| 27 |
+
args: [--fix, --unsafe-fixes]
|
| 28 |
+
- id: ruff-format
|
README.md
CHANGED
|
@@ -5,9 +5,9 @@ colorFrom: blue
|
|
| 5 |
colorTo: purple
|
| 6 |
sdk: gradio
|
| 7 |
python_version: 3.12
|
| 8 |
-
sdk_version: 6.
|
| 9 |
app_file: app.py
|
| 10 |
pinned: false
|
| 11 |
---
|
| 12 |
|
| 13 |
-
Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
|
|
|
|
| 5 |
colorTo: purple
|
| 6 |
sdk: gradio
|
| 7 |
python_version: 3.12
|
| 8 |
+
sdk_version: 6.29.0
|
| 9 |
app_file: app.py
|
| 10 |
pinned: false
|
| 11 |
---
|
| 12 |
|
| 13 |
+
Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
|
app.py
CHANGED
|
@@ -1,193 +1,314 @@
|
|
| 1 |
-
# Adapted from https://huggingface.co/spaces/hlydecker/MegaDetector_v5
|
| 2 |
# Adapted from https://huggingface.co/spaces/sofmi/MegaDetector_DLClive/blob/main/app.py
|
| 3 |
-
# Adapted from https://huggingface.co/spaces/Neslihan/megadetector_dlcmodels/blob/main/app.py
|
| 4 |
# Adapted from https://huggingface.co/spaces/DeepLabCut/MegaDetector_DeepLabCut
|
| 5 |
|
| 6 |
import os
|
| 7 |
-
import
|
| 8 |
-
import numpy as np
|
| 9 |
-
from matplotlib import cm
|
| 10 |
-
import gradio as gr
|
| 11 |
-
import deeplabcut
|
| 12 |
-
import dlclibrary
|
| 13 |
-
import dlclive
|
| 14 |
-
# import transformers
|
| 15 |
-
|
| 16 |
-
from PIL import Image, ImageColor, ImageFont, ImageDraw
|
| 17 |
-
import requests
|
| 18 |
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 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
|
| 25 |
from dlclibrary.dlcmodelzoo.modelzoo_download import (
|
| 26 |
download_huggingface_model,
|
| 27 |
-
MODELOPTIONS,
|
| 28 |
)
|
| 29 |
-
from dlclive import
|
| 30 |
|
|
|
|
|
|
|
| 31 |
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
|
| 37 |
-
#
|
| 38 |
-
|
| 39 |
-
|
|
|
|
| 40 |
|
| 41 |
# megadetector and dlc model look up
|
| 42 |
-
MD_models_dict = {
|
| 43 |
-
|
|
|
|
|
|
|
| 44 |
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
|
|
|
|
|
|
|
|
|
| 49 |
|
| 50 |
|
| 51 |
#####################################################
|
| 52 |
-
def
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
|
| 65 |
if not flag_dlc_only:
|
| 66 |
-
############################################################
|
| 67 |
# ### Run Megadetector
|
| 68 |
-
md_results = predict_md(
|
| 69 |
-
|
| 70 |
-
|
|
|
|
|
|
|
| 71 |
|
| 72 |
################################################################
|
| 73 |
-
# Obtain animal crops
|
| 74 |
-
list_crops = crop_animal_detections(img_input,
|
| 75 |
-
md_results,
|
| 76 |
-
bbox_likelihood_th)
|
| 77 |
|
| 78 |
############################################################
|
| 79 |
|
| 80 |
-
## Get DLC model and label map
|
| 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 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
# Run DLC and visualize results
|
| 102 |
-
dlc_proc = Processor()
|
| 103 |
|
| 104 |
# if required: ignore MD crops and run DLC on full image [mostly for testing]
|
| 105 |
if flag_dlc_only:
|
| 106 |
# compute kpts on input img
|
| 107 |
-
list_kpts_per_crop = predict_dlc([np.asarray(img_input)],
|
| 108 |
-
kpts_likelihood_th,
|
| 109 |
-
path_to_DLCmodel,
|
| 110 |
-
dlc_proc)
|
| 111 |
# draw kpts on input img #fix!
|
| 112 |
-
draw_keypoints_on_image(
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 125 |
|
| 126 |
else:
|
| 127 |
# Compute kpts for each crop
|
| 128 |
-
list_kpts_per_crop = predict_dlc(list_crops,
|
| 129 |
-
|
| 130 |
-
path_to_DLCmodel,
|
| 131 |
-
dlc_proc)
|
| 132 |
-
|
| 133 |
# resize input image to match megadetector output
|
| 134 |
-
img_background = img_input.resize((md_results.ims[0].shape[1],
|
| 135 |
-
md_results.ims[0].shape[0]))
|
| 136 |
-
|
| 137 |
-
# draw keypoints on each crop and paste to background img
|
| 138 |
-
for ic, (np_crop, kpts_crop) in enumerate(zip(list_crops,
|
| 139 |
-
list_kpts_per_crop)):
|
| 140 |
|
|
|
|
|
|
|
| 141 |
img_crop = Image.fromarray(np_crop)
|
| 142 |
|
| 143 |
# Draw keypts on crop
|
| 144 |
-
draw_keypoints_on_image(
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 153 |
|
| 154 |
# Paste crop in original image
|
| 155 |
-
img_background.paste(img_crop,
|
| 156 |
-
box = tuple([int(t) for t in md_results.xyxy[0][ic,:2]]))
|
| 157 |
|
| 158 |
# Plot bbox
|
| 159 |
-
bb_per_animal =
|
| 160 |
-
pred = md_results.xyxy[0].tolist()[ic][4]
|
| 161 |
-
if bbox_likelihood_th < pred:
|
| 162 |
-
draw_bbox_w_text(img_background,
|
| 163 |
-
bb_per_animal,
|
| 164 |
-
font_size=font_size) # TODO: add selectable color for bbox?
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
# Save detection results as json
|
| 168 |
-
download_file = save_results_as_json(md_results,list_kpts_per_crop,map_label_id_to_str, bbox_likelihood_th,dlc_model_input_str,mega_model_input)
|
| 169 |
|
| 170 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 171 |
|
|
|
|
|
|
|
|
|
|
| 172 |
|
| 173 |
|
| 174 |
#########################################################
|
| 175 |
# Define user interface and launch
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Adapted from https://huggingface.co/spaces/hlydecker/MegaDetector_v5
|
| 2 |
# Adapted from https://huggingface.co/spaces/sofmi/MegaDetector_DLClive/blob/main/app.py
|
| 3 |
+
# Adapted from https://huggingface.co/spaces/Neslihan/megadetector_dlcmodels/blob/main/app.py
|
| 4 |
# Adapted from https://huggingface.co/spaces/DeepLabCut/MegaDetector_DeepLabCut
|
| 5 |
|
| 6 |
import os
|
| 7 |
+
import threading
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
|
| 9 |
+
import gradio as gr
|
| 10 |
+
import numpy as np
|
| 11 |
+
import yaml
|
|
|
|
|
|
|
|
|
|
| 12 |
from dlclibrary.dlcmodelzoo.modelzoo_download import (
|
| 13 |
download_huggingface_model,
|
|
|
|
| 14 |
)
|
| 15 |
+
from dlclive import Processor
|
| 16 |
|
| 17 |
+
# import transformers
|
| 18 |
+
from PIL import Image
|
| 19 |
|
| 20 |
+
from detection_utils import crop_animal_detections, predict_md
|
| 21 |
+
from dlc_utils import predict_dlc
|
| 22 |
+
from pytorch_utils import PYTORCH_MODELS, load_superanimal, predict_superanimal
|
| 23 |
+
from ui_utils import (
|
| 24 |
+
confidence_legend_html,
|
| 25 |
+
dlc_theme,
|
| 26 |
+
gradio_description_and_examples,
|
| 27 |
+
gradio_inputs_for_MD_DLC,
|
| 28 |
+
gradio_outputs_for_MD_DLC,
|
| 29 |
+
)
|
| 30 |
+
from viz_utils import (
|
| 31 |
+
draw_bbox_w_text,
|
| 32 |
+
draw_keypoints_on_image,
|
| 33 |
+
keypoint_confidence_rows,
|
| 34 |
+
save_annotated_image,
|
| 35 |
+
save_results_as_json,
|
| 36 |
+
save_results_only_dlc,
|
| 37 |
+
save_results_pytorch,
|
| 38 |
+
)
|
| 39 |
|
| 40 |
+
# TESTING (passes) download the SuperAnimal models:
|
| 41 |
+
# model = 'superanimal_topviewmouse'
|
| 42 |
+
# train_dir = 'DLC_models/sa-tvm'
|
| 43 |
+
# download_huggingface_model(model, train_dir)
|
| 44 |
|
| 45 |
# megadetector and dlc model look up
|
| 46 |
+
MD_models_dict = {
|
| 47 |
+
"md_v5a": "MD_models/md_v5a.0.0.pt", #
|
| 48 |
+
"md_v5b": "MD_models/md_v5b.0.0.pt",
|
| 49 |
+
}
|
| 50 |
|
| 51 |
+
BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
|
| 52 |
+
|
| 53 |
+
# TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
|
| 54 |
+
DLC_models_dict = {
|
| 55 |
+
"superanimal_topviewmouse": ("superanimal_topviewmouse_dlcrnet", "DLC_models/sa-tvm"),
|
| 56 |
+
"superanimal_quadruped": ("superanimal_quadruped_dlcrnet", "DLC_models/sa-q"),
|
| 57 |
+
}
|
| 58 |
|
| 59 |
|
| 60 |
#####################################################
|
| 61 |
+
def finalize_outputs(img_output, download_file, kpts_per_animal, map_label_id_to_str, color_by_confidence, colormap):
|
| 62 |
+
annotated_file = save_annotated_image(img_output)
|
| 63 |
+
confidence_rows = keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str)
|
| 64 |
+
legend = confidence_legend_html(colormap) if color_by_confidence else ""
|
| 65 |
+
return img_output, legend, download_file, annotated_file, confidence_rows
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
#####################################################
|
| 69 |
+
def predict_pipeline_pytorch(
|
| 70 |
+
img_input,
|
| 71 |
+
superanimal,
|
| 72 |
+
flag_dlc_only,
|
| 73 |
+
flag_show_str_labels,
|
| 74 |
+
bbox_likelihood_th,
|
| 75 |
+
kpts_likelihood_th,
|
| 76 |
+
font_style,
|
| 77 |
+
font_size,
|
| 78 |
+
keypt_color,
|
| 79 |
+
marker_size,
|
| 80 |
+
flag_color_by_confidence,
|
| 81 |
+
colormap,
|
| 82 |
+
bbox_color,
|
| 83 |
+
):
|
| 84 |
+
# detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
|
| 85 |
+
img_output, animals, bodyparts = predict_superanimal(
|
| 86 |
+
img_input, superanimal, bbox_likelihood_th, kpts_likelihood_th, full_image=flag_dlc_only
|
| 87 |
+
)
|
| 88 |
+
map_label_id_to_str = dict(enumerate(bodyparts))
|
| 89 |
+
|
| 90 |
+
for animal in animals:
|
| 91 |
+
draw_keypoints_on_image(
|
| 92 |
+
img_output,
|
| 93 |
+
animal["kpts"],
|
| 94 |
+
map_label_id_to_str,
|
| 95 |
+
flag_show_str_labels,
|
| 96 |
+
use_normalized_coordinates=False,
|
| 97 |
+
font_style=font_style,
|
| 98 |
+
font_size=font_size,
|
| 99 |
+
keypt_color=keypt_color,
|
| 100 |
+
marker_size=marker_size,
|
| 101 |
+
color_by_confidence=flag_color_by_confidence,
|
| 102 |
+
colormap=colormap,
|
| 103 |
+
)
|
| 104 |
+
if not flag_dlc_only:
|
| 105 |
+
draw_bbox_w_text(img_output, animal["bbox"], font_size=font_size, bbox_color=bbox_color)
|
| 106 |
+
|
| 107 |
+
pose_model, detector = PYTORCH_MODELS[superanimal]
|
| 108 |
+
download_file = save_results_pytorch(
|
| 109 |
+
animals,
|
| 110 |
+
map_label_id_to_str,
|
| 111 |
+
superanimal,
|
| 112 |
+
pose_model,
|
| 113 |
+
None if flag_dlc_only else detector,
|
| 114 |
+
image_size=img_input.size,
|
| 115 |
+
annotated_size=img_output.size,
|
| 116 |
+
)
|
| 117 |
+
return finalize_outputs(
|
| 118 |
+
img_output,
|
| 119 |
+
download_file,
|
| 120 |
+
[animal["kpts"] for animal in animals],
|
| 121 |
+
map_label_id_to_str,
|
| 122 |
+
flag_color_by_confidence,
|
| 123 |
+
colormap,
|
| 124 |
+
)
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
#####################################################
|
| 128 |
+
def predict_pipeline(
|
| 129 |
+
img_input,
|
| 130 |
+
backend,
|
| 131 |
+
mega_model_input,
|
| 132 |
+
dlc_model_input_str,
|
| 133 |
+
flag_dlc_only,
|
| 134 |
+
flag_show_str_labels,
|
| 135 |
+
bbox_likelihood_th,
|
| 136 |
+
kpts_likelihood_th,
|
| 137 |
+
font_style,
|
| 138 |
+
font_size,
|
| 139 |
+
keypt_color,
|
| 140 |
+
marker_size,
|
| 141 |
+
flag_color_by_confidence,
|
| 142 |
+
colormap,
|
| 143 |
+
bbox_color,
|
| 144 |
+
):
|
| 145 |
+
|
| 146 |
+
if backend == "PyTorch":
|
| 147 |
+
return predict_pipeline_pytorch(
|
| 148 |
+
img_input,
|
| 149 |
+
dlc_model_input_str,
|
| 150 |
+
flag_dlc_only,
|
| 151 |
+
flag_show_str_labels,
|
| 152 |
+
bbox_likelihood_th,
|
| 153 |
+
kpts_likelihood_th,
|
| 154 |
+
font_style,
|
| 155 |
+
font_size,
|
| 156 |
+
keypt_color,
|
| 157 |
+
marker_size,
|
| 158 |
+
flag_color_by_confidence,
|
| 159 |
+
colormap,
|
| 160 |
+
bbox_color,
|
| 161 |
+
)
|
| 162 |
+
|
| 163 |
+
# TensorFlow (legacy): MegaDetector crops + DLCLive
|
| 164 |
+
dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
|
| 165 |
|
| 166 |
if not flag_dlc_only:
|
| 167 |
+
############################################################
|
| 168 |
# ### Run Megadetector
|
| 169 |
+
md_results = predict_md(
|
| 170 |
+
img_input,
|
| 171 |
+
MD_models_dict[mega_model_input], # mega_model_input,
|
| 172 |
+
size=640,
|
| 173 |
+
) # Image.fromarray(results.imgs[0])
|
| 174 |
|
| 175 |
################################################################
|
| 176 |
+
# Obtain animal crops (and their bboxes) with confidence above th
|
| 177 |
+
list_crops, list_bboxes = crop_animal_detections(img_input, md_results, bbox_likelihood_th)
|
|
|
|
|
|
|
| 178 |
|
| 179 |
############################################################
|
| 180 |
|
| 181 |
+
## Get DLC model and label map
|
| 182 |
+
|
| 183 |
# If model is found: do not download (previous execution is likely within same day)
|
| 184 |
# TODO: can we ask the user whether to reload dlc model if a directory is found?
|
| 185 |
+
path_to_DLCmodel = dlc_model_dir
|
| 186 |
+
if not (os.path.isdir(dlc_model_dir) and len(os.listdir(dlc_model_dir)) > 0):
|
| 187 |
+
download_huggingface_model(dlc_model_name, path_to_DLCmodel)
|
|
|
|
|
|
|
|
|
|
| 188 |
|
| 189 |
# extract map label ids to strings
|
| 190 |
+
pose_cfg_path = os.path.join(dlc_model_dir, "pose_cfg.yaml")
|
| 191 |
+
with open(pose_cfg_path) as stream:
|
| 192 |
+
pose_cfg_dict = yaml.safe_load(stream)
|
| 193 |
+
map_label_id_to_str = dict(
|
| 194 |
+
[
|
| 195 |
+
(k, v)
|
| 196 |
+
for k, v in zip(
|
| 197 |
+
[
|
| 198 |
+
el[0] for el in pose_cfg_dict["all_joints"]
|
| 199 |
+
], # pose_cfg_dict['all_joints'] is a list of one-element lists,
|
| 200 |
+
pose_cfg_dict["all_joints_names"],
|
| 201 |
+
strict=True,
|
| 202 |
+
)
|
| 203 |
+
]
|
| 204 |
+
)
|
| 205 |
+
|
| 206 |
+
##############################################################
|
| 207 |
# Run DLC and visualize results
|
| 208 |
+
dlc_proc = Processor() # TODO: update deeplabcut.video_inference_superanimal() once merged
|
| 209 |
|
| 210 |
# if required: ignore MD crops and run DLC on full image [mostly for testing]
|
| 211 |
if flag_dlc_only:
|
| 212 |
# compute kpts on input img
|
| 213 |
+
list_kpts_per_crop = predict_dlc([np.asarray(img_input)], kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
|
|
|
|
|
|
|
|
|
|
| 214 |
# draw kpts on input img #fix!
|
| 215 |
+
draw_keypoints_on_image(
|
| 216 |
+
img_input,
|
| 217 |
+
list_kpts_per_crop[0], # a numpy array with shape [num_keypoints, 2].
|
| 218 |
+
map_label_id_to_str,
|
| 219 |
+
flag_show_str_labels,
|
| 220 |
+
use_normalized_coordinates=False,
|
| 221 |
+
font_style=font_style,
|
| 222 |
+
font_size=font_size,
|
| 223 |
+
keypt_color=keypt_color,
|
| 224 |
+
marker_size=marker_size,
|
| 225 |
+
color_by_confidence=flag_color_by_confidence,
|
| 226 |
+
colormap=colormap,
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
donw_file = save_results_only_dlc(
|
| 230 |
+
list_kpts_per_crop[0], map_label_id_to_str, dlc_model_name, image_size=img_input.size
|
| 231 |
+
)
|
| 232 |
+
|
| 233 |
+
return finalize_outputs(
|
| 234 |
+
img_input, donw_file, [list_kpts_per_crop[0]], map_label_id_to_str, flag_color_by_confidence, colormap
|
| 235 |
+
)
|
| 236 |
|
| 237 |
else:
|
| 238 |
# Compute kpts for each crop
|
| 239 |
+
list_kpts_per_crop = predict_dlc(list_crops, kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
|
| 240 |
+
|
|
|
|
|
|
|
|
|
|
| 241 |
# resize input image to match megadetector output
|
| 242 |
+
img_background = img_input.resize((md_results.ims[0].shape[1], md_results.ims[0].shape[0]))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 243 |
|
| 244 |
+
# draw keypoints on each crop and paste to background img
|
| 245 |
+
for np_crop, kpts_crop, bb_per_animal in zip(list_crops, list_kpts_per_crop, list_bboxes, strict=True):
|
| 246 |
img_crop = Image.fromarray(np_crop)
|
| 247 |
|
| 248 |
# Draw keypts on crop
|
| 249 |
+
draw_keypoints_on_image(
|
| 250 |
+
img_crop,
|
| 251 |
+
kpts_crop, # a numpy array with shape [num_keypoints, 2].
|
| 252 |
+
map_label_id_to_str,
|
| 253 |
+
flag_show_str_labels,
|
| 254 |
+
use_normalized_coordinates=False, # if True, then I should use md_results.xyxyn for list_kpts_crop
|
| 255 |
+
font_style=font_style,
|
| 256 |
+
font_size=font_size,
|
| 257 |
+
keypt_color=keypt_color,
|
| 258 |
+
marker_size=marker_size,
|
| 259 |
+
color_by_confidence=flag_color_by_confidence,
|
| 260 |
+
colormap=colormap,
|
| 261 |
+
)
|
| 262 |
|
| 263 |
# Paste crop in original image
|
| 264 |
+
img_background.paste(img_crop, box=tuple([int(t) for t in bb_per_animal[:2]]))
|
|
|
|
| 265 |
|
| 266 |
# Plot bbox
|
| 267 |
+
draw_bbox_w_text(img_background, bb_per_animal, font_size=font_size, bbox_color=bbox_color)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 268 |
|
| 269 |
+
# Save detection results as json
|
| 270 |
+
download_file = save_results_as_json(
|
| 271 |
+
md_results,
|
| 272 |
+
list_kpts_per_crop,
|
| 273 |
+
list_bboxes,
|
| 274 |
+
map_label_id_to_str,
|
| 275 |
+
dlc_model_name,
|
| 276 |
+
mega_model_input,
|
| 277 |
+
image_size=img_input.size,
|
| 278 |
+
)
|
| 279 |
|
| 280 |
+
return finalize_outputs(
|
| 281 |
+
img_background, download_file, list_kpts_per_crop, map_label_id_to_str, flag_color_by_confidence, colormap
|
| 282 |
+
)
|
| 283 |
|
| 284 |
|
| 285 |
#########################################################
|
| 286 |
# Define user interface and launch
|
| 287 |
+
[gr_title, gr_description, examples] = gradio_description_and_examples()
|
| 288 |
+
|
| 289 |
+
with gr.Blocks(title=gr_title) as demo:
|
| 290 |
+
gr.Markdown(f"# {gr_title}\n{gr_description}")
|
| 291 |
+
with gr.Row():
|
| 292 |
+
with gr.Column():
|
| 293 |
+
inputs = gradio_inputs_for_MD_DLC(BACKENDS, list(MD_models_dict.keys()), list(DLC_models_dict.keys()))
|
| 294 |
+
run_button = gr.Button("Run", variant="primary")
|
| 295 |
+
with gr.Column():
|
| 296 |
+
outputs = gradio_outputs_for_MD_DLC()
|
| 297 |
+
|
| 298 |
+
# the MegaDetector choice only applies to the TensorFlow (legacy) backend
|
| 299 |
+
gr_backend_input, gr_mega_model_input = inputs[1], inputs[2]
|
| 300 |
+
gr_backend_input.change(
|
| 301 |
+
lambda backend: gr.update(visible=backend != "PyTorch"), inputs=gr_backend_input, outputs=gr_mega_model_input
|
| 302 |
+
)
|
| 303 |
+
|
| 304 |
+
run_button.click(predict_pipeline, inputs=inputs, outputs=outputs, api_name="predict")
|
| 305 |
+
|
| 306 |
+
# cached on first click, so a failing download cannot block startup
|
| 307 |
+
gr.Examples(examples, inputs=inputs, outputs=outputs, fn=predict_pipeline, cache_examples=True, cache_mode="lazy")
|
| 308 |
+
|
| 309 |
+
# download and build the default model while the app starts; a request arriving
|
| 310 |
+
# earlier waits on the same lock instead of downloading again
|
| 311 |
+
threading.Thread(target=load_superanimal, args=("superanimal_quadruped",), daemon=True).start()
|
| 312 |
+
|
| 313 |
+
demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
|
| 314 |
+
demo.launch(theme=dlc_theme())
|
detection_utils.py
CHANGED
|
@@ -1,92 +1,74 @@
|
|
| 1 |
-
|
| 2 |
-
from tkinter import W
|
| 3 |
-
import gradio as gr
|
| 4 |
-
from matplotlib import cm
|
| 5 |
-
import torch
|
| 6 |
-
import torchvision
|
| 7 |
-
import matplotlib
|
| 8 |
-
import PIL
|
| 9 |
-
from PIL import Image, ImageColor, ImageFont, ImageDraw
|
| 10 |
-
import numpy as np
|
| 11 |
import math
|
| 12 |
|
|
|
|
|
|
|
|
|
|
| 13 |
|
| 14 |
-
import yaml
|
| 15 |
-
import pdb
|
| 16 |
|
| 17 |
############################################
|
| 18 |
# Predict detections with MegaDetector v5a model
|
| 19 |
-
def predict_md(
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
|
|
|
|
|
|
| 23 |
# resize image
|
| 24 |
-
g =
|
| 25 |
-
im = im.resize((int(x * g) for x in im.size),
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 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 |
-
|
| 49 |
-
|
| 50 |
-
|
|
|
|
| 51 |
|
| 52 |
|
| 53 |
##########################################
|
| 54 |
-
def crop_animal_detections(img_in,
|
| 55 |
-
yolo_results,
|
| 56 |
-
likelihood_th):
|
| 57 |
|
| 58 |
## Extract animal crops
|
| 59 |
-
list_labels_as_str = [i for i in yolo_results.names.values()] # ['animal', 'person', 'vehicle']
|
| 60 |
list_np_animal_crops = []
|
|
|
|
| 61 |
|
| 62 |
# image to crop (scale as input for megadetector)
|
| 63 |
-
img_in = img_in.resize((yolo_results.ims[0].shape[1],
|
| 64 |
-
|
| 65 |
-
# for every detection in the img
|
| 66 |
for det_array in yolo_results.xyxy:
|
| 67 |
-
|
| 68 |
# for every detection
|
| 69 |
for j in range(det_array.shape[0]):
|
| 70 |
-
|
| 71 |
# compute coords around bbox rounded to the nearest integer (for pasting later)
|
| 72 |
-
xmin_rd = int(math.floor(det_array[j,0]))
|
| 73 |
-
ymin_rd = int(math.floor(det_array[j,1]))
|
| 74 |
|
| 75 |
-
xmax_rd = int(math.ceil(det_array[j,2]))
|
| 76 |
-
ymax_rd = int(math.ceil(det_array[j,3]))
|
| 77 |
|
| 78 |
-
pred_llk = det_array[j,4]
|
| 79 |
-
pred_label = det_array[j,5]
|
| 80 |
# keep animal crops above threshold
|
| 81 |
-
if (pred_label == list_labels_as_str.index(
|
| 82 |
-
(pred_llk >= likelihood_th):
|
| 83 |
area = (xmin_rd, ymin_rd, xmax_rd, ymax_rd)
|
| 84 |
|
| 85 |
-
#pdb.set_trace()
|
| 86 |
-
crop = img_in.crop(area)
|
| 87 |
crop_np = np.asarray(crop)
|
| 88 |
|
| 89 |
# add to list
|
| 90 |
list_np_animal_crops.append(crop_np)
|
|
|
|
| 91 |
|
| 92 |
-
return list_np_animal_crops
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import math
|
| 2 |
|
| 3 |
+
import numpy as np
|
| 4 |
+
import PIL
|
| 5 |
+
import torch
|
| 6 |
|
|
|
|
|
|
|
| 7 |
|
| 8 |
############################################
|
| 9 |
# Predict detections with MegaDetector v5a model
|
| 10 |
+
def predict_md(
|
| 11 |
+
im,
|
| 12 |
+
megadetector_model, # Megadet_Models[mega_model_input]
|
| 13 |
+
size=640,
|
| 14 |
+
):
|
| 15 |
+
|
| 16 |
# resize image
|
| 17 |
+
g = size / max(im.size) # multipl factor to make max size of the image equal to input size
|
| 18 |
+
im = im.resize((int(x * g) for x in im.size), PIL.Image.Resampling.LANCZOS) # resize
|
| 19 |
+
# device: yolov5's select_device expects a CUDA index ('0') or 'cpu', not 'cuda'
|
| 20 |
+
md_device = "0" if torch.cuda.is_available() else "cpu"
|
| 21 |
+
|
| 22 |
+
# megadetector
|
| 23 |
+
MD_model = torch.hub.load(
|
| 24 |
+
"ultralytics/yolov5", # repo_or_dir
|
| 25 |
+
"custom", # model
|
| 26 |
+
megadetector_model, # args for callable model
|
| 27 |
+
skip_validation=True, # avoid GitHub API rate limit (403)
|
| 28 |
+
device=md_device,
|
| 29 |
+
trust_repo=True,
|
| 30 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
## detect objects
|
| 33 |
+
# vars(results).keys(): imgs, pred, names, files, times, xyxy, xywh, xyxyn, xywhn, n, t, s
|
| 34 |
+
results = MD_model(im)
|
| 35 |
+
|
| 36 |
+
return results
|
| 37 |
|
| 38 |
|
| 39 |
##########################################
|
| 40 |
+
def crop_animal_detections(img_in, yolo_results, likelihood_th):
|
|
|
|
|
|
|
| 41 |
|
| 42 |
## Extract animal crops
|
| 43 |
+
list_labels_as_str = [i for i in yolo_results.names.values()] # ['animal', 'person', 'vehicle']
|
| 44 |
list_np_animal_crops = []
|
| 45 |
+
list_animal_bboxes = [] # detection rows [x1,y1,x2,y2,conf,label] matching each crop
|
| 46 |
|
| 47 |
# image to crop (scale as input for megadetector)
|
| 48 |
+
img_in = img_in.resize((yolo_results.ims[0].shape[1], yolo_results.ims[0].shape[0]))
|
| 49 |
+
# for every detection in the img
|
|
|
|
| 50 |
for det_array in yolo_results.xyxy:
|
|
|
|
| 51 |
# for every detection
|
| 52 |
for j in range(det_array.shape[0]):
|
|
|
|
| 53 |
# compute coords around bbox rounded to the nearest integer (for pasting later)
|
| 54 |
+
xmin_rd = int(math.floor(det_array[j, 0])) # int() should suffice?
|
| 55 |
+
ymin_rd = int(math.floor(det_array[j, 1]))
|
| 56 |
|
| 57 |
+
xmax_rd = int(math.ceil(det_array[j, 2]))
|
| 58 |
+
ymax_rd = int(math.ceil(det_array[j, 3]))
|
| 59 |
|
| 60 |
+
pred_llk = det_array[j, 4]
|
| 61 |
+
pred_label = det_array[j, 5]
|
| 62 |
# keep animal crops above threshold
|
| 63 |
+
if (pred_label == list_labels_as_str.index("animal")) and (pred_llk >= likelihood_th):
|
|
|
|
| 64 |
area = (xmin_rd, ymin_rd, xmax_rd, ymax_rd)
|
| 65 |
|
| 66 |
+
# pdb.set_trace()
|
| 67 |
+
crop = img_in.crop(area) # Image.fromarray(img_in).crop(area)
|
| 68 |
crop_np = np.asarray(crop)
|
| 69 |
|
| 70 |
# add to list
|
| 71 |
list_np_animal_crops.append(crop_np)
|
| 72 |
+
list_animal_bboxes.append(det_array[j, :].tolist())
|
| 73 |
|
| 74 |
+
return list_np_animal_crops, list_animal_bboxes
|
dlc_utils.py
CHANGED
|
@@ -1,32 +1,28 @@
|
|
| 1 |
-
import deeplabcut
|
| 2 |
-
from tkinter import W
|
| 3 |
-
import gradio as gr
|
| 4 |
import numpy as np
|
| 5 |
-
from dlclive import DLCLive
|
| 6 |
|
| 7 |
|
| 8 |
##########################################
|
| 9 |
-
def predict_dlc(list_np_crops,
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
# run dlc thru list of crops
|
| 15 |
dlc_live = DLCLive(dlc_model_folder, processor=dlc_proc)
|
| 16 |
dlc_live.init_inference(list_np_crops[0])
|
| 17 |
|
| 18 |
list_kpts_per_crop = []
|
| 19 |
-
all_kypts = []
|
| 20 |
-
np_aux = np.empty((1,3)) # can I avoid hardcoding here?
|
| 21 |
for crop in list_np_crops:
|
| 22 |
-
|
| 23 |
-
keypts_xyp = dlc_live.get_pose(crop) # third column is llk!
|
| 24 |
# set kpts below threhsold to nan
|
| 25 |
-
|
| 26 |
-
#
|
| 27 |
-
keypts_xyp[keypts_xyp[:,-1] < kpts_likelihood_th,:] = np_aux.fill(np.nan)
|
| 28 |
-
# add kpts of this crop to list
|
| 29 |
list_kpts_per_crop.append(keypts_xyp)
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
return list_kpts_per_crop
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import numpy as np
|
| 2 |
+
from dlclive import DLCLive
|
| 3 |
|
| 4 |
|
| 5 |
##########################################
|
| 6 |
+
def predict_dlc(list_np_crops, kpts_likelihood_th, dlc_model_folder, dlc_proc):
|
| 7 |
+
|
| 8 |
+
# no animal detected: nothing to run
|
| 9 |
+
if len(list_np_crops) == 0:
|
| 10 |
+
return []
|
| 11 |
+
|
| 12 |
+
# DLCLive always converts 3-channel frames BGR->RGB, but our crops are
|
| 13 |
+
# already RGB (PIL), so pass them as BGR for the model to see RGB
|
| 14 |
+
list_np_crops = [np.ascontiguousarray(crop[..., ::-1]) for crop in list_np_crops]
|
| 15 |
+
|
| 16 |
# run dlc thru list of crops
|
| 17 |
dlc_live = DLCLive(dlc_model_folder, processor=dlc_proc)
|
| 18 |
dlc_live.init_inference(list_np_crops[0])
|
| 19 |
|
| 20 |
list_kpts_per_crop = []
|
|
|
|
|
|
|
| 21 |
for crop in list_np_crops:
|
| 22 |
+
keypts_xyp = dlc_live.get_pose(crop) # third column is llk!
|
|
|
|
| 23 |
# set kpts below threhsold to nan
|
| 24 |
+
keypts_xyp[keypts_xyp[:, -1] < kpts_likelihood_th, :] = np.nan
|
| 25 |
+
# add kpts of this crop to list
|
|
|
|
|
|
|
| 26 |
list_kpts_per_crop.append(keypts_xyp)
|
| 27 |
+
|
| 28 |
+
return list_kpts_per_crop
|
|
|
pre-requirements.txt
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
pip>=26.2
|
pyproject.toml
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[project]
|
| 2 |
+
name = "deeplabcut-modelzoo-superanimals"
|
| 3 |
+
version = "0.1.0"
|
| 4 |
+
description = "Gradio Space for the DeepLabCut Model Zoo SuperAnimal models"
|
| 5 |
+
readme = "README.md"
|
| 6 |
+
requires-python = ">=3.12,<3.13"
|
| 7 |
+
classifiers = [
|
| 8 |
+
"Programming Language :: Python :: 3 :: Only",
|
| 9 |
+
"Programming Language :: Python :: 3.12",
|
| 10 |
+
]
|
| 11 |
+
# Hugging Face Spaces install from requirements.txt only: keep both lists in sync.
|
| 12 |
+
dependencies = [
|
| 13 |
+
"deeplabcut[modelzoo,tf]==3.0.2",
|
| 14 |
+
"deeplabcut-live",
|
| 15 |
+
"dlclibrary",
|
| 16 |
+
"gitpython>=3.1.30",
|
| 17 |
+
"gradio==6.29",
|
| 18 |
+
"humanfriendly",
|
| 19 |
+
"psutil",
|
| 20 |
+
"ruamel-yaml==0.17.21",
|
| 21 |
+
"seaborn",
|
| 22 |
+
"ultralytics",
|
| 23 |
+
]
|
| 24 |
+
|
| 25 |
+
[tool.ruff]
|
| 26 |
+
target-version = "py312"
|
| 27 |
+
line-length = 120
|
| 28 |
+
format.docstring-code-format = true
|
| 29 |
+
lint.select = [ "B", "E", "F", "I", "PIE", "UP" ]
|
| 30 |
+
lint.ignore = [ "B007", "E741" ]
|
| 31 |
+
|
| 32 |
+
[tool.pyproject-fmt]
|
| 33 |
+
max_supported_python = "3.12"
|
pytorch_utils.py
ADDED
|
@@ -0,0 +1,118 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import threading
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
|
| 4 |
+
import numpy as np
|
| 5 |
+
import PIL
|
| 6 |
+
from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners
|
| 7 |
+
from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
|
| 8 |
+
from deeplabcut.pose_estimation_pytorch.modelzoo.utils import MODEL_FILENAME_MAPPING
|
| 9 |
+
from dlclibrary import download_huggingface_model
|
| 10 |
+
|
| 11 |
+
# SuperAnimal (pose model, detector) used by the PyTorch backend
|
| 12 |
+
PYTORCH_MODELS = {
|
| 13 |
+
"superanimal_quadruped": ("hrnet_w32", "fasterrcnn_resnet50_fpn_v2"),
|
| 14 |
+
"superanimal_topviewmouse": ("hrnet_w32", "fasterrcnn_resnet50_fpn_v2"),
|
| 15 |
+
}
|
| 16 |
+
|
| 17 |
+
MAX_INDIVIDUALS = 10
|
| 18 |
+
MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
|
| 19 |
+
# next to the TF models, not deeplabcut/modelzoo/checkpoints: site-packages is read-only for the Space's non-root user
|
| 20 |
+
WEIGHTS_DIR = Path(__file__).parent / "DLC_models" / "pytorch"
|
| 21 |
+
|
| 22 |
+
_runners = {}
|
| 23 |
+
_build_lock = threading.Lock()
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
##########################################
|
| 27 |
+
def snapshot_path(superanimal, model_name):
|
| 28 |
+
"""Path to a SuperAnimal snapshot in WEIGHTS_DIR, downloaded on first use (as deeplabcut does)."""
|
| 29 |
+
name = f"{superanimal}_{model_name}"
|
| 30 |
+
path = WEIGHTS_DIR / f"{name}.pt"
|
| 31 |
+
if not path.exists():
|
| 32 |
+
source = MODEL_FILENAME_MAPPING.get(name, path.name)
|
| 33 |
+
rename = None if source == path.name else {source: path.name}
|
| 34 |
+
download_huggingface_model(name, target_dir=str(WEIGHTS_DIR), rename_mapping=rename)
|
| 35 |
+
return path
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
##########################################
|
| 39 |
+
def load_superanimal(superanimal, device="auto"):
|
| 40 |
+
"""Build (once) the detector and pose runners for a SuperAnimal model; weights are downloaded on first use."""
|
| 41 |
+
with _build_lock:
|
| 42 |
+
if superanimal not in _runners:
|
| 43 |
+
pose_model, detector = PYTORCH_MODELS[superanimal]
|
| 44 |
+
cfg = PoseConfig.build_for_superanimal_inference(
|
| 45 |
+
superanimal,
|
| 46 |
+
model_name=pose_model,
|
| 47 |
+
detector_name=detector,
|
| 48 |
+
max_individuals=MAX_INDIVIDUALS,
|
| 49 |
+
device=device,
|
| 50 |
+
)
|
| 51 |
+
# keep low-score boxes: the UI threshold filters them afterwards
|
| 52 |
+
cfg["detector"]["model"]["box_score_thresh"] = 0.05
|
| 53 |
+
pose_runner, detector_runner = get_inference_runners(
|
| 54 |
+
cfg,
|
| 55 |
+
snapshot_path=snapshot_path(superanimal, pose_model),
|
| 56 |
+
detector_path=snapshot_path(superanimal, detector),
|
| 57 |
+
max_individuals=MAX_INDIVIDUALS,
|
| 58 |
+
inference_cfg={"multithreading": {"enabled": False}},
|
| 59 |
+
)
|
| 60 |
+
_runners[superanimal] = {
|
| 61 |
+
"pose": pose_runner,
|
| 62 |
+
"detector": detector_runner,
|
| 63 |
+
"bodyparts": list(cfg["metadata"]["bodyparts"]),
|
| 64 |
+
"lock": threading.Lock(),
|
| 65 |
+
} # runners are not thread-safe
|
| 66 |
+
return _runners[superanimal]
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
##########################################
|
| 70 |
+
def resize_max_side(img, max_size=MAX_IMAGE_SIZE):
|
| 71 |
+
scale = max_size / max(img.size)
|
| 72 |
+
if scale >= 1:
|
| 73 |
+
return img
|
| 74 |
+
return img.resize([int(x * scale) for x in img.size], PIL.Image.Resampling.LANCZOS)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
##########################################
|
| 78 |
+
def predict_superanimal(img_input, superanimal, bbox_likelihood_th, kpts_likelihood_th, full_image=False):
|
| 79 |
+
"""Detect animals and estimate their pose with a PyTorch SuperAnimal model.
|
| 80 |
+
|
| 81 |
+
Returns the (resized) RGB image the predictions refer to, the list of animals
|
| 82 |
+
as dicts {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3) array of x,y,llk
|
| 83 |
+
in image pixels, NaN below kpts_likelihood_th}, and the bodypart names.
|
| 84 |
+
"""
|
| 85 |
+
img = resize_max_side(img_input.convert("RGB"))
|
| 86 |
+
img_np = np.asarray(img)
|
| 87 |
+
runners = load_superanimal(superanimal)
|
| 88 |
+
|
| 89 |
+
with runners["lock"]:
|
| 90 |
+
if full_image:
|
| 91 |
+
# skip the detector and treat the whole image as one animal
|
| 92 |
+
h, w = img_np.shape[:2]
|
| 93 |
+
detections = {
|
| 94 |
+
"bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
|
| 95 |
+
"bbox_scores": np.array([1.0], dtype=np.float32),
|
| 96 |
+
}
|
| 97 |
+
else:
|
| 98 |
+
detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
|
| 99 |
+
keep = detections["bbox_scores"] >= bbox_likelihood_th
|
| 100 |
+
detections = {"bboxes": detections["bboxes"][keep], "bbox_scores": detections["bbox_scores"][keep]}
|
| 101 |
+
|
| 102 |
+
if len(detections["bboxes"]) == 0:
|
| 103 |
+
return img, [], runners["bodyparts"]
|
| 104 |
+
|
| 105 |
+
predictions = runners["pose"].inference([(img_np, detections)])[0]
|
| 106 |
+
|
| 107 |
+
animals = []
|
| 108 |
+
# outputs are padded to MAX_INDIVIDUALS with -1
|
| 109 |
+
for kpts, (x, y, w, h), score in zip(
|
| 110 |
+
predictions["bodyparts"], predictions["bboxes"], predictions["bbox_scores"], strict=True
|
| 111 |
+
):
|
| 112 |
+
if score < 0:
|
| 113 |
+
continue
|
| 114 |
+
kpts = kpts.astype(float)
|
| 115 |
+
kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
|
| 116 |
+
animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)], "kpts": kpts})
|
| 117 |
+
|
| 118 |
+
return img, animals, runners["bodyparts"]
|
requirements.txt
CHANGED
|
@@ -1,10 +1,13 @@
|
|
| 1 |
-
|
|
|
|
|
|
|
|
|
|
| 2 |
gitpython>=3.1.30
|
| 3 |
seaborn
|
| 4 |
-
deeplabcut[modelzoo,tf]
|
| 5 |
deeplabcut-live
|
| 6 |
ruamel.yaml==0.17.21
|
| 7 |
dlclibrary
|
| 8 |
humanfriendly
|
| 9 |
psutil
|
| 10 |
-
ultralytics
|
|
|
|
| 1 |
+
# Hugging Face Spaces install from this file; keep in sync with pyproject.toml dependencies.
|
| 2 |
+
# CPU-only PyTorch wheels: the Space has no GPU, and the default Linux wheels pull ~2-3 GB of CUDA libraries
|
| 3 |
+
--extra-index-url https://download.pytorch.org/whl/cpu
|
| 4 |
+
gradio==6.29.0
|
| 5 |
gitpython>=3.1.30
|
| 6 |
seaborn
|
| 7 |
+
deeplabcut[modelzoo,tf]==3.0.2
|
| 8 |
deeplabcut-live
|
| 9 |
ruamel.yaml==0.17.21
|
| 10 |
dlclibrary
|
| 11 |
humanfriendly
|
| 12 |
psutil
|
| 13 |
+
ultralytics
|
save_results.py
DELETED
|
@@ -1,56 +0,0 @@
|
|
| 1 |
-
import json
|
| 2 |
-
import numpy as np
|
| 3 |
-
import pdb
|
| 4 |
-
|
| 5 |
-
dict_pred = {0: 'animal', 1: 'person', 2: 'vehicle'}
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
def save_results(md_results, dlc_outputs,map_label_id_to_str,thr,output_file = 'dowload_predictions.json'):
|
| 9 |
-
|
| 10 |
-
"""
|
| 11 |
-
|
| 12 |
-
write json
|
| 13 |
-
|
| 14 |
-
"""
|
| 15 |
-
info = {}
|
| 16 |
-
## info megaDetector
|
| 17 |
-
info['file']= md_results.files[0]
|
| 18 |
-
number_bb = len(md_results.xyxy[0].tolist())
|
| 19 |
-
info['number_of_bb'] = number_bb
|
| 20 |
-
number_bb_thr = len(dlc_outputs)
|
| 21 |
-
labels = [n for n in map_label_id_to_str.values()]
|
| 22 |
-
#pdb.set_trace()
|
| 23 |
-
new_index = []
|
| 24 |
-
for i in range(number_bb):
|
| 25 |
-
corner_x1,corner_y1,corner_x2,corner_y2,confidence, _ = md_results.xyxy[0].tolist()[i]
|
| 26 |
-
|
| 27 |
-
if confidence > thr:
|
| 28 |
-
new_index.append(i)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
for i in range(number_bb_thr):
|
| 32 |
-
aux={}
|
| 33 |
-
corner_x1,corner_y1,corner_x2,corner_y2,confidence, _ = md_results.xyxy[0].tolist()[new_index[i]]
|
| 34 |
-
aux['corner_1'] = (corner_x1,corner_y1)
|
| 35 |
-
aux['corner_2'] = (corner_x2,corner_y2)
|
| 36 |
-
aux['predict MD'] = md_results.names[0]
|
| 37 |
-
aux['confidence MD'] = confidence
|
| 38 |
-
|
| 39 |
-
## info dlc
|
| 40 |
-
kypts = []
|
| 41 |
-
for s in dlc_outputs[i]:
|
| 42 |
-
aux1 = []
|
| 43 |
-
for j in s:
|
| 44 |
-
aux1.append(float(j))
|
| 45 |
-
|
| 46 |
-
kypts.append(aux1)
|
| 47 |
-
aux['dlc_pred'] = dict(zip(labels,kypts))
|
| 48 |
-
info['bb_' + str(new_index[i]) ]=aux
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
with open(output_file, 'w') as f:
|
| 52 |
-
json.dump(info, f, indent=1)
|
| 53 |
-
print('Output file saved at {}'.format(output_file))
|
| 54 |
-
|
| 55 |
-
return output_file
|
| 56 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
ui_utils.py
CHANGED
|
@@ -1,21 +1,78 @@
|
|
| 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,12 +80,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(
|
| 30 |
-
value=True,
|
| 31 |
-
label="Show bodypart labels?",
|
| 32 |
)
|
| 33 |
|
| 34 |
# Gradio Slider signature is (minimum, maximum, value, step, ...)
|
|
@@ -49,36 +101,60 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
|
|
| 49 |
)
|
| 50 |
|
| 51 |
# Data viz
|
| 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 |
return [
|
| 81 |
gr_image_input,
|
|
|
|
| 82 |
gr_mega_model_input,
|
| 83 |
gr_dlc_model_input,
|
| 84 |
gr_dlc_only_checkbox,
|
|
@@ -89,38 +165,64 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
|
|
| 89 |
gr_slider_font_size,
|
| 90 |
gr_keypt_color,
|
| 91 |
gr_slider_marker_size,
|
|
|
|
|
|
|
|
|
|
| 92 |
]
|
| 93 |
|
| 94 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
def gradio_outputs_for_MD_DLC():
|
| 96 |
gr_image_output = gr.Image(type="pil", label="Output Image")
|
| 97 |
-
|
| 98 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
|
| 100 |
|
| 101 |
def gradio_description_and_examples():
|
| 102 |
-
title = "DeepLabCut Model Zoo SuperAnimals"
|
| 103 |
description = (
|
| 104 |
-
"
|
| 105 |
-
"
|
| 106 |
-
"
|
| 107 |
-
"
|
| 108 |
-
"Want to run on videos on the cloud or locally? See the "
|
| 109 |
-
"<a href='http://www.mackenziemathislab.org/dlc-modelzoo'>DeepLabCut ModelZoo</a>."
|
| 110 |
)
|
| 111 |
|
| 112 |
-
examples = [
|
| 113 |
-
"
|
| 114 |
-
"
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
3,
|
| 124 |
-
]]
|
| 125 |
-
|
| 126 |
-
return [title, description, examples]
|
|
|
|
| 1 |
import gradio as gr
|
| 2 |
+
from matplotlib import colormaps
|
| 3 |
+
from matplotlib.colors import to_hex
|
| 4 |
+
from PIL import Image
|
| 5 |
+
|
| 6 |
+
from pytorch_utils import MAX_IMAGE_SIZE
|
| 7 |
+
from viz_utils import COLORMAPS
|
| 8 |
+
|
| 9 |
+
# shades built around the DeepLabCut docs palette (DeepLabCut/docs/_static/custom.css)
|
| 10 |
+
DLC_PURPLE = gr.themes.Color(
|
| 11 |
+
name="dlc_purple",
|
| 12 |
+
c50="#f5edff",
|
| 13 |
+
c100="#ead7ff",
|
| 14 |
+
c200="#d9b8ff",
|
| 15 |
+
c300="#c084fc",
|
| 16 |
+
c400="#ac72f0",
|
| 17 |
+
c500="#9b5de5",
|
| 18 |
+
c600="#8550c4",
|
| 19 |
+
c700="#73439a",
|
| 20 |
+
c800="#5c357c",
|
| 21 |
+
c900="#4b236f",
|
| 22 |
+
c950="#2e1546",
|
| 23 |
+
)
|
| 24 |
+
DLC_TEAL = gr.themes.Color(
|
| 25 |
+
name="dlc_teal",
|
| 26 |
+
c50="#e8f8f7",
|
| 27 |
+
c100="#c9efec",
|
| 28 |
+
c200="#9fe2dd",
|
| 29 |
+
c300="#7ad3ce",
|
| 30 |
+
c400="#57c4be",
|
| 31 |
+
c500="#2fb0a8",
|
| 32 |
+
c600="#21a197",
|
| 33 |
+
c700="#1a8078",
|
| 34 |
+
c800="#16655f",
|
| 35 |
+
c900="#124f4a",
|
| 36 |
+
c950="#0a2e2b",
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def dlc_theme():
|
| 41 |
+
# white text on purple 500 is 4.1:1, so filled buttons use 700 (7.0:1) and 600 (5.3:1)
|
| 42 |
+
return gr.themes.Default(primary_hue=DLC_PURPLE, secondary_hue=DLC_TEAL, neutral_hue="slate").set(
|
| 43 |
+
button_primary_background_fill="*primary_700",
|
| 44 |
+
button_primary_background_fill_hover="*primary_600",
|
| 45 |
+
button_primary_background_fill_dark="*primary_600",
|
| 46 |
+
button_primary_background_fill_hover_dark="*primary_500",
|
| 47 |
+
button_primary_text_color="white",
|
| 48 |
+
button_primary_text_color_dark="white",
|
| 49 |
+
button_primary_border_color="*primary_700",
|
| 50 |
+
button_primary_border_color_dark="*primary_600",
|
| 51 |
+
)
|
| 52 |
|
| 53 |
|
| 54 |
+
def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
|
| 55 |
# Input image
|
| 56 |
gr_image_input = gr.Image(type="pil", label="Input Image")
|
| 57 |
|
| 58 |
# Models
|
| 59 |
+
gr_backend_input = gr.Radio(
|
| 60 |
+
choices=backends_list,
|
| 61 |
+
value=backends_list[0],
|
| 62 |
+
label="Select backend",
|
| 63 |
+
)
|
| 64 |
+
|
| 65 |
gr_mega_model_input = gr.Dropdown(
|
| 66 |
choices=md_models_list,
|
| 67 |
value="md_v5a",
|
| 68 |
type="value",
|
| 69 |
+
label="Select Detector model (TensorFlow legacy only)",
|
| 70 |
+
visible=gr_backend_input.value != "PyTorch",
|
| 71 |
)
|
| 72 |
|
| 73 |
gr_dlc_model_input = gr.Dropdown(
|
| 74 |
choices=dlc_models_list,
|
| 75 |
+
value="superanimal_quadruped",
|
| 76 |
type="value",
|
| 77 |
label="Select DeepLabCut model",
|
| 78 |
)
|
|
|
|
| 80 |
# Other inputs
|
| 81 |
gr_dlc_only_checkbox = gr.Checkbox(
|
| 82 |
value=False,
|
| 83 |
+
label="Run DeepLabCut only, directly on input image?",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
)
|
| 85 |
|
| 86 |
# Gradio Slider signature is (minimum, maximum, value, step, ...)
|
|
|
|
| 101 |
)
|
| 102 |
|
| 103 |
# Data viz
|
| 104 |
+
with gr.Accordion("Display options", open=False):
|
| 105 |
+
gr_str_labels_checkbox = gr.Checkbox(
|
| 106 |
+
value=True,
|
| 107 |
+
label="Show bodypart labels?",
|
| 108 |
+
)
|
| 109 |
+
|
| 110 |
+
gr_color_by_confidence_checkbox = gr.Checkbox(
|
| 111 |
+
value=True,
|
| 112 |
+
label="Color keypoints by confidence? (otherwise by bodypart)",
|
| 113 |
+
)
|
| 114 |
+
|
| 115 |
+
gr_colormap = gr.Dropdown(
|
| 116 |
+
choices=COLORMAPS,
|
| 117 |
+
value="viridis",
|
| 118 |
+
type="value",
|
| 119 |
+
label="Keypoint colormap",
|
| 120 |
+
)
|
| 121 |
+
|
| 122 |
+
gr_keypt_color = gr.ColorPicker(
|
| 123 |
+
value="#862db7",
|
| 124 |
+
label="Choose color for keypoint label",
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
gr_bbox_color = gr.ColorPicker(
|
| 128 |
+
value="#ff0000",
|
| 129 |
+
label="Choose color for bounding boxes",
|
| 130 |
+
)
|
| 131 |
+
|
| 132 |
+
gr_labels_font_style = gr.Dropdown(
|
| 133 |
+
choices=["amiko", "animals", "nature", "painter", "zen"],
|
| 134 |
+
value="amiko",
|
| 135 |
+
type="value",
|
| 136 |
+
label="Select keypoint label font",
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
gr_slider_font_size = gr.Slider(
|
| 140 |
+
minimum=5,
|
| 141 |
+
maximum=30,
|
| 142 |
+
value=18,
|
| 143 |
+
step=1,
|
| 144 |
+
label="Set font size",
|
| 145 |
+
)
|
| 146 |
+
|
| 147 |
+
gr_slider_marker_size = gr.Slider(
|
| 148 |
+
minimum=1,
|
| 149 |
+
maximum=20,
|
| 150 |
+
value=6,
|
| 151 |
+
step=1,
|
| 152 |
+
label="Set marker size",
|
| 153 |
+
)
|
| 154 |
|
| 155 |
return [
|
| 156 |
gr_image_input,
|
| 157 |
+
gr_backend_input,
|
| 158 |
gr_mega_model_input,
|
| 159 |
gr_dlc_model_input,
|
| 160 |
gr_dlc_only_checkbox,
|
|
|
|
| 165 |
gr_slider_font_size,
|
| 166 |
gr_keypt_color,
|
| 167 |
gr_slider_marker_size,
|
| 168 |
+
gr_color_by_confidence_checkbox,
|
| 169 |
+
gr_colormap,
|
| 170 |
+
gr_bbox_color,
|
| 171 |
]
|
| 172 |
|
| 173 |
|
| 174 |
+
def confidence_legend_html(colormap="viridis"):
|
| 175 |
+
# the colormap as a CSS gradient, matching the keypoint fill in draw_keypoints_on_image
|
| 176 |
+
gradient = ", ".join(f"{to_hex(colormaps[colormap](i / 10))} {i * 10}%" for i in range(11))
|
| 177 |
+
# no leading newline: gradio prefixes "'" to cached example values starting with one (CSV injection guard)
|
| 178 |
+
return f"""<div style="display:flex; align-items:flex-start; gap:12px; flex-wrap:wrap; font-size:var(--text-sm);">
|
| 179 |
+
<span style="white-space:nowrap; line-height:12px;">Keypoint confidence</span>
|
| 180 |
+
<div style="flex:1; min-width:160px; max-width:360px;">
|
| 181 |
+
<div style="height:12px; border-radius:6px; background:linear-gradient(to right, {gradient});"></div>
|
| 182 |
+
<div style="display:flex; justify-content:space-between; margin-top:2px; font-variant-numeric:tabular-nums;">
|
| 183 |
+
<span>0</span><span>0.5</span><span>1</span>
|
| 184 |
+
</div>
|
| 185 |
+
</div>
|
| 186 |
+
</div>"""
|
| 187 |
+
|
| 188 |
+
|
| 189 |
def gradio_outputs_for_MD_DLC():
|
| 190 |
gr_image_output = gr.Image(type="pil", label="Output Image")
|
| 191 |
+
gr_confidence_legend = gr.HTML("")
|
| 192 |
+
with gr.Row():
|
| 193 |
+
gr_file_download = gr.File(label="Download JSON file")
|
| 194 |
+
gr_image_download = gr.File(label="Download annotated image")
|
| 195 |
+
gr_confidence_table = gr.Dataframe(
|
| 196 |
+
headers=["animal", "bodypart", "confidence"],
|
| 197 |
+
label="Keypoint confidence (lowest first)",
|
| 198 |
+
interactive=False,
|
| 199 |
+
)
|
| 200 |
+
return [gr_image_output, gr_confidence_legend, gr_file_download, gr_image_download, gr_confidence_table]
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def example_sizes(path):
|
| 204 |
+
# font and marker sizes for the resolution the PyTorch backend draws on
|
| 205 |
+
side = min(max(Image.open(path).size), MAX_IMAGE_SIZE)
|
| 206 |
+
return round(side / 64), round(side / 200)
|
| 207 |
|
| 208 |
|
| 209 |
def gradio_description_and_examples():
|
| 210 |
+
title = "DeepLabCut Model Zoo: SuperAnimals"
|
| 211 |
description = (
|
| 212 |
+
"Estimate animal poses with the SuperAnimal models from the "
|
| 213 |
+
"[DeepLabCut Model Zoo](http://www.mackenziemathislab.org/dlc-modelzoo) "
|
| 214 |
+
"([paper](https://arxiv.org/abs/2203.07436)). "
|
| 215 |
+
"Upload an image or pick an example below; to run on videos, see the Model Zoo page."
|
|
|
|
|
|
|
| 216 |
)
|
| 217 |
|
| 218 |
+
examples = [
|
| 219 |
+
[image, "PyTorch", "md_v5a", "superanimal_quadruped", False, True, 0.5, 0.4, "amiko"]
|
| 220 |
+
+ [font_size, "#ff0000", marker_size, True, "viridis", "#ff0000"]
|
| 221 |
+
for image in (
|
| 222 |
+
"examples/dog.jpeg",
|
| 223 |
+
"examples/cat.jpg",
|
| 224 |
+
)
|
| 225 |
+
for font_size, marker_size in [example_sizes(image)]
|
| 226 |
+
]
|
| 227 |
+
|
| 228 |
+
return [title, description, examples]
|
|
|
|
|
|
|
|
|
|
|
|
viz_utils.py
CHANGED
|
@@ -1,31 +1,38 @@
|
|
| 1 |
-
import json
|
| 2 |
-
import
|
|
|
|
| 3 |
|
| 4 |
-
from matplotlib import cm
|
| 5 |
-
import matplotlib
|
| 6 |
-
from PIL import Image, ImageColor, ImageFont, ImageDraw
|
| 7 |
import numpy as np
|
| 8 |
-
import
|
| 9 |
-
from
|
|
|
|
| 10 |
today = date.today()
|
| 11 |
-
FONTS = {
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
|
| 17 |
#########################################
|
| 18 |
# Draw keypoints on image
|
| 19 |
-
def draw_keypoints_on_image(
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
|
|
|
|
|
|
|
|
|
| 29 |
"""Draws keypoints on an image.
|
| 30 |
Modified from:
|
| 31 |
https://www.programcreek.com/python/?code=fjchange%2Fobject_centric_VAD%2Fobject_centric_VAD-master%2Fobject_detection%2Futils%2Fvisualization_utils.py
|
|
@@ -39,159 +46,232 @@ def draw_keypoints_on_image(image,
|
|
| 39 |
use_normalized_coordinates: if True (default), treat keypoint values as
|
| 40 |
relative to the image. Otherwise treat them as absolute.
|
| 41 |
|
| 42 |
-
|
| 43 |
"""
|
| 44 |
# get a drawing context
|
| 45 |
-
draw = ImageDraw.Draw(image,"RGBA")
|
| 46 |
|
| 47 |
im_width, im_height = image.size
|
| 48 |
keypoints_x = [k[0] for k in keypoints]
|
| 49 |
keypoints_y = [k[1] for k in keypoints]
|
| 50 |
-
|
| 51 |
-
norm = matplotlib.colors.Normalize(vmin=0, vmax=255)
|
| 52 |
-
|
| 53 |
-
# debugging keypoints
|
| 54 |
-
print (keypoints)
|
| 55 |
-
|
| 56 |
-
names_for_color = [i for i in map_label_id_to_str.keys()]
|
| 57 |
-
colores = np.linspace(0, 255, num=len(names_for_color),dtype= int)
|
| 58 |
|
| 59 |
# adjust keypoints coords if required
|
| 60 |
if use_normalized_coordinates:
|
| 61 |
keypoints_x = tuple([im_width * x for x in keypoints_x])
|
| 62 |
keypoints_y = tuple([im_height * y for y in keypoints_y])
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
cmap2 = matplotlib.cm.get_cmap('Greys')
|
| 66 |
# draw ellipses around keypoints
|
| 67 |
-
for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y)):
|
| 68 |
-
round_fill = list(cm.viridis(norm(colores[i]),bytes=True))#[round(num*255) for num in list(cmap(i))[:3]] #check!
|
| 69 |
# handling potential nans in the keypoints
|
| 70 |
if np.isnan(keypoint_x).any():
|
| 71 |
continue
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 80 |
|
| 81 |
# add string labels around keypoints
|
| 82 |
if flag_show_str_labels:
|
| 83 |
-
font = ImageFont.truetype(FONTS[font_style],
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 89 |
|
| 90 |
#########################################
|
| 91 |
# Draw bboxes on image
|
| 92 |
-
def draw_bbox_w_text(img,
|
| 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 |
-
Output detections as json file
|
| 120 |
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
number_bb_thr = len(dlc_outputs)
|
| 132 |
-
labels = [n for n in map_dlc_label_id_to_str.values()]
|
| 133 |
-
|
| 134 |
-
# create list of bboxes above th
|
| 135 |
-
new_index = []
|
| 136 |
-
for i in range(number_bb):
|
| 137 |
-
corner_x1,corner_y1,corner_x2,corner_y2,confidence, _ = md_results.xyxy[0].tolist()[i]
|
| 138 |
-
|
| 139 |
-
if confidence > thr:
|
| 140 |
-
new_index.append(i)
|
| 141 |
-
|
| 142 |
-
# define aux dict for every bounding box above threshold
|
| 143 |
-
for i in range(number_bb_thr):
|
| 144 |
-
aux={}
|
| 145 |
-
# MD output
|
| 146 |
-
corner_x1,corner_y1,corner_x2,corner_y2,confidence, _ = md_results.xyxy[0].tolist()[new_index[i]]
|
| 147 |
-
aux['corner_1'] = (corner_x1,corner_y1)
|
| 148 |
-
aux['corner_2'] = (corner_x2,corner_y2)
|
| 149 |
-
aux['predict MD'] = md_results.names[0]
|
| 150 |
-
aux['confidence MD'] = confidence
|
| 151 |
-
|
| 152 |
-
# DLC output
|
| 153 |
-
info['dlc_model'] = model
|
| 154 |
-
kypts = []
|
| 155 |
-
for s in dlc_outputs[i]:
|
| 156 |
-
aux1 = []
|
| 157 |
-
for j in s:
|
| 158 |
-
aux1.append(float(j))
|
| 159 |
-
|
| 160 |
-
kypts.append(aux1)
|
| 161 |
-
aux['dlc_pred'] = dict(zip(labels,kypts))
|
| 162 |
-
info['bb_' + str(new_index[i]) ]=aux
|
| 163 |
-
|
| 164 |
-
# save dict as json
|
| 165 |
-
with open(path_to_output_file, 'w') as f:
|
| 166 |
-
json.dump(info, f, indent=1)
|
| 167 |
-
print('Output file saved at {}'.format(path_to_output_file))
|
| 168 |
|
|
|
|
|
|
|
|
|
|
| 169 |
return path_to_output_file
|
| 170 |
|
| 171 |
|
| 172 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 173 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 174 |
"""
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
info =
|
| 178 |
-
info[
|
| 179 |
-
|
| 180 |
-
info[
|
| 181 |
-
|
| 182 |
-
for s in dlc_outputs:
|
| 183 |
-
aux1 = []
|
| 184 |
-
for j in s:
|
| 185 |
-
aux1.append(float(j))
|
| 186 |
|
| 187 |
-
|
| 188 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 189 |
|
| 190 |
-
with open(output_file, 'w') as f:
|
| 191 |
-
json.dump(info, f, indent=1)
|
| 192 |
-
print('Output file saved at {}'.format(output_file))
|
| 193 |
|
| 194 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 195 |
|
| 196 |
|
| 197 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import math
|
| 3 |
+
from datetime import date
|
| 4 |
|
|
|
|
|
|
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
+
from matplotlib import colormaps
|
| 7 |
+
from PIL import ImageColor, ImageDraw, ImageFont
|
| 8 |
+
|
| 9 |
today = date.today()
|
| 10 |
+
FONTS = {
|
| 11 |
+
"amiko": "fonts/Amiko-Regular.ttf",
|
| 12 |
+
"nature": "fonts/LoveNature.otf",
|
| 13 |
+
"painter": "fonts/PainterDecorator.otf",
|
| 14 |
+
"animals": "fonts/UncialAnimals.ttf",
|
| 15 |
+
"zen": "fonts/ZEN.TTF",
|
| 16 |
+
}
|
| 17 |
+
# perceptually uniform maps first; turbo separates neighbouring bodyparts best
|
| 18 |
+
COLORMAPS = ["viridis", "plasma", "magma", "cividis", "turbo"]
|
| 19 |
+
|
| 20 |
|
| 21 |
#########################################
|
| 22 |
# Draw keypoints on image
|
| 23 |
+
def draw_keypoints_on_image(
|
| 24 |
+
image,
|
| 25 |
+
keypoints,
|
| 26 |
+
map_label_id_to_str,
|
| 27 |
+
flag_show_str_labels,
|
| 28 |
+
use_normalized_coordinates=True,
|
| 29 |
+
font_style="amiko",
|
| 30 |
+
font_size=8,
|
| 31 |
+
keypt_color="#ff0000",
|
| 32 |
+
marker_size=2,
|
| 33 |
+
color_by_confidence=True,
|
| 34 |
+
colormap="viridis",
|
| 35 |
+
):
|
| 36 |
"""Draws keypoints on an image.
|
| 37 |
Modified from:
|
| 38 |
https://www.programcreek.com/python/?code=fjchange%2Fobject_centric_VAD%2Fobject_centric_VAD-master%2Fobject_detection%2Futils%2Fvisualization_utils.py
|
|
|
|
| 46 |
use_normalized_coordinates: if True (default), treat keypoint values as
|
| 47 |
relative to the image. Otherwise treat them as absolute.
|
| 48 |
|
| 49 |
+
|
| 50 |
"""
|
| 51 |
# get a drawing context
|
| 52 |
+
draw = ImageDraw.Draw(image, "RGBA")
|
| 53 |
|
| 54 |
im_width, im_height = image.size
|
| 55 |
keypoints_x = [k[0] for k in keypoints]
|
| 56 |
keypoints_y = [k[1] for k in keypoints]
|
| 57 |
+
confidences = [k[2] for k in keypoints]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 58 |
|
| 59 |
# adjust keypoints coords if required
|
| 60 |
if use_normalized_coordinates:
|
| 61 |
keypoints_x = tuple([im_width * x for x in keypoints_x])
|
| 62 |
keypoints_y = tuple([im_height * y for y in keypoints_y])
|
| 63 |
+
|
| 64 |
+
cmap = colormaps[colormap]
|
|
|
|
| 65 |
# draw ellipses around keypoints
|
| 66 |
+
for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y, strict=True)):
|
|
|
|
| 67 |
# handling potential nans in the keypoints
|
| 68 |
if np.isnan(keypoint_x).any():
|
| 69 |
continue
|
| 70 |
+
|
| 71 |
+
confidence = float(np.clip(confidences[i], 0, 1))
|
| 72 |
+
if color_by_confidence:
|
| 73 |
+
# fill color encodes the keypoint confidence (see confidence_legend_html in ui_utils)
|
| 74 |
+
round_fill = cmap(confidence, bytes=True)
|
| 75 |
+
else:
|
| 76 |
+
# one color per bodypart, transparency encodes the confidence
|
| 77 |
+
round_fill = list(cmap(i / max(len(keypoints) - 1, 1), bytes=True))
|
| 78 |
+
round_fill[3] = round(confidence * 255)
|
| 79 |
+
round_fill = tuple(round_fill)
|
| 80 |
+
draw.ellipse(
|
| 81 |
+
[
|
| 82 |
+
(keypoint_x - marker_size, keypoint_y - marker_size),
|
| 83 |
+
(keypoint_x + marker_size, keypoint_y + marker_size),
|
| 84 |
+
],
|
| 85 |
+
fill=tuple(round_fill),
|
| 86 |
+
outline="black",
|
| 87 |
+
width=1,
|
| 88 |
+
) # fill and outline: [0,255]
|
| 89 |
|
| 90 |
# add string labels around keypoints
|
| 91 |
if flag_show_str_labels:
|
| 92 |
+
font = ImageFont.truetype(FONTS[font_style], font_size)
|
| 93 |
+
draw.text(
|
| 94 |
+
(keypoint_x + marker_size, keypoint_y + marker_size), # (0.5*im_width, 0.5*im_height), #-------
|
| 95 |
+
display_bodypart(map_label_id_to_str[i]),
|
| 96 |
+
ImageColor.getcolor(keypt_color, "RGB"), # rgb #
|
| 97 |
+
font=font,
|
| 98 |
+
)
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
#########################################
|
| 102 |
+
# Bodypart names for display
|
| 103 |
+
# display names where the SuperAnimal definitions misspell (quadruped "thai") or read oddly
|
| 104 |
+
# (top-view mouse "backend"); the JSON output keeps the model's names
|
| 105 |
+
BODYPART_DISPLAY_NAMES = {
|
| 106 |
+
"front_left_thai": "front left thigh",
|
| 107 |
+
"front_right_thai": "front right thigh",
|
| 108 |
+
"back_left_thai": "back left thigh",
|
| 109 |
+
"back_right_thai": "back right thigh",
|
| 110 |
+
"mid_backend": "mid back end",
|
| 111 |
+
"mid_backend2": "mid back end 2",
|
| 112 |
+
"mid_backend3": "mid back end 3",
|
| 113 |
+
}
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def display_bodypart(name):
|
| 117 |
+
return BODYPART_DISPLAY_NAMES.get(name, name.replace("_", " "))
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
#########################################
|
| 121 |
+
# Keypoint confidences as table rows
|
| 122 |
+
def keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str):
|
| 123 |
+
"""(animal, bodypart, confidence) for every keypoint kept (not NaN), lowest confidence first."""
|
| 124 |
+
rows = []
|
| 125 |
+
for i_animal, kpts in enumerate(kpts_per_animal):
|
| 126 |
+
for i_kpt, kpt in enumerate(kpts):
|
| 127 |
+
if not np.isnan(kpt[2]):
|
| 128 |
+
rows.append([i_animal, display_bodypart(map_label_id_to_str[i_kpt]), round(float(kpt[2]), 3)])
|
| 129 |
+
return sorted(rows, key=lambda row: row[2])
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
#########################################
|
| 133 |
+
# Save the annotated image for download
|
| 134 |
+
def save_annotated_image(image, path_to_output_file="download_annotated.png"):
|
| 135 |
+
image.save(path_to_output_file)
|
| 136 |
+
return path_to_output_file
|
| 137 |
+
|
| 138 |
|
| 139 |
#########################################
|
| 140 |
# Draw bboxes on image
|
| 141 |
+
def draw_bbox_w_text(img, results, font_style="amiko", font_size=8, bbox_color="#ff0000"):
|
| 142 |
+
x1, y1, x2, y2, confidence = results[:5]
|
| 143 |
+
draw = ImageDraw.Draw(img)
|
| 144 |
+
draw.rectangle([(x1, y1), (x2, y2)], outline=bbox_color, width=max(2, round(font_size / 5)))
|
| 145 |
+
|
| 146 |
+
label = f"animal {confidence:.2f}"
|
| 147 |
+
font = ImageFont.truetype(FONTS[font_style], font_size)
|
| 148 |
+
left, top, right, bottom = draw.textbbox((0, 0), label, font=font)
|
| 149 |
+
pad = max(2, font_size // 5)
|
| 150 |
+
label_w, label_h = right - left + 2 * pad, bottom - top + 2 * pad
|
| 151 |
+
# label above the box, or inside it when the box touches the top of the image
|
| 152 |
+
label_y = y1 - label_h if y1 >= label_h else y1
|
| 153 |
+
draw.rectangle([(x1, label_y), (x1 + label_w, label_y + label_h)], fill=bbox_color)
|
| 154 |
+
draw.text((x1 + pad - left, label_y + pad - top), label, font=font, fill=label_text_color(bbox_color))
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def label_text_color(background):
|
| 158 |
+
# black or white, whichever contrasts more with the background (WCAG relative luminance)
|
| 159 |
+
channels = [c / 255 for c in ImageColor.getrgb(background)[:3]]
|
| 160 |
+
r, g, b = [c / 12.92 if c <= 0.03928 else ((c + 0.055) / 1.055) ** 2.4 for c in channels]
|
| 161 |
+
return "black" if 0.2126 * r + 0.7152 * g + 0.0722 * b > 0.179 else "white"
|
| 162 |
+
|
| 163 |
|
| 164 |
###########################################
|
| 165 |
+
# JSON outputs: pixel coordinates in the input image, hidden keypoints as null
|
| 166 |
+
COORDINATES = "pixels in the input image (after EXIF orientation), origin top-left, x right, y down"
|
| 167 |
|
|
|
|
|
|
|
| 168 |
|
| 169 |
+
def keypoints_to_json(kpts, offset=(0, 0), scale=(1.0, 1.0)):
|
| 170 |
+
"""[x, y, confidence] per keypoint, mapped by (k + offset) * scale; NaN (below threshold) becomes null."""
|
| 171 |
+
out = []
|
| 172 |
+
for x, y, conf in np.asarray(kpts, dtype=float)[:, :3]:
|
| 173 |
+
if np.isnan(conf):
|
| 174 |
+
out.append(None)
|
| 175 |
+
else:
|
| 176 |
+
out.append([(x + offset[0]) * scale[0], (y + offset[1]) * scale[1], conf])
|
| 177 |
+
return out
|
| 178 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 179 |
|
| 180 |
+
def write_json(info, path_to_output_file):
|
| 181 |
+
with open(path_to_output_file, "w") as f:
|
| 182 |
+
json.dump(info, f, indent=1, allow_nan=False)
|
| 183 |
return path_to_output_file
|
| 184 |
|
| 185 |
|
| 186 |
+
def json_header(image_size, annotated_size):
|
| 187 |
+
return {
|
| 188 |
+
"date": str(today),
|
| 189 |
+
"coordinates": COORDINATES,
|
| 190 |
+
"image_size": list(image_size),
|
| 191 |
+
"annotated_image_size": list(annotated_size),
|
| 192 |
+
}
|
| 193 |
|
| 194 |
+
|
| 195 |
+
def save_results_as_json(
|
| 196 |
+
md_results,
|
| 197 |
+
dlc_outputs,
|
| 198 |
+
animal_bboxes,
|
| 199 |
+
map_dlc_label_id_to_str,
|
| 200 |
+
model,
|
| 201 |
+
mega_model_input,
|
| 202 |
+
image_size,
|
| 203 |
+
path_to_output_file="download_predictions.json",
|
| 204 |
+
):
|
| 205 |
+
"""TF (legacy) MegaDetector + DLC results.
|
| 206 |
+
|
| 207 |
+
animal_bboxes: detection rows [x1,y1,x2,y2,conf,label] in the MegaDetector frame, one per entry of dlc_outputs
|
| 208 |
+
dlc_outputs: keypoints relative to each crop (crops start at floor(x1), floor(y1), see crop_animal_detections)
|
| 209 |
"""
|
| 210 |
+
md_h, md_w = md_results.ims[0].shape[:2]
|
| 211 |
+
scale = (image_size[0] / md_w, image_size[1] / md_h)
|
| 212 |
+
info = json_header(image_size, (md_w, md_h))
|
| 213 |
+
info["MD_model"] = str(mega_model_input)
|
| 214 |
+
info["number_of_bb"] = len(dlc_outputs)
|
| 215 |
+
info["dlc_model"] = model
|
| 216 |
+
labels = list(map_dlc_label_id_to_str.values())
|
|
|
|
|
|
|
|
|
|
|
|
|
| 217 |
|
| 218 |
+
for i, kpts in enumerate(dlc_outputs):
|
| 219 |
+
x1, y1, x2, y2, confidence, _ = animal_bboxes[i]
|
| 220 |
+
info["bb_" + str(i)] = {
|
| 221 |
+
"corner_1": (x1 * scale[0], y1 * scale[1]),
|
| 222 |
+
"corner_2": (x2 * scale[0], y2 * scale[1]),
|
| 223 |
+
"predict MD": md_results.names[0],
|
| 224 |
+
"confidence MD": float(confidence),
|
| 225 |
+
"dlc_pred": dict(
|
| 226 |
+
zip(labels, keypoints_to_json(kpts, offset=(math.floor(x1), math.floor(y1)), scale=scale), strict=True)
|
| 227 |
+
),
|
| 228 |
+
}
|
| 229 |
+
return write_json(info, path_to_output_file)
|
| 230 |
|
|
|
|
|
|
|
|
|
|
| 231 |
|
| 232 |
+
def save_results_only_dlc(
|
| 233 |
+
dlc_outputs, map_label_id_to_str, model, image_size, output_file="dowload_predictions_dlc.json"
|
| 234 |
+
):
|
| 235 |
+
"""TF (legacy) DLC run on the whole input image (keypoints already in input pixels)."""
|
| 236 |
+
info = json_header(image_size, image_size)
|
| 237 |
+
info["dlc_model"] = model
|
| 238 |
+
info["dlc_pred"] = dict(zip(map_label_id_to_str.values(), keypoints_to_json(dlc_outputs), strict=True))
|
| 239 |
+
return write_json(info, output_file)
|
| 240 |
|
| 241 |
|
| 242 |
+
def save_results_pytorch(
|
| 243 |
+
animals,
|
| 244 |
+
map_label_id_to_str,
|
| 245 |
+
model,
|
| 246 |
+
pose_model,
|
| 247 |
+
detector,
|
| 248 |
+
image_size,
|
| 249 |
+
annotated_size,
|
| 250 |
+
path_to_output_file="download_predictions.json",
|
| 251 |
+
):
|
| 252 |
+
"""PyTorch SuperAnimal results (same layout as save_results_as_json).
|
| 253 |
+
|
| 254 |
+
animals: list of {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3)}, in the annotated (resized) image
|
| 255 |
+
detector: None if the detector was skipped (whole image used as one animal)
|
| 256 |
+
"""
|
| 257 |
+
scale = (image_size[0] / annotated_size[0], image_size[1] / annotated_size[1])
|
| 258 |
+
info = json_header(image_size, annotated_size)
|
| 259 |
+
info["backend"] = "pytorch"
|
| 260 |
+
info["dlc_model"] = model
|
| 261 |
+
info["pose_model"] = pose_model
|
| 262 |
+
info["detector"] = detector
|
| 263 |
+
info["number_of_bb"] = len(animals)
|
| 264 |
+
labels = list(map_label_id_to_str.values())
|
| 265 |
+
|
| 266 |
+
for i, animal in enumerate(animals):
|
| 267 |
+
x1, y1, x2, y2, confidence = animal["bbox"]
|
| 268 |
+
info["bb_" + str(i)] = {
|
| 269 |
+
"corner_1": (x1 * scale[0], y1 * scale[1]),
|
| 270 |
+
"corner_2": (x2 * scale[0], y2 * scale[1]),
|
| 271 |
+
"confidence": confidence,
|
| 272 |
+
"dlc_pred": dict(zip(labels, keypoints_to_json(animal["kpts"], scale=scale), strict=True)),
|
| 273 |
+
}
|
| 274 |
+
return write_json(info, path_to_output_file)
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
###########################################
|