Download dinac3/encoder.py from data-archetype/dinac3_96: direct link, hf CLI and curl.
- Browser
- Download file 5.47 kB
-
https://huggingface.co/data-archetype/dinac3_96/resolve/main/dinac3/encoder.py
- Command line
-
hf download hf://data-archetype/dinac3_96/dinac3/encoder.py
-
curl -L -o encoder.py https://huggingface.co/data-archetype/dinac3_96/resolve/main/dinac3/encoder.py
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 | |