C-Achard's picture
Add confidence-based keypoint outputs
e189412
Raw History Blame
13.3 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 yaml
import numpy as np
from matplotlib import cm
import gradio as gr
import deeplabcut
import dlclibrary
import dlclive
# import transformers
from PIL import Image, ImageColor, ImageFont, ImageDraw
from viz_utils import save_results_as_json, draw_keypoints_on_image, draw_bbox_w_text, save_results_only_dlc, save_results_pytorch
from viz_utils import add_confidence_legend, keypoint_confidence_rows, save_annotated_image
from detection_utils import predict_md, crop_animal_detections
from dlc_utils import predict_dlc
from pytorch_utils import predict_superanimal, load_superanimal, PYTORCH_MODELS
from ui_utils import gradio_inputs_for_MD_DLC, gradio_outputs_for_MD_DLC, gradio_description_and_examples
from deeplabcut.utils import auxiliaryfunctions
from dlclibrary.dlcmodelzoo.modelzoo_download import (
download_huggingface_model,
MODELOPTIONS,
)
from dlclive import DLCLive, Processor
# 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):
# confidence legend, annotated image for download and per-keypoint confidence table
if color_by_confidence:
img_output = add_confidence_legend(img_output)
annotated_file = save_annotated_image(img_output)
confidence_rows = keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str)
return img_output, 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,
):
# 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)
if not flag_dlc_only:
draw_bbox_w_text(img_output,
animal['bbox'],
font_size=font_size)
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)
return finalize_outputs(img_output, download_file,
[animal['kpts'] for animal in animals], map_label_id_to_str,
flag_color_by_confidence)
#####################################################
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,
):
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)
# 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, "r") 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'])])
##############################################################
# 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)
donw_file = save_results_only_dlc(list_kpts_per_crop[0], map_label_id_to_str,dlc_model_name)
return finalize_outputs(img_input, donw_file,
[list_kpts_per_crop[0]], map_label_id_to_str,
flag_color_by_confidence)
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):
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)
# 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) # TODO: add selectable color for bbox?
# 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)
return finalize_outputs(img_background, download_file,
list_kpts_per_crop, map_label_id_to_str,
flag_color_by_confidence)
#########################################################
# 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=gr.themes.Default())