Download pytorch_utils.py from DeepLabCut/DeepLabCutModelZoo-SuperAnimals: direct link, hf CLI and curl.
- Browser
- Download file 5.02 kB
-
https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/main/pytorch_utils.py
- Command line
-
hf download hf://spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/pytorch_utils.py
-
curl -L -o pytorch_utils.py https://huggingface.co/spaces/DeepLabCut/DeepLabCutModelZoo-SuperAnimals/resolve/main/pytorch_utils.py
5.02 kB
| import threading | |
| from pathlib import Path | |
| 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 MODEL_FILENAME_MAPPING | |
| from dlclibrary import download_huggingface_model | |
| # 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) | |
| # next to the TF models, not deeplabcut/modelzoo/checkpoints: site-packages is read-only for the Space's non-root user | |
| WEIGHTS_DIR = Path(__file__).parent / "DLC_models" / "pytorch" | |
| _runners = {} | |
| _build_lock = threading.Lock() | |
| ########################################## | |
| def snapshot_path(superanimal, model_name): | |
| """Path to a SuperAnimal snapshot in WEIGHTS_DIR, downloaded on first use (as deeplabcut does).""" | |
| name = f"{superanimal}_{model_name}" | |
| path = WEIGHTS_DIR / f"{name}.pt" | |
| if not path.exists(): | |
| source = MODEL_FILENAME_MAPPING.get(name, path.name) | |
| rename = None if source == path.name else {source: path.name} | |
| download_huggingface_model(name, target_dir=str(WEIGHTS_DIR), rename_mapping=rename) | |
| return path | |
| ########################################## | |
| def load_superanimal(superanimal, device="auto"): | |
| """Build (once) the detector and pose runners for a SuperAnimal model; weights are downloaded on first use.""" | |
| 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=snapshot_path(superanimal, pose_model), | |
| detector_path=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"], strict=True | |
| ): | |
| 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"] | |