dinac3_96 / dinac3 /encoder.py
data-archetype's picture
dinac3_96 v1.0
75ff4df
Raw History Blame Contribute Delete
5.47 kB
"""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