C-Achard commited on
Commit
e189412
Β·
1 Parent(s): 734addc

Add confidence-based keypoint outputs

Browse files

Extend the app and UI to surface keypoint confidence more clearly. This adds optional confidence-based coloring with a legend, an annotated image download, and a confidence table sorted by lowest scores first. It also preloads the default SuperAnimal PyTorch model in a background thread so the first request is less likely to block on model download.

Files changed (3) hide show
  1. app.py +36 -8
  2. ui_utils.py +21 -10
  3. viz_utils.py +61 -14
app.py CHANGED
@@ -4,6 +4,7 @@
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
@@ -16,9 +17,10 @@ import dlclive
16
  from PIL import Image, ImageColor, ImageFont, ImageDraw
17
 
18
  from viz_utils import save_results_as_json, draw_keypoints_on_image, draw_bbox_w_text, save_results_only_dlc, save_results_pytorch
 
19
  from detection_utils import predict_md, crop_animal_detections
20
  from dlc_utils import predict_dlc
21
- from pytorch_utils import predict_superanimal, PYTORCH_MODELS
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
@@ -46,6 +48,16 @@ DLC_models_dict = {'superanimal_topviewmouse': ('superanimal_topviewmouse_dlcrne
46
 
47
 
48
 
 
 
 
 
 
 
 
 
 
 
49
  #####################################################
50
  def predict_pipeline_pytorch(img_input,
51
  superanimal,
@@ -57,6 +69,7 @@ def predict_pipeline_pytorch(img_input,
57
  font_size,
58
  keypt_color,
59
  marker_size,
 
60
  ):
61
  # detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
62
  img_output, animals, bodyparts = predict_superanimal(img_input,
@@ -75,7 +88,8 @@ def predict_pipeline_pytorch(img_input,
75
  font_style=font_style,
76
  font_size=font_size,
77
  keypt_color=keypt_color,
78
- marker_size=marker_size)
 
79
  if not flag_dlc_only:
80
  draw_bbox_w_text(img_output,
81
  animal['bbox'],
@@ -84,7 +98,9 @@ def predict_pipeline_pytorch(img_input,
84
  pose_model, detector = PYTORCH_MODELS[superanimal]
85
  download_file = save_results_pytorch(animals, map_label_id_to_str, superanimal,
86
  pose_model, None if flag_dlc_only else detector)
87
- return img_output, download_file
 
 
88
 
89
 
90
  #####################################################
@@ -100,6 +116,7 @@ def predict_pipeline(img_input,
100
  font_size,
101
  keypt_color,
102
  marker_size,
 
103
  ):
104
 
105
  if backend == "PyTorch":
@@ -112,7 +129,8 @@ def predict_pipeline(img_input,
112
  font_style,
113
  font_size,
114
  keypt_color,
115
- marker_size)
 
116
 
117
  # TensorFlow (legacy): MegaDetector crops + DLCLive
118
  dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
@@ -169,11 +187,14 @@ def predict_pipeline(img_input,
169
  font_style=font_style,
170
  font_size=font_size,
171
  keypt_color=keypt_color,
172
- marker_size=marker_size)
 
173
 
174
  donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_name)
175
 
176
- return img_input, donw_file
 
 
177
 
178
  else:
179
  # Compute kpts for each crop
@@ -202,7 +223,8 @@ def predict_pipeline(img_input,
202
  font_style=font_style,
203
  font_size=font_size,
204
  keypt_color=keypt_color,
205
- marker_size=marker_size)
 
206
 
207
  # Paste crop in original image
208
  img_background.paste(img_crop,
@@ -217,7 +239,9 @@ def predict_pipeline(img_input,
217
  # Save detection results as json
218
  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)
219
 
220
- return img_background, download_file
 
 
221
 
222
 
223
 
@@ -254,5 +278,9 @@ with gr.Blocks(title=gr_title) as demo:
254
  cache_examples=True,
255
  cache_mode="lazy")
256
 
 
 
 
 
257
  demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
258
  demo.launch(theme=gr.themes.Default())
 
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
 
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
 
48
 
49
 
50
 
51
+ #####################################################
52
+ def finalize_outputs(img_output, download_file, kpts_per_animal, map_label_id_to_str, color_by_confidence):
53
+ # confidence legend, annotated image for download and per-keypoint confidence table
54
+ if color_by_confidence:
55
+ img_output = add_confidence_legend(img_output)
56
+ annotated_file = save_annotated_image(img_output)
57
+ confidence_rows = keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str)
58
+ return img_output, download_file, annotated_file, confidence_rows
59
+
60
+
61
  #####################################################
62
  def predict_pipeline_pytorch(img_input,
63
  superanimal,
 
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,
 
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'],
 
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
  #####################################################
 
116
  font_size,
117
  keypt_color,
118
  marker_size,
119
+ flag_color_by_confidence,
120
  ):
121
 
122
  if backend == "PyTorch":
 
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]
 
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
 
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,
 
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
 
 
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
283
+ threading.Thread(target=load_superanimal, args=("superanimal_quadruped",), daemon=True).start()
284
+
285
  demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
286
  demo.launch(theme=gr.themes.Default())
ui_utils.py CHANGED
@@ -57,6 +57,11 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
57
  label="Show bodypart labels?",
58
  )
59
 
 
 
 
 
 
60
  gr_keypt_color = gr.ColorPicker(
61
  value="#862db7",
62
  label="Choose color for keypoint label",
@@ -98,28 +103,34 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
98
  gr_slider_font_size,
99
  gr_keypt_color,
100
  gr_slider_marker_size,
 
101
  ]
102
 
103
 
104
  def gradio_outputs_for_MD_DLC():
105
  gr_image_output = gr.Image(type="pil", label="Output Image")
106
- gr_file_download = gr.File(label="Download JSON file")
107
- return [gr_image_output, gr_file_download]
 
 
 
 
 
 
 
108
 
109
 
110
  def gradio_description_and_examples():
111
- title = "DeepLabCut Model Zoo SuperAnimals"
112
  description = (
113
- "Test the SuperAnimal models from the "
114
- "<a href='http://www.mackenziemathislab.org/dlc-modelzoo'>"
115
- "DeepLabCut ModelZoo Project</a>, and read more on arXiv: "
116
- "https://arxiv.org/abs/2203.07436! Simply upload an image and see how it does. "
117
- "Want to run on videos on the cloud or locally? See the "
118
- "<a href='http://www.mackenziemathislab.org/dlc-modelzoo'>DeepLabCut ModelZoo</a>."
119
  )
120
 
121
  examples = [
122
- [image, "PyTorch", "md_v5a", "superanimal_quadruped", False, True, 0.5, 0.4, "amiko", 10, "#ff0000", 5]
123
  for image in ("examples/dog.jpeg", "examples/cat.jpg")
124
  ]
125
 
 
57
  label="Show bodypart labels?",
58
  )
59
 
60
+ gr_color_by_confidence_checkbox = gr.Checkbox(
61
+ value=True,
62
+ label="Color keypoints by confidence? (otherwise by bodypart)",
63
+ )
64
+
65
  gr_keypt_color = gr.ColorPicker(
66
  value="#862db7",
67
  label="Choose color for keypoint label",
 
103
  gr_slider_font_size,
104
  gr_keypt_color,
105
  gr_slider_marker_size,
106
+ gr_color_by_confidence_checkbox,
107
  ]
108
 
109
 
110
  def gradio_outputs_for_MD_DLC():
111
  gr_image_output = gr.Image(type="pil", label="Output Image")
112
+ with gr.Row():
113
+ gr_file_download = gr.File(label="Download JSON file")
114
+ gr_image_download = gr.File(label="Download annotated image")
115
+ gr_confidence_table = gr.Dataframe(
116
+ headers=["animal", "bodypart", "confidence"],
117
+ label="Keypoint confidence (lowest first)",
118
+ interactive=False,
119
+ )
120
+ return [gr_image_output, gr_file_download, gr_image_download, gr_confidence_table]
121
 
122
 
123
  def gradio_description_and_examples():
124
+ title = "DeepLabCut Model Zoo: SuperAnimals"
125
  description = (
126
+ "Estimate animal poses with the SuperAnimal models from the "
127
+ "[DeepLabCut Model Zoo](http://www.mackenziemathislab.org/dlc-modelzoo) "
128
+ "([paper](https://arxiv.org/abs/2203.07436)). "
129
+ "Upload an image or pick an example below; to run on videos, see the Model Zoo page."
 
 
130
  )
131
 
132
  examples = [
133
+ [image, "PyTorch", "md_v5a", "superanimal_quadruped", False, True, 0.5, 0.4, "amiko", 10, "#ff0000", 5, True]
134
  for image in ("examples/dog.jpeg", "examples/cat.jpg")
135
  ]
136
 
viz_utils.py CHANGED
@@ -25,6 +25,7 @@ def draw_keypoints_on_image(image,
25
  font_size=8,
26
  keypt_color="#ff0000",
27
  marker_size=2,
 
28
  ):
29
  """Draws keypoints on an image.
30
  Modified from:
@@ -47,14 +48,7 @@ def draw_keypoints_on_image(image,
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:
@@ -64,15 +58,19 @@ def draw_keypoints_on_image(image,
64
  #cmap = matplotlib.cm.get_cmap('hsv')
65
  # draw ellipses around keypoints
66
  for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y)):
67
- round_fill = list(cm.viridis(norm(colores[i]),bytes=True))#[round(num*255) for num in list(cmap(i))[:3]] #check!
68
  # handling potential nans in the keypoints
69
  if np.isnan(keypoint_x).any():
70
  continue
71
-
72
- if np.isnan(alpha[i]) == False :
73
- round_fill[3] = round(alpha[i] *255)
74
- #print(round_fill)
75
- #round_outline = [round(num*255) for num in list(cmap2(alpha[i]))[:3]]
 
 
 
 
 
76
  draw.ellipse([(keypoint_x - marker_size, keypoint_y - marker_size),
77
  (keypoint_x + marker_size, keypoint_y + marker_size)],
78
  fill=tuple(round_fill), outline= 'black', width=1) #fill and outline: [0,255]
@@ -86,6 +84,55 @@ def draw_keypoints_on_image(image,
86
  ImageColor.getcolor(keypt_color, "RGB"), # rgb #
87
  font=font)
88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89
  #########################################
90
  # Draw bboxes on image
91
  def draw_bbox_w_text(img,
 
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:
 
48
  im_width, im_height = image.size
49
  keypoints_x = [k[0] for k in keypoints]
50
  keypoints_y = [k[1] for k in keypoints]
51
+ confidences = [k[2] for k in keypoints]
 
 
 
 
 
 
 
52
 
53
  # adjust keypoints coords if required
54
  if use_normalized_coordinates:
 
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
64
+
65
+ confidence = float(np.clip(confidences[i], 0, 1))
66
+ if color_by_confidence:
67
+ # fill color encodes the keypoint confidence (see add_confidence_legend)
68
+ round_fill = cm.viridis(confidence, bytes=True)
69
+ else:
70
+ # one color per bodypart, transparency encodes the confidence
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]
 
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
93
+ coordinates are unchanged.
94
+ """
95
+ im_width, im_height = image.size
96
+ strip_w = max(60, im_width // 5)
97
+ strip_h = max(6, im_height // 60)
98
+ font = ImageFont.truetype(FONTS[font_style], max(10, strip_h * 2))
99
+ margin = strip_h
100
+ band_h = strip_h + font.size + 3 * margin
101
+
102
+ out = Image.new("RGB", (im_width, im_height + band_h), "white")
103
+ out.paste(image.convert("RGB"), (0, 0))
104
+ draw = ImageDraw.Draw(out)
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")
113
+ draw.text((x0 + strip_w / 2, label_y), "confidence", fill="black", font=font, anchor="ma")
114
+ return out
115
+
116
+
117
+ #########################################
118
+ # Keypoint confidences as table rows
119
+ def keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str):
120
+ """(animal, bodypart, confidence) for every keypoint kept (not NaN), lowest confidence first."""
121
+ rows = []
122
+ for i_animal, kpts in enumerate(kpts_per_animal):
123
+ for i_kpt, kpt in enumerate(kpts):
124
+ if not np.isnan(kpt[2]):
125
+ rows.append([i_animal, map_label_id_to_str[i_kpt], round(float(kpt[2]), 3)])
126
+ return sorted(rows, key=lambda row: row[2])
127
+
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,