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

Add PyTorch SuperAnimal backend

Browse files

Introduce a PyTorch inference path for SuperAnimal models, including backend selection in the UI, cached pose/detector runners, and JSON export of predictions. The legacy TensorFlow flow now maps model names separately from local directories, and MegaDetector device selection was corrected for yolov5 loading. The DeepLabCut dependency is also pinned to 3.0.2 to match the new backend integration.

Files changed (6) hide show
  1. app.py +73 -15
  2. detection_utils.py +3 -11
  3. pytorch_utils.py +104 -0
  4. requirements.txt +1 -1
  5. ui_utils.py +13 -5
  6. viz_utils.py +33 -0
app.py CHANGED
@@ -16,9 +16,10 @@ import dlclive
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
20
  from detection_utils import predict_md, crop_animal_detections
21
  from dlc_utils import predict_dlc
 
22
  from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
23
 
24
  from deeplabcut.utils import auxiliaryfunctions
@@ -42,14 +43,58 @@ image = Image.open(requests.get(url, stream=True).raw)
42
  MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
43
  'md_v5b': "MD_models/md_v5b.0.0.pt"}
44
 
45
- # DLC models target dirs
46
- DLC_models_dict = {'superanimal_topviewmouse_dlcrnet': "DLC_models/sa-tvm",
47
- 'superanimal_quadruped_dlcrnet': "DLC_models/sa-q"}
48
-
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49
 
50
 
51
  #####################################################
52
  def predict_pipeline(img_input,
 
53
  mega_model_input,
54
  dlc_model_input_str,
55
  flag_dlc_only,
@@ -62,6 +107,21 @@ def predict_pipeline(img_input,
62
  marker_size,
63
  ):
64
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
65
  if not flag_dlc_only:
66
  ############################################################
67
  # ### Run Megadetector
@@ -81,15 +141,12 @@ def predict_pipeline(img_input,
81
 
82
  # If model is found: do not download (previous execution is likely within same day)
83
  # TODO: can we ask the user whether to reload dlc model if a directory is found?
84
- if os.path.isdir(DLC_models_dict[dlc_model_input_str]) and \
85
- len(os.listdir(DLC_models_dict[dlc_model_input_str])) > 0:
86
- path_to_DLCmodel = DLC_models_dict[dlc_model_input_str]
87
- else:
88
- path_to_DLCmodel = DLC_models_dict[dlc_model_input_str]
89
- download_huggingface_model(dlc_model_input_str, path_to_DLCmodel)
90
 
91
  # extract map label ids to strings
92
- pose_cfg_path = os.path.join(DLC_models_dict[dlc_model_input_str],
93
  'pose_cfg.yaml')
94
  with open(pose_cfg_path, "r") as stream:
95
  pose_cfg_dict = yaml.safe_load(stream)
@@ -119,7 +176,7 @@ def predict_pipeline(img_input,
119
  keypt_color=keypt_color,
120
  marker_size=marker_size)
121
 
122
- donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_input_str)
123
 
124
  return img_input, donw_file
125
 
@@ -163,7 +220,7 @@ def predict_pipeline(img_input,
163
 
164
 
165
  # Save detection results as json
166
- download_file = save_results_as_json(md_results,list_kpts_per_crop,list_bboxes,map_label_id_to_str,dlc_model_input_str,mega_model_input)
167
 
168
  return img_background, download_file
169
 
@@ -171,7 +228,8 @@ def predict_pipeline(img_input,
171
 
172
  #########################################################
173
  # Define user interface and launch
174
- inputs = gradio_inputs_for_MD_DLC(list(MD_models_dict.keys()),
 
175
  list(DLC_models_dict.keys()))
176
  outputs = gradio_outputs_for_MD_DLC()
177
  [gr_title,
 
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
21
  from dlc_utils import predict_dlc
22
+ from pytorch_utils import predict_superanimal, PYTORCH_MODELS
23
  from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
24
 
25
  from deeplabcut.utils import auxiliaryfunctions
 
43
  MD_models_dict = {'md_v5a': "MD_models/md_v5a.0.0.pt", #
44
  'md_v5b': "MD_models/md_v5b.0.0.pt"}
45
 
46
+ BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
47
+
48
+ # TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
49
+ DLC_models_dict = {'superanimal_topviewmouse': ('superanimal_topviewmouse_dlcrnet', "DLC_models/sa-tvm"),
50
+ 'superanimal_quadruped': ('superanimal_quadruped_dlcrnet', "DLC_models/sa-q")}
51
+
52
+
53
+
54
+ #####################################################
55
+ def predict_pipeline_pytorch(img_input,
56
+ superanimal,
57
+ flag_dlc_only,
58
+ flag_show_str_labels,
59
+ bbox_likelihood_th,
60
+ kpts_likelihood_th,
61
+ font_style,
62
+ font_size,
63
+ keypt_color,
64
+ marker_size,
65
+ ):
66
+ # detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
67
+ img_output, animals, bodyparts = predict_superanimal(img_input,
68
+ superanimal,
69
+ bbox_likelihood_th,
70
+ kpts_likelihood_th,
71
+ full_image=flag_dlc_only)
72
+ map_label_id_to_str = dict(enumerate(bodyparts))
73
+
74
+ for animal in animals:
75
+ draw_keypoints_on_image(img_output,
76
+ animal['kpts'],
77
+ map_label_id_to_str,
78
+ flag_show_str_labels,
79
+ use_normalized_coordinates=False,
80
+ font_style=font_style,
81
+ font_size=font_size,
82
+ keypt_color=keypt_color,
83
+ marker_size=marker_size)
84
+ if not flag_dlc_only:
85
+ draw_bbox_w_text(img_output,
86
+ animal['bbox'],
87
+ font_size=font_size)
88
+
89
+ pose_model, detector = PYTORCH_MODELS[superanimal]
90
+ download_file = save_results_pytorch(animals, map_label_id_to_str, superanimal,
91
+ pose_model, None if flag_dlc_only else detector)
92
+ return img_output, download_file
93
 
94
 
95
  #####################################################
96
  def predict_pipeline(img_input,
97
+ backend,
98
  mega_model_input,
99
  dlc_model_input_str,
100
  flag_dlc_only,
 
107
  marker_size,
108
  ):
109
 
110
+ if backend == "PyTorch":
111
+ return predict_pipeline_pytorch(img_input,
112
+ dlc_model_input_str,
113
+ flag_dlc_only,
114
+ flag_show_str_labels,
115
+ bbox_likelihood_th,
116
+ kpts_likelihood_th,
117
+ font_style,
118
+ font_size,
119
+ keypt_color,
120
+ marker_size)
121
+
122
+ # TensorFlow (legacy): MegaDetector crops + DLCLive
123
+ dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
124
+
125
  if not flag_dlc_only:
126
  ############################################################
127
  # ### Run Megadetector
 
141
 
142
  # If model is found: do not download (previous execution is likely within same day)
143
  # TODO: can we ask the user whether to reload dlc model if a directory is found?
144
+ path_to_DLCmodel = dlc_model_dir
145
+ if not (os.path.isdir(dlc_model_dir) and len(os.listdir(dlc_model_dir)) > 0):
146
+ download_huggingface_model(dlc_model_name, path_to_DLCmodel)
 
 
 
147
 
148
  # extract map label ids to strings
149
+ pose_cfg_path = os.path.join(dlc_model_dir,
150
  'pose_cfg.yaml')
151
  with open(pose_cfg_path, "r") as stream:
152
  pose_cfg_dict = yaml.safe_load(stream)
 
176
  keypt_color=keypt_color,
177
  marker_size=marker_size)
178
 
179
+ donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_name)
180
 
181
  return img_input, donw_file
182
 
 
220
 
221
 
222
  # Save detection results as json
223
+ download_file = save_results_as_json(md_results,list_kpts_per_crop,list_bboxes,map_label_id_to_str,dlc_model_name,mega_model_input)
224
 
225
  return img_background, download_file
226
 
 
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,
detection_utils.py CHANGED
@@ -24,11 +24,8 @@ def predict_md(im,
24
  g = (size / max(im.size)) # multipl factor to make max size of the image equal to input size
25
  im = im.resize((int(x * g) for x in im.size),
26
  PIL.Image.Resampling.LANCZOS) # resize
27
- # device
28
- if torch.cuda.is_available():
29
- md_device = torch.device('cuda')
30
- else:
31
- md_device = torch.device('cpu')
32
 
33
  # megadetector
34
  MD_model = torch.hub.load('ultralytics/yolov5', # repo_or_dir
@@ -37,12 +34,7 @@ def predict_md(im,
37
  skip_validation=True, # avoid GitHub API rate limit (403)
38
  device=md_device,
39
  trust_repo=True
40
- )
41
-
42
- # send model to gpu if possible
43
- if (md_device == torch.device('cuda')):
44
- print('Sending model to GPU')
45
- MD_model.to(md_device)
46
 
47
  ## detect objects
48
  results = MD_model(im) # inference # vars(results).keys()= dict_keys(['imgs', 'pred', 'names', 'files', 'times', 'xyxy', 'xywh', 'xyxyn', 'xywhn', 'n', 't', 's'])
 
24
  g = (size / max(im.size)) # multipl factor to make max size of the image equal to input size
25
  im = im.resize((int(x * g) for x in im.size),
26
  PIL.Image.Resampling.LANCZOS) # resize
27
+ # device: yolov5's select_device expects a CUDA index ('0') or 'cpu', not 'cuda'
28
+ md_device = '0' if torch.cuda.is_available() else 'cpu'
 
 
 
29
 
30
  # megadetector
31
  MD_model = torch.hub.load('ultralytics/yolov5', # repo_or_dir
 
34
  skip_validation=True, # avoid GitHub API rate limit (403)
35
  device=md_device,
36
  trust_repo=True
37
+ )
 
 
 
 
 
38
 
39
  ## detect objects
40
  results = MD_model(im) # inference # vars(results).keys()= dict_keys(['imgs', 'pred', 'names', 'files', 'times', 'xyxy', 'xywh', 'xyxyn', 'xywhn', 'n', 't', 's'])
pytorch_utils.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import threading
2
+
3
+ import numpy as np
4
+ import PIL
5
+ from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners
6
+ from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig
7
+ from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path
8
+
9
+
10
+ # SuperAnimal (pose model, detector) used by the PyTorch backend
11
+ PYTORCH_MODELS = {'superanimal_quadruped': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2'),
12
+ 'superanimal_topviewmouse': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2')}
13
+
14
+ MAX_INDIVIDUALS = 10
15
+ MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on)
16
+
17
+ _runners = {}
18
+ _build_lock = threading.Lock()
19
+
20
+
21
+ ##########################################
22
+ def load_superanimal(superanimal, device="auto"):
23
+ """Build (once) the detector and pose runners for a SuperAnimal model.
24
+
25
+ Weights are downloaded on first use to deeplabcut/modelzoo/checkpoints.
26
+ """
27
+ with _build_lock:
28
+ if superanimal not in _runners:
29
+ pose_model, detector = PYTORCH_MODELS[superanimal]
30
+ cfg = PoseConfig.build_for_superanimal_inference(superanimal,
31
+ model_name=pose_model,
32
+ detector_name=detector,
33
+ max_individuals=MAX_INDIVIDUALS,
34
+ device=device)
35
+ # keep low-score boxes: the UI threshold filters them afterwards
36
+ cfg["detector"]["model"]["box_score_thresh"] = 0.05
37
+ pose_runner, detector_runner = get_inference_runners(
38
+ cfg,
39
+ snapshot_path=get_super_animal_snapshot_path(superanimal, pose_model),
40
+ detector_path=get_super_animal_snapshot_path(superanimal, detector),
41
+ max_individuals=MAX_INDIVIDUALS,
42
+ inference_cfg={"multithreading": {"enabled": False}},
43
+ )
44
+ _runners[superanimal] = {"pose": pose_runner,
45
+ "detector": detector_runner,
46
+ "bodyparts": list(cfg["metadata"]["bodyparts"]),
47
+ "lock": threading.Lock()} # runners are not thread-safe
48
+ return _runners[superanimal]
49
+
50
+
51
+ ##########################################
52
+ def resize_max_side(img, max_size=MAX_IMAGE_SIZE):
53
+ scale = max_size / max(img.size)
54
+ if scale >= 1:
55
+ return img
56
+ return img.resize([int(x * scale) for x in img.size], PIL.Image.Resampling.LANCZOS)
57
+
58
+
59
+ ##########################################
60
+ def predict_superanimal(img_input,
61
+ superanimal,
62
+ bbox_likelihood_th,
63
+ kpts_likelihood_th,
64
+ full_image=False):
65
+ """Detect animals and estimate their pose with a PyTorch SuperAnimal model.
66
+
67
+ Returns the (resized) RGB image the predictions refer to, the list of animals
68
+ as dicts {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3) array of x,y,llk
69
+ in image pixels, NaN below kpts_likelihood_th}, and the bodypart names.
70
+ """
71
+ img = resize_max_side(img_input.convert("RGB"))
72
+ img_np = np.asarray(img)
73
+ runners = load_superanimal(superanimal)
74
+
75
+ with runners["lock"]:
76
+ if full_image:
77
+ # skip the detector and treat the whole image as one animal
78
+ h, w = img_np.shape[:2]
79
+ detections = {"bboxes": np.array([[0, 0, w, h]], dtype=np.float32),
80
+ "bbox_scores": np.array([1.0], dtype=np.float32)}
81
+ else:
82
+ detections = runners["detector"].inference([img_np])[0] # bboxes in xywh
83
+ keep = detections["bbox_scores"] >= bbox_likelihood_th
84
+ detections = {"bboxes": detections["bboxes"][keep],
85
+ "bbox_scores": detections["bbox_scores"][keep]}
86
+
87
+ if len(detections["bboxes"]) == 0:
88
+ return img, [], runners["bodyparts"]
89
+
90
+ predictions = runners["pose"].inference([(img_np, detections)])[0]
91
+
92
+ animals = []
93
+ # outputs are padded to MAX_INDIVIDUALS with -1
94
+ for kpts, (x, y, w, h), score in zip(predictions["bodyparts"],
95
+ predictions["bboxes"],
96
+ predictions["bbox_scores"]):
97
+ if score < 0:
98
+ continue
99
+ kpts = kpts.astype(float)
100
+ kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan
101
+ animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)],
102
+ "kpts": kpts})
103
+
104
+ return img, animals, runners["bodyparts"]
requirements.txt CHANGED
@@ -1,7 +1,7 @@
1
  gradio
2
  gitpython>=3.1.30
3
  seaborn
4
- deeplabcut[modelzoo,tf]>=3.0.0rc14
5
  deeplabcut-live
6
  ruamel.yaml==0.17.21
7
  dlclibrary
 
1
  gradio
2
  gitpython>=3.1.30
3
  seaborn
4
+ deeplabcut[modelzoo,tf]==3.0.2
5
  deeplabcut-live
6
  ruamel.yaml==0.17.21
7
  dlclibrary
ui_utils.py CHANGED
@@ -1,21 +1,27 @@
1
  import gradio as gr
2
 
3
 
4
- def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
5
  # Input image
6
  gr_image_input = gr.Image(type="pil", label="Input Image")
7
 
8
  # Models
 
 
 
 
 
 
9
  gr_mega_model_input = gr.Dropdown(
10
  choices=md_models_list,
11
  value="md_v5a",
12
  type="value",
13
- label="Select Detector model",
14
  )
15
 
16
  gr_dlc_model_input = gr.Dropdown(
17
  choices=dlc_models_list,
18
- value="superanimal_quadruped_dlcrnet",
19
  type="value",
20
  label="Select DeepLabCut model",
21
  )
@@ -23,7 +29,7 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
23
  # Other inputs
24
  gr_dlc_only_checkbox = gr.Checkbox(
25
  value=False,
26
- label="Run DLClive only, directly on input image?",
27
  )
28
 
29
  gr_str_labels_checkbox = gr.Checkbox(
@@ -79,6 +85,7 @@ def gradio_inputs_for_MD_DLC(md_models_list, dlc_models_list):
79
 
80
  return [
81
  gr_image_input,
 
82
  gr_mega_model_input,
83
  gr_dlc_model_input,
84
  gr_dlc_only_checkbox,
@@ -111,8 +118,9 @@ def gradio_description_and_examples():
111
 
112
  examples = [[
113
  "examples/dog.jpeg",
 
114
  "md_v5a",
115
- "superanimal_quadruped_dlcrnet",
116
  False,
117
  True,
118
  0.5,
 
1
  import gradio as gr
2
 
3
 
4
+ def gradio_inputs_for_MD_DLC(backends_list, md_models_list, dlc_models_list):
5
  # Input image
6
  gr_image_input = gr.Image(type="pil", label="Input Image")
7
 
8
  # Models
9
+ gr_backend_input = gr.Radio(
10
+ choices=backends_list,
11
+ value=backends_list[0],
12
+ label="Select backend",
13
+ )
14
+
15
  gr_mega_model_input = gr.Dropdown(
16
  choices=md_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(
23
  choices=dlc_models_list,
24
+ value="superanimal_quadruped",
25
  type="value",
26
  label="Select DeepLabCut model",
27
  )
 
29
  # Other inputs
30
  gr_dlc_only_checkbox = gr.Checkbox(
31
  value=False,
32
+ label="Run DeepLabCut only, directly on input image?",
33
  )
34
 
35
  gr_str_labels_checkbox = gr.Checkbox(
 
85
 
86
  return [
87
  gr_image_input,
88
+ gr_backend_input,
89
  gr_mega_model_input,
90
  gr_dlc_model_input,
91
  gr_dlc_only_checkbox,
 
118
 
119
  examples = [[
120
  "examples/dog.jpeg",
121
+ "PyTorch",
122
  "md_v5a",
123
+ "superanimal_quadruped",
124
  False,
125
  True,
126
  0.5,
viz_utils.py CHANGED
@@ -185,4 +185,37 @@ def save_results_only_dlc(dlc_outputs,map_label_id_to_str,model,output_file = 'd
185
  return output_file
186
 
187
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
188
  ###########################################
 
185
  return output_file
186
 
187
 
188
+ def save_results_pytorch(animals, map_label_id_to_str, model, pose_model, detector, path_to_output_file = 'download_predictions.json'):
189
+
190
+ """
191
+ Output PyTorch SuperAnimal predictions as json file (same layout as save_results_as_json)
192
+
193
+ animals: list of {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3)}, in image coords
194
+ detector: None if the detector was skipped (whole image used as one animal)
195
+ """
196
+ info = {}
197
+ info['date'] = str(today)
198
+ info['backend'] = 'pytorch'
199
+ info['dlc_model'] = model
200
+ info['pose_model'] = pose_model
201
+ info['detector'] = detector
202
+ info['number_of_bb'] = len(animals)
203
+ labels = [n for n in map_label_id_to_str.values()]
204
+
205
+ for i, animal in enumerate(animals):
206
+ corner_x1, corner_y1, corner_x2, corner_y2, confidence = animal['bbox']
207
+ aux = {}
208
+ aux['corner_1'] = (corner_x1, corner_y1)
209
+ aux['corner_2'] = (corner_x2, corner_y2)
210
+ aux['confidence'] = confidence
211
+ aux['dlc_pred'] = dict(zip(labels, [[float(v) for v in kpt] for kpt in animal['kpts']]))
212
+ info['bb_' + str(i)] = aux
213
+
214
+ with open(path_to_output_file, 'w') as f:
215
+ json.dump(info, f, indent=1)
216
+ print('Output file saved at {}'.format(path_to_output_file))
217
+
218
+ return path_to_output_file
219
+
220
+
221
  ###########################################