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
|