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]