Instructions to use copperscout1/vitpose-plus-base with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use copperscout1/vitpose-plus-base with Transformers:
# Load model directly from transformers import AutoImageProcessor, VitPoseForPoseEstimation processor = AutoImageProcessor.from_pretrained("copperscout1/vitpose-plus-base") model = VitPoseForPoseEstimation.from_pretrained("copperscout1/vitpose-plus-base", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download handler.py from copperscout1/vitpose-plus-base: direct link, hf CLI and curl.
- Browser
- Download file 9.09 kB
-
https://huggingface.co/copperscout1/vitpose-plus-base/resolve/main/handler.py
- Command line
-
hf download hf://copperscout1/vitpose-plus-base/handler.py
-
curl -L -o handler.py https://huggingface.co/copperscout1/vitpose-plus-base/resolve/main/handler.py
9.09 kB
| """Custom Inference Endpoint handler for ViTPose+ (top-down pose). | |
| ViTPose is not a catalog / pipeline task, so Hugging Face Inference Endpoints | |
| need this EndpointHandler. It detects people (RT-DETR) then estimates COCO-17 | |
| keypoints. Pass `boxes` to skip the detector. | |
| Request JSON: | |
| {"inputs": "<base64 or URL>", "parameters": {"threshold": 0.3, "dataset_index": 0}} | |
| {"inputs": "<base64>", "boxes": [[x, y, w, h], ...]} # COCO xywh, skips detector | |
| dataset_index (ViTPose+ MoE experts): | |
| 0 COCO, 1 AIC, 2 MPII, 3 AP-10K, 4 APT-36K, 5 COCO-WholeBody | |
| """ | |
| from __future__ import annotations | |
| import base64 | |
| import io | |
| from typing import Any | |
| from urllib.parse import urlparse | |
| import numpy as np | |
| import requests | |
| import torch | |
| from PIL import Image | |
| from transformers import AutoProcessor, RTDetrForObjectDetection, VitPoseForPoseEstimation | |
| DETECTOR_ID = "PekingU/rtdetr_r50vd_coco_o365" | |
| POSE_FALLBACK_ID = "usyd-community/vitpose-plus-base" | |
| def _as_float(value: Any) -> float: | |
| if hasattr(value, "item"): | |
| return float(value.item()) | |
| return float(value) | |
| def _as_list(value: Any) -> list: | |
| if hasattr(value, "detach"): | |
| value = value.detach().cpu().numpy() | |
| if hasattr(value, "tolist"): | |
| return value.tolist() | |
| return list(value) | |
| class EndpointHandler: | |
| def __init__(self, path: str = "") -> None: | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.dtype = torch.float16 if self.device.type == "cuda" else torch.float32 | |
| pose_id = path or POSE_FALLBACK_ID | |
| self.processor = AutoProcessor.from_pretrained(pose_id) | |
| self.model = VitPoseForPoseEstimation.from_pretrained(pose_id) | |
| self.model.to(device=self.device, dtype=self.dtype) | |
| self.model.eval() | |
| backbone = getattr(self.model.config, "backbone_config", None) | |
| self.num_experts = int(getattr(backbone, "num_experts", 1) or 1) | |
| self.id2label = {int(k): v for k, v in self.model.config.id2label.items()} | |
| self.det_processor = AutoProcessor.from_pretrained(DETECTOR_ID) | |
| self.det_model = RTDetrForObjectDetection.from_pretrained(DETECTOR_ID) | |
| self.det_model.to(device=self.device, dtype=self.dtype) | |
| self.det_model.eval() | |
| self.person_label_ids = { | |
| int(i) | |
| for i, name in self.det_model.config.id2label.items() | |
| if str(name).lower() == "person" | |
| } or {0} | |
| def __call__(self, data: dict[str, Any]) -> dict[str, Any]: | |
| payload = dict(data or {}) | |
| parameters = payload.pop("parameters", None) or {} | |
| if not isinstance(parameters, dict): | |
| parameters = {} | |
| raw = payload.pop("inputs", payload) | |
| boxes = payload.pop("boxes", parameters.get("boxes")) | |
| boxes_format = str(payload.pop("boxes_format", parameters.get("boxes_format", "xywh"))).lower() | |
| threshold = float(payload.pop("threshold", parameters.get("threshold", 0.3))) | |
| detect_threshold = float( | |
| payload.pop("detect_threshold", parameters.get("detect_threshold", 0.3)) | |
| ) | |
| dataset_index = int(payload.pop("dataset_index", parameters.get("dataset_index", 0))) | |
| if isinstance(raw, dict): | |
| boxes = raw.get("boxes", boxes) | |
| boxes_format = str(raw.get("boxes_format", boxes_format)).lower() | |
| threshold = float(raw.get("threshold", threshold)) | |
| detect_threshold = float(raw.get("detect_threshold", detect_threshold)) | |
| dataset_index = int(raw.get("dataset_index", dataset_index)) | |
| raw = raw.get("image", raw.get("inputs", raw)) | |
| image = self._load_image(raw) | |
| person_boxes = self._resolve_boxes(image, boxes, boxes_format, detect_threshold) | |
| if person_boxes.shape[0] == 0: | |
| return { | |
| "people": [], | |
| "width": image.width, | |
| "height": image.height, | |
| "dataset_index": dataset_index, | |
| } | |
| inputs = self.processor(image, boxes=[person_boxes], return_tensors="pt") | |
| inputs = { | |
| k: v.to(self.device, dtype=self.dtype) if torch.is_floating_point(v) else v.to(self.device) | |
| for k, v in inputs.items() | |
| } | |
| if self.num_experts > 1: | |
| inputs["dataset_index"] = torch.tensor([dataset_index], device=self.device) | |
| with torch.inference_mode(): | |
| outputs = self.model(**inputs) | |
| pose_results = self.processor.post_process_pose_estimation( | |
| outputs, boxes=[person_boxes], threshold=threshold | |
| ) | |
| image_pose_result = pose_results[0] if pose_results else [] | |
| people: list[dict[str, Any]] = [] | |
| for i, person_pose in enumerate(image_pose_result): | |
| box = person_boxes[i].tolist() if i < len(person_boxes) else None | |
| keypoints = [] | |
| for keypoint, label, score in zip( | |
| person_pose["keypoints"], person_pose["labels"], person_pose["scores"] | |
| ): | |
| label_id = int(_as_float(label)) | |
| xy = _as_list(keypoint) | |
| keypoints.append( | |
| { | |
| "name": self.id2label.get(label_id, str(label_id)), | |
| "label": label_id, | |
| "x": float(xy[0]), | |
| "y": float(xy[1]), | |
| "score": _as_float(score), | |
| } | |
| ) | |
| people.append({"box": box, "keypoints": keypoints}) | |
| return { | |
| "people": people, | |
| "width": image.width, | |
| "height": image.height, | |
| "dataset_index": dataset_index, | |
| } | |
| def _load_image(self, image_input: Any) -> Image.Image: | |
| if isinstance(image_input, Image.Image): | |
| return image_input.convert("RGB") | |
| if isinstance(image_input, (bytes, bytearray, memoryview)): | |
| return Image.open(io.BytesIO(bytes(image_input))).convert("RGB") | |
| if isinstance(image_input, np.ndarray): | |
| if image_input.ndim == 3: | |
| return Image.fromarray(image_input.astype("uint8")).convert("RGB") | |
| raise ValueError("ndarray image must be HWC uint8") | |
| if isinstance(image_input, list) and image_input and isinstance(image_input[0], int): | |
| return Image.open(io.BytesIO(bytes(image_input))).convert("RGB") | |
| if not isinstance(image_input, str): | |
| raise ValueError("inputs must be a PIL image, base64 string, URL, or bytes") | |
| text = image_input.strip() | |
| if not text: | |
| raise ValueError("empty image input") | |
| parsed = urlparse(text) | |
| if parsed.scheme in ("http", "https"): | |
| response = requests.get(text, timeout=30) | |
| response.raise_for_status() | |
| return Image.open(io.BytesIO(response.content)).convert("RGB") | |
| if text.startswith("data:") and "," in text: | |
| text = text.split(",", 1)[1] | |
| try: | |
| raw = base64.b64decode(text, validate=False) | |
| except Exception as exc: | |
| raise ValueError("inputs string is not a valid image URL or base64 payload") from exc | |
| return Image.open(io.BytesIO(raw)).convert("RGB") | |
| def _resolve_boxes( | |
| self, | |
| image: Image.Image, | |
| boxes: Any, | |
| boxes_format: str, | |
| detect_threshold: float, | |
| ) -> np.ndarray: | |
| if boxes is not None: | |
| arr = np.asarray(boxes, dtype=np.float32) | |
| if arr.size == 0: | |
| return np.zeros((0, 4), dtype=np.float32) | |
| if arr.ndim == 1: | |
| arr = arr.reshape(1, 4) | |
| if arr.shape[-1] != 4: | |
| raise ValueError("boxes must be [x, y, w, h] or [x1, y1, x2, y2]") | |
| if boxes_format in ("xyxy", "voc"): | |
| arr = arr.copy() | |
| arr[:, 2] = arr[:, 2] - arr[:, 0] | |
| arr[:, 3] = arr[:, 3] - arr[:, 1] | |
| return arr | |
| det_inputs = self.det_processor(images=image, return_tensors="pt") | |
| det_inputs = { | |
| k: v.to(self.device, dtype=self.dtype) if torch.is_floating_point(v) else v.to(self.device) | |
| for k, v in det_inputs.items() | |
| } | |
| with torch.inference_mode(): | |
| det_outputs = self.det_model(**det_inputs) | |
| results = self.det_processor.post_process_object_detection( | |
| det_outputs, | |
| target_sizes=torch.tensor([(image.height, image.width)], device=self.device), | |
| threshold=detect_threshold, | |
| ) | |
| result = results[0] | |
| labels = result["labels"] | |
| mask = torch.zeros_like(labels, dtype=torch.bool) | |
| for person_id in self.person_label_ids: | |
| mask |= labels == person_id | |
| person_boxes = result["boxes"][mask].detach().cpu().numpy().astype(np.float32) | |
| if person_boxes.size == 0: | |
| return np.zeros((0, 4), dtype=np.float32) | |
| person_boxes[:, 2] = person_boxes[:, 2] - person_boxes[:, 0] | |
| person_boxes[:, 3] = person_boxes[:, 3] - person_boxes[:, 1] | |
| return person_boxes | |