# 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())