Add PyTorch SuperAnimal backend and refresh the Space

#14
by C-Achard - opened
Files changed (13) hide show
  1. .gitignore +14 -1
  2. .pre-commit-config.yaml +28 -0
  3. README.md +2 -2
  4. app.py +261 -140
  5. detection_utils.py +43 -61
  6. dlc_utils.py +16 -20
  7. pre-requirements.txt +1 -0
  8. pyproject.toml +33 -0
  9. pytorch_utils.py +118 -0
  10. requirements.txt +6 -3
  11. save_results.py +0 -56
  12. ui_utils.py +162 -60
  13. 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.11.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
 
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 yaml
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
- 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
25
  from dlclibrary.dlcmodelzoo.modelzoo_download import (
26
  download_huggingface_model,
27
- MODELOPTIONS,
28
  )
29
- from dlclive import DLCLive, Processor
30
 
 
 
31
 
32
- # TESTING (passes) download the SuperAnimal models:
33
- #model = 'superanimal_topviewmouse'
34
- #train_dir = 'DLC_models/sa-tvm'
35
- #download_huggingface_model(model, train_dir)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
 
37
- # grab demo data cooco cat:
38
- url = "http://images.cocodataset.org/val2017/000000039769.jpg"
39
- image = Image.open(requests.get(url, stream=True).raw)
 
40
 
41
  # megadetector and dlc model look up
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
- # DLC models target dirs
46
- DLC_models_dict = {'superanimal_topviewmouse_dlcrnet': "DLC_models/sa-tvm",
47
- 'superanimal_quadruped_dlcrnet': "DLC_models/sa-q"}
48
-
 
 
 
49
 
50
 
51
  #####################################################
52
- def predict_pipeline(img_input,
53
- mega_model_input,
54
- dlc_model_input_str,
55
- flag_dlc_only,
56
- flag_show_str_labels,
57
- bbox_likelihood_th,
58
- kpts_likelihood_th,
59
- font_style,
60
- font_size,
61
- keypt_color,
62
- marker_size,
63
- ):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
 
65
  if not flag_dlc_only:
66
- ############################################################
67
  # ### Run Megadetector
68
- md_results = predict_md(img_input,
69
- MD_models_dict[mega_model_input], #mega_model_input,
70
- size=640) #Image.fromarray(results.imgs[0])
 
 
71
 
72
  ################################################################
73
- # Obtain animal crops for bboxes with confidence above th
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
- if os.path.isdir(DLC_models_dict[dlc_model_input_str]) and \
85
- len(os.listdir(DLC_models_dict[dlc_model_input_str])) > 0:
86
- path_to_DLCmodel = DLC_models_dict[dlc_model_input_str]
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(DLC_models_dict[dlc_model_input_str],
93
- 'pose_cfg.yaml')
94
- with open(pose_cfg_path, "r") as stream:
95
- pose_cfg_dict = yaml.safe_load(stream)
96
- map_label_id_to_str = dict([(k,v) for k,v in zip([el[0] for el in pose_cfg_dict['all_joints']], # pose_cfg_dict['all_joints'] is a list of one-element lists,
97
- pose_cfg_dict['all_joints_names'])])
98
-
99
-
100
- ##############################################################
 
 
 
 
 
 
 
 
101
  # Run DLC and visualize results
102
- dlc_proc = Processor() #TODO: update deeplabcut.video_inference_superanimal() once merged
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(img_input,
113
- list_kpts_per_crop[0], # a numpy array with shape [num_keypoints, 2].
114
- map_label_id_to_str,
115
- flag_show_str_labels,
116
- use_normalized_coordinates=False,
117
- font_style=font_style,
118
- font_size=font_size,
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,dlc_model_input_str)
123
-
124
- return img_input, donw_file
 
 
 
 
 
 
 
 
125
 
126
  else:
127
  # Compute kpts for each crop
128
- list_kpts_per_crop = predict_dlc(list_crops,
129
- kpts_likelihood_th,
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(img_crop,
145
- kpts_crop, # a numpy array with shape [num_keypoints, 2].
146
- map_label_id_to_str,
147
- flag_show_str_labels,
148
- use_normalized_coordinates=False, # if True, then I should use md_results.xyxyn for list_kpts_crop
149
- font_style=font_style,
150
- font_size=font_size,
151
- keypt_color=keypt_color,
152
- marker_size=marker_size)
 
 
 
 
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 = md_results.xyxy[0].tolist()[ic]
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
- return img_background, download_file
 
 
 
 
 
 
 
 
 
171
 
 
 
 
172
 
173
 
174
  #########################################################
175
  # Define user interface and launch
176
- inputs = gradio_inputs_for_MD_DLC(list(MD_models_dict.keys()),
177
- list(DLC_models_dict.keys()))
178
- outputs = gradio_outputs_for_MD_DLC()
179
- [gr_title,
180
- gr_description,
181
- examples] = gradio_description_and_examples()
182
-
183
- # launch
184
- demo = gr.Interface(predict_pipeline,
185
- inputs=inputs,
186
- outputs=outputs,
187
- title=gr_title,
188
- description=gr_description,
189
- examples=examples,
190
- )
191
-
192
- demo.queue()
193
- demo.launch(theme=gr.themes.Default())
 
 
 
 
 
 
 
 
 
 
 
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(im,
20
- megadetector_model, #Megadet_Models[mega_model_input]
21
- size=640):
22
-
 
 
23
  # resize image
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
35
- 'custom', #model
36
- megadetector_model, # args for callable model
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'])
49
-
50
- return results
 
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
- yolo_results.ims[0].shape[0]))
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])) # int() should suffice?
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('animal')) and \
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) #Image.fromarray(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, Processor
6
 
7
 
8
  ##########################################
9
- def predict_dlc(list_np_crops,
10
- kpts_likelihood_th,
11
- dlc_model_folder,
12
- dlc_proc):
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
- # scale crop here?
23
- keypts_xyp = dlc_live.get_pose(crop) # third column is llk!
24
  # set kpts below threhsold to nan
25
-
26
- #pdb.set_trace()
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
- all_kypts.append(keypts_xyp)
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
- gradio
 
 
 
2
  gitpython>=3.1.30
3
  seaborn
4
- deeplabcut[modelzoo,tf]>=3.0.0rc14
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="superanimal_quadruped_dlcrnet",
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 DLClive only, directly on input image?",
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
- gr_keypt_color = gr.ColorPicker(
53
- value="#862db7",
54
- label="Choose color for keypoint label",
55
- )
56
-
57
- gr_labels_font_style = gr.Dropdown(
58
- choices=["amiko", "animals", "nature", "painter", "zen"],
59
- value="amiko",
60
- type="value",
61
- label="Select keypoint label font",
62
- )
63
-
64
- gr_slider_font_size = gr.Slider(
65
- minimum=5,
66
- maximum=30,
67
- value=8,
68
- step=1,
69
- label="Set font size",
70
- )
71
-
72
- gr_slider_marker_size = gr.Slider(
73
- minimum=1,
74
- maximum=20,
75
- value=9,
76
- step=1,
77
- label="Set marker size",
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
- gr_file_download = gr.File(label="Download JSON file")
98
- return [gr_image_output, gr_file_download]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99
 
100
 
101
  def gradio_description_and_examples():
102
- title = "DeepLabCut Model Zoo SuperAnimals"
103
  description = (
104
- "Test the SuperAnimal models from the "
105
- "<a href='http://www.mackenziemathislab.org/dlc-modelzoo'>"
106
- "DeepLabCut ModelZoo Project</a>, and read more on arXiv: "
107
- "https://arxiv.org/abs/2203.07436! Simply upload an image and see how it does. "
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
- "examples/dog.jpeg",
114
- "md_v5a",
115
- "superanimal_quadruped_dlcrnet",
116
- False,
117
- True,
118
- 0.5,
119
- 0.0,
120
- "amiko",
121
- 9,
122
- "#ff0000",
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 numpy as np
 
3
 
4
- from matplotlib import cm
5
- import matplotlib
6
- from PIL import Image, ImageColor, ImageFont, ImageDraw
7
  import numpy as np
8
- import pdb
9
- from datetime import date
 
10
  today = date.today()
11
- FONTS = {'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
  #########################################
18
  # Draw keypoints on image
19
- def draw_keypoints_on_image(image,
20
- keypoints,
21
- map_label_id_to_str,
22
- flag_show_str_labels,
23
- use_normalized_coordinates=True,
24
- font_style='amiko',
25
- font_size=8,
26
- keypt_color="#ff0000",
27
- marker_size=2,
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
- alpha = [k[2] for k in keypoints]
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
- #cmap = matplotlib.cm.get_cmap('hsv')
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
- if np.isnan(alpha[i]) == False :
74
- round_fill[3] = round(alpha[i] *255)
75
- #print(round_fill)
76
- #round_outline = [round(num*255) for num in list(cmap2(alpha[i]))[:3]]
77
- draw.ellipse([(keypoint_x - marker_size, keypoint_y - marker_size),
78
- (keypoint_x + marker_size, keypoint_y + marker_size)],
79
- fill=tuple(round_fill), outline= 'black', width=1) #fill and outline: [0,255]
 
 
 
 
 
 
 
 
 
 
 
80
 
81
  # add string labels around keypoints
82
  if flag_show_str_labels:
83
- font = ImageFont.truetype(FONTS[font_style],
84
- font_size)
85
- draw.text((keypoint_x + marker_size, keypoint_y + marker_size),#(0.5*im_width, 0.5*im_height), #-------
86
- map_label_id_to_str[i],
87
- ImageColor.getcolor(keypt_color, "RGB"), # rgb #
88
- font=font)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89
 
90
  #########################################
91
  # Draw bboxes on image
92
- def draw_bbox_w_text(img,
93
- results,
94
- font_style='amiko',
95
- font_size=8): #TODO: select color too?
96
- #pdb.set_trace()
97
- bbxyxy = results
98
- w, h = bbxyxy[2], bbxyxy[3]
99
- shape = [(bbxyxy[0], bbxyxy[1]), (w , h)]
100
- imgR = ImageDraw.Draw(img)
101
- imgR.rectangle(shape, outline ="red",width=5) ##bb for animal
102
-
103
- confidence = bbxyxy[4]
104
- string_bb = 'animal ' + str(round(confidence, 2))
105
- font = ImageFont.truetype(FONTS[font_style], font_size)
106
-
107
- text_size = font.getbbox(string_bb) # (h,w)
108
- position = (bbxyxy[0],bbxyxy[1] - text_size[1] -2 )
109
- left, top, right, bottom = imgR.textbbox(position, string_bb, font=font)
110
- imgR.rectangle((left, top-5, right+5, bottom+5), fill="red")
111
- imgR.text((bbxyxy[0] + 3 ,bbxyxy[1] - text_size[1] -2 ), string_bb, font=font, fill="black")
112
-
113
- return imgR
114
 
115
  ###########################################
116
- def save_results_as_json(md_results, dlc_outputs, map_dlc_label_id_to_str, thr,model,mega_model_input, path_to_output_file = 'download_predictions.json'):
 
117
 
118
- """
119
- Output detections as json file
120
 
121
- """
122
- # initialise dict to save to json
123
- info = {}
124
- info['date'] = str(today)
125
- info['MD_model'] = str(mega_model_input)
126
- # info from megaDetector
127
- info['file']= md_results.files[0]
128
- number_bb = len(md_results.xyxy[0].tolist())
129
- info['number_of_bb'] = number_bb
130
- # info from DLC
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 save_results_only_dlc(dlc_outputs,map_label_id_to_str,model,output_file = 'dowload_predictions_dlc.json'):
 
 
 
 
 
 
173
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
174
  """
175
- write json dlc output
176
- """
177
- info = {}
178
- info['date'] = str(today)
179
- labels = [n for n in map_label_id_to_str.values()]
180
- info['dlc_model'] = model
181
- kypts = []
182
- for s in dlc_outputs:
183
- aux1 = []
184
- for j in s:
185
- aux1.append(float(j))
186
 
187
- kypts.append(aux1)
188
- info['dlc_pred'] = dict(zip(labels,kypts))
 
 
 
 
 
 
 
 
 
 
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
- return output_file
 
 
 
 
 
 
 
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
+ ###########################################