File size: 5,468 Bytes
75ff4df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
"""Frozen DINOv3-B behind a trained patch embedding, and the layer projection.

Every block's patch tokens (under the final LayerNorm) are standardized per
channel with fixed statistics, concatenated (12 x 768) and mapped by one linear
projection to the latent channels.
"""

from __future__ import annotations

from typing import TYPE_CHECKING, cast

import timm
import torch
from timm.layers.pos_embed_sincos import RotaryEmbeddingDinoV3
from torch import Tensor, nn

from .config import PATCH

if TYPE_CHECKING:
    from timm.models.eva import Eva

    from .config import Dinac3Config

IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)


def _require_supported_backbone(backbone: Eva) -> None:
    """Fail fast if timm builds a backbone the explicit forward does not cover.

    Raises:
        ValueError: On absolute position embeddings, patch dropout, a pre-norm,
            missing CLS/register tokens or a RoPE other than DINOv3's without
            coordinate augmentation.
    """

    supported = (
        backbone.pos_embed is None
        and backbone.patch_drop is None
        and isinstance(backbone.norm_pre, nn.Identity)
        and backbone.cls_token is not None
        and backbone.reg_token is not None
        and isinstance(backbone.rope, RotaryEmbeddingDinoV3)
        and not backbone.rope.aug_active
    )
    if not supported:
        raise ValueError("Unsupported timm DINOv3 backbone structure for dinac3")


class Encoder(nn.Module):
    """Images in [-1, 1] to raw latents ``[B, C, H / 16, W / 16]``."""

    pixel_mean: Tensor
    pixel_std: Tensor
    layer_mean: Tensor
    layer_std: Tensor

    def __init__(self, config: Dinac3Config) -> None:
        """Build the architecture only; every tensor comes from the artifact."""

        super().__init__()
        self.backbone = cast(
            "Eva",
            timm.create_model(
                config.backbone.value,
                pretrained=False,
                num_classes=0,
                dynamic_img_size=True,
                dynamic_img_pad=True,
            ),
        )
        _require_supported_backbone(self.backbone)
        # timm keeps RoPE periods non-persistent; the artifact stores the exact
        # values the training run used.
        rope = cast("nn.Module", self.backbone.rope)
        rope.register_buffer("periods", cast("Tensor", rope.periods), persistent=True)
        self.prefix_tokens = self.backbone.num_prefix_tokens
        width = self.backbone.embed_dim * len(self.backbone.blocks)
        self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1))
        self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1))
        self.register_buffer("layer_mean", torch.zeros(width))
        self.register_buffer("layer_std", torch.ones(width))
        self.projection = nn.Linear(width, config.latent_channels)
        self._rope_cache: dict[tuple[int, int, torch.device, torch.dtype], Tensor] = {}

    def rope_embed(self, height: int, width: int) -> Tensor:
        """The DINOv3 RoPE table of an image size, built once per token grid
        (outside any compiled graph: timm builds coordinates on the CPU)."""

        rope = cast("RotaryEmbeddingDinoV3", self.backbone.rope)
        periods = cast("Tensor", rope.periods)
        key = (height // PATCH, width // PATCH, periods.device, periods.dtype)
        cached = self._rope_cache.get(key)
        if cached is None:
            with torch.inference_mode(False), torch.no_grad():
                cached = rope.get_embed(shape=[key[0], key[1]])
            self._rope_cache[key] = cached
        return cached

    def forward(self, images: Tensor, rope: Tensor) -> Tensor:
        """Raw latents in the activation dtype; ``rope`` from :meth:`rope_embed`."""

        dtype = self._prefix_tokens()[0].dtype
        x = ((images.float().add(1.0).mul(0.5) - self.pixel_mean) / self.pixel_std).to(
            dtype=dtype
        )
        features = self._block_features(x, rope)
        tokens = torch.cat([feature.to(dtype=dtype) for feature in features], dim=-1)
        standardized = (tokens.float() - self.layer_mean) / self.layer_std
        latents = self.projection(standardized)
        b, _, height, width = images.shape
        return latents.transpose(1, 2).reshape(b, -1, height // PATCH, width // PATCH)

    def _block_features(self, x: Tensor, rope: Tensor) -> list[Tensor]:
        """Final-normed patch tokens of every block (timm's intermediates path)."""

        backbone = self.backbone
        tokens = backbone.patch_embed(x)
        b, _, _, c = tokens.shape
        tokens = tokens.view(b, -1, c)
        cls_token, reg_token = self._prefix_tokens()
        tokens = torch.cat(
            [cls_token.expand(b, -1, -1), reg_token.expand(b, -1, -1), tokens], dim=1
        )
        features: list[Tensor] = []
        for block in backbone.blocks:
            tokens = block(tokens, rope=rope)
            features.append(backbone.norm(tokens)[:, self.prefix_tokens :])
        return features

    def _prefix_tokens(self) -> tuple[Tensor, Tensor]:
        """The CLS and register tokens (present: checked at construction)."""

        cls_token, reg_token = self.backbone.cls_token, self.backbone.reg_token
        if cls_token is None or reg_token is None:
            raise RuntimeError("The DINOv3 backbone lost its CLS or register tokens")
        return cls_token, reg_token