"""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