Image Feature Extraction
timm
Safetensors
re-identification
metric-learning
wildlife
dinov2
gorilla
Eval Results (legacy)
Instructions to use gorilla-watch/GorillaWatch-DINOv2-Giant with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use gorilla-watch/GorillaWatch-DINOv2-Giant with timm:
import timm model = timm.create_model("hf-hub:gorilla-watch/GorillaWatch-DINOv2-Giant", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Add files using upload-large-folder tool
Browse files- README.md +191 -0
- config.json +24 -0
- model.safetensors +3 -0
- modeling.py +255 -0
README.md
ADDED
|
@@ -0,0 +1,191 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: cc-by-4.0
|
| 3 |
+
library_name: timm
|
| 4 |
+
pipeline_tag: image-feature-extraction
|
| 5 |
+
tags:
|
| 6 |
+
- image-feature-extraction
|
| 7 |
+
- re-identification
|
| 8 |
+
- metric-learning
|
| 9 |
+
- wildlife
|
| 10 |
+
- dinov2
|
| 11 |
+
- gorilla
|
| 12 |
+
datasets:
|
| 13 |
+
- gorilla-watch/Gorilla-SPAC-Wild
|
| 14 |
+
model-index:
|
| 15 |
+
- name: GorillaWatch-DINOv2-Giant
|
| 16 |
+
results:
|
| 17 |
+
- task:
|
| 18 |
+
type: image-feature-extraction
|
| 19 |
+
name: facial gorilla re-identification
|
| 20 |
+
dataset:
|
| 21 |
+
type: gorilla-watch/Gorilla-SPAC-Wild
|
| 22 |
+
name: Gorilla-SPAC-Wild
|
| 23 |
+
config: face_with_body
|
| 24 |
+
split: test
|
| 25 |
+
metrics:
|
| 26 |
+
- name: Micro Accuracy
|
| 27 |
+
type: accuracy
|
| 28 |
+
value: 0.5554
|
| 29 |
+
- name: Macro Accuracy
|
| 30 |
+
type: accuracy
|
| 31 |
+
value: 0.4629
|
| 32 |
+
- name: Tracklet Micro Accuracy
|
| 33 |
+
type: accuracy
|
| 34 |
+
value: 0.6121
|
| 35 |
+
- name: Tracklet Macro Accuracy
|
| 36 |
+
type: accuracy
|
| 37 |
+
value: 0.4451
|
| 38 |
+
- task:
|
| 39 |
+
type: image-feature-extraction
|
| 40 |
+
name: facial gorilla re-identification
|
| 41 |
+
dataset:
|
| 42 |
+
type: gorilla-watch/Gorilla-Zoo-Berlin
|
| 43 |
+
name: Gorilla-Zoo-Berlin
|
| 44 |
+
config: face_with_body
|
| 45 |
+
split: test
|
| 46 |
+
metrics:
|
| 47 |
+
- name: Micro Accuracy
|
| 48 |
+
type: accuracy
|
| 49 |
+
value: 0.7657
|
| 50 |
+
- name: Macro Accuracy
|
| 51 |
+
type: accuracy
|
| 52 |
+
value: 0.759
|
| 53 |
+
- name: Tracklet Micro Accuracy
|
| 54 |
+
type: accuracy
|
| 55 |
+
value: 0.8218
|
| 56 |
+
- name: Tracklet Macro Accuracy
|
| 57 |
+
type: accuracy
|
| 58 |
+
value: 0.8044
|
| 59 |
+
---
|
| 60 |
+
|
| 61 |
+
# GorillaWatch-DINOv2-Giant
|
| 62 |
+
|
| 63 |
+
Gorilla re-identification model from **[GorillaWatch: An Automated System for In-the-Wild Gorilla Re-Identification and Population Monitoring](https://arxiv.org/abs/2512.07776)** (WACV 2026). Further project details can be found [here](https://gorilla-watch.github.io/).
|
| 64 |
+
|
| 65 |
+
A `vit_giant_patch14_dinov2.lvd142m` DINOv2 backbone fine-tuned with hard-mining triplet loss on
|
| 66 |
+
[Gorilla-SPAC-Wild](https://huggingface.co/datasets/gorilla-watch/Gorilla-SPAC-Wild), projecting to a
|
| 67 |
+
**256-dimensional embedding**. Identification is done by k-NN retrieval against a
|
| 68 |
+
gallery of embeddings, not by classification. The model has no fixed identity vocabulary, to enable
|
| 69 |
+
generalisation to individuals unseen during training.
|
| 70 |
+
|
| 71 |
+
| | |
|
| 72 |
+
|---|---|
|
| 73 |
+
| Backbone | `vit_giant_patch14_dinov2.lvd142m` |
|
| 74 |
+
| Input resolution | 518×518 |
|
| 75 |
+
| Embedding dimension | 256 |
|
| 76 |
+
| Parameters | 1136.9M |
|
| 77 |
+
| Training data | Gorilla-SPAC-Wild (`face_with_body`) |
|
| 78 |
+
|
| 79 |
+
## Preprocessing
|
| 80 |
+
|
| 81 |
+
> [!IMPORTANT]
|
| 82 |
+
> This model does **not** use timm's default DINOv2 transform. It expects a **square resize**
|
| 83 |
+
> (which only preserves the aspect ratio when the input images are already squared, which is the case in our datasets) and normalization with **mean = std = 0.5**, not the ImageNet
|
| 84 |
+
> statistics reported in the backbone's `default_cfg`. Using timm's default transform produces
|
| 85 |
+
> incorrect embeddings.
|
| 86 |
+
|
| 87 |
+
```python
|
| 88 |
+
transforms.Compose([
|
| 89 |
+
transforms.Resize((518, 518)),
|
| 90 |
+
transforms.ToTensor(),
|
| 91 |
+
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
|
| 92 |
+
])
|
| 93 |
+
```
|
| 94 |
+
|
| 95 |
+
`modeling.py` in this repository exposes this as `model.get_transform()`.
|
| 96 |
+
|
| 97 |
+
## Usage
|
| 98 |
+
|
| 99 |
+
Here we provide a minimal setup to use the model for feature extraction.
|
| 100 |
+
|
| 101 |
+
Requires `torch`, `timm`, `safetensors`, `huggingface_hub` and `torchvision`.
|
| 102 |
+
|
| 103 |
+
```python
|
| 104 |
+
import sys, torch
|
| 105 |
+
from huggingface_hub import snapshot_download
|
| 106 |
+
from PIL import Image
|
| 107 |
+
|
| 108 |
+
# Fetch weights, config and the self-contained modeling.py in one go
|
| 109 |
+
local_dir = snapshot_download("gorilla-watch/GorillaWatch-DINOv2-Giant")
|
| 110 |
+
sys.path.insert(0, local_dir)
|
| 111 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 112 |
+
from modeling import load_model
|
| 113 |
+
|
| 114 |
+
model = load_model(local_dir, device=device) # already in eval mode
|
| 115 |
+
transform = model.get_transform()
|
| 116 |
+
|
| 117 |
+
image = Image.open("gorilla.png").convert("RGB")
|
| 118 |
+
with torch.no_grad():
|
| 119 |
+
embedding = model(transform(image).unsqueeze(0).to(model.device)) # (1, 256)
|
| 120 |
+
```
|
| 121 |
+
|
| 122 |
+
`load_model` also accepts the repo id directly (`load_model("gorilla-watch/GorillaWatch-DINOv2-Giant")`) if you would rather not
|
| 123 |
+
manage a local directory.
|
| 124 |
+
|
| 125 |
+
Identity assignment uses **k-NN with k=5 under Euclidean distance** against a gallery of embeddings.
|
| 126 |
+
The paper's protocol masks out gallery entries from the same encounter (same camera on the same
|
| 127 |
+
date) to avoid trivially easy matches. The full evaluation code can be found in our [GitHub Repo](https://github.com/gorilla-watch/gorillawatch).
|
| 128 |
+
|
| 129 |
+
## Training
|
| 130 |
+
|
| 131 |
+
Fine-tuned from the upstream `vit_giant_patch14_dinov2.lvd142m` DINOv2 checkpoint.
|
| 132 |
+
|
| 133 |
+
| Hyperparameter | Value |
|
| 134 |
+
|---|---|
|
| 135 |
+
| Loss | Online triplet, hard mining, Euclidean, margin 0.647 |
|
| 136 |
+
| Optimizer | AdamW (β=0.9/0.999, ε=1e-7) |
|
| 137 |
+
| Learning rate | 1.9e-7, cosine annealing to 1e-7 |
|
| 138 |
+
| Batch size | 8 (effective 48 via 6 gradient accumulation steps) |
|
| 139 |
+
| Regularization | L2 = 0.0059, L2-SP = 1.3e-5 |
|
| 140 |
+
| Epochs | 100 max, best-validation-loss checkpoint retained |
|
| 141 |
+
| Precision | AMP (fp16 autocast, fp32 master weights) |
|
| 142 |
+
| Seed | 42 |
|
| 143 |
+
|
| 144 |
+
The code used to train these models can be found in our [Github Repository](https://github.com/gorilla-watch/gorillawatch).
|
| 145 |
+
|
| 146 |
+
## Results
|
| 147 |
+
|
| 148 |
+
k-NN retrieval accuracy (k=5, Euclidean distance). Gallery entries from the same encounter (same camera on the same date) are masked out, so every match is made across encounters. Macro accuracy averages over identities and is the harder number: it weights rarely-seen individuals equally with frequently-seen ones.
|
| 149 |
+
|
| 150 |
+
### In-domain: Gorilla-SPAC-Wild
|
| 151 |
+
|
| 152 |
+
Test split of [Gorilla-SPAC-Wild](https://huggingface.co/datasets/gorilla-watch/Gorilla-SPAC-Wild), the distribution the model was fine-tuned on.
|
| 153 |
+
|
| 154 |
+
| Protocol | Micro accuracy | Macro accuracy |
|
| 155 |
+
|---|---|---|
|
| 156 |
+
| Per image | 0.5554 | 0.4629 |
|
| 157 |
+
| Per tracklet (average pooling) | 0.6121 | 0.4451 |
|
| 158 |
+
|
| 159 |
+
### Out-of-distribution: Gorilla-Zoo-Berlin
|
| 160 |
+
|
| 161 |
+
[Gorilla-Zoo-Berlin](https://huggingface.co/datasets/gorilla-watch/Gorilla-Zoo-Berlin) is a **zero-shot domain-transfer test**: the model is applied to footage recorded in the Berlin Zoo, with no fine-tuning on it, so enclosure, lighting, camera hardware and the individuals themselves are all unseen. The numbers are still higher, since the amount of individuals is much lower than in the SPAC dataset. This evaluation clearly shows that the model is able to generalize to new, unseen populations.
|
| 162 |
+
|
| 163 |
+
| Protocol | Micro accuracy | Macro accuracy |
|
| 164 |
+
|---|---|---|
|
| 165 |
+
| Per image | 0.7657 | 0.7590 |
|
| 166 |
+
| Per tracklet (average pooling) | 0.8218 | 0.8044 |
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
## Provenance
|
| 170 |
+
|
| 171 |
+
These weights are bit-identical conversions from the `.pth` files created in the training process. They were converted to the `model.safetensors` format for better integration with HuggingFace.
|
| 172 |
+
|
| 173 |
+
## License
|
| 174 |
+
|
| 175 |
+
This model is released under the **CC-BY-4.0 License**.
|
| 176 |
+
|
| 177 |
+
## Citation
|
| 178 |
+
|
| 179 |
+
```bibtex
|
| 180 |
+
@inproceedings{GorillaWatch2026,
|
| 181 |
+
title={GorillaWatch: An Automated System for In-the-Wild Gorilla Re-Identification and Population Monitoring},
|
| 182 |
+
booktitle={Proceedings of the IEEE/CVF Winter Conference on Applications of Computer Vision (WACV)},
|
| 183 |
+
author={Maximilian Schall and Felix Leonard Kn\"ofel and Noah Elias K\"onig and Jan Jonas Kubeler and Maximilian von Klinski and Joan Wilhelm Linnemann and Xiaoshi Liu and Iven Jelle Schlegelmilch and Ole Woyciniuk and Alexandra Schild and Dante Wasmuht and Magdalena Bermejo Espinet and German Illera Basas and Gerard de Melo},
|
| 184 |
+
year={2026},
|
| 185 |
+
archivePrefix={arXiv},
|
| 186 |
+
eprint={2512.07776}
|
| 187 |
+
}
|
| 188 |
+
```
|
| 189 |
+
|
| 190 |
+
## Acknowledgements
|
| 191 |
+
The project on which this report is based was funded by the Federal Ministry of Research, Technology and Space under the funding code “KI-Servicezentrum Berlin-Brandenburg” 16IS22092. We acknowledge the support of Sabine Plattner African Charities (SPAC) for their funding to this research. We are grateful to Zoo Berlin for their expert assistance and facility access. This collaboration enabled the development of AI tools capable of being deployed in the wild to directly support gorilla conservation. The responsibility for the content of this publication remains with the authors.
|
config.json
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"backbone_name": "vit_giant_patch14_dinov2.lvd142m",
|
| 3 |
+
"img_size": 518,
|
| 4 |
+
"embedding_size": 256,
|
| 5 |
+
"embedding_id": "linear",
|
| 6 |
+
"pool_mode": "none",
|
| 7 |
+
"dropout_p": 0.0,
|
| 8 |
+
"image_mean": [
|
| 9 |
+
0.5,
|
| 10 |
+
0.5,
|
| 11 |
+
0.5
|
| 12 |
+
],
|
| 13 |
+
"image_std": [
|
| 14 |
+
0.5,
|
| 15 |
+
0.5,
|
| 16 |
+
0.5
|
| 17 |
+
],
|
| 18 |
+
"resize_mode": "square",
|
| 19 |
+
"interpolation": "bilinear",
|
| 20 |
+
"timm_version": "1.0.15",
|
| 21 |
+
"torch_version": "2.9.1+cu128",
|
| 22 |
+
"source_checkpoint": "vit_giant_patch14_dinov2.lvd142m_fine_tuned.pth",
|
| 23 |
+
"source_sha256": "dc6807a766a7a8d070afdc9d05f25bebc3256203d9eff6985f573dc529c98009"
|
| 24 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:17a06000a498afafeaff3247c337669ea438f7a446117d06e83629c307157527
|
| 3 |
+
size 4547548872
|
modeling.py
ADDED
|
@@ -0,0 +1,255 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Standalone model definition for the published GorillaWatch re-ID checkpoints.
|
| 2 |
+
|
| 3 |
+
This file is intentionally self-contained: it is copied verbatim into every
|
| 4 |
+
HuggingFace model repository so the weights are usable without installing
|
| 5 |
+
anything from the GorillaWatch source tree.
|
| 6 |
+
|
| 7 |
+
The forward pass mirrors ``TimmWrapper`` in
|
| 8 |
+
``gorillawatch/src/gorillawatch/model/basemodel.py`` exactly.
|
| 9 |
+
|
| 10 |
+
Usage::
|
| 11 |
+
|
| 12 |
+
from modeling import load_model
|
| 13 |
+
|
| 14 |
+
model = load_model("gorilla-watch/GorillaWatch-DINOv2-Large")
|
| 15 |
+
tf = model.get_transform()
|
| 16 |
+
embeddings = model(tf(image).unsqueeze(0)) # (1, 256)
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
import json
|
| 20 |
+
from dataclasses import asdict, dataclass, field, fields
|
| 21 |
+
from pathlib import Path
|
| 22 |
+
from typing import Any, Optional, Sequence, Union
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
|
| 27 |
+
CONFIG_NAME = "config.json"
|
| 28 |
+
WEIGHTS_NAME = "model.safetensors"
|
| 29 |
+
|
| 30 |
+
# Fields of GorillaWatchConfig that determine the module topology. Everything
|
| 31 |
+
# else in the config is preprocessing or provenance metadata.
|
| 32 |
+
ARCH_FIELDS = (
|
| 33 |
+
"backbone_name",
|
| 34 |
+
"img_size",
|
| 35 |
+
"embedding_size",
|
| 36 |
+
"embedding_id",
|
| 37 |
+
"pool_mode",
|
| 38 |
+
"dropout_p",
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
try: # optional; only needed for the idiomatic from_pretrained path
|
| 42 |
+
from huggingface_hub import PyTorchModelHubMixin
|
| 43 |
+
|
| 44 |
+
_HUB_MIXIN_AVAILABLE = True
|
| 45 |
+
except ImportError: # pragma: no cover - exercised only in minimal envs
|
| 46 |
+
_HUB_MIXIN_AVAILABLE = False
|
| 47 |
+
|
| 48 |
+
class PyTorchModelHubMixin: # type: ignore[no-redef]
|
| 49 |
+
"""No-op stand-in when huggingface_hub is not installed."""
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
@dataclass
|
| 53 |
+
class GorillaWatchConfig:
|
| 54 |
+
"""Everything needed to rebuild a published checkpoint and preprocess for it.
|
| 55 |
+
|
| 56 |
+
The training pipeline stores none of this in the ``.pth`` file, which is the
|
| 57 |
+
whole reason this dataclass exists. ``img_size`` in particular is not
|
| 58 |
+
recoverable from the checkpoint filename and silently changes the
|
| 59 |
+
positional-embedding shape.
|
| 60 |
+
"""
|
| 61 |
+
|
| 62 |
+
backbone_name: str
|
| 63 |
+
img_size: int = 518
|
| 64 |
+
embedding_size: int = 256
|
| 65 |
+
embedding_id: str = "linear"
|
| 66 |
+
pool_mode: str = "none"
|
| 67 |
+
dropout_p: float = 0.0
|
| 68 |
+
|
| 69 |
+
# Preprocessing. NOTE: mean/std are 0.5, *not* the ImageNet statistics that
|
| 70 |
+
# timm's DINOv2 default_cfg reports. Using timm's default transform against
|
| 71 |
+
# these weights produces wrong embeddings.
|
| 72 |
+
image_mean: Sequence[float] = field(default_factory=lambda: [0.5, 0.5, 0.5])
|
| 73 |
+
image_std: Sequence[float] = field(default_factory=lambda: [0.5, 0.5, 0.5])
|
| 74 |
+
resize_mode: str = "square" # Resize((S, S)); requires squared input to preserve aspect ratio
|
| 75 |
+
interpolation: str = "bilinear"
|
| 76 |
+
|
| 77 |
+
timm_version: Optional[str] = None
|
| 78 |
+
torch_version: Optional[str] = None
|
| 79 |
+
source_checkpoint: Optional[str] = None
|
| 80 |
+
source_sha256: Optional[str] = None
|
| 81 |
+
|
| 82 |
+
def to_dict(self) -> dict[str, Any]:
|
| 83 |
+
return asdict(self)
|
| 84 |
+
|
| 85 |
+
@classmethod
|
| 86 |
+
def from_dict(cls, data: dict[str, Any]) -> "GorillaWatchConfig":
|
| 87 |
+
known = {f.name for f in fields(cls)}
|
| 88 |
+
unknown = set(data) - known
|
| 89 |
+
if unknown:
|
| 90 |
+
raise ValueError(
|
| 91 |
+
f"Unrecognised keys in config: {sorted(unknown)}. "
|
| 92 |
+
"Refusing to load rather than silently ignoring architecture settings."
|
| 93 |
+
)
|
| 94 |
+
return cls(**data)
|
| 95 |
+
|
| 96 |
+
def save(self, path: Union[str, Path]) -> None:
|
| 97 |
+
Path(path).write_text(json.dumps(self.to_dict(), indent=2) + "\n")
|
| 98 |
+
|
| 99 |
+
@classmethod
|
| 100 |
+
def load(cls, path: Union[str, Path]) -> "GorillaWatchConfig":
|
| 101 |
+
return cls.from_dict(json.loads(Path(path).read_text()))
|
| 102 |
+
|
| 103 |
+
@property
|
| 104 |
+
def arch_kwargs(self) -> dict[str, Any]:
|
| 105 |
+
return {name: getattr(self, name) for name in ARCH_FIELDS}
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def build_transform(
|
| 109 |
+
img_size: int = 518,
|
| 110 |
+
image_mean: Sequence[float] = (0.5, 0.5, 0.5),
|
| 111 |
+
image_std: Sequence[float] = (0.5, 0.5, 0.5),
|
| 112 |
+
):
|
| 113 |
+
"""The exact eval transform used by the training pipeline.
|
| 114 |
+
|
| 115 |
+
Mirrors ``get_transform`` in
|
| 116 |
+
``gorillawatch/src/gorillawatch/data_hf/data_loading.py``.
|
| 117 |
+
"""
|
| 118 |
+
from torchvision import transforms
|
| 119 |
+
|
| 120 |
+
return transforms.Compose(
|
| 121 |
+
[
|
| 122 |
+
transforms.Resize((img_size, img_size)),
|
| 123 |
+
transforms.ToTensor(),
|
| 124 |
+
transforms.Normalize(mean=list(image_mean), std=list(image_std)),
|
| 125 |
+
]
|
| 126 |
+
)
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def _build_embedding_layer(embedding_id: str, feature_dim: int, embedding_dim: int) -> nn.Module:
|
| 130 |
+
if embedding_id == "linear":
|
| 131 |
+
return nn.Linear(feature_dim, embedding_dim)
|
| 132 |
+
if embedding_id == "identity":
|
| 133 |
+
return nn.Identity()
|
| 134 |
+
raise NotImplementedError(
|
| 135 |
+
f"embedding_id={embedding_id!r} is not supported by the published wrapper. "
|
| 136 |
+
"All released checkpoints use 'linear'."
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
class GorillaWatchViT(nn.Module, PyTorchModelHubMixin):
|
| 141 |
+
"""DINOv2 ViT backbone with a linear projection to a 256-d re-ID embedding.
|
| 142 |
+
|
| 143 |
+
Constructor arguments mirror ``config.json`` one-for-one so that
|
| 144 |
+
``from_pretrained`` round-trips without losing fields.
|
| 145 |
+
"""
|
| 146 |
+
|
| 147 |
+
def __init__(
|
| 148 |
+
self,
|
| 149 |
+
backbone_name: str,
|
| 150 |
+
img_size: int = 518,
|
| 151 |
+
embedding_size: int = 256,
|
| 152 |
+
embedding_id: str = "linear",
|
| 153 |
+
pool_mode: str = "none",
|
| 154 |
+
dropout_p: float = 0.0,
|
| 155 |
+
image_mean: Optional[Sequence[float]] = None,
|
| 156 |
+
image_std: Optional[Sequence[float]] = None,
|
| 157 |
+
resize_mode: str = "square",
|
| 158 |
+
interpolation: str = "bilinear",
|
| 159 |
+
pretrained: bool = False,
|
| 160 |
+
**_ignored: Any,
|
| 161 |
+
) -> None:
|
| 162 |
+
super().__init__()
|
| 163 |
+
|
| 164 |
+
if pool_mode != "none":
|
| 165 |
+
raise NotImplementedError(
|
| 166 |
+
f"pool_mode={pool_mode!r} is not supported by the published wrapper. "
|
| 167 |
+
"All released checkpoints are ViTs trained with pool_mode='none'."
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
import timm
|
| 171 |
+
|
| 172 |
+
# pretrained=False by default: the fine-tuned weights overwrite the
|
| 173 |
+
# upstream DINOv2 weights anyway, so downloading them first is pure
|
| 174 |
+
# waste (4.5 GB for the giant).
|
| 175 |
+
self.model = timm.create_model(
|
| 176 |
+
backbone_name,
|
| 177 |
+
pretrained=pretrained,
|
| 178 |
+
drop_rate=0.0,
|
| 179 |
+
img_size=img_size,
|
| 180 |
+
)
|
| 181 |
+
self.num_features = self.model.num_features
|
| 182 |
+
self.embedding_layer = _build_embedding_layer(
|
| 183 |
+
embedding_id, self.num_features, embedding_size
|
| 184 |
+
)
|
| 185 |
+
|
| 186 |
+
self.backbone_name = backbone_name
|
| 187 |
+
self.img_size = img_size
|
| 188 |
+
self.embedding_size = embedding_size
|
| 189 |
+
self.embedding_id = embedding_id
|
| 190 |
+
self.pool_mode = pool_mode
|
| 191 |
+
self.dropout_p = dropout_p
|
| 192 |
+
self.image_mean = list(image_mean) if image_mean is not None else [0.5, 0.5, 0.5]
|
| 193 |
+
self.image_std = list(image_std) if image_std is not None else [0.5, 0.5, 0.5]
|
| 194 |
+
self.resize_mode = resize_mode
|
| 195 |
+
self.interpolation = interpolation
|
| 196 |
+
|
| 197 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 198 |
+
x = self.model.forward_features(x)
|
| 199 |
+
x = self.model.forward_head(x, pre_logits=True)
|
| 200 |
+
if x.dim() == 3: # VisionTransformer: take the CLS token
|
| 201 |
+
x = x[:, 0, :]
|
| 202 |
+
return self.embedding_layer(x)
|
| 203 |
+
|
| 204 |
+
@property
|
| 205 |
+
def device(self) -> torch.device:
|
| 206 |
+
return next(self.parameters()).device
|
| 207 |
+
|
| 208 |
+
def get_transform(self):
|
| 209 |
+
return build_transform(self.img_size, self.image_mean, self.image_std)
|
| 210 |
+
|
| 211 |
+
@classmethod
|
| 212 |
+
def from_config(cls, config: GorillaWatchConfig, pretrained: bool = False) -> "GorillaWatchViT":
|
| 213 |
+
return cls(pretrained=pretrained, **config.to_dict())
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def load_model(
|
| 217 |
+
model_id_or_path: Union[str, Path],
|
| 218 |
+
device: Union[str, torch.device] = "cpu",
|
| 219 |
+
revision: Optional[str] = None,
|
| 220 |
+
token: Optional[str] = None,
|
| 221 |
+
cache_dir: Optional[str] = None,
|
| 222 |
+
) -> GorillaWatchViT:
|
| 223 |
+
"""Load a published checkpoint from a local directory or the HuggingFace Hub.
|
| 224 |
+
|
| 225 |
+
This is the guaranteed-stable loader; it does not depend on
|
| 226 |
+
``PyTorchModelHubMixin`` behaviour, which varies across huggingface_hub
|
| 227 |
+
versions.
|
| 228 |
+
"""
|
| 229 |
+
from safetensors.torch import load_file
|
| 230 |
+
|
| 231 |
+
local = Path(model_id_or_path)
|
| 232 |
+
if local.is_dir():
|
| 233 |
+
config_path = local / CONFIG_NAME
|
| 234 |
+
weights_path = local / WEIGHTS_NAME
|
| 235 |
+
missing = [p.name for p in (config_path, weights_path) if not p.exists()]
|
| 236 |
+
if missing:
|
| 237 |
+
raise FileNotFoundError(f"{local} is missing {missing}")
|
| 238 |
+
else:
|
| 239 |
+
from huggingface_hub import hf_hub_download
|
| 240 |
+
|
| 241 |
+
download_kwargs = dict(
|
| 242 |
+
repo_id=str(model_id_or_path),
|
| 243 |
+
repo_type="model",
|
| 244 |
+
revision=revision,
|
| 245 |
+
token=token,
|
| 246 |
+
cache_dir=cache_dir,
|
| 247 |
+
)
|
| 248 |
+
config_path = Path(hf_hub_download(filename=CONFIG_NAME, **download_kwargs))
|
| 249 |
+
weights_path = Path(hf_hub_download(filename=WEIGHTS_NAME, **download_kwargs))
|
| 250 |
+
|
| 251 |
+
config = GorillaWatchConfig.load(config_path)
|
| 252 |
+
model = GorillaWatchViT.from_config(config)
|
| 253 |
+
state_dict = load_file(str(weights_path))
|
| 254 |
+
model.load_state_dict(state_dict, strict=True)
|
| 255 |
+
return model.to(device).eval()
|