C-Achard's picture
Add confidence-based keypoint outputs
e189412
Raw History Blame
4 kB
import gradio as gr
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_keypt_color = gr.ColorPicker(
value="#862db7",
label="Choose color for keypoint label",
)
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=8,
step=1,
label="Set font size",
)
gr_slider_marker_size = gr.Slider(
minimum=1,
maximum=20,
value=9,
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,
]
def gradio_outputs_for_MD_DLC():
gr_image_output = gr.Image(type="pil", label="Output Image")
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_file_download, gr_image_download, gr_confidence_table]
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", 10, "#ff0000", 5, True]
for image in ("examples/dog.jpeg", "examples/cat.jpg")
]
return [title, description, examples]