Add PyTorch SuperAnimal backend and refresh the Space
#14
by C-Achard - opened
Scope
New PyTorch-based SuperAnimal inference (new default), fixes to the TensorFlow pipeline (kept as "TensorFlow (legacy)"), a new UI, a revised JSON output, and dependency changes for the CPU Space.
Motivation
The TensorFlow pipeline (MegaDetector + TF ResNet-50 through DLCLive) is an older pipeline: the pose model runs on crops much smaller than its training crops, DLCLive 1.1.0 swaps RGB/BGR (DeepLabCut-live#175), and grayscale input performs poorly. The SuperAnimal PyTorch models (Faster R-CNN detector, HRNet-w32 top-down) give more confident and more stable keypoints on RGB and grayscale in the images tried.
Main changes
- The PyTorch backend (
pytorch_utils.py) is the default: SuperAnimal Faster R-CNN + HRNet-w32, inputs resized to 1280 px on the longest side, weights downloaded on first use intoDLC_models/pytorch. - The TF legacy path passes BGR crops to work around DeepLabCut-live#175, aligns boxes with crops, and handles images without detections; the MegaDetector choice is shown only for this backend.
- The app uses
gr.Blockswith a named/predictendpoint, lazily cached examples, and the default model preloaded in the background at startup; the COCO image fetched at startup was removed. - Keypoints are coloured by confidence (legend below the image) or by bodypart, with colormap and bbox-colour pickers, a per-keypoint confidence table, an annotated PNG download and a DeepLabCut colour theme; labels are cleaned for display (e.g.
thaiβthigh) while the JSON keeps the model's names. - The JSON uses input-image pixels for all backends, adds
image_size,annotated_image_sizeandcoordinates, and writesnullfor keypoints below threshold; it previously wroteNaN(invalid JSON), crop-relative keypoints on the TF + MegaDetector path, and anumber_of_bbcounting all detections. - Dependencies:
sdk_versionandgradio6.29.0,deeplabcut==3.0.2, CPU-only torch wheels,pip>=26.2inpre-requirements.txt;pyproject.tomland a pre-commit config are added and the unusedsave_results.pyis removed.
Additional context
- On CPU, PyTorch is about 4Γ slower per image with models loaded, but faster per request than the TF path, which reloads its models on every request. Requests run one at a time because the PyTorch runners are not thread-safe.
- The first use of each species downloads about 285 MB; the Space disk is not persistent, so this repeats after each restart.
- Validated manually on Windows with a GPU (both backends, detector and full-image modes, strict JSON parsing) and in a Linux CPU container (install and startup); not run on Spaces hardware.
- Rollback if needed: revert the merge commit on
main.
C-Achard changed pull request status to open
mwmathis changed pull request status to merged