C-Achard commited on
Commit
752f8f0
Β·
1 Parent(s): 935b48a

run pre commit hooks

Browse files
Files changed (8) hide show
  1. README.md +1 -1
  2. app.py +174 -177
  3. detection_utils.py +41 -53
  4. dlc_utils.py +6 -12
  5. pytorch_utils.py +27 -26
  6. requirements.txt +2 -1
  7. ui_utils.py +1 -1
  8. viz_utils.py +111 -97
README.md CHANGED
@@ -10,4 +10,4 @@ app_file: app.py
10
  pinned: false
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
10
  pinned: false
11
  ---
12
 
13
+ Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
app.py CHANGED
@@ -1,51 +1,55 @@
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
- import yaml
9
- import numpy as np
10
- from matplotlib import cm
11
- import gradio as gr
12
- import deeplabcut
13
- import dlclibrary
14
- import dlclive
15
- # import transformers
16
-
17
- from PIL import Image, ImageColor, ImageFont, ImageDraw
18
-
19
- from viz_utils import save_results_as_json, draw_keypoints_on_image, draw_bbox_w_text, save_results_only_dlc, save_results_pytorch
20
- from viz_utils import add_confidence_legend, keypoint_confidence_rows, save_annotated_image
21
- from detection_utils import predict_md, crop_animal_detections
22
- from dlc_utils import predict_dlc
23
- from pytorch_utils import predict_superanimal, load_superanimal, PYTORCH_MODELS
24
- from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
25
 
26
- from deeplabcut.utils import auxiliaryfunctions
 
 
27
  from dlclibrary.dlcmodelzoo.modelzoo_download import (
28
  download_huggingface_model,
29
- MODELOPTIONS,
30
  )
31
- from dlclive import DLCLive, Processor
 
 
 
32
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
 
34
  # TESTING (passes) download the SuperAnimal models:
35
- #model = 'superanimal_topviewmouse'
36
- #train_dir = 'DLC_models/sa-tvm'
37
- #download_huggingface_model(model, train_dir)
38
 
39
  # megadetector and dlc model look up
40
- MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
41
- 'md_v5b': "MD_models/md_v5b.0.0.pt"}
 
 
42
 
43
  BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
44
 
45
  # TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
46
- DLC_models_dict = {'superanimal_topviewmouse': ('superanimal_topviewmouse_dlcrnet', "DLC_models/sa-tvm"),
47
- 'superanimal_quadruped': ('superanimal_quadruped_dlcrnet', "DLC_models/sa-q")}
48
-
 
49
 
50
 
51
  #####################################################
@@ -59,99 +63,102 @@ def finalize_outputs(img_output, download_file, kpts_per_animal, map_label_id_to
59
 
60
 
61
  #####################################################
62
- def predict_pipeline_pytorch(img_input,
63
- superanimal,
64
- flag_dlc_only,
65
- flag_show_str_labels,
66
- bbox_likelihood_th,
67
- kpts_likelihood_th,
68
- font_style,
69
- font_size,
70
- keypt_color,
71
- marker_size,
72
- flag_color_by_confidence,
73
- ):
 
74
  # detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
75
- img_output, animals, bodyparts = predict_superanimal(img_input,
76
- superanimal,
77
- bbox_likelihood_th,
78
- kpts_likelihood_th,
79
- full_image=flag_dlc_only)
80
  map_label_id_to_str = dict(enumerate(bodyparts))
81
 
82
  for animal in animals:
83
- draw_keypoints_on_image(img_output,
84
- animal['kpts'],
85
- map_label_id_to_str,
86
- flag_show_str_labels,
87
- use_normalized_coordinates=False,
88
- font_style=font_style,
89
- font_size=font_size,
90
- keypt_color=keypt_color,
91
- marker_size=marker_size,
92
- color_by_confidence=flag_color_by_confidence)
 
 
93
  if not flag_dlc_only:
94
- draw_bbox_w_text(img_output,
95
- animal['bbox'],
96
- font_size=font_size)
97
 
98
  pose_model, detector = PYTORCH_MODELS[superanimal]
99
- download_file = save_results_pytorch(animals, map_label_id_to_str, superanimal,
100
- pose_model, None if flag_dlc_only else detector)
101
- return finalize_outputs(img_output, download_file,
102
- [animal['kpts'] for animal in animals], map_label_id_to_str,
103
- flag_color_by_confidence)
 
104
 
105
 
106
  #####################################################
107
- def predict_pipeline(img_input,
108
- backend,
109
- mega_model_input,
110
- dlc_model_input_str,
111
- flag_dlc_only,
112
- flag_show_str_labels,
113
- bbox_likelihood_th,
114
- kpts_likelihood_th,
115
- font_style,
116
- font_size,
117
- keypt_color,
118
- marker_size,
119
- flag_color_by_confidence,
120
- ):
 
121
 
122
  if backend == "PyTorch":
123
- return predict_pipeline_pytorch(img_input,
124
- dlc_model_input_str,
125
- flag_dlc_only,
126
- flag_show_str_labels,
127
- bbox_likelihood_th,
128
- kpts_likelihood_th,
129
- font_style,
130
- font_size,
131
- keypt_color,
132
- marker_size,
133
- flag_color_by_confidence)
 
 
134
 
135
  # TensorFlow (legacy): MegaDetector crops + DLCLive
136
  dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
137
 
138
  if not flag_dlc_only:
139
- ############################################################
140
  # ### Run Megadetector
141
- md_results = predict_md(img_input,
142
- MD_models_dict[mega_model_input], #mega_model_input,
143
- size=640) #Image.fromarray(results.imgs[0])
 
 
144
 
145
  ################################################################
146
  # Obtain animal crops (and their bboxes) with confidence above th
147
- list_crops, list_bboxes = crop_animal_detections(img_input,
148
- md_results,
149
- bbox_likelihood_th)
150
 
151
  ############################################################
152
 
153
- ## Get DLC model and label map
154
-
155
  # If model is found: do not download (previous execution is likely within same day)
156
  # TODO: can we ask the user whether to reload dlc model if a directory is found?
157
  path_to_DLCmodel = dlc_model_dir
@@ -159,124 +166,114 @@ def predict_pipeline(img_input,
159
  download_huggingface_model(dlc_model_name, path_to_DLCmodel)
160
 
161
  # extract map label ids to strings
162
- pose_cfg_path = os.path.join(dlc_model_dir,
163
- 'pose_cfg.yaml')
164
- with open(pose_cfg_path, "r") as stream:
165
- pose_cfg_dict = yaml.safe_load(stream)
166
- 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,
167
- pose_cfg_dict['all_joints_names'])])
168
-
169
-
170
- ##############################################################
 
 
 
 
 
 
 
 
171
  # Run DLC and visualize results
172
- dlc_proc = Processor() #TODO: update deeplabcut.video_inference_superanimal() once merged
173
 
174
  # if required: ignore MD crops and run DLC on full image [mostly for testing]
175
  if flag_dlc_only:
176
  # compute kpts on input img
177
- list_kpts_per_crop = predict_dlc([np.asarray(img_input)],
178
- kpts_likelihood_th,
179
- path_to_DLCmodel,
180
- dlc_proc)
181
  # draw kpts on input img #fix!
182
- draw_keypoints_on_image(img_input,
183
- list_kpts_per_crop[0], # a numpy array with shape [num_keypoints, 2].
184
- map_label_id_to_str,
185
- flag_show_str_labels,
186
- use_normalized_coordinates=False,
187
- font_style=font_style,
188
- font_size=font_size,
189
- keypt_color=keypt_color,
190
- marker_size=marker_size,
191
- color_by_confidence=flag_color_by_confidence)
192
-
193
- donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_name)
194
-
195
- return finalize_outputs(img_input, donw_file,
196
- [list_kpts_per_crop[0]], map_label_id_to_str,
197
- flag_color_by_confidence)
 
 
198
 
199
  else:
200
  # Compute kpts for each crop
201
- list_kpts_per_crop = predict_dlc(list_crops,
202
- kpts_likelihood_th,
203
- path_to_DLCmodel,
204
- dlc_proc)
205
-
206
  # resize input image to match megadetector output
207
- img_background = img_input.resize((md_results.ims[0].shape[1],
208
- md_results.ims[0].shape[0]))
209
-
210
- # draw keypoints on each crop and paste to background img
211
- for np_crop, kpts_crop, bb_per_animal in zip(list_crops,
212
- list_kpts_per_crop,
213
- list_bboxes):
214
 
 
 
215
  img_crop = Image.fromarray(np_crop)
216
 
217
  # Draw keypts on crop
218
- draw_keypoints_on_image(img_crop,
219
- kpts_crop, # a numpy array with shape [num_keypoints, 2].
220
- map_label_id_to_str,
221
- flag_show_str_labels,
222
- use_normalized_coordinates=False, # if True, then I should use md_results.xyxyn for list_kpts_crop
223
- font_style=font_style,
224
- font_size=font_size,
225
- keypt_color=keypt_color,
226
- marker_size=marker_size,
227
- color_by_confidence=flag_color_by_confidence)
 
 
228
 
229
  # Paste crop in original image
230
- img_background.paste(img_crop,
231
- box = tuple([int(t) for t in bb_per_animal[:2]]))
232
 
233
  # Plot bbox
234
- draw_bbox_w_text(img_background,
235
- bb_per_animal,
236
- font_size=font_size) # TODO: add selectable color for bbox?
237
-
238
 
239
  # Save detection results as json
240
- download_file = save_results_as_json(md_results,list_kpts_per_crop,list_bboxes,map_label_id_to_str,dlc_model_name,mega_model_input)
241
-
242
- return finalize_outputs(img_background, download_file,
243
- list_kpts_per_crop, map_label_id_to_str,
244
- flag_color_by_confidence)
245
 
 
 
 
246
 
247
 
248
  #########################################################
249
  # Define user interface and launch
250
- [gr_title,
251
- gr_description,
252
- examples] = gradio_description_and_examples()
253
 
254
  with gr.Blocks(title=gr_title) as demo:
255
  gr.Markdown(f"# {gr_title}\n{gr_description}")
256
  with gr.Row():
257
  with gr.Column():
258
- inputs = gradio_inputs_for_MD_DLC(BACKENDS,
259
- list(MD_models_dict.keys()),
260
- list(DLC_models_dict.keys()))
261
  run_button = gr.Button("Run", variant="primary")
262
  with gr.Column():
263
  outputs = gradio_outputs_for_MD_DLC()
264
 
265
  # the MegaDetector choice only applies to the TensorFlow (legacy) backend
266
  gr_backend_input, gr_mega_model_input = inputs[1], inputs[2]
267
- gr_backend_input.change(lambda backend: gr.update(visible=backend != "PyTorch"),
268
- inputs=gr_backend_input,
269
- outputs=gr_mega_model_input)
270
 
271
  run_button.click(predict_pipeline, inputs=inputs, outputs=outputs, api_name="predict")
272
 
273
  # cached on first click, so a failing download cannot block startup
274
- gr.Examples(examples,
275
- inputs=inputs,
276
- outputs=outputs,
277
- fn=predict_pipeline,
278
- cache_examples=True,
279
- cache_mode="lazy")
280
 
281
  # download and build the default model while the app starts; a request arriving
282
  # earlier waits on the same lock instead of downloading again
 
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 gradio_description_and_examples, gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC
24
+ from viz_utils import (
25
+ add_confidence_legend,
26
+ draw_bbox_w_text,
27
+ draw_keypoints_on_image,
28
+ keypoint_confidence_rows,
29
+ save_annotated_image,
30
+ save_results_as_json,
31
+ save_results_only_dlc,
32
+ save_results_pytorch,
33
+ )
34
 
35
  # TESTING (passes) download the SuperAnimal models:
36
+ # model = 'superanimal_topviewmouse'
37
+ # train_dir = 'DLC_models/sa-tvm'
38
+ # download_huggingface_model(model, train_dir)
39
 
40
  # megadetector and dlc model look up
41
+ MD_models_dict = {
42
+ "md_v5a": "MD_models/md_v5a.0.0.pt", #
43
+ "md_v5b": "MD_models/md_v5b.0.0.pt",
44
+ }
45
 
46
  BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
47
 
48
  # TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
49
+ DLC_models_dict = {
50
+ "superanimal_topviewmouse": ("superanimal_topviewmouse_dlcrnet", "DLC_models/sa-tvm"),
51
+ "superanimal_quadruped": ("superanimal_quadruped_dlcrnet", "DLC_models/sa-q"),
52
+ }
53
 
54
 
55
  #####################################################
 
63
 
64
 
65
  #####################################################
66
+ def predict_pipeline_pytorch(
67
+ img_input,
68
+ superanimal,
69
+ flag_dlc_only,
70
+ flag_show_str_labels,
71
+ bbox_likelihood_th,
72
+ kpts_likelihood_th,
73
+ font_style,
74
+ font_size,
75
+ keypt_color,
76
+ marker_size,
77
+ flag_color_by_confidence,
78
+ ):
79
  # detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
80
+ img_output, animals, bodyparts = predict_superanimal(
81
+ img_input, superanimal, bbox_likelihood_th, kpts_likelihood_th, full_image=flag_dlc_only
82
+ )
 
 
83
  map_label_id_to_str = dict(enumerate(bodyparts))
84
 
85
  for animal in animals:
86
+ draw_keypoints_on_image(
87
+ img_output,
88
+ animal["kpts"],
89
+ map_label_id_to_str,
90
+ flag_show_str_labels,
91
+ use_normalized_coordinates=False,
92
+ font_style=font_style,
93
+ font_size=font_size,
94
+ keypt_color=keypt_color,
95
+ marker_size=marker_size,
96
+ color_by_confidence=flag_color_by_confidence,
97
+ )
98
  if not flag_dlc_only:
99
+ draw_bbox_w_text(img_output, animal["bbox"], font_size=font_size)
 
 
100
 
101
  pose_model, detector = PYTORCH_MODELS[superanimal]
102
+ download_file = save_results_pytorch(
103
+ animals, map_label_id_to_str, superanimal, pose_model, None if flag_dlc_only else detector
104
+ )
105
+ return finalize_outputs(
106
+ img_output, download_file, [animal["kpts"] for animal in animals], map_label_id_to_str, flag_color_by_confidence
107
+ )
108
 
109
 
110
  #####################################################
111
+ def predict_pipeline(
112
+ img_input,
113
+ backend,
114
+ mega_model_input,
115
+ dlc_model_input_str,
116
+ flag_dlc_only,
117
+ flag_show_str_labels,
118
+ bbox_likelihood_th,
119
+ kpts_likelihood_th,
120
+ font_style,
121
+ font_size,
122
+ keypt_color,
123
+ marker_size,
124
+ flag_color_by_confidence,
125
+ ):
126
 
127
  if backend == "PyTorch":
128
+ return predict_pipeline_pytorch(
129
+ img_input,
130
+ dlc_model_input_str,
131
+ flag_dlc_only,
132
+ flag_show_str_labels,
133
+ bbox_likelihood_th,
134
+ kpts_likelihood_th,
135
+ font_style,
136
+ font_size,
137
+ keypt_color,
138
+ marker_size,
139
+ flag_color_by_confidence,
140
+ )
141
 
142
  # TensorFlow (legacy): MegaDetector crops + DLCLive
143
  dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
144
 
145
  if not flag_dlc_only:
146
+ ############################################################
147
  # ### Run Megadetector
148
+ md_results = predict_md(
149
+ img_input,
150
+ MD_models_dict[mega_model_input], # mega_model_input,
151
+ size=640,
152
+ ) # Image.fromarray(results.imgs[0])
153
 
154
  ################################################################
155
  # Obtain animal crops (and their bboxes) with confidence above th
156
+ list_crops, list_bboxes = crop_animal_detections(img_input, md_results, bbox_likelihood_th)
 
 
157
 
158
  ############################################################
159
 
160
+ ## Get DLC model and label map
161
+
162
  # If model is found: do not download (previous execution is likely within same day)
163
  # TODO: can we ask the user whether to reload dlc model if a directory is found?
164
  path_to_DLCmodel = dlc_model_dir
 
166
  download_huggingface_model(dlc_model_name, path_to_DLCmodel)
167
 
168
  # extract map label ids to strings
169
+ pose_cfg_path = os.path.join(dlc_model_dir, "pose_cfg.yaml")
170
+ with open(pose_cfg_path) as stream:
171
+ pose_cfg_dict = yaml.safe_load(stream)
172
+ map_label_id_to_str = dict(
173
+ [
174
+ (k, v)
175
+ for k, v in zip(
176
+ [
177
+ el[0] for el in pose_cfg_dict["all_joints"]
178
+ ], # pose_cfg_dict['all_joints'] is a list of one-element lists,
179
+ pose_cfg_dict["all_joints_names"],
180
+ strict=True,
181
+ )
182
+ ]
183
+ )
184
+
185
+ ##############################################################
186
  # Run DLC and visualize results
187
+ dlc_proc = Processor() # TODO: update deeplabcut.video_inference_superanimal() once merged
188
 
189
  # if required: ignore MD crops and run DLC on full image [mostly for testing]
190
  if flag_dlc_only:
191
  # compute kpts on input img
192
+ list_kpts_per_crop = predict_dlc([np.asarray(img_input)], kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
 
 
 
193
  # draw kpts on input img #fix!
194
+ draw_keypoints_on_image(
195
+ img_input,
196
+ list_kpts_per_crop[0], # a numpy array with shape [num_keypoints, 2].
197
+ map_label_id_to_str,
198
+ flag_show_str_labels,
199
+ use_normalized_coordinates=False,
200
+ font_style=font_style,
201
+ font_size=font_size,
202
+ keypt_color=keypt_color,
203
+ marker_size=marker_size,
204
+ color_by_confidence=flag_color_by_confidence,
205
+ )
206
+
207
+ donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str, dlc_model_name)
208
+
209
+ return finalize_outputs(
210
+ img_input, donw_file, [list_kpts_per_crop[0]], map_label_id_to_str, flag_color_by_confidence
211
+ )
212
 
213
  else:
214
  # Compute kpts for each crop
215
+ list_kpts_per_crop = predict_dlc(list_crops, kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
216
+
 
 
 
217
  # resize input image to match megadetector output
218
+ img_background = img_input.resize((md_results.ims[0].shape[1], md_results.ims[0].shape[0]))
 
 
 
 
 
 
219
 
220
+ # draw keypoints on each crop and paste to background img
221
+ for np_crop, kpts_crop, bb_per_animal in zip(list_crops, list_kpts_per_crop, list_bboxes, strict=True):
222
  img_crop = Image.fromarray(np_crop)
223
 
224
  # Draw keypts on crop
225
+ draw_keypoints_on_image(
226
+ img_crop,
227
+ kpts_crop, # a numpy array with shape [num_keypoints, 2].
228
+ map_label_id_to_str,
229
+ flag_show_str_labels,
230
+ use_normalized_coordinates=False, # if True, then I should use md_results.xyxyn for list_kpts_crop
231
+ font_style=font_style,
232
+ font_size=font_size,
233
+ keypt_color=keypt_color,
234
+ marker_size=marker_size,
235
+ color_by_confidence=flag_color_by_confidence,
236
+ )
237
 
238
  # Paste crop in original image
239
+ img_background.paste(img_crop, box=tuple([int(t) for t in bb_per_animal[:2]]))
 
240
 
241
  # Plot bbox
242
+ draw_bbox_w_text(img_background, bb_per_animal, font_size=font_size) # TODO: add selectable color for bbox?
 
 
 
243
 
244
  # Save detection results as json
245
+ download_file = save_results_as_json(
246
+ md_results, list_kpts_per_crop, list_bboxes, map_label_id_to_str, dlc_model_name, mega_model_input
247
+ )
 
 
248
 
249
+ return finalize_outputs(
250
+ img_background, download_file, list_kpts_per_crop, map_label_id_to_str, flag_color_by_confidence
251
+ )
252
 
253
 
254
  #########################################################
255
  # Define user interface and launch
256
+ [gr_title, gr_description, examples] = gradio_description_and_examples()
 
 
257
 
258
  with gr.Blocks(title=gr_title) as demo:
259
  gr.Markdown(f"# {gr_title}\n{gr_description}")
260
  with gr.Row():
261
  with gr.Column():
262
+ inputs = gradio_inputs_for_MD_DLC(BACKENDS, list(MD_models_dict.keys()), list(DLC_models_dict.keys()))
 
 
263
  run_button = gr.Button("Run", variant="primary")
264
  with gr.Column():
265
  outputs = gradio_outputs_for_MD_DLC()
266
 
267
  # the MegaDetector choice only applies to the TensorFlow (legacy) backend
268
  gr_backend_input, gr_mega_model_input = inputs[1], inputs[2]
269
+ gr_backend_input.change(
270
+ lambda backend: gr.update(visible=backend != "PyTorch"), inputs=gr_backend_input, outputs=gr_mega_model_input
271
+ )
272
 
273
  run_button.click(predict_pipeline, inputs=inputs, outputs=outputs, api_name="predict")
274
 
275
  # cached on first click, so a failing download cannot block startup
276
+ gr.Examples(examples, inputs=inputs, outputs=outputs, fn=predict_pipeline, cache_examples=True, cache_mode="lazy")
 
 
 
 
 
277
 
278
  # download and build the default model while the app starts; a request arriving
279
  # earlier waits on the same lock instead of downloading again
detection_utils.py CHANGED
@@ -1,86 +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: yolov5's select_device expects a CUDA index ('0') or 'cpu', not 'cuda'
28
- md_device = '0' if torch.cuda.is_available() else 'cpu'
29
-
30
- # megadetector
31
- MD_model = torch.hub.load('ultralytics/yolov5', # repo_or_dir
32
- 'custom', #model
33
- megadetector_model, # args for callable model
34
- skip_validation=True, # avoid GitHub API rate limit (403)
35
- device=md_device,
36
- trust_repo=True
37
- )
 
38
 
39
  ## detect objects
40
- results = MD_model(im) # inference # vars(results).keys()= dict_keys(['imgs', 'pred', 'names', 'files', 'times', 'xyxy', 'xywh', 'xyxyn', 'xywhn', 'n', 't', 's'])
41
-
42
- return results
 
43
 
44
 
45
  ##########################################
46
- def crop_animal_detections(img_in,
47
- yolo_results,
48
- likelihood_th):
49
 
50
  ## Extract animal crops
51
  list_labels_as_str = [i for i in yolo_results.names.values()] # ['animal', 'person', 'vehicle']
52
  list_np_animal_crops = []
53
- list_animal_bboxes = [] # detection rows [x1,y1,x2,y2,conf,label] matching each crop
54
 
55
  # image to crop (scale as input for megadetector)
56
- img_in = img_in.resize((yolo_results.ims[0].shape[1],
57
- yolo_results.ims[0].shape[0]))
58
- # for every detection in the img
59
  for det_array in yolo_results.xyxy:
60
-
61
  # for every detection
62
  for j in range(det_array.shape[0]):
63
-
64
  # compute coords around bbox rounded to the nearest integer (for pasting later)
65
- xmin_rd = int(math.floor(det_array[j,0])) # int() should suffice?
66
- ymin_rd = int(math.floor(det_array[j,1]))
67
 
68
- xmax_rd = int(math.ceil(det_array[j,2]))
69
- ymax_rd = int(math.ceil(det_array[j,3]))
70
 
71
- pred_llk = det_array[j,4]
72
- pred_label = det_array[j,5]
73
  # keep animal crops above threshold
74
- if (pred_label == list_labels_as_str.index('animal')) and \
75
- (pred_llk >= likelihood_th):
76
  area = (xmin_rd, ymin_rd, xmax_rd, ymax_rd)
77
 
78
- #pdb.set_trace()
79
- crop = img_in.crop(area) #Image.fromarray(img_in).crop(area)
80
  crop_np = np.asarray(crop)
81
 
82
  # add to list
83
  list_np_animal_crops.append(crop_np)
84
- list_animal_bboxes.append(det_array[j,:].tolist())
85
 
86
- return list_np_animal_crops, list_animal_bboxes
 
 
 
 
 
 
 
 
 
 
 
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,16 +1,10 @@
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
  # no animal detected: nothing to run
15
  if len(list_np_crops) == 0:
16
  return []
@@ -25,10 +19,10 @@ def predict_dlc(list_np_crops,
25
 
26
  list_kpts_per_crop = []
27
  for crop in list_np_crops:
28
- keypts_xyp = dlc_live.get_pose(crop) # third column is llk!
29
  # set kpts below threhsold to nan
30
- keypts_xyp[keypts_xyp[:,-1] < kpts_likelihood_th,:] = np.nan
31
  # add kpts of this crop to list
32
  list_kpts_per_crop.append(keypts_xyp)
33
 
34
- 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 []
 
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
pytorch_utils.py CHANGED
@@ -6,10 +6,11 @@ from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners
6
  from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
7
  from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path
8
 
9
-
10
  # SuperAnimal (pose model, detector) used by the PyTorch backend
11
- PYTORCH_MODELS = {'superanimal_quadruped': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2'),
12
- 'superanimal_topviewmouse': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2')}
 
 
13
 
14
  MAX_INDIVIDUALS = 10
15
  MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
@@ -27,11 +28,13 @@ def load_superanimal(superanimal, device="auto"):
27
  with _build_lock:
28
  if superanimal not in _runners:
29
  pose_model, detector = PYTORCH_MODELS[superanimal]
30
- cfg = PoseConfig.build_for_superanimal_inference(superanimal,
31
- model_name=pose_model,
32
- detector_name=detector,
33
- max_individuals=MAX_INDIVIDUALS,
34
- device=device)
 
 
35
  # keep low-score boxes: the UI threshold filters them afterwards
36
  cfg["detector"]["model"]["box_score_thresh"] = 0.05
37
  pose_runner, detector_runner = get_inference_runners(
@@ -41,10 +44,12 @@ def load_superanimal(superanimal, device="auto"):
41
  max_individuals=MAX_INDIVIDUALS,
42
  inference_cfg={"multithreading": {"enabled": False}},
43
  )
44
- _runners[superanimal] = {"pose": pose_runner,
45
- "detector": detector_runner,
46
- "bodyparts": list(cfg["metadata"]["bodyparts"]),
47
- "lock": threading.Lock()} # runners are not thread-safe
 
 
48
  return _runners[superanimal]
49
 
50
 
@@ -57,11 +62,7 @@ def resize_max_side(img, max_size=MAX_IMAGE_SIZE):
57
 
58
 
59
  ##########################################
60
- def predict_superanimal(img_input,
61
- superanimal,
62
- bbox_likelihood_th,
63
- kpts_likelihood_th,
64
- full_image=False):
65
  """Detect animals and estimate their pose with a PyTorch SuperAnimal model.
66
 
67
  Returns the (resized) RGB image the predictions refer to, the list of animals
@@ -76,13 +77,14 @@ def predict_superanimal(img_input,
76
  if full_image:
77
  # skip the detector and treat the whole image as one animal
78
  h, w = img_np.shape[:2]
79
- detections = {"bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
80
- "bbox_scores": np.array([1.0], dtype=np.float32)}
 
 
81
  else:
82
  detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
83
  keep = detections["bbox_scores"] >= bbox_likelihood_th
84
- detections = {"bboxes": detections["bboxes"][keep],
85
- "bbox_scores": detections["bbox_scores"][keep]}
86
 
87
  if len(detections["bboxes"]) == 0:
88
  return img, [], runners["bodyparts"]
@@ -91,14 +93,13 @@ def predict_superanimal(img_input,
91
 
92
  animals = []
93
  # outputs are padded to MAX_INDIVIDUALS with -1
94
- for kpts, (x, y, w, h), score in zip(predictions["bodyparts"],
95
- predictions["bboxes"],
96
- predictions["bbox_scores"]):
97
  if score < 0:
98
  continue
99
  kpts = kpts.astype(float)
100
  kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
101
- animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)],
102
- "kpts": kpts})
103
 
104
  return img, animals, runners["bodyparts"]
 
6
  from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
7
  from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path
8
 
 
9
  # SuperAnimal (pose model, detector) used by the PyTorch backend
10
+ PYTORCH_MODELS = {
11
+ "superanimal_quadruped": ("hrnet_w32", "fasterrcnn_resnet50_fpn_v2"),
12
+ "superanimal_topviewmouse": ("hrnet_w32", "fasterrcnn_resnet50_fpn_v2"),
13
+ }
14
 
15
  MAX_INDIVIDUALS = 10
16
  MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
 
28
  with _build_lock:
29
  if superanimal not in _runners:
30
  pose_model, detector = PYTORCH_MODELS[superanimal]
31
+ cfg = PoseConfig.build_for_superanimal_inference(
32
+ superanimal,
33
+ model_name=pose_model,
34
+ detector_name=detector,
35
+ max_individuals=MAX_INDIVIDUALS,
36
+ device=device,
37
+ )
38
  # keep low-score boxes: the UI threshold filters them afterwards
39
  cfg["detector"]["model"]["box_score_thresh"] = 0.05
40
  pose_runner, detector_runner = get_inference_runners(
 
44
  max_individuals=MAX_INDIVIDUALS,
45
  inference_cfg={"multithreading": {"enabled": False}},
46
  )
47
+ _runners[superanimal] = {
48
+ "pose": pose_runner,
49
+ "detector": detector_runner,
50
+ "bodyparts": list(cfg["metadata"]["bodyparts"]),
51
+ "lock": threading.Lock(),
52
+ } # runners are not thread-safe
53
  return _runners[superanimal]
54
 
55
 
 
62
 
63
 
64
  ##########################################
65
+ def predict_superanimal(img_input, superanimal, bbox_likelihood_th, kpts_likelihood_th, full_image=False):
 
 
 
 
66
  """Detect animals and estimate their pose with a PyTorch SuperAnimal model.
67
 
68
  Returns the (resized) RGB image the predictions refer to, the list of animals
 
77
  if full_image:
78
  # skip the detector and treat the whole image as one animal
79
  h, w = img_np.shape[:2]
80
+ detections = {
81
+ "bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
82
+ "bbox_scores": np.array([1.0], dtype=np.float32),
83
+ }
84
  else:
85
  detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
86
  keep = detections["bbox_scores"] >= bbox_likelihood_th
87
+ detections = {"bboxes": detections["bboxes"][keep], "bbox_scores": detections["bbox_scores"][keep]}
 
88
 
89
  if len(detections["bboxes"]) == 0:
90
  return img, [], runners["bodyparts"]
 
93
 
94
  animals = []
95
  # outputs are padded to MAX_INDIVIDUALS with -1
96
+ for kpts, (x, y, w, h), score in zip(
97
+ predictions["bodyparts"], predictions["bboxes"], predictions["bbox_scores"], strict=True
98
+ ):
99
  if score < 0:
100
  continue
101
  kpts = kpts.astype(float)
102
  kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
103
+ animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)], "kpts": kpts})
 
104
 
105
  return img, animals, runners["bodyparts"]
requirements.txt CHANGED
@@ -1,3 +1,4 @@
 
1
  gradio
2
  gitpython>=3.1.30
3
  seaborn
@@ -7,4 +8,4 @@ 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
  gradio
3
  gitpython>=3.1.30
4
  seaborn
 
8
  dlclibrary
9
  humanfriendly
10
  psutil
11
+ ultralytics
ui_utils.py CHANGED
@@ -134,4 +134,4 @@ def gradio_description_and_examples():
134
  for image in ("examples/dog.jpeg", "examples/cat.jpg")
135
  ]
136
 
137
- return [title, description, examples]
 
134
  for image in ("examples/dog.jpeg", "examples/cat.jpg")
135
  ]
136
 
137
+ return [title, description, examples]
viz_utils.py CHANGED
@@ -1,32 +1,34 @@
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
- color_by_confidence=True,
29
- ):
 
30
  """Draws keypoints on an image.
31
  Modified from:
32
  https://www.programcreek.com/python/?code=fjchange%2Fobject_centric_VAD%2Fobject_centric_VAD-master%2Fobject_detection%2Futils%2Fvisualization_utils.py
@@ -40,10 +42,10 @@ def draw_keypoints_on_image(image,
40
  use_normalized_coordinates: if True (default), treat keypoint values as
41
  relative to the image. Otherwise treat them as absolute.
42
 
43
-
44
  """
45
  # get a drawing context
46
- draw = ImageDraw.Draw(image,"RGBA")
47
 
48
  im_width, im_height = image.size
49
  keypoints_x = [k[0] for k in keypoints]
@@ -54,10 +56,10 @@ def draw_keypoints_on_image(image,
54
  if use_normalized_coordinates:
55
  keypoints_x = tuple([im_width * x for x in keypoints_x])
56
  keypoints_y = tuple([im_height * y for y in keypoints_y])
57
-
58
- #cmap = matplotlib.cm.get_cmap('hsv')
59
  # draw ellipses around keypoints
60
- for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y)):
61
  # handling potential nans in the keypoints
62
  if np.isnan(keypoint_x).any():
63
  continue
@@ -71,22 +73,30 @@ def draw_keypoints_on_image(image,
71
  round_fill = list(cm.viridis(i / max(len(keypoints) - 1, 1), bytes=True))
72
  round_fill[3] = round(confidence * 255)
73
  round_fill = tuple(round_fill)
74
- draw.ellipse([(keypoint_x - marker_size, keypoint_y - marker_size),
75
- (keypoint_x + marker_size, keypoint_y + marker_size)],
76
- fill=tuple(round_fill), outline= 'black', width=1) #fill and outline: [0,255]
 
 
 
 
 
 
77
 
78
  # add string labels around keypoints
79
  if flag_show_str_labels:
80
- font = ImageFont.truetype(FONTS[font_style],
81
- font_size)
82
- draw.text((keypoint_x + marker_size, keypoint_y + marker_size),#(0.5*im_width, 0.5*im_height), #-------
83
- map_label_id_to_str[i],
84
- ImageColor.getcolor(keypt_color, "RGB"), # rgb #
85
- font=font)
 
 
86
 
87
  #########################################
88
  # Legend for the keypoint confidence colors
89
- def add_confidence_legend(image, font_style='amiko'):
90
  """Returns the image with a white band below it holding a viridis strip (confidence 0 to 1).
91
 
92
  The band is added below the image, so the legend never covers it and keypoint
@@ -105,8 +115,7 @@ def add_confidence_legend(image, font_style='amiko'):
105
  x0 = im_width - strip_w - 2 * margin
106
  y0 = im_height + margin
107
  for dx in range(strip_w):
108
- draw.line([(x0 + dx, y0), (x0 + dx, y0 + strip_h)],
109
- fill=cm.viridis(dx / (strip_w - 1), bytes=True))
110
  label_y = y0 + strip_h + margin // 2
111
  draw.text((x0, label_y), "0", fill="black", font=font)
112
  draw.text((x0 + strip_w, label_y), "1", fill="black", font=font, anchor="ra")
@@ -128,39 +137,44 @@ def keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str):
128
 
129
  #########################################
130
  # Save the annotated image for download
131
- def save_annotated_image(image, path_to_output_file='download_annotated.png'):
132
  image.save(path_to_output_file)
133
  return path_to_output_file
134
 
135
 
136
  #########################################
137
  # Draw bboxes on image
138
- def draw_bbox_w_text(img,
139
- results,
140
- font_style='amiko',
141
- font_size=8): #TODO: select color too?
142
- #pdb.set_trace()
143
  bbxyxy = results
144
  w, h = bbxyxy[2], bbxyxy[3]
145
- shape = [(bbxyxy[0], bbxyxy[1]), (w , h)]
146
- imgR = ImageDraw.Draw(img)
147
- imgR.rectangle(shape, outline ="red",width=5) ##bb for animal
148
 
149
  confidence = bbxyxy[4]
150
- string_bb = 'animal ' + str(round(confidence, 2))
151
- font = ImageFont.truetype(FONTS[font_style], font_size)
152
 
153
- text_size = font.getbbox(string_bb) # (h,w)
154
- position = (bbxyxy[0],bbxyxy[1] - text_size[1] -2 )
155
  left, top, right, bottom = imgR.textbbox(position, string_bb, font=font)
156
- imgR.rectangle((left, top-5, right+5, bottom+5), fill="red")
157
- imgR.text((bbxyxy[0] + 3 ,bbxyxy[1] - text_size[1] -2 ), string_bb, font=font, fill="black")
158
 
159
  return imgR
160
 
161
- ###########################################
162
- def save_results_as_json(md_results, dlc_outputs, animal_bboxes, map_dlc_label_id_to_str, model,mega_model_input, path_to_output_file = 'download_predictions.json'):
163
 
 
 
 
 
 
 
 
 
 
 
164
  """
165
  Output detections as json file
166
 
@@ -168,26 +182,26 @@ def save_results_as_json(md_results, dlc_outputs, animal_bboxes, map_dlc_label_i
168
  """
169
  # initialise dict to save to json
170
  info = {}
171
- info['date'] = str(today)
172
- info['MD_model'] = str(mega_model_input)
173
  # info from megaDetector
174
- info['file']= md_results.files[0]
175
  number_bb = len(md_results.xyxy[0].tolist())
176
- info['number_of_bb'] = number_bb
177
  # info from DLC
178
- info['dlc_model'] = model
179
  labels = [n for n in map_dlc_label_id_to_str.values()]
180
 
181
  # define aux dict for every animal bounding box above threshold
182
  for i in range(len(dlc_outputs)):
183
- aux={}
184
  # MD output
185
- corner_x1,corner_y1,corner_x2,corner_y2,confidence, _ = animal_bboxes[i]
186
- aux['corner_1'] = (corner_x1,corner_y1)
187
- aux['corner_2'] = (corner_x2,corner_y2)
188
- aux['predict MD'] = md_results.names[0]
189
- aux['confidence MD'] = confidence
190
-
191
  # DLC output
192
  kypts = []
193
  for s in dlc_outputs[i]:
@@ -196,26 +210,25 @@ def save_results_as_json(md_results, dlc_outputs, animal_bboxes, map_dlc_label_i
196
  aux1.append(float(j))
197
 
198
  kypts.append(aux1)
199
- aux['dlc_pred'] = dict(zip(labels,kypts))
200
- info['bb_' + str(i) ]=aux
201
 
202
  # save dict as json
203
- with open(path_to_output_file, 'w') as f:
204
  json.dump(info, f, indent=1)
205
- print('Output file saved at {}'.format(path_to_output_file))
206
 
207
  return path_to_output_file
208
 
209
 
210
- def save_results_only_dlc(dlc_outputs,map_label_id_to_str,model,output_file = 'dowload_predictions_dlc.json'):
211
-
212
  """
213
  write json dlc output
214
  """
215
  info = {}
216
- info['date'] = str(today)
217
  labels = [n for n in map_label_id_to_str.values()]
218
- info['dlc_model'] = model
219
  kypts = []
220
  for s in dlc_outputs:
221
  aux1 = []
@@ -223,17 +236,18 @@ def save_results_only_dlc(dlc_outputs,map_label_id_to_str,model,output_file = 'd
223
  aux1.append(float(j))
224
 
225
  kypts.append(aux1)
226
- info['dlc_pred'] = dict(zip(labels,kypts))
227
 
228
- with open(output_file, 'w') as f:
229
  json.dump(info, f, indent=1)
230
- print('Output file saved at {}'.format(output_file))
231
 
232
  return output_file
233
 
234
 
235
- def save_results_pytorch(animals, map_label_id_to_str, model, pose_model, detector, path_to_output_file = 'download_predictions.json'):
236
-
 
237
  """
238
  Output PyTorch SuperAnimal predictions as json file (same layout as save_results_as_json)
239
 
@@ -241,28 +255,28 @@ def save_results_pytorch(animals, map_label_id_to_str, model, pose_model, detect
241
  detector: None if the detector was skipped (whole image used as one animal)
242
  """
243
  info = {}
244
- info['date'] = str(today)
245
- info['backend'] = 'pytorch'
246
- info['dlc_model'] = model
247
- info['pose_model'] = pose_model
248
- info['detector'] = detector
249
- info['number_of_bb'] = len(animals)
250
  labels = [n for n in map_label_id_to_str.values()]
251
 
252
  for i, animal in enumerate(animals):
253
- corner_x1, corner_y1, corner_x2, corner_y2, confidence = animal['bbox']
254
  aux = {}
255
- aux['corner_1'] = (corner_x1, corner_y1)
256
- aux['corner_2'] = (corner_x2, corner_y2)
257
- aux['confidence'] = confidence
258
- aux['dlc_pred'] = dict(zip(labels, [[float(v) for v in kpt] for kpt in animal['kpts']]))
259
- info['bb_' + str(i)] = aux
260
 
261
- with open(path_to_output_file, 'w') as f:
262
  json.dump(info, f, indent=1)
263
- print('Output file saved at {}'.format(path_to_output_file))
264
 
265
  return path_to_output_file
266
 
267
 
268
- ###########################################
 
1
+ import json
2
+ from datetime import date
3
 
 
 
 
4
  import numpy as np
5
+ from matplotlib import cm
6
+ from PIL import Image, ImageColor, ImageDraw, ImageFont
7
+
8
  today = date.today()
9
+ FONTS = {
10
+ "amiko": "fonts/Amiko-Regular.ttf",
11
+ "nature": "fonts/LoveNature.otf",
12
+ "painter": "fonts/PainterDecorator.otf",
13
+ "animals": "fonts/UncialAnimals.ttf",
14
+ "zen": "fonts/ZEN.TTF",
15
+ }
16
+
17
 
18
  #########################################
19
  # Draw keypoints on image
20
+ def draw_keypoints_on_image(
21
+ image,
22
+ keypoints,
23
+ map_label_id_to_str,
24
+ flag_show_str_labels,
25
+ use_normalized_coordinates=True,
26
+ font_style="amiko",
27
+ font_size=8,
28
+ keypt_color="#ff0000",
29
+ marker_size=2,
30
+ color_by_confidence=True,
31
+ ):
32
  """Draws keypoints on an image.
33
  Modified from:
34
  https://www.programcreek.com/python/?code=fjchange%2Fobject_centric_VAD%2Fobject_centric_VAD-master%2Fobject_detection%2Futils%2Fvisualization_utils.py
 
42
  use_normalized_coordinates: if True (default), treat keypoint values as
43
  relative to the image. Otherwise treat them as absolute.
44
 
45
+
46
  """
47
  # get a drawing context
48
+ draw = ImageDraw.Draw(image, "RGBA")
49
 
50
  im_width, im_height = image.size
51
  keypoints_x = [k[0] for k in keypoints]
 
56
  if use_normalized_coordinates:
57
  keypoints_x = tuple([im_width * x for x in keypoints_x])
58
  keypoints_y = tuple([im_height * y for y in keypoints_y])
59
+
60
+ # cmap = matplotlib.cm.get_cmap('hsv')
61
  # draw ellipses around keypoints
62
+ for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y, strict=True)):
63
  # handling potential nans in the keypoints
64
  if np.isnan(keypoint_x).any():
65
  continue
 
73
  round_fill = list(cm.viridis(i / max(len(keypoints) - 1, 1), bytes=True))
74
  round_fill[3] = round(confidence * 255)
75
  round_fill = tuple(round_fill)
76
+ draw.ellipse(
77
+ [
78
+ (keypoint_x - marker_size, keypoint_y - marker_size),
79
+ (keypoint_x + marker_size, keypoint_y + marker_size),
80
+ ],
81
+ fill=tuple(round_fill),
82
+ outline="black",
83
+ width=1,
84
+ ) # fill and outline: [0,255]
85
 
86
  # add string labels around keypoints
87
  if flag_show_str_labels:
88
+ font = ImageFont.truetype(FONTS[font_style], font_size)
89
+ draw.text(
90
+ (keypoint_x + marker_size, keypoint_y + marker_size), # (0.5*im_width, 0.5*im_height), #-------
91
+ map_label_id_to_str[i],
92
+ ImageColor.getcolor(keypt_color, "RGB"), # rgb #
93
+ font=font,
94
+ )
95
+
96
 
97
  #########################################
98
  # Legend for the keypoint confidence colors
99
+ def add_confidence_legend(image, font_style="amiko"):
100
  """Returns the image with a white band below it holding a viridis strip (confidence 0 to 1).
101
 
102
  The band is added below the image, so the legend never covers it and keypoint
 
115
  x0 = im_width - strip_w - 2 * margin
116
  y0 = im_height + margin
117
  for dx in range(strip_w):
118
+ draw.line([(x0 + dx, y0), (x0 + dx, y0 + strip_h)], fill=cm.viridis(dx / (strip_w - 1), bytes=True))
 
119
  label_y = y0 + strip_h + margin // 2
120
  draw.text((x0, label_y), "0", fill="black", font=font)
121
  draw.text((x0 + strip_w, label_y), "1", fill="black", font=font, anchor="ra")
 
137
 
138
  #########################################
139
  # Save the annotated image for download
140
+ def save_annotated_image(image, path_to_output_file="download_annotated.png"):
141
  image.save(path_to_output_file)
142
  return path_to_output_file
143
 
144
 
145
  #########################################
146
  # Draw bboxes on image
147
+ def draw_bbox_w_text(img, results, font_style="amiko", font_size=8): # TODO: select color too?
148
+ # pdb.set_trace()
 
 
 
149
  bbxyxy = results
150
  w, h = bbxyxy[2], bbxyxy[3]
151
+ shape = [(bbxyxy[0], bbxyxy[1]), (w, h)]
152
+ imgR = ImageDraw.Draw(img)
153
+ imgR.rectangle(shape, outline="red", width=5) ##bb for animal
154
 
155
  confidence = bbxyxy[4]
156
+ string_bb = "animal " + str(round(confidence, 2))
157
+ font = ImageFont.truetype(FONTS[font_style], font_size)
158
 
159
+ text_size = font.getbbox(string_bb) # (h,w)
160
+ position = (bbxyxy[0], bbxyxy[1] - text_size[1] - 2)
161
  left, top, right, bottom = imgR.textbbox(position, string_bb, font=font)
162
+ imgR.rectangle((left, top - 5, right + 5, bottom + 5), fill="red")
163
+ imgR.text((bbxyxy[0] + 3, bbxyxy[1] - text_size[1] - 2), string_bb, font=font, fill="black")
164
 
165
  return imgR
166
 
 
 
167
 
168
+ ###########################################
169
+ def save_results_as_json(
170
+ md_results,
171
+ dlc_outputs,
172
+ animal_bboxes,
173
+ map_dlc_label_id_to_str,
174
+ model,
175
+ mega_model_input,
176
+ path_to_output_file="download_predictions.json",
177
+ ):
178
  """
179
  Output detections as json file
180
 
 
182
  """
183
  # initialise dict to save to json
184
  info = {}
185
+ info["date"] = str(today)
186
+ info["MD_model"] = str(mega_model_input)
187
  # info from megaDetector
188
+ info["file"] = md_results.files[0]
189
  number_bb = len(md_results.xyxy[0].tolist())
190
+ info["number_of_bb"] = number_bb
191
  # info from DLC
192
+ info["dlc_model"] = model
193
  labels = [n for n in map_dlc_label_id_to_str.values()]
194
 
195
  # define aux dict for every animal bounding box above threshold
196
  for i in range(len(dlc_outputs)):
197
+ aux = {}
198
  # MD output
199
+ corner_x1, corner_y1, corner_x2, corner_y2, confidence, _ = animal_bboxes[i]
200
+ aux["corner_1"] = (corner_x1, corner_y1)
201
+ aux["corner_2"] = (corner_x2, corner_y2)
202
+ aux["predict MD"] = md_results.names[0]
203
+ aux["confidence MD"] = confidence
204
+
205
  # DLC output
206
  kypts = []
207
  for s in dlc_outputs[i]:
 
210
  aux1.append(float(j))
211
 
212
  kypts.append(aux1)
213
+ aux["dlc_pred"] = dict(zip(labels, kypts, strict=True))
214
+ info["bb_" + str(i)] = aux
215
 
216
  # save dict as json
217
+ with open(path_to_output_file, "w") as f:
218
  json.dump(info, f, indent=1)
219
+ print(f"Output file saved at {path_to_output_file}")
220
 
221
  return path_to_output_file
222
 
223
 
224
+ def save_results_only_dlc(dlc_outputs, map_label_id_to_str, model, output_file="dowload_predictions_dlc.json"):
 
225
  """
226
  write json dlc output
227
  """
228
  info = {}
229
+ info["date"] = str(today)
230
  labels = [n for n in map_label_id_to_str.values()]
231
+ info["dlc_model"] = model
232
  kypts = []
233
  for s in dlc_outputs:
234
  aux1 = []
 
236
  aux1.append(float(j))
237
 
238
  kypts.append(aux1)
239
+ info["dlc_pred"] = dict(zip(labels, kypts, strict=True))
240
 
241
+ with open(output_file, "w") as f:
242
  json.dump(info, f, indent=1)
243
+ print(f"Output file saved at {output_file}")
244
 
245
  return output_file
246
 
247
 
248
+ def save_results_pytorch(
249
+ animals, map_label_id_to_str, model, pose_model, detector, path_to_output_file="download_predictions.json"
250
+ ):
251
  """
252
  Output PyTorch SuperAnimal predictions as json file (same layout as save_results_as_json)
253
 
 
255
  detector: None if the detector was skipped (whole image used as one animal)
256
  """
257
  info = {}
258
+ info["date"] = str(today)
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 = [n for n in map_label_id_to_str.values()]
265
 
266
  for i, animal in enumerate(animals):
267
+ corner_x1, corner_y1, corner_x2, corner_y2, confidence = animal["bbox"]
268
  aux = {}
269
+ aux["corner_1"] = (corner_x1, corner_y1)
270
+ aux["corner_2"] = (corner_x2, corner_y2)
271
+ aux["confidence"] = confidence
272
+ aux["dlc_pred"] = dict(zip(labels, [[float(v) for v in kpt] for kpt in animal["kpts"]], strict=True))
273
+ info["bb_" + str(i)] = aux
274
 
275
+ with open(path_to_output_file, "w") as f:
276
  json.dump(info, f, indent=1)
277
+ print(f"Output file saved at {path_to_output_file}")
278
 
279
  return path_to_output_file
280
 
281
 
282
+ ###########################################