File size: 7,135 Bytes
7206ed3 e0fc43c 7206ed3 50ce3c9 d62a8dc 7206ed3 50ce3c9 d62a8dc 50ce3c9 734addc d62a8dc 50ce3c9 d62a8dc 7206ed3 d62a8dc 50ce3c9 d62a8dc 734addc e189412 e0fc43c 734addc e0fc43c 734addc e0fc43c 734addc e0fc43c 734addc d62a8dc 50ce3c9 d62a8dc e189412 e0fc43c d62a8dc e0fc43c 7206ed3 d62a8dc e0fc43c e189412 e0fc43c d62a8dc 7206ed3 e189412 d62a8dc e189412 d62a8dc 7206ed3 734addc e0fc43c 734addc 7206ed3 752f8f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 | import gradio as gr
from matplotlib import colormaps
from matplotlib.colors import to_hex
from PIL import Image
from pytorch_utils import MAX_IMAGE_SIZE
from viz_utils import COLORMAPS
# shades built around the DeepLabCut docs palette (DeepLabCut/docs/_static/custom.css)
DLC_PURPLE = gr.themes.Color(
name="dlc_purple",
c50="#f5edff",
c100="#ead7ff",
c200="#d9b8ff",
c300="#c084fc",
c400="#ac72f0",
c500="#9b5de5",
c600="#8550c4",
c700="#73439a",
c800="#5c357c",
c900="#4b236f",
c950="#2e1546",
)
DLC_TEAL = gr.themes.Color(
name="dlc_teal",
c50="#e8f8f7",
c100="#c9efec",
c200="#9fe2dd",
c300="#7ad3ce",
c400="#57c4be",
c500="#2fb0a8",
c600="#21a197",
c700="#1a8078",
c800="#16655f",
c900="#124f4a",
c950="#0a2e2b",
)
def dlc_theme():
# white text on purple 500 is 4.1:1, so filled buttons use 700 (7.0:1) and 600 (5.3:1)
return gr.themes.Default(primary_hue=DLC_PURPLE, secondary_hue=DLC_TEAL, neutral_hue="slate").set(
button_primary_background_fill="*primary_700",
button_primary_background_fill_hover="*primary_600",
button_primary_background_fill_dark="*primary_600",
button_primary_background_fill_hover_dark="*primary_500",
button_primary_text_color="white",
button_primary_text_color_dark="white",
button_primary_border_color="*primary_700",
button_primary_border_color_dark="*primary_600",
)
def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
# Input image
gr_image_input = gr.Image(type="pil", label="Input Image")
# Models
gr_backend_input = gr.Radio(
choices=backends_list,
value=backends_list[0],
label="Select backend",
)
gr_mega_model_input = gr.Dropdown(
choices=md_models_list,
value="md_v5a",
type="value",
label="Select Detector model (TensorFlow legacy only)",
visible=gr_backend_input.value != "PyTorch",
)
gr_dlc_model_input = gr.Dropdown(
choices=dlc_models_list,
value="superanimal_quadruped",
type="value",
label="Select DeepLabCut model",
)
# Other inputs
gr_dlc_only_checkbox = gr.Checkbox(
value=False,
label="Run DeepLabCut only, directly on input image?",
)
# Gradio Slider signature is (minimum, maximum, value, step, ...)
gr_slider_conf_bboxes = gr.Slider(
minimum=0,
maximum=1,
value=0.2,
step=0.05,
label="Set confidence threshold for animal detections",
)
gr_slider_conf_keypoints = gr.Slider(
minimum=0,
maximum=1,
value=0.4,
step=0.05,
label="Set confidence threshold for keypoints",
)
# Data viz
with gr.Accordion("Display options", open=False):
gr_str_labels_checkbox = gr.Checkbox(
value=True,
label="Show bodypart labels?",
)
gr_color_by_confidence_checkbox = gr.Checkbox(
value=True,
label="Color keypoints by confidence? (otherwise by bodypart)",
)
gr_colormap = gr.Dropdown(
choices=COLORMAPS,
value="viridis",
type="value",
label="Keypoint colormap",
)
gr_keypt_color = gr.ColorPicker(
value="#862db7",
label="Choose color for keypoint label",
)
gr_bbox_color = gr.ColorPicker(
value="#ff0000",
label="Choose color for bounding boxes",
)
gr_labels_font_style = gr.Dropdown(
choices=["amiko", "animals", "nature", "painter", "zen"],
value="amiko",
type="value",
label="Select keypoint label font",
)
gr_slider_font_size = gr.Slider(
minimum=5,
maximum=30,
value=18,
step=1,
label="Set font size",
)
gr_slider_marker_size = gr.Slider(
minimum=1,
maximum=20,
value=6,
step=1,
label="Set marker size",
)
return [
gr_image_input,
gr_backend_input,
gr_mega_model_input,
gr_dlc_model_input,
gr_dlc_only_checkbox,
gr_str_labels_checkbox,
gr_slider_conf_bboxes,
gr_slider_conf_keypoints,
gr_labels_font_style,
gr_slider_font_size,
gr_keypt_color,
gr_slider_marker_size,
gr_color_by_confidence_checkbox,
gr_colormap,
gr_bbox_color,
]
def confidence_legend_html(colormap="viridis"):
# the colormap as a CSS gradient, matching the keypoint fill in draw_keypoints_on_image
gradient = ", ".join(f"{to_hex(colormaps[colormap](i / 10))} {i * 10}%" for i in range(11))
# no leading newline: gradio prefixes "'" to cached example values starting with one (CSV injection guard)
return f"""<div style="display:flex; align-items:flex-start; gap:12px; flex-wrap:wrap; font-size:var(--text-sm);">
<span style="white-space:nowrap; line-height:12px;">Keypoint confidence</span>
<div style="flex:1; min-width:160px; max-width:360px;">
<div style="height:12px; border-radius:6px; background:linear-gradient(to right, {gradient});"></div>
<div style="display:flex; justify-content:space-between; margin-top:2px; font-variant-numeric:tabular-nums;">
<span>0</span><span>0.5</span><span>1</span>
</div>
</div>
</div>"""
def gradio_outputs_for_MD_DLC():
gr_image_output = gr.Image(type="pil", label="Output Image")
gr_confidence_legend = gr.HTML("")
with gr.Row():
gr_file_download = gr.File(label="Download JSON file")
gr_image_download = gr.File(label="Download annotated image")
gr_confidence_table = gr.Dataframe(
headers=["animal", "bodypart", "confidence"],
label="Keypoint confidence (lowest first)",
interactive=False,
)
return [gr_image_output, gr_confidence_legend, gr_file_download, gr_image_download, gr_confidence_table]
def example_sizes(path):
# font and marker sizes for the resolution the PyTorch backend draws on
side = min(max(Image.open(path).size), MAX_IMAGE_SIZE)
return round(side / 64), round(side / 200)
def gradio_description_and_examples():
title = "DeepLabCut Model Zoo: SuperAnimals"
description = (
"Estimate animal poses with the SuperAnimal models from the "
"[DeepLabCut Model Zoo](http://www.mackenziemathislab.org/dlc-modelzoo) "
"([paper](https://arxiv.org/abs/2203.07436)). "
"Upload an image or pick an example below; to run on videos, see the Model Zoo page."
)
examples = [
[image, "PyTorch", "md_v5a", "superanimal_quadruped", False, True, 0.5, 0.4, "amiko"]
+ [font_size, "#ff0000", marker_size, True, "viridis", "#ff0000"]
for image in (
"examples/dog.jpeg",
"examples/cat.jpg",
)
for font_size, marker_size in [example_sizes(image)]
]
return [title, description, examples]
|