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