C-Achard commited on
Commit
734addc
Β·
1 Parent(s): 50ce3c9

Refine Gradio app startup and layout

Browse files

Switch the UI from `gr.Interface` to a `gr.Blocks` layout with an explicit Run button, lazy-cached examples, and serialized queueing to avoid startup/download issues and PyTorch concurrency problems. The MegaDetector selector is now shown only for non-PyTorch backends, display settings are grouped under an accordion, and the bundled examples were cleaned up and expanded.

Files changed (2) hide show
  1. app.py +31 -22
  2. ui_utils.py +38 -46
app.py CHANGED
@@ -13,8 +13,7 @@ import dlclibrary
13
  import dlclive
14
  # import transformers
15
 
16
- from PIL import Image, ImageColor, ImageFont, ImageDraw
17
- import requests
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 detection_utils import predict_md, crop_animal_detections
@@ -35,10 +34,6 @@ from dlclive import DLCLive, Processor
35
  #train_dir = 'DLC_models/sa-tvm'
36
  #download_huggingface_model(model, train_dir)
37
 
38
- # grab demo data cooco cat:
39
- url = "http://images.cocodataset.org/val2017/000000039769.jpg"
40
- image = Image.open(requests.get(url, stream=True).raw)
41
-
42
  # megadetector and dlc model look up
43
  MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
44
  'md_v5b': "MD_models/md_v5b.0.0.pt"}
@@ -228,22 +223,36 @@ def predict_pipeline(img_input,
228
 
229
  #########################################################
230
  # Define user interface and launch
231
- inputs = gradio_inputs_for_MD_DLC(BACKENDS,
232
- list(MD_models_dict.keys()),
233
- list(DLC_models_dict.keys()))
234
- outputs = gradio_outputs_for_MD_DLC()
235
- [gr_title,
236
- gr_description,
237
  examples] = gradio_description_and_examples()
238
 
239
- # launch
240
- demo = gr.Interface(predict_pipeline,
241
- inputs=inputs,
242
- outputs=outputs,
243
- title=gr_title,
244
- description=gr_description,
245
- examples=examples,
246
- )
247
-
248
- demo.queue()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
249
  demo.launch(theme=gr.themes.Default())
 
13
  import dlclive
14
  # import transformers
15
 
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
 
34
  #train_dir = 'DLC_models/sa-tvm'
35
  #download_huggingface_model(model, train_dir)
36
 
 
 
 
 
37
  # megadetector and dlc model look up
38
  MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
39
  'md_v5b': "MD_models/md_v5b.0.0.pt"}
 
223
 
224
  #########################################################
225
  # Define user interface and launch
226
+ [gr_title,
227
+ gr_description,
 
 
 
 
228
  examples] = gradio_description_and_examples()
229
 
230
+ with gr.Blocks(title=gr_title) as demo:
231
+ gr.Markdown(f"# {gr_title}\n{gr_description}")
232
+ with gr.Row():
233
+ with gr.Column():
234
+ inputs = gradio_inputs_for_MD_DLC(BACKENDS,
235
+ list(MD_models_dict.keys()),
236
+ list(DLC_models_dict.keys()))
237
+ run_button = gr.Button("Run", variant="primary")
238
+ with gr.Column():
239
+ outputs = gradio_outputs_for_MD_DLC()
240
+
241
+ # the MegaDetector choice only applies to the TensorFlow (legacy) backend
242
+ gr_backend_input, gr_mega_model_input = inputs[1], inputs[2]
243
+ gr_backend_input.change(lambda backend: gr.update(visible=backend != "PyTorch"),
244
+ inputs=gr_backend_input,
245
+ outputs=gr_mega_model_input)
246
+
247
+ run_button.click(predict_pipeline, inputs=inputs, outputs=outputs, api_name="predict")
248
+
249
+ # cached on first click, so a failing download cannot block startup
250
+ gr.Examples(examples,
251
+ inputs=inputs,
252
+ outputs=outputs,
253
+ fn=predict_pipeline,
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())
ui_utils.py CHANGED
@@ -17,6 +17,7 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
17
  value="md_v5a",
18
  type="value",
19
  label="Select Detector model (TensorFlow legacy only)",
 
20
  )
21
 
22
  gr_dlc_model_input = gr.Dropdown(
@@ -32,11 +33,6 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
32
  label="Run DeepLabCut only, directly on input image?",
33
  )
34
 
35
- gr_str_labels_checkbox = gr.Checkbox(
36
- value=True,
37
- label="Show bodypart labels?",
38
- )
39
-
40
  # Gradio Slider signature is (minimum, maximum, value, step, ...)
41
  gr_slider_conf_bboxes = gr.Slider(
42
  minimum=0,
@@ -55,33 +51,39 @@ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
55
  )
56
 
57
  # Data viz
58
- gr_keypt_color = gr.ColorPicker(
59
- value="#862db7",
60
- label="Choose color for keypoint label",
61
- )
62
-
63
- gr_labels_font_style = gr.Dropdown(
64
- choices=["amiko", "animals", "nature", "painter", "zen"],
65
- value="amiko",
66
- type="value",
67
- label="Select keypoint label font",
68
- )
69
-
70
- gr_slider_font_size = gr.Slider(
71
- minimum=5,
72
- maximum=30,
73
- value=8,
74
- step=1,
75
- label="Set font size",
76
- )
77
-
78
- gr_slider_marker_size = gr.Slider(
79
- minimum=1,
80
- maximum=20,
81
- value=9,
82
- step=1,
83
- label="Set marker size",
84
- )
 
 
 
 
 
 
85
 
86
  return [
87
  gr_image_input,
@@ -116,19 +118,9 @@ def gradio_description_and_examples():
116
  "<a href='http://www.mackenziemathislab.org/dlc-modelzoo'>DeepLabCut ModelZoo</a>."
117
  )
118
 
119
- examples = [[
120
- "examples/dog.jpeg",
121
- "PyTorch",
122
- "md_v5a",
123
- "superanimal_quadruped",
124
- False,
125
- True,
126
- 0.5,
127
- 0.0,
128
- "amiko",
129
- 9,
130
- "#ff0000",
131
- 3,
132
- ]]
133
 
134
  return [title, description, examples]
 
17
  value="md_v5a",
18
  type="value",
19
  label="Select Detector model (TensorFlow legacy only)",
20
+ visible=gr_backend_input.value != "PyTorch",
21
  )
22
 
23
  gr_dlc_model_input = gr.Dropdown(
 
33
  label="Run DeepLabCut only, directly on input image?",
34
  )
35
 
 
 
 
 
 
36
  # Gradio Slider signature is (minimum, maximum, value, step, ...)
37
  gr_slider_conf_bboxes = gr.Slider(
38
  minimum=0,
 
51
  )
52
 
53
  # Data viz
54
+ with gr.Accordion("Display options", open=False):
55
+ gr_str_labels_checkbox = gr.Checkbox(
56
+ value=True,
57
+ label="Show bodypart labels?",
58
+ )
59
+
60
+ gr_keypt_color = gr.ColorPicker(
61
+ value="#862db7",
62
+ label="Choose color for keypoint label",
63
+ )
64
+
65
+ gr_labels_font_style = gr.Dropdown(
66
+ choices=["amiko", "animals", "nature", "painter", "zen"],
67
+ value="amiko",
68
+ type="value",
69
+ label="Select keypoint label font",
70
+ )
71
+
72
+ gr_slider_font_size = gr.Slider(
73
+ minimum=5,
74
+ maximum=30,
75
+ value=8,
76
+ step=1,
77
+ label="Set font size",
78
+ )
79
+
80
+ gr_slider_marker_size = gr.Slider(
81
+ minimum=1,
82
+ maximum=20,
83
+ value=9,
84
+ step=1,
85
+ label="Set marker size",
86
+ )
87
 
88
  return [
89
  gr_image_input,
 
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
 
126
  return [title, description, examples]