Download app.py from DeepLabCut/DeepLabCutModelZoo-SuperAnimals: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/main/app.py
- Command line
-
hf download hf://spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/app.py
-
curl -L -o app.py https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/main/app.py
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()) | |