"""Loading and inference API of dinac3 autoencoders.""" from __future__ import annotations import os from collections.abc import Callable from pathlib import Path from typing import NoReturn import torch from safetensors.torch import load_file from torch import Tensor, nn from torch.fx.experimental import _config as fx_config from .config import PATCH, Dinac3Config from .decoder import Decoder from .encoder import Encoder from .precision import require_storage_policy CONFIG_FILENAME = "config.json" WEIGHTS_FILENAME = "model.safetensors" SUPPORTED_DTYPES = (torch.bfloat16, torch.float32) def _looks_like_local_path(name: str) -> bool: """Whether a string names a filesystem path rather than a Hub ``org/name`` id: it starts with ``.``, ``/`` or ``~``, uses a path separator other than the id's single ``/``, or has more than one ``/``.""" other_separator = os.sep != "/" and os.sep in name return name.startswith((".", "/", "~")) or other_separator or name.count("/") > 1 def resolve_model_dir(path_or_repo: str | Path, *, revision: str | None) -> Path: """A local artifact directory, or a Hugging Face Hub snapshot of a repo id. A ``Path`` is always local. A ``str`` is local when it names an existing directory; a string that looks like a path but names none is an error; any other string is a Hub repository id, of which only ``config.json`` and ``model.safetensors`` are downloaded. Raises: FileNotFoundError: If a local path does not name a directory. ValueError: If ``revision`` is given for a local directory. """ match path_or_repo: case Path() as directory: local = directory.expanduser() case str() as name if Path(name).expanduser().is_dir(): local = Path(name).expanduser() case str() as name if _looks_like_local_path(name): raise FileNotFoundError(f"Local model path not found: {name}") case str() as repo_id: from huggingface_hub import snapshot_download return Path( snapshot_download( repo_id, revision=revision, allow_patterns=[CONFIG_FILENAME, WEIGHTS_FILENAME], ) ) case other: raise TypeError(f"Expected a str or Path, got {type(other)}") if not local.is_dir(): raise FileNotFoundError(f"Model directory does not exist: {local}") if revision is not None: raise ValueError(f"revision applies to Hub repositories, not to {local}") return local def compile_dynamic(call: Callable[..., Tensor]) -> Callable[..., Tensor]: """``torch.compile`` with dynamic shapes and without duck sizing. Dynamo otherwise assumes sizes that happen to be equal on the first call (an image width and a token count, a batch and a grid side) stay equal, and recompiles at the first shape that breaks the coincidence. Dimensions of size 1 are still specialized: a batch of 1, or a 16-pixel image side (a latent grid side of 1), compiles one more graph per such pattern. Relies on PyTorch's private ``torch.fx.experimental._config.use_duck_shape`` (tested with PyTorch 2.13). """ compiled = torch.compile(call, dynamic=True, fullgraph=True) def traced(*args: Tensor) -> Tensor: """Run the compiled call with duck sizing off (a per-thread setting).""" with fx_config.patch(use_duck_shape=False): # ty: ignore[unresolved-attribute] return compiled(*args) return traced class Dinac3(nn.Module): """Deterministic image autoencoder with a DINOv3-aligned latent. Images are RGB in [-1, 1] with height and width multiples of 16. The latent has ``config.latent_channels`` channels at stride 16: free (detail) channels first, the ``config.semantic_channels`` semantic channels last. ``encode``/``decode`` use the whitened latent (zero mean, unit variance per channel under the training distribution); ``encode_raw``/``decode_raw`` the raw one. """ latent_mean: Tensor latent_var: Tensor def __init__(self, config: Dinac3Config) -> None: """Build the architecture without weights (no network access).""" super().__init__() self.config = config self.encoder = Encoder(config) self.decoder = Decoder(config) self.register_buffer("latent_mean", torch.zeros(config.latent_channels)) self.register_buffer("latent_var", torch.ones(config.latent_channels)) self.compute_dtype = torch.bfloat16 self._encode_call: Callable[[Tensor, Tensor], Tensor] = self.encoder.forward self._decode_call: Callable[[Tensor], Tensor] = self.decoder.forward @classmethod def from_pretrained( cls, path_or_repo: str | Path, *, device: torch.device | str, dtype: torch.dtype = torch.bfloat16, compile_encoder: bool = True, compile_decoder: bool = True, revision: str | None = None, ) -> Dinac3: """Load an artifact strictly and prepare it for CUDA inference. Args: path_or_repo: Local directory (``Path``, or ``str`` naming one) or Hugging Face Hub repository id. device: CUDA device. dtype: ``torch.bfloat16`` (the stored mixed precision, BF16 autocast: the training precision) or ``torch.float32`` (every tensor upcast, no autocast). compile_encoder: ``torch.compile`` the encoder with dynamic shapes. compile_decoder: ``torch.compile`` the decoder with dynamic shapes. revision: Hub revision, for repository ids only. Raises: ValueError: On an unsupported dtype or device, or tensors stored in another dtype than the storage policy. """ directory = resolve_model_dir(path_or_repo, revision=revision) config = Dinac3Config.load(directory / CONFIG_FILENAME) state = load_file(str(directory / WEIGHTS_FILENAME)) require_storage_policy(state) # Every tensor comes from the artifact: build on the meta device and # let the loaded tensors become the parameters and buffers. with torch.device("meta"): model = cls(config) model.load_state_dict(state, strict=True, assign=True) model.prepare( device=torch.device(device), dtype=dtype, compile_encoder=compile_encoder, compile_decoder=compile_decoder, ) return model def prepare( self, *, device: torch.device, dtype: torch.dtype, compile_encoder: bool, compile_decoder: bool, ) -> None: """Move to ``device``, set the compute precision and compile once. Raises: ValueError: On a non-CUDA device or an unsupported dtype. """ if device.type != "cuda": raise ValueError("dinac3 inference requires a CUDA device") if dtype not in SUPPORTED_DTYPES: raise ValueError(f"dtype must be one of {SUPPORTED_DTYPES}, got {dtype}") self.to(device=device) if dtype == torch.float32: for tensor in (*self.parameters(), *self.buffers()): tensor.data = tensor.data.float() self.compute_dtype = dtype self.eval().requires_grad_(False) self.decoder.conv_up_head.fold_kernels() self._encode_call = ( compile_dynamic(self.encoder.forward) if compile_encoder else self.encoder.forward ) self._decode_call = ( compile_dynamic(self.decoder.forward) if compile_decoder else self.decoder.forward ) @torch.inference_mode() def encode(self, images: Tensor) -> Tensor: """Whitened FP32 latents ``[B, C, H / 16, W / 16]`` of [-1, 1] images.""" return self.whiten(self.encode_raw(images)) @torch.inference_mode() def encode_raw(self, images: Tensor) -> Tensor: """Raw (unwhitened) FP32 latents of [-1, 1] images.""" self._require_images(images) rope = self.encoder.rope_embed(images.shape[-2], images.shape[-1]) with self._autocast(): return self._encode_call(images, rope).float() @torch.inference_mode() def decode(self, latents: Tensor, height: int, width: int) -> Tensor: """FP32 RGB ``[B, 3, height, width]`` (unclamped, about [-1, 1]) from whitened latents, in one deterministic pass.""" if (height, width) != (latents.shape[-2] * PATCH, latents.shape[-1] * PATCH): raise ValueError( f"{height}x{width} is not the 16x image of a " f"{latents.shape[-2]}x{latents.shape[-1]} latent grid" ) return self.decode_raw(self.dewhiten(latents)) @torch.inference_mode() def decode_raw(self, latents: Tensor) -> Tensor: """FP32 RGB from raw (unwhitened) latents, in one deterministic pass.""" self._require_latents(latents) with self._autocast(): images = self._decode_call(latents.float()) return images.float().contiguous() def whiten(self, latents: Tensor) -> Tensor: """Raw to whitened latents, in FP32.""" mean, std = self._latent_stats() return (latents.float() - mean) / std def dewhiten(self, latents: Tensor) -> Tensor: """Whitened to raw latents, in FP32.""" mean, std = self._latent_stats() return latents.float() * std + mean def semantic_channels(self, latents: Tensor) -> Tensor: """The DINOv3-aligned channels: the last ``semantic_channels``.""" return latents[:, self.config.free_channels :] def free_channels(self, latents: Tensor) -> Tensor: """The free (detail) channels: all but the semantic ones.""" return latents[:, : self.config.free_channels] def to(self, *args: object, **kwargs: object) -> Dinac3: """Move to a device; dtype casts are rejected (see :meth:`_reject_cast`).""" casts = ( "dtype" in kwargs or "tensor" in kwargs or any(isinstance(arg, torch.dtype | Tensor) for arg in args) ) if casts: self._reject_cast() return super().to(*args, **kwargs) # ty: ignore[no-matching-overload] def half(self) -> NoReturn: """Rejected: see :meth:`_reject_cast`.""" self._reject_cast() def bfloat16(self) -> NoReturn: """Rejected: see :meth:`_reject_cast`.""" self._reject_cast() def float(self) -> NoReturn: """Rejected: see :meth:`_reject_cast`.""" self._reject_cast() def double(self) -> NoReturn: """Rejected: see :meth:`_reject_cast`.""" self._reject_cast() def type(self, dst_type: object) -> NoReturn: """Rejected: see :meth:`_reject_cast`.""" del dst_type self._reject_cast() def _reject_cast(self) -> NoReturn: """Module-wide dtype casts would break the mixed storage policy. Raises: TypeError: Always, pointing to ``from_pretrained(dtype=...)``. """ raise TypeError( "dinac3 keeps a mixed storage policy (BF16 weights, FP32 residual-path, " "statistics and RoPE tensors); a module-wide dtype cast would break it. " "Choose the precision with Dinac3.from_pretrained(..., " "dtype=torch.bfloat16 | torch.float32)." ) def _latent_stats(self) -> tuple[Tensor, Tensor]: """Per-channel ``(mean, std)`` ``[1, C, 1, 1]`` in FP32.""" mean = self.latent_mean.float().view(1, -1, 1, 1) var = self.latent_var.float().view(1, -1, 1, 1) return mean, torch.sqrt(var + self.config.latent_stats_eps) def _autocast(self) -> torch.autocast: """BF16 autocast for BF16 inference, disabled for FP32.""" return torch.autocast( "cuda", dtype=torch.bfloat16, enabled=self.compute_dtype == torch.bfloat16 ) def _require_images(self, images: Tensor) -> None: """Validate images at the API boundary, outside compiled graphs. Raises: ValueError: On a wrong shape, size, dtype or device. """ if images.ndim != 4 or images.shape[1] != 3 or images.shape[0] < 1: raise ValueError("images must have shape [B, 3, H, W] with B >= 1") height, width = images.shape[-2:] if min(height, width) < PATCH or height % PATCH or width % PATCH: raise ValueError(f"Image sides must be positive multiples of {PATCH}") self._require_tensor(images) def _require_latents(self, latents: Tensor) -> None: """Validate latents at the API boundary. Raises: ValueError: On a wrong shape, dtype or device. """ if latents.ndim != 4 or latents.shape[1] != self.config.latent_channels: raise ValueError( f"latents must have shape [B, {self.config.latent_channels}, h, w]" ) if min(latents.shape[0], *latents.shape[-2:]) < 1: raise ValueError("Latent batch and grid sides must be positive") self._require_tensor(latents) def _require_tensor(self, tensor: Tensor) -> None: """Require a floating tensor on the model's CUDA device. Raises: ValueError: If the tensor is elsewhere or not floating point. """ device = self.latent_mean.device if tensor.device != device: raise ValueError(f"Inputs must be on the model device {device}") if not tensor.is_floating_point(): raise ValueError("Inputs must be floating-point tensors")