mwmathis's picture C-Achard's picture
Color keypoints by bodypart by default (#15)
5149d74
Raw History Blame Contribute Delete
7.14 kB
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=False,
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, False, "viridis", "#ff0000"]
for image in (
"examples/dog.jpeg",
"examples/cat.jpg",
)
for font_size, marker_size in [example_sizes(image)]
]
return [title, description, examples]