Download pytorch_utils.py from DeepLabCut/DeepLabCutModelZoo-SuperAnimals: direct link, hf CLI and curl.
- Browser
- Download file 4.68 kB
-
https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/935b48a0cbdab550059c9246881d528d812d7486/pytorch_utils.py
- Command line
-
hf download hf://spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals@935b48a0cbdab550059c9246881d528d812d7486/pytorch_utils.py
-
curl -L -o pytorch_utils.py https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/935b48a0cbdab550059c9246881d528d812d7486/pytorch_utils.py
4.68 kB
| import threading | |
| import numpy as np | |
| import PIL | |
| from deeplabcut.pose_estimation_pytorch.apis.utils import get_inference_runners | |
| from deeplabcut.pose_estimation_pytorch.config.pose import PoseConfig | |
| from deeplabcut.pose_estimation_pytorch.modelzoo.utils import get_super_animal_snapshot_path | |
| # SuperAnimal (pose model, detector) used by the PyTorch backend | |
| PYTORCH_MODELS = {'superanimal_quadruped': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2'), | |
| 'superanimal_topviewmouse': ('hrnet_w32', 'fasterrcnn_resnet50_fpn_v2')} | |
| MAX_INDIVIDUALS = 10 | |
| MAX_IMAGE_SIZE = 1280 # longest side fed to the models (and drawn on) | |
| _runners = {} | |
| _build_lock = threading.Lock() | |
| ########################################## | |
| def load_superanimal(superanimal, device="auto"): | |
| """Build (once) the detector and pose runners for a SuperAnimal model. | |
| Weights are downloaded on first use to deeplabcut/modelzoo/checkpoints. | |
| """ | |
| with _build_lock: | |
| if superanimal not in _runners: | |
| pose_model, detector = PYTORCH_MODELS[superanimal] | |
| cfg = PoseConfig.build_for_superanimal_inference(superanimal, | |
| model_name=pose_model, | |
| detector_name=detector, | |
| max_individuals=MAX_INDIVIDUALS, | |
| device=device) | |
| # keep low-score boxes: the UI threshold filters them afterwards | |
| cfg["detector"]["model"]["box_score_thresh"] = 0.05 | |
| pose_runner, detector_runner = get_inference_runners( | |
| cfg, | |
| snapshot_path=get_super_animal_snapshot_path(superanimal, pose_model), | |
| detector_path=get_super_animal_snapshot_path(superanimal, detector), | |
| max_individuals=MAX_INDIVIDUALS, | |
| inference_cfg={"multithreading": {"enabled": False}}, | |
| ) | |
| _runners[superanimal] = {"pose": pose_runner, | |
| "detector": detector_runner, | |
| "bodyparts": list(cfg["metadata"]["bodyparts"]), | |
| "lock": threading.Lock()} # runners are not thread-safe | |
| return _runners[superanimal] | |
| ########################################## | |
| def resize_max_side(img, max_size=MAX_IMAGE_SIZE): | |
| scale = max_size / max(img.size) | |
| if scale >= 1: | |
| return img | |
| return img.resize([int(x * scale) for x in img.size], PIL.Image.Resampling.LANCZOS) | |
| ########################################## | |
| def predict_superanimal(img_input, | |
| superanimal, | |
| bbox_likelihood_th, | |
| kpts_likelihood_th, | |
| full_image=False): | |
| """Detect animals and estimate their pose with a PyTorch SuperAnimal model. | |
| Returns the (resized) RGB image the predictions refer to, the list of animals | |
| as dicts {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3) array of x,y,llk | |
| in image pixels, NaN below kpts_likelihood_th}, and the bodypart names. | |
| """ | |
| img = resize_max_side(img_input.convert("RGB")) | |
| img_np = np.asarray(img) | |
| runners = load_superanimal(superanimal) | |
| with runners["lock"]: | |
| if full_image: | |
| # skip the detector and treat the whole image as one animal | |
| h, w = img_np.shape[:2] | |
| detections = {"bboxes": np.array([[0, 0, w, h]], dtype=np.float32), | |
| "bbox_scores": np.array([1.0], dtype=np.float32)} | |
| else: | |
| detections = runners["detector"].inference([img_np])[0] # bboxes in xywh | |
| keep = detections["bbox_scores"] >= bbox_likelihood_th | |
| detections = {"bboxes": detections["bboxes"][keep], | |
| "bbox_scores": detections["bbox_scores"][keep]} | |
| if len(detections["bboxes"]) == 0: | |
| return img, [], runners["bodyparts"] | |
| predictions = runners["pose"].inference([(img_np, detections)])[0] | |
| animals = [] | |
| # outputs are padded to MAX_INDIVIDUALS with -1 | |
| for kpts, (x, y, w, h), score in zip(predictions["bodyparts"], | |
| predictions["bboxes"], | |
| predictions["bbox_scores"]): | |
| if score < 0: | |
| continue | |
| kpts = kpts.astype(float) | |
| kpts[kpts[:, 2] < kpts_likelihood_th, :] = np.nan | |
| animals.append({"bbox": [float(x), float(y), float(x + w), float(y + h), float(score)], | |
| "kpts": kpts}) | |
| return img, animals, runners["bodyparts"] | |