Add confidence-based keypoint outputs
Browse filesExtend 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.
- app.py +36 -8
- ui_utils.py +21 -10
- 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 |
-
|
| 107 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 108 |
|
| 109 |
|
| 110 |
def gradio_description_and_examples():
|
| 111 |
-
title = "DeepLabCut Model Zoo SuperAnimals"
|
| 112 |
description = (
|
| 113 |
-
"
|
| 114 |
-
"
|
| 115 |
-
"
|
| 116 |
-
"
|
| 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 |
-
|
| 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 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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,
|