multimodalart HF Staff commited on
Commit
b06cf47
·
verified ·
1 Parent(s): 72b879b

Upload folder using huggingface_hub

Browse files
Files changed (6) hide show
  1. .gitattributes +2 -0
  2. README.md +32 -7
  3. app.py +552 -0
  4. city_with_cars.png +3 -0
  5. fruit_store.png +3 -0
  6. requirements.txt +9 -0
.gitattributes CHANGED
@@ -33,3 +33,5 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ city_with_cars.png filter=lfs diff=lfs merge=lfs -text
37
+ fruit_store.png filter=lfs diff=lfs merge=lfs -text
README.md CHANGED
@@ -1,13 +1,38 @@
1
  ---
2
- title: Vit Up
3
- emoji: 🌍
4
- colorFrom: pink
5
  colorTo: gray
6
  sdk: gradio
7
- sdk_version: 6.19.0
8
- python_version: '3.12'
9
  app_file: app.py
10
- pinned: false
 
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: ViT-Up Feature Upsampler
3
+ emoji: 🔼
4
+ colorFrom: blue
5
  colorTo: gray
6
  sdk: gradio
7
+ sdk_version: "5.50.0"
 
8
  app_file: app.py
9
+ short_description: DINOv3 feature upsampling with ViT-Up
10
+ python_version: "3.12"
11
+ startup_duration_timeout: 900
12
  ---
13
 
14
+ # ViT-Up: Faithful Feature Upsampling for Vision Transformers
15
+
16
+ This Space demonstrates **ViT-Up**, an implicit feature upsampler for Vision
17
+ Transformers that predicts backbone-aligned features at arbitrary continuous
18
+ image coordinates.
19
+
20
+ ## How it works
21
+
22
+ 1. **Input**: An image is padded to square, resized to 448×448, and normalised
23
+ with ImageNet statistics.
24
+ 2. **Backbone**: A DINOv3-S+ ViT backbone (loaded from the non-gated
25
+ `timm/vit_small_plus_patch16_dinov3.lvd1689m` mirror) extracts multi-layer
26
+ hidden states. LoRA adapters from the ViT-Up checkpoint are applied.
27
+ 3. **Upsampling**: ViT-Up queries features at a dense grid of user-selected
28
+ resolution (e.g. 112×112), producing high-resolution feature maps aligned
29
+ with the backbone.
30
+ 4. **Visualization**: The 3 principal components of the upsampled features are
31
+ projected to RGB via PCA, showing the semantic structure learned by ViT-Up.
32
+
33
+ ## Model
34
+
35
+ - **Paper**: [ViT-Up: Faithful Feature Upsampling for Vision Transformers](https://huggingface.co/papers/2606.14024)
36
+ - **Weights**: [Krispin/vit-up](https://huggingface.co/Krispin/vit-up)
37
+ - **Code**: [GitHub](https://github.com/krispinwandel/vit-up)
38
+ - **License**: CC-BY-NC-SA-4.0
app.py ADDED
@@ -0,0 +1,552 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ViT-Up: Faithful Feature Upsampling for Vision Transformers.
2
+
3
+ Interactive demo that loads the ViT-Up feature upsampler, extracts dense
4
+ features from an input image at a user-selected output resolution, and
5
+ visualises them via a 3-component PCA projection to RGB.
6
+
7
+ The DINOv3 backbone checkpoint on Hugging Face is gated, so this demo
8
+ loads the equivalent pretrained weights from the non-gated timm mirror
9
+ (`timm/vit_small_plus_patch16_dinov3.lvd1689m`) and maps them into the
10
+ same ``DINOv3ViT`` module structure the ViT-Up code expects. The ViT-Up
11
+ LoRA adapters and upsampler head are then loaded from ``Krispin/vit-up``.
12
+ """
13
+
14
+ import os
15
+
16
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
17
+
18
+ import spaces # MUST come before torch / any CUDA-touching import
19
+ import sys
20
+ import math
21
+ from pathlib import Path
22
+ from typing import Any, Dict, List, Optional
23
+
24
+ import torch
25
+ import torch.nn as nn
26
+ import torch.nn.functional as F
27
+ import numpy as np
28
+ from PIL import Image, ImageOps
29
+ import gradio as gr
30
+ from huggingface_hub import hf_hub_download
31
+ from safetensors.torch import load_file as load_safetensors
32
+ import torchvision.transforms.v2 as T
33
+
34
+ # ---------------------------------------------------------------------------
35
+ # Config constants — DINOv3-S+ variant
36
+ # ---------------------------------------------------------------------------
37
+ BACKBONE_TIMM_REPO = "timm/vit_small_plus_patch16_dinov3.lvd1689m"
38
+ VITUP_WEIGHTS_REPO = "Krispin/vit-up"
39
+ VITUP_WEIGHTS_FILE = "vit_up_dinov3_splus.safetensors"
40
+ HIDDEN_SIZE = 384
41
+ NUM_LAYERS = 12
42
+ NUM_HEADS = 6
43
+ INTERMEDIATE_SIZE = 1536
44
+ PATCH_SIZE = 16
45
+ NUM_REGISTER_TOKENS = 4
46
+ IMAGE_SIZE = 448
47
+ LAYER_INDICES = [0, 2, 4, 6, 8, 10, 12]
48
+ RESNET_MEAN = torch.tensor([0.485, 0.456, 0.406])
49
+ RESNET_STD = torch.tensor([0.229, 0.224, 0.225])
50
+
51
+ # ---------------------------------------------------------------------------
52
+ # Shallow-copy the vit_up package from the cloned repo so we can import it
53
+ # ---------------------------------------------------------------------------
54
+ _REPO_ROOT = Path("/tmp/hugging-demos-build-paper_2606.14024-g169ewzl/vit-up")
55
+ if str(_REPO_ROOT) not in sys.path:
56
+ sys.path.insert(0, str(_REPO_ROOT))
57
+
58
+ from transformers import DINOv3ViTConfig
59
+ from vit_up.layers.backbones.dinov3_vit import DINOv3ViT
60
+ from vit_up.model.vit_up import ViTUp
61
+ from vit_up.utils.state_dict_migration import migrate_vit_up_state_dict_keys
62
+ from peft import LoraConfig, get_peft_model
63
+
64
+
65
+ # ---------------------------------------------------------------------------
66
+ # Weight mapping: timm -> DINOv3ViT
67
+ # ---------------------------------------------------------------------------
68
+ def _map_timm_to_dinov3(timm_sd: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
69
+ """Convert a timm ViT state-dict to the DINOv3ViT module key names."""
70
+
71
+ mapped: Dict[str, torch.Tensor] = {}
72
+ for key, val in timm_sd.items():
73
+ if key == "cls_token":
74
+ mapped["embeddings.cls_token"] = val
75
+ elif key == "reg_token":
76
+ mapped["embeddings.register_tokens"] = val
77
+ elif key == "patch_embed.proj.weight":
78
+ mapped["embeddings.patch_embeddings.weight"] = val
79
+ elif key == "patch_embed.proj.bias":
80
+ mapped["embeddings.patch_embeddings.bias"] = val
81
+ elif key.startswith("blocks.") and key.endswith(".attn.qkv.weight"):
82
+ idx = int(key.split(".")[1])
83
+ qkv = val # (3*hidden, hidden) but timm uses fused qkv
84
+ q, k, v = qkv.chunk(3, dim=0)
85
+ mapped[f"layer.{idx}.attention.q_proj.weight"] = q
86
+ mapped[f"layer.{idx}.attention.k_proj.weight"] = k
87
+ mapped[f"layer.{idx}.attention.v_proj.weight"] = v
88
+ elif key.startswith("blocks.") and key.endswith(".attn.qkv.bias"):
89
+ idx = int(key.split(".")[1])
90
+ qkv = val
91
+ if val is not None and val.numel() > 0:
92
+ q, k, v = qkv.chunk(3, dim=0)
93
+ mapped[f"layer.{idx}.attention.q_proj.bias"] = q
94
+ mapped[f"layer.{idx}.attention.k_proj.bias"] = k
95
+ mapped[f"layer.{idx}.attention.v_proj.bias"] = v
96
+ elif key.startswith("blocks.") and ".attn.proj." in key:
97
+ idx = int(key.split(".")[1])
98
+ suffix = key.split(".attn.proj.")[-1] # weight or bias
99
+ mapped[f"layer.{idx}.attention.o_proj.{suffix}"] = val
100
+ elif key.startswith("blocks.") and ".norm1." in key:
101
+ idx = int(key.split(".")[1])
102
+ suffix = key.split(".norm1.")[-1]
103
+ mapped[f"layer.{idx}.norm1.{suffix}"] = val
104
+ elif key.startswith("blocks.") and ".norm2." in key:
105
+ idx = int(key.split(".")[1])
106
+ suffix = key.split(".norm2.")[-1]
107
+ mapped[f"layer.{idx}.norm2.{suffix}"] = val
108
+ elif key.startswith("blocks.") and ".mlp.fc1_g." in key:
109
+ idx = int(key.split(".")[1])
110
+ suffix = key.split(".mlp.fc1_g.")[-1]
111
+ mapped[f"layer.{idx}.mlp.gate_proj.{suffix}"] = val
112
+ elif key.startswith("blocks.") and ".mlp.fc1_x." in key:
113
+ idx = int(key.split(".")[1])
114
+ suffix = key.split(".mlp.fc1_x.")[-1]
115
+ mapped[f"layer.{idx}.mlp.up_proj.{suffix}"] = val
116
+ elif key.startswith("blocks.") and ".mlp.fc2." in key:
117
+ idx = int(key.split(".")[1])
118
+ suffix = key.split(".mlp.fc2.")[-1]
119
+ mapped[f"layer.{idx}.mlp.down_proj.{suffix}"] = val
120
+ elif key.startswith("blocks.") and key.endswith(".gamma_1"):
121
+ idx = int(key.split(".")[1])
122
+ mapped[f"layer.{idx}.layer_scale1.lambda1"] = val
123
+ elif key.startswith("blocks.") and key.endswith(".gamma_2"):
124
+ idx = int(key.split(".")[1])
125
+ mapped[f"layer.{idx}.layer_scale2.lambda1"] = val
126
+ elif key == "norm.weight":
127
+ mapped["norm.weight"] = val
128
+ elif key == "norm.bias":
129
+ mapped["norm.bias"] = val
130
+ # pos_embed is handled by RoPE — skip
131
+ return mapped
132
+
133
+
134
+ # ---------------------------------------------------------------------------
135
+ # Build the backbone from config + timm weights + LoRA
136
+ # ---------------------------------------------------------------------------
137
+ def _build_backbone(device: str, dtype: torch.dtype) -> DINOv3ViT:
138
+ config = DINOv3ViTConfig(
139
+ hidden_size=HIDDEN_SIZE,
140
+ num_hidden_layers=NUM_LAYERS,
141
+ num_attention_heads=NUM_HEADS,
142
+ intermediate_size=INTERMEDIATE_SIZE,
143
+ patch_size=PATCH_SIZE,
144
+ image_size=IMAGE_SIZE,
145
+ num_register_tokens=NUM_REGISTER_TOKENS,
146
+ use_gated_mlp=True,
147
+ layerscale_value=1e-5,
148
+ query_bias=True,
149
+ key_bias=False,
150
+ value_bias=True,
151
+ proj_bias=True,
152
+ mlp_bias=True,
153
+ )
154
+ backbone = DINOv3ViT(config)
155
+
156
+ # Load timm weights
157
+ timm_safetensors_path = hf_hub_download(
158
+ BACKBONE_TIMM_REPO, "model.safetensors"
159
+ )
160
+ timm_sd = load_safetensors(timm_safetensors_path, device="cpu")
161
+ mapped_sd = _map_timm_to_dinov3(timm_sd)
162
+ missing, unexpected = backbone.load_state_dict(mapped_sd, strict=False)
163
+ # embeddings.mask_token won't be in timm weights — that's fine
164
+ real_missing = [k for k in missing if "mask_token" not in k]
165
+ if real_missing:
166
+ print(f"[WARNING] Missing backbone keys after timm load: {real_missing[:10]}")
167
+ print(f"[INFO] Loaded backbone from timm: {len(mapped_sd)} tensors mapped")
168
+
169
+ # Apply LoRA
170
+ lora_config = LoraConfig(
171
+ r=16,
172
+ lora_alpha=32,
173
+ lora_dropout=0.05,
174
+ bias="none",
175
+ target_modules=[
176
+ "patch_embeddings",
177
+ "q_proj",
178
+ "k_proj",
179
+ "v_proj",
180
+ "o_proj",
181
+ ],
182
+ )
183
+ backbone = get_peft_model(backbone, lora_config)
184
+ backbone = backbone.to(device=device, dtype=dtype).eval()
185
+ return backbone
186
+
187
+
188
+ # ---------------------------------------------------------------------------
189
+ # Build the ViT-Up model from config
190
+ # ---------------------------------------------------------------------------
191
+ def _build_vit_up(device: str, dtype: torch.dtype) -> ViTUp:
192
+ """Instantiate the ViTUp upsampler from the config tree (same as the repo)."""
193
+ from vit_up.layers.query_encoder import QueryEncoder
194
+ from vit_up.layers.pos_enc import FourierPositionalEncoding
195
+ from vit_up.layers.continuous_rope import ContinuousRoPE2D
196
+ from vit_up.layers.smart_module_list import SmartModuleList
197
+ from vit_up.layers.mlp import SimpleMLP
198
+ from vit_up.layers.cross_attention import CrossAttention
199
+ from vit_up.layers.film import SimpleFiLMV2
200
+
201
+ dim = HIDDEN_SIZE
202
+
203
+ query_embedding = QueryEncoder(
204
+ layer_index=0,
205
+ img_in_size=3584,
206
+ window_size=0,
207
+ out_proj_module=None,
208
+ )
209
+
210
+ rel_pos_enc = FourierPositionalEncoding(num_bands=16, max_resolution=10.0)
211
+
212
+ q_rope_embeddings = ContinuousRoPE2D(dim=64, base=100.0, scale=2 * math.pi)
213
+
214
+ vit_up_blocks = SmartModuleList(
215
+ n_blocks=6,
216
+ block_class_path="vit_up.model.vit_up.ViTUpBlock",
217
+ block_init_args={
218
+ "dim": dim,
219
+ "dim_h": dim,
220
+ "transition_mlp": SimpleMLP(
221
+ dims=[dim, dim * 2, dim],
222
+ activation="gelu",
223
+ input_layernorm=True,
224
+ use_residual=True,
225
+ ),
226
+ "cross_attention": CrossAttention(
227
+ dim=dim,
228
+ num_heads=NUM_HEADS,
229
+ cross_attn_window_size=32,
230
+ qkv_bias=True,
231
+ attn_dropout=0.0,
232
+ proj_dropout=0.0,
233
+ ),
234
+ "featx": SimpleFiLMV2(
235
+ input_module=nn.LayerNorm(dim),
236
+ gamma_beta_mlp=SimpleMLP(
237
+ dims=[66, dim, dim * 2],
238
+ activation="gelu",
239
+ zero_init_last=True,
240
+ ),
241
+ post_mlp=SimpleMLP(
242
+ dims=[dim, dim * 4, dim],
243
+ activation="gelu",
244
+ input_layernorm=True,
245
+ use_residual=False,
246
+ ),
247
+ ),
248
+ "mlp": SimpleMLP(
249
+ dims=[dim, dim * 4, dim],
250
+ activation="gelu",
251
+ ),
252
+ },
253
+ )
254
+
255
+ decoder_mlp = SmartModuleList(
256
+ n_blocks=7,
257
+ block_class_path="vit_up.layers.mlp.SimpleMLP",
258
+ block_init_args={
259
+ "input_layernorm": True,
260
+ "dims": [dim, dim],
261
+ },
262
+ )
263
+
264
+ vit_up = ViTUp(
265
+ layer_indices=LAYER_INDICES,
266
+ query_embedding=query_embedding,
267
+ rel_pos_enc=rel_pos_enc,
268
+ vit_up_blocks=vit_up_blocks,
269
+ decoder_mlp=decoder_mlp,
270
+ q_rope_embeddings=q_rope_embeddings,
271
+ )
272
+ vit_up = vit_up.to(device=device, dtype=dtype).eval()
273
+ return vit_up
274
+
275
+
276
+ # ---------------------------------------------------------------------------
277
+ # Load ViT-Up + LoRA weights from the safetensors checkpoint
278
+ # ---------------------------------------------------------------------------
279
+ def _load_vit_up_weights(
280
+ backbone: nn.Module,
281
+ vit_up: ViTUp,
282
+ device: str,
283
+ ) -> None:
284
+ """Load the combined LoRA + ViT-Up weights from Krispin/vit-up."""
285
+ weights_path = hf_hub_download(VITUP_WEIGHTS_REPO, VITUP_WEIGHTS_FILE)
286
+ state_dict = load_safetensors(weights_path, device="cpu")
287
+
288
+ backbone_sd: Dict[str, torch.Tensor] = {}
289
+ vit_up_sd: Dict[str, torch.Tensor] = {}
290
+ for key, val in state_dict.items():
291
+ if key.startswith("backbone."):
292
+ backbone_sd[key.removeprefix("backbone.")] = val
293
+ else:
294
+ vit_up_sd[key] = val
295
+
296
+ # Load backbone LoRA weights
297
+ missing_b, unexpected_b = backbone.load_state_dict(backbone_sd, strict=False)
298
+ print(f"[INFO] Loaded backbone LoRA: {len(backbone_sd)} tensors, "
299
+ f"missing={len(missing_b)}, unexpected={len(unexpected_b)}")
300
+
301
+ # Load ViT-Up weights (with key migration)
302
+ migrated_vit_up_sd = migrate_vit_up_state_dict_keys(vit_up_sd)
303
+ missing_v, unexpected_v = vit_up.load_state_dict(migrated_vit_up_sd, strict=False)
304
+ print(f"[INFO] Loaded ViT-Up: {len(migrated_vit_up_sd)} tensors, "
305
+ f"missing={len(missing_v)}, unexpected={len(unexpected_v)}")
306
+ if missing_v:
307
+ print(f" Missing ViT-Up keys: {missing_v[:10]}")
308
+
309
+
310
+ # ---------------------------------------------------------------------------
311
+ # PCA utilities (from the repo's correspondence.py)
312
+ # ---------------------------------------------------------------------------
313
+ def _fit_pca(tokens_nc: torch.Tensor, k: int = 3) -> dict:
314
+ """Fit a simple PCA on (N, C) feature tokens."""
315
+ tokens = tokens_nc.float()
316
+ mean = tokens.mean(dim=0)
317
+ centered = tokens - mean
318
+ _, singular_values, vh = torch.linalg.svd(centered, full_matrices=False)
319
+ components = vh[:k].T
320
+ projected = centered @ components
321
+ color_min = projected.amin(dim=0)
322
+ color_max = projected.amax(dim=0)
323
+ flat = torch.isclose(color_max, color_min)
324
+ color_max = torch.where(flat, color_min + 1.0, color_max)
325
+ return {
326
+ "pca_eig": components,
327
+ "pca_mean": mean,
328
+ "pca_color_min": color_min,
329
+ "pca_color_max": color_max,
330
+ }
331
+
332
+
333
+ def _apply_pca_rgb(feats_hwc: torch.Tensor, pca_data: dict) -> torch.Tensor:
334
+ h, w, c = feats_hwc.shape
335
+ tokens = feats_hwc.float().reshape(-1, c)
336
+ mean = pca_data["pca_mean"].to(device=tokens.device, dtype=tokens.dtype)
337
+ components = pca_data["pca_eig"].to(device=tokens.device, dtype=tokens.dtype)
338
+ color_min = pca_data["pca_color_min"].to(device=tokens.device, dtype=tokens.dtype)
339
+ color_max = pca_data["pca_color_max"].to(device=tokens.device, dtype=tokens.dtype)
340
+ projected = (tokens - mean) @ components
341
+ rgb = (projected - color_min.view(1, -1)) / (color_max - color_min).view(1, -1).add(1e-8)
342
+ rgb = rgb.clamp(0.0, 1.0).mul(255.0).to(torch.uint8)
343
+ return rgb.reshape(h, w, 3)
344
+
345
+
346
+ # ---------------------------------------------------------------------------
347
+ # Image utilities
348
+ # ---------------------------------------------------------------------------
349
+ def pad_image_to_square(img: Image.Image) -> Image.Image:
350
+ w, h = img.size
351
+ if w == h:
352
+ return img
353
+ max_side = max(w, h)
354
+ if w > h:
355
+ py = (w - h) // 2
356
+ return ImageOps.expand(img, border=(0, py), fill=0)
357
+ else:
358
+ px = (h - w) // 2
359
+ return ImageOps.expand(img, border=(px, 0), fill=0)
360
+
361
+
362
+ def crop_feature_square_to_image_aspect(
363
+ feat_img: Image.Image,
364
+ original_size: tuple,
365
+ ) -> Image.Image:
366
+ width, height = original_size
367
+ max_size = max(width, height)
368
+ px, py = (0, 0)
369
+ if width > height:
370
+ py = (width - height) // 2
371
+ elif height > width:
372
+ px = (height - width) // 2
373
+ scale_x = feat_img.width / max_size
374
+ scale_y = feat_img.height / max_size
375
+ left = int(round(px * scale_x))
376
+ top = int(round(py * scale_y))
377
+ right = int(round((px + width) * scale_x))
378
+ bottom = int(round((py + height) * scale_y))
379
+ return feat_img.crop((left, top, right, bottom))
380
+
381
+
382
+ # ---------------------------------------------------------------------------
383
+ # Build the full model at module scope
384
+ # ---------------------------------------------------------------------------
385
+ print("[INFO] Building ViT-Up model...")
386
+ DEVICE = "cuda"
387
+ DTYPE = torch.bfloat16
388
+
389
+ backbone = _build_backbone(DEVICE, DTYPE)
390
+ vit_up = _build_vit_up(DEVICE, DTYPE)
391
+ _load_vit_up_weights(backbone, vit_up, DEVICE)
392
+ backbone = backbone.eval()
393
+ vit_up = vit_up.eval()
394
+ print("[INFO] Model ready.")
395
+
396
+
397
+ # ---------------------------------------------------------------------------
398
+ # Inference
399
+ # ---------------------------------------------------------------------------
400
+ def _prepare_image(img: Image.Image) -> torch.Tensor:
401
+ """Pad to square, resize, normalise — return (1, 3, H, W) on device."""
402
+ img_square = pad_image_to_square(img.convert("RGB"))
403
+ transform = T.Compose([
404
+ T.ToImage(),
405
+ T.Resize((IMAGE_SIZE, IMAGE_SIZE), interpolation=T.InterpolationMode.BILINEAR, antialias=True),
406
+ T.ToDtype(torch.float32, scale=True),
407
+ T.Normalize(mean=RESNET_MEAN, std=RESNET_STD),
408
+ ])
409
+ return transform(img_square).unsqueeze(0).to(DEVICE)
410
+
411
+
412
+ def _compute_query_coords(out_size: int) -> torch.Tensor:
413
+ coords = torch.linspace(0.5, out_size - 0.5, out_size) / out_size
414
+ grid_y, grid_x = torch.meshgrid(coords, coords, indexing="ij")
415
+ return torch.stack((grid_x, grid_y), dim=-1).reshape(1, -1, 2)
416
+
417
+
418
+ @spaces.GPU(duration=120)
419
+ def extract_and_visualize(
420
+ input_image: Image.Image,
421
+ output_resolution: int,
422
+ ) -> tuple[Image.Image, Image.Image, str]:
423
+ """Extract dense ViT-Up features and visualise them via PCA.
424
+
425
+ Args:
426
+ input_image: Input PIL image.
427
+ output_resolution: Output feature map resolution (pixels per side).
428
+
429
+ Returns:
430
+ Tuple of (pca_visualization, input_resized, info_text).
431
+ """
432
+ if input_image is None:
433
+ return None, None, "Please provide an input image."
434
+
435
+ out_size = int(output_resolution)
436
+ orig_w, orig_h = input_image.size
437
+
438
+ # Prepare input
439
+ pixel_values = _prepare_image(input_image)
440
+
441
+ # Compute cache data (backbone hidden states)
442
+ with torch.no_grad(), torch.autocast(device_type="cuda", dtype=DTYPE):
443
+ cache_data = vit_up.compute_cache_data(
444
+ pixel_values=pixel_values,
445
+ backbone=backbone,
446
+ hidden_layer_img_size=IMAGE_SIZE,
447
+ )
448
+
449
+ # Query coords for dense output
450
+ query_coords = _compute_query_coords(out_size).to(DEVICE, dtype=DTYPE)
451
+
452
+ # Extract features
453
+ chunk_size = 4096
454
+ q_chunks = []
455
+ for q_start in range(0, query_coords.shape[1], chunk_size):
456
+ q_end = min(q_start + chunk_size, query_coords.shape[1])
457
+ q_chunk = vit_up(
458
+ pixel_values=None,
459
+ q_xy_normalized=query_coords[:, q_start:q_end, :],
460
+ cache_data=cache_data,
461
+ )
462
+ q_chunks.append(q_chunk[-1]) # final layer
463
+
464
+ features = torch.cat(q_chunks, dim=1) # (1, out_size*out_size, D)
465
+ features_hwc = features[0].reshape(out_size, out_size, -1).float().cpu()
466
+
467
+ # PCA
468
+ pca_data = _fit_pca(features_hwc.reshape(-1, features_hwc.shape[-1]), k=3)
469
+ pca_rgb = _apply_pca_rgb(features_hwc, pca_data)
470
+ pca_img = Image.fromarray(pca_rgb.numpy().astype(np.uint8), mode="RGB")
471
+
472
+ # Crop to original aspect ratio
473
+ pca_img = crop_feature_square_to_image_aspect(pca_img, (orig_w, orig_h))
474
+
475
+ # Resize for display
476
+ display_w, display_h = orig_w, orig_h
477
+ max_display = 512
478
+ if max(display_w, display_h) > max_display:
479
+ scale = max_display / max(display_w, display_h)
480
+ display_w = int(display_w * scale)
481
+ display_h = int(display_h * scale)
482
+ pca_display = pca_img.resize((display_w, display_h), Image.Resampling.NEAREST)
483
+
484
+ # Also create a resized input for side-by-side comparison
485
+ input_display = input_image.convert("RGB").resize((display_w, display_h), Image.Resampling.LANCZOS)
486
+
487
+ info = (f"Feature dim: {features_hwc.shape[-1]} | "
488
+ f"Output resolution: {out_size}x{out_size} | "
489
+ f"Total query points: {out_size * out_size}")
490
+
491
+ return pca_display, input_display, info
492
+
493
+
494
+ # ---------------------------------------------------------------------------
495
+ # Gradio UI
496
+ # ---------------------------------------------------------------------------
497
+ CSS = """
498
+ #col-container { max-width: 1100px; margin: 0 auto; }
499
+ .dark .gradio-container { color: var(--body-text-color); }
500
+ """
501
+
502
+ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo:
503
+ gr.Markdown("# ViT-Up: Faithful Feature Upsampling for Vision Transformers")
504
+ gr.Markdown(
505
+ "Upload an image to extract dense DINOv3 features at arbitrary resolution "
506
+ "via the ViT-Up feature upsampler. The PCA visualization shows the "
507
+ "3 principal components of the upsampled feature map as RGB."
508
+ )
509
+ gr.Markdown(
510
+ "[Paper](https://huggingface.co/papers/2606.14024) | "
511
+ "[GitHub](https://github.com/krispinwandel/vit-up) | "
512
+ "[Model Weights](https://huggingface.co/Krispin/vit-up)"
513
+ )
514
+
515
+ with gr.Row():
516
+ with gr.Column():
517
+ input_img = gr.Image(label="Input Image", type="pil")
518
+ with gr.Accordion("Advanced settings", open=False):
519
+ out_res = gr.Slider(
520
+ label="Output resolution (pixels per side)",
521
+ minimum=28,
522
+ maximum=224,
523
+ value=112,
524
+ step=28,
525
+ )
526
+ run_btn = gr.Button("Extract Features", variant="primary")
527
+ with gr.Column():
528
+ pca_out = gr.Image(label="PCA Feature Visualization")
529
+ input_display = gr.Image(label="Input (resized)")
530
+
531
+ info_text = gr.Textbox(label="Info", interactive=False)
532
+
533
+ run_btn.click(
534
+ fn=extract_and_visualize,
535
+ inputs=[input_img, out_res],
536
+ outputs=[pca_out, input_display, info_text],
537
+ api_name="extract_features",
538
+ )
539
+
540
+ gr.Examples(
541
+ examples=[
542
+ ["city_with_cars.png", 112],
543
+ ["fruit_store.png", 112],
544
+ ],
545
+ inputs=[input_img, out_res],
546
+ outputs=[pca_out, input_display, info_text],
547
+ fn=extract_and_visualize,
548
+ cache_examples=True,
549
+ cache_mode="lazy",
550
+ )
551
+
552
+ demo.launch(mcp_server=True)
city_with_cars.png ADDED

Git LFS Details

  • SHA256: 7635ea2c8c10d1935632b2c368901e6842092b7663cf5ad1d9d80a061e9db547
  • Pointer size: 131 Bytes
  • Size of remote file: 254 kB
fruit_store.png ADDED

Git LFS Details

  • SHA256: 2c82aef91dcbe8af7851b1520744ab097bba4f0d9b4f3b9a71277392d5ede9f9
  • Pointer size: 131 Bytes
  • Size of remote file: 508 kB
requirements.txt ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ transformers
2
+ peft
3
+ omegaconf
4
+ safetensors
5
+ torchvision
6
+ einops
7
+ pillow
8
+ numpy
9
+ scikit-learn