jkubuni commited on
Commit
753f0ec
·
verified ·
1 Parent(s): 0c901ce

Add files using upload-large-folder tool

Browse files
Files changed (4) hide show
  1. README.md +191 -0
  2. config.json +24 -0
  3. model.safetensors +3 -0
  4. 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()