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