mwmathis's picture C-Achard's picture
Add PyTorch SuperAnimal backend and refresh the Space (#14)
2e1b62d
Raw History Blame Contribute Delete
11.2 kB
# Adapted from https://huggingface.co/spaces/hlydecker/MegaDetector_v5
# Adapted from https://huggingface.co/spaces/sofmi/MegaDetector_DLClive/blob/main/app.py
# Adapted from https://huggingface.co/spaces/Neslihan/megadetector_dlcmodels/blob/main/app.py
# Adapted from https://huggingface.co/spaces/DeepLabCut/MegaDetector_DeepLabCut
import os
import threading
import gradio as gr
import numpy as np
import yaml
from dlclibrary.dlcmodelzoo.modelzoo_download import (
download_huggingface_model,
)
from dlclive import Processor
# import transformers
from PIL import Image
from detection_utils import crop_animal_detections, predict_md
from dlc_utils import predict_dlc
from pytorch_utils import PYTORCH_MODELS, load_superanimal, predict_superanimal
from ui_utils import (
confidence_legend_html,
dlc_theme,
gradio_description_and_examples,
gradio_inputs_for_MD_DLC,
gradio_outputs_for_MD_DLC,
)
from viz_utils import (
draw_bbox_w_text,
draw_keypoints_on_image,
keypoint_confidence_rows,
save_annotated_image,
save_results_as_json,
save_results_only_dlc,
save_results_pytorch,
)
# TESTING (passes) download the SuperAnimal models:
# model = 'superanimal_topviewmouse'
# train_dir = 'DLC_models/sa-tvm'
# download_huggingface_model(model, train_dir)
# megadetector and dlc model look up
MD_models_dict = {
"md_v5a": "MD_models/md_v5a.0.0.pt", #
"md_v5b": "MD_models/md_v5b.0.0.pt",
}
BACKENDS = ["PyTorch", "TensorFlow (legacy)"]
# TF (legacy) DLC models: model zoo name and target dir, per SuperAnimal
DLC_models_dict = {
"superanimal_topviewmouse": ("superanimal_topviewmouse_dlcrnet", "DLC_models/sa-tvm"),
"superanimal_quadruped": ("superanimal_quadruped_dlcrnet", "DLC_models/sa-q"),
}
#####################################################
def finalize_outputs(img_output, download_file, kpts_per_animal, map_label_id_to_str, color_by_confidence, colormap):
annotated_file = save_annotated_image(img_output)
confidence_rows = keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str)
legend = confidence_legend_html(colormap) if color_by_confidence else ""
return img_output, legend, download_file, annotated_file, confidence_rows
#####################################################
def predict_pipeline_pytorch(
img_input,
superanimal,
flag_dlc_only,
flag_show_str_labels,
bbox_likelihood_th,
kpts_likelihood_th,
font_style,
font_size,
keypt_color,
marker_size,
flag_color_by_confidence,
colormap,
bbox_color,
):
# detection + pose with the SuperAnimal PyTorch models (keypoints in image coords)
img_output, animals, bodyparts = predict_superanimal(
img_input, superanimal, bbox_likelihood_th, kpts_likelihood_th, full_image=flag_dlc_only
)
map_label_id_to_str = dict(enumerate(bodyparts))
for animal in animals:
draw_keypoints_on_image(
img_output,
animal["kpts"],
map_label_id_to_str,
flag_show_str_labels,
use_normalized_coordinates=False,
font_style=font_style,
font_size=font_size,
keypt_color=keypt_color,
marker_size=marker_size,
color_by_confidence=flag_color_by_confidence,
colormap=colormap,
)
if not flag_dlc_only:
draw_bbox_w_text(img_output, animal["bbox"], font_size=font_size, bbox_color=bbox_color)
pose_model, detector = PYTORCH_MODELS[superanimal]
download_file = save_results_pytorch(
animals,
map_label_id_to_str,
superanimal,
pose_model,
None if flag_dlc_only else detector,
image_size=img_input.size,
annotated_size=img_output.size,
)
return finalize_outputs(
img_output,
download_file,
[animal["kpts"] for animal in animals],
map_label_id_to_str,
flag_color_by_confidence,
colormap,
)
#####################################################
def predict_pipeline(
img_input,
backend,
mega_model_input,
dlc_model_input_str,
flag_dlc_only,
flag_show_str_labels,
bbox_likelihood_th,
kpts_likelihood_th,
font_style,
font_size,
keypt_color,
marker_size,
flag_color_by_confidence,
colormap,
bbox_color,
):
if backend == "PyTorch":
return predict_pipeline_pytorch(
img_input,
dlc_model_input_str,
flag_dlc_only,
flag_show_str_labels,
bbox_likelihood_th,
kpts_likelihood_th,
font_style,
font_size,
keypt_color,
marker_size,
flag_color_by_confidence,
colormap,
bbox_color,
)
# TensorFlow (legacy): MegaDetector crops + DLCLive
dlc_model_name, dlc_model_dir = DLC_models_dict[dlc_model_input_str]
if not flag_dlc_only:
############################################################
# ### Run Megadetector
md_results = predict_md(
img_input,
MD_models_dict[mega_model_input], # mega_model_input,
size=640,
) # Image.fromarray(results.imgs[0])
################################################################
# Obtain animal crops (and their bboxes) with confidence above th
list_crops, list_bboxes = crop_animal_detections(img_input, md_results, bbox_likelihood_th)
############################################################
## Get DLC model and label map
# If model is found: do not download (previous execution is likely within same day)
# TODO: can we ask the user whether to reload dlc model if a directory is found?
path_to_DLCmodel = dlc_model_dir
if not (os.path.isdir(dlc_model_dir) and len(os.listdir(dlc_model_dir)) > 0):
download_huggingface_model(dlc_model_name, path_to_DLCmodel)
# extract map label ids to strings
pose_cfg_path = os.path.join(dlc_model_dir, "pose_cfg.yaml")
with open(pose_cfg_path) as stream:
pose_cfg_dict = yaml.safe_load(stream)
map_label_id_to_str = dict(
[
(k, v)
for k, v in zip(
[
el[0] for el in pose_cfg_dict["all_joints"]
], # pose_cfg_dict['all_joints'] is a list of one-element lists,
pose_cfg_dict["all_joints_names"],
strict=True,
)
]
)
##############################################################
# Run DLC and visualize results
dlc_proc = Processor() # TODO: update deeplabcut.video_inference_superanimal() once merged
# if required: ignore MD crops and run DLC on full image [mostly for testing]
if flag_dlc_only:
# compute kpts on input img
list_kpts_per_crop = predict_dlc([np.asarray(img_input)], kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
# draw kpts on input img #fix!
draw_keypoints_on_image(
img_input,
list_kpts_per_crop[0], # a numpy array with shape [num_keypoints, 2].
map_label_id_to_str,
flag_show_str_labels,
use_normalized_coordinates=False,
font_style=font_style,
font_size=font_size,
keypt_color=keypt_color,
marker_size=marker_size,
color_by_confidence=flag_color_by_confidence,
colormap=colormap,
)
donw_file = save_results_only_dlc(
list_kpts_per_crop[0], map_label_id_to_str, dlc_model_name, image_size=img_input.size
)
return finalize_outputs(
img_input, donw_file, [list_kpts_per_crop[0]], map_label_id_to_str, flag_color_by_confidence, colormap
)
else:
# Compute kpts for each crop
list_kpts_per_crop = predict_dlc(list_crops, kpts_likelihood_th, path_to_DLCmodel, dlc_proc)
# resize input image to match megadetector output
img_background = img_input.resize((md_results.ims[0].shape[1], md_results.ims[0].shape[0]))
# draw keypoints on each crop and paste to background img
for np_crop, kpts_crop, bb_per_animal in zip(list_crops, list_kpts_per_crop, list_bboxes, strict=True):
img_crop = Image.fromarray(np_crop)
# Draw keypts on crop
draw_keypoints_on_image(
img_crop,
kpts_crop, # a numpy array with shape [num_keypoints, 2].
map_label_id_to_str,
flag_show_str_labels,
use_normalized_coordinates=False, # if True, then I should use md_results.xyxyn for list_kpts_crop
font_style=font_style,
font_size=font_size,
keypt_color=keypt_color,
marker_size=marker_size,
color_by_confidence=flag_color_by_confidence,
colormap=colormap,
)
# Paste crop in original image
img_background.paste(img_crop, box=tuple([int(t) for t in bb_per_animal[:2]]))
# Plot bbox
draw_bbox_w_text(img_background, bb_per_animal, font_size=font_size, bbox_color=bbox_color)
# Save detection results as json
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,
image_size=img_input.size,
)
return finalize_outputs(
img_background, download_file, list_kpts_per_crop, map_label_id_to_str, flag_color_by_confidence, colormap
)
#########################################################
# Define user interface and launch
[gr_title, gr_description, examples] = gradio_description_and_examples()
with gr.Blocks(title=gr_title) as demo:
gr.Markdown(f"# {gr_title}\n{gr_description}")
with gr.Row():
with gr.Column():
inputs = gradio_inputs_for_MD_DLC(BACKENDS, list(MD_models_dict.keys()), list(DLC_models_dict.keys()))
run_button = gr.Button("Run", variant="primary")
with gr.Column():
outputs = gradio_outputs_for_MD_DLC()
# the MegaDetector choice only applies to the TensorFlow (legacy) backend
gr_backend_input, gr_mega_model_input = inputs[1], inputs[2]
gr_backend_input.change(
lambda backend: gr.update(visible=backend != "PyTorch"), inputs=gr_backend_input, outputs=gr_mega_model_input
)
run_button.click(predict_pipeline, inputs=inputs, outputs=outputs, api_name="predict")
# cached on first click, so a failing download cannot block startup
gr.Examples(examples, inputs=inputs, outputs=outputs, fn=predict_pipeline, cache_examples=True, cache_mode="lazy")
# download and build the default model while the app starts; a request arriving
# earlier waits on the same lock instead of downloading again
threading.Thread(target=load_superanimal, args=("superanimal_quadruped",), daemon=True).start()
demo.queue(default_concurrency_limit=1) # PyTorch runners are not thread-safe
demo.launch(theme=dlc_theme())