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 into DLC_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.Blocks with a named /predict endpoint, 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_size and coordinates, and writes null for keypoints below threshold; it previously wrote NaN (invalid JSON), crop-relative keypoints on the TF + MegaDetector path, and a number_of_bb counting all detections.
  • Dependencies: sdk_version and gradio 6.29.0, deeplabcut==3.0.2, CPU-only torch wheels, pip>=26.2 in pre-requirements.txt; pyproject.toml and a pre-commit config are added and the unused save_results.py is 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

Sign up or log in to comment