Polish pose visualization UI
Browse filesRefresh 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.
- app.py +30 -13
- ui_utils.py +99 -5
- 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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=
|
|
|
|
| 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=
|
| 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=
|
| 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"
|
| 134 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 6 |
-
from PIL import
|
| 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 |
-
|
| 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
|
| 70 |
-
round_fill =
|
| 71 |
else:
|
| 72 |
# one color per bodypart, transparency encodes the confidence
|
| 73 |
-
round_fill = list(
|
| 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 |
-
#
|
| 99 |
-
|
| 100 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
|
| 102 |
-
|
| 103 |
-
|
| 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):
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
###########################################
|