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

Polish pose visualization UI

Browse files

Refresh the Gradio app with a DeepLabCut-themed UI, configurable keypoint colormaps and bounding-box colors, and a standalone HTML confidence legend instead of drawing it into the output image. This also improves example defaults for different image sizes, cleans up displayed bodypart names, and makes bbox labels more readable with contrast-aware text.

Files changed (3) hide show
  1. app.py +30 -13
  2. ui_utils.py +99 -5
  3. viz_utils.py +46 -51
app.py CHANGED
@@ -20,9 +20,14 @@ from PIL import Image
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,
@@ -53,13 +58,11 @@ DLC_models_dict = {
53
 
54
 
55
  #####################################################
56
- def finalize_outputs(img_output, download_file, kpts_per_animal, map_label_id_to_str, color_by_confidence):
57
- # confidence legend, annotated image for download and per-keypoint confidence table
58
- if color_by_confidence:
59
- img_output = add_confidence_legend(img_output)
60
  annotated_file = save_annotated_image(img_output)
61
  confidence_rows = keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str)
62
- return img_output, download_file, annotated_file, confidence_rows
 
63
 
64
 
65
  #####################################################
@@ -75,6 +78,8 @@ def predict_pipeline_pytorch(
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(
@@ -94,16 +99,22 @@ def predict_pipeline_pytorch(
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
 
@@ -122,6 +133,8 @@ def predict_pipeline(
122
  keypt_color,
123
  marker_size,
124
  flag_color_by_confidence,
 
 
125
  ):
126
 
127
  if backend == "PyTorch":
@@ -137,6 +150,8 @@ def predict_pipeline(
137
  keypt_color,
138
  marker_size,
139
  flag_color_by_confidence,
 
 
140
  )
141
 
142
  # TensorFlow (legacy): MegaDetector crops + DLCLive
@@ -202,12 +217,13 @@ def predict_pipeline(
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:
@@ -233,13 +249,14 @@ def predict_pipeline(
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(
@@ -247,7 +264,7 @@ def predict_pipeline(
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
 
@@ -280,4 +297,4 @@ with gr.Blocks(title=gr_title) as demo:
280
  threading.Thread(target=load_superanimal, args=("superanimal_quadruped",), daemon=True).start()
281
 
282
  demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
283
- demo.launch(theme=gr.themes.Default())
 
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,
 
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
  #####################################################
 
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(
 
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, map_label_id_to_str, superanimal, pose_model, None if flag_dlc_only else detector
110
  )
111
  return finalize_outputs(
112
+ img_output,
113
+ download_file,
114
+ [animal["kpts"] for animal in animals],
115
+ map_label_id_to_str,
116
+ flag_color_by_confidence,
117
+ colormap,
118
  )
119
 
120
 
 
133
  keypt_color,
134
  marker_size,
135
  flag_color_by_confidence,
136
+ colormap,
137
+ bbox_color,
138
  ):
139
 
140
  if backend == "PyTorch":
 
150
  keypt_color,
151
  marker_size,
152
  flag_color_by_confidence,
153
+ colormap,
154
+ bbox_color,
155
  )
156
 
157
  # TensorFlow (legacy): MegaDetector crops + DLCLive
 
217
  keypt_color=keypt_color,
218
  marker_size=marker_size,
219
  color_by_confidence=flag_color_by_confidence,
220
+ colormap=colormap,
221
  )
222
 
223
  donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str, dlc_model_name)
224
 
225
  return finalize_outputs(
226
+ img_input, donw_file, [list_kpts_per_crop[0]], map_label_id_to_str, flag_color_by_confidence, colormap
227
  )
228
 
229
  else:
 
249
  keypt_color=keypt_color,
250
  marker_size=marker_size,
251
  color_by_confidence=flag_color_by_confidence,
252
+ colormap=colormap,
253
  )
254
 
255
  # Paste crop in original image
256
  img_background.paste(img_crop, box=tuple([int(t) for t in bb_per_animal[:2]]))
257
 
258
  # Plot bbox
259
+ draw_bbox_w_text(img_background, bb_per_animal, font_size=font_size, bbox_color=bbox_color)
260
 
261
  # Save detection results as json
262
  download_file = save_results_as_json(
 
264
  )
265
 
266
  return finalize_outputs(
267
+ img_background, download_file, list_kpts_per_crop, map_label_id_to_str, flag_color_by_confidence, colormap
268
  )
269
 
270
 
 
297
  threading.Thread(target=load_superanimal, args=("superanimal_quadruped",), daemon=True).start()
298
 
299
  demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
300
+ demo.launch(theme=dlc_theme())
ui_utils.py CHANGED
@@ -1,4 +1,54 @@
1
  import gradio as gr
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
 
3
 
4
  def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
@@ -62,11 +112,23 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
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",
68
  )
69
 
 
 
 
 
 
70
  gr_labels_font_style = gr.Dropdown(
71
  choices=["amiko", "animals", "nature", "painter", "zen"],
72
  value="amiko",
@@ -77,7 +139,7 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
77
  gr_slider_font_size = gr.Slider(
78
  minimum=5,
79
  maximum=30,
80
- value=8,
81
  step=1,
82
  label="Set font size",
83
  )
@@ -85,7 +147,7 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
85
  gr_slider_marker_size = gr.Slider(
86
  minimum=1,
87
  maximum=20,
88
- value=9,
89
  step=1,
90
  label="Set marker size",
91
  )
@@ -104,11 +166,29 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
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")
@@ -117,7 +197,13 @@ def gradio_outputs_for_MD_DLC():
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():
@@ -130,8 +216,16 @@ def gradio_description_and_examples():
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
 
137
  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):
 
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",
 
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
  )
 
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
  )
 
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")
 
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():
 
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
+ "examples/lynx.jpg",
225
+ "examples/goat.jpg",
226
+ "examples/giraffe.jpg",
227
+ )
228
+ for font_size, marker_size in [example_sizes(image)]
229
  ]
230
 
231
  return [title, description, examples]
viz_utils.py CHANGED
@@ -2,8 +2,8 @@ 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 = {
@@ -13,6 +13,8 @@ FONTS = {
13
  "animals": "fonts/UncialAnimals.ttf",
14
  "zen": "fonts/ZEN.TTF",
15
  }
 
 
16
 
17
 
18
  #########################################
@@ -28,6 +30,7 @@ def draw_keypoints_on_image(
28
  keypt_color="#ff0000",
29
  marker_size=2,
30
  color_by_confidence=True,
 
31
  ):
32
  """Draws keypoints on an image.
33
  Modified from:
@@ -57,7 +60,7 @@ def draw_keypoints_on_image(
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
@@ -66,11 +69,11 @@ def draw_keypoints_on_image(
66
 
67
  confidence = float(np.clip(confidences[i], 0, 1))
68
  if color_by_confidence:
69
- # fill color encodes the keypoint confidence (see add_confidence_legend)
70
- round_fill = cm.viridis(confidence, bytes=True)
71
  else:
72
  # one color per bodypart, transparency encodes the confidence
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(
@@ -88,39 +91,29 @@ def draw_keypoints_on_image(
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
103
- coordinates are unchanged.
104
- """
105
- im_width, im_height = image.size
106
- strip_w = max(60, im_width // 5)
107
- strip_h = max(6, im_height // 60)
108
- font = ImageFont.truetype(FONTS[font_style], max(10, strip_h * 2))
109
- margin = strip_h
110
- band_h = strip_h + font.size + 3 * margin
111
-
112
- out = Image.new("RGB", (im_width, im_height + band_h), "white")
113
- out.paste(image.convert("RGB"), (0, 0))
114
- draw = ImageDraw.Draw(out)
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")
122
- draw.text((x0 + strip_w / 2, label_y), "confidence", fill="black", font=font, anchor="ma")
123
- return out
124
 
125
 
126
  #########################################
@@ -131,7 +124,7 @@ def keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str):
131
  for i_animal, kpts in enumerate(kpts_per_animal):
132
  for i_kpt, kpt in enumerate(kpts):
133
  if not np.isnan(kpt[2]):
134
- rows.append([i_animal, map_label_id_to_str[i_kpt], round(float(kpt[2]), 3)])
135
  return sorted(rows, key=lambda row: row[2])
136
 
137
 
@@ -144,25 +137,27 @@ def save_annotated_image(image, path_to_output_file="download_annotated.png"):
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
  ###########################################
 
2
  from datetime import date
3
 
4
  import numpy as np
5
+ from matplotlib import colormaps
6
+ from PIL import ImageColor, ImageDraw, ImageFont
7
 
8
  today = date.today()
9
  FONTS = {
 
13
  "animals": "fonts/UncialAnimals.ttf",
14
  "zen": "fonts/ZEN.TTF",
15
  }
16
+ # perceptually uniform maps first; turbo separates neighbouring bodyparts best
17
+ COLORMAPS = ["viridis", "plasma", "magma", "cividis", "turbo"]
18
 
19
 
20
  #########################################
 
30
  keypt_color="#ff0000",
31
  marker_size=2,
32
  color_by_confidence=True,
33
+ colormap="viridis",
34
  ):
35
  """Draws keypoints on an image.
36
  Modified from:
 
60
  keypoints_x = tuple([im_width * x for x in keypoints_x])
61
  keypoints_y = tuple([im_height * y for y in keypoints_y])
62
 
63
+ cmap = colormaps[colormap]
64
  # draw ellipses around keypoints
65
  for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y, strict=True)):
66
  # handling potential nans in the keypoints
 
69
 
70
  confidence = float(np.clip(confidences[i], 0, 1))
71
  if color_by_confidence:
72
+ # fill color encodes the keypoint confidence (see confidence_legend_html in ui_utils)
73
+ round_fill = cmap(confidence, bytes=True)
74
  else:
75
  # one color per bodypart, transparency encodes the confidence
76
+ round_fill = list(cmap(i / max(len(keypoints) - 1, 1), bytes=True))
77
  round_fill[3] = round(confidence * 255)
78
  round_fill = tuple(round_fill)
79
  draw.ellipse(
 
91
  font = ImageFont.truetype(FONTS[font_style], font_size)
92
  draw.text(
93
  (keypoint_x + marker_size, keypoint_y + marker_size), # (0.5*im_width, 0.5*im_height), #-------
94
+ display_bodypart(map_label_id_to_str[i]),
95
  ImageColor.getcolor(keypt_color, "RGB"), # rgb #
96
  font=font,
97
  )
98
 
99
 
100
  #########################################
101
+ # Bodypart names for display
102
+ # display names where the SuperAnimal definitions misspell (quadruped "thai") or read oddly
103
+ # (top-view mouse "backend"); the JSON output keeps the model's names
104
+ BODYPART_DISPLAY_NAMES = {
105
+ "front_left_thai": "front left thigh",
106
+ "front_right_thai": "front right thigh",
107
+ "back_left_thai": "back left thigh",
108
+ "back_right_thai": "back right thigh",
109
+ "mid_backend": "mid back end",
110
+ "mid_backend2": "mid back end 2",
111
+ "mid_backend3": "mid back end 3",
112
+ }
113
 
114
+
115
+ def display_bodypart(name):
116
+ return BODYPART_DISPLAY_NAMES.get(name, name.replace("_", " "))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
117
 
118
 
119
  #########################################
 
124
  for i_animal, kpts in enumerate(kpts_per_animal):
125
  for i_kpt, kpt in enumerate(kpts):
126
  if not np.isnan(kpt[2]):
127
+ rows.append([i_animal, display_bodypart(map_label_id_to_str[i_kpt]), round(float(kpt[2]), 3)])
128
  return sorted(rows, key=lambda row: row[2])
129
 
130
 
 
137
 
138
  #########################################
139
  # Draw bboxes on image
140
+ def draw_bbox_w_text(img, results, font_style="amiko", font_size=8, bbox_color="#ff0000"):
141
+ x1, y1, x2, y2, confidence = results[:5]
142
+ draw = ImageDraw.Draw(img)
143
+ draw.rectangle([(x1, y1), (x2, y2)], outline=bbox_color, width=max(2, round(font_size / 5)))
 
 
 
 
 
 
 
 
 
 
 
 
 
144
 
145
+ label = f"animal {confidence:.2f}"
146
+ font = ImageFont.truetype(FONTS[font_style], font_size)
147
+ left, top, right, bottom = draw.textbbox((0, 0), label, font=font)
148
+ pad = max(2, font_size // 5)
149
+ label_w, label_h = right - left + 2 * pad, bottom - top + 2 * pad
150
+ # label above the box, or inside it when the box touches the top of the image
151
+ label_y = y1 - label_h if y1 >= label_h else y1
152
+ draw.rectangle([(x1, label_y), (x1 + label_w, label_y + label_h)], fill=bbox_color)
153
+ draw.text((x1 + pad - left, label_y + pad - top), label, font=font, fill=label_text_color(bbox_color))
154
+
155
+
156
+ def label_text_color(background):
157
+ # black or white, whichever contrasts more with the background (WCAG relative luminance)
158
+ channels = [c / 255 for c in ImageColor.getrgb(background)[:3]]
159
+ r, g, b = [c / 12.92 if c <= 0.03928 else ((c + 0.055) / 1.055) ** 2.4 for c in channels]
160
+ return "black" if 0.2126 * r + 0.7152 * g + 0.0722 * b > 0.179 else "white"
161
 
162
 
163
  ###########################################