File size: 3,145 Bytes
f66bbd0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""

backbone.py — Frozen DINOv3 feature extractor.



GazeRefine never updates the backbone. We load a pretrained DINOv3 ViT via

`timm`, freeze every parameter, and pull out raw patch tokens from one or

more transformer blocks using forward hooks. No adapters, no projections,

no fine-tuning.

"""

from __future__ import annotations

import torch
import torch.nn as nn


class FrozenDINOv3(nn.Module):
    """Completely frozen DINOv3 ViT backbone that exposes raw patch tokens.



    Parameters

    ----------

    model_name : str

        Any DINOv3 variant available in `timm` (e.g.

        ``"vit_large_patch16_dinov3.lvd1689m"``,

        ``"vit_base_patch16_dinov3.lvd1689m"``). Larger backbones generally

        give cleaner semantic separation but cost more memory/compute.

    extract_mode : {"last", "all"}

        - ``"last"``: use only the final block's patch tokens (fast, the

          default used in our reported results).

        - ``"all"``: pool patch tokens from 4 evenly-spaced blocks

          (0, n/4, 3n/4, n-1) and average the resulting similarity maps in

          ``GazeRefine``. This sometimes helps on harder modalities at the

          cost of ~4x compute.

    """

    def __init__(

        self,

        model_name: str = "vit_large_patch16_dinov3.lvd1689m",

        extract_mode: str = "last",

    ):
        super().__init__()
        import timm  # local import: keeps `timm` optional for users who only read the code

        print(f"[GazeRefine] Loading frozen backbone: {model_name} (extract_mode={extract_mode})")
        bb = timm.create_model(model_name, pretrained=True, num_classes=0)
        for p in bb.parameters():
            p.requires_grad_(False)
        bb.eval()

        self.backbone = bb
        self.embed_dim = bb.embed_dim
        self.num_blocks = len(bb.blocks)

        if extract_mode == "last":
            self.levels = [-1]
        elif extract_mode == "all":
            self.levels = [0, self.num_blocks // 4, (self.num_blocks * 3) // 4, self.num_blocks - 1]
        else:
            raise ValueError(f"extract_mode must be 'last' or 'all', got {extract_mode!r}")

        print(f"[GazeRefine] Hooking transformer blocks at levels: {self.levels}")
        self._feats: dict[int, torch.Tensor] = {}
        for lvl in self.levels:
            bb.blocks[lvl].register_forward_hook(self._make_hook(lvl))

    def _make_hook(self, lvl: int):
        def _hook_fn(module, inp, out):
            # DINOv3 token layout: [CLS, register_1..register_k, patch_1..patch_N]
            n_prefix = 1 + getattr(self.backbone, "num_register_tokens", 4)
            self._feats[lvl] = out[:, n_prefix:, :]  # (B, N, D) — patch tokens only
        return _hook_fn

    @torch.no_grad()
    def forward(self, x: torch.Tensor) -> list[torch.Tensor]:
        """Run the frozen backbone and return a list of (B, N, D) patch-token tensors,

        one per hooked level, in the order given by ``self.levels``."""
        self._feats.clear()
        self.backbone(x)
        return [self._feats[lvl] for lvl in self.levels]