"""Finite Scalar Quantization from https://arxiv.org/abs/2309.15505.""" from __future__ import annotations import torch from einops import pack, rearrange, unpack from torch import Tensor, int32, nn # helper functions def _pack_single( tensor: Tensor, pattern: str, ) -> tuple[Tensor, list[torch.Size]]: return pack([tensor], pattern) def _unpack_single( tensor: Tensor, packed_shape: list[torch.Size], pattern: str, ) -> Tensor: return unpack(tensor, packed_shape, pattern)[0] # tensor helpers def straight_through_round(inputs: Tensor) -> Tensor: """Round with straight through gradients.""" rounded = inputs.round() return inputs + (rounded - inputs).detach() # main class class FiniteScalarQuantizer(nn.Module): """Quantize continuous features into a product of scalar codebooks.""" def __init__( self, levels: list[int] | tuple[int, ...], input_dimension: int | None = None, output_dimension: int | None = None, num_codebooks: int = 1, keep_codebook_dimension: bool | None = None, scale: float | None = None, jitter_spread: float = 0.0, ) -> None: super().__init__() if not levels or any(level < 2 for level in levels): raise ValueError("levels must contain integers greater than one.") if num_codebooks <= 0: raise ValueError("num_codebooks must be positive.") if jitter_spread < 0: raise ValueError("jitter_spread must be non-negative.") _levels = torch.tensor(levels, dtype=int32) self.register_buffer("_levels", _levels, persistent=False) _basis = torch.cumprod(torch.tensor([1] + levels[:-1]), dim=0, dtype=int32) self.register_buffer("_basis", _basis, persistent=False) self.scale = scale codebook_dimension = len(levels) self.codebook_dimension = codebook_dimension self.jitter_spread = jitter_spread effective_codebook_dimension = codebook_dimension * num_codebooks self.num_codebooks = num_codebooks self.effective_codebook_dimension = effective_codebook_dimension if keep_codebook_dimension is None: keep_codebook_dimension = num_codebooks > 1 if num_codebooks > 1 and not keep_codebook_dimension: raise ValueError( "keep_codebook_dimension must be true with multiple codebooks." ) self.keep_codebook_dimension = keep_codebook_dimension self.input_dimension = ( input_dimension if input_dimension is not None else len(_levels) * num_codebooks ) if self.input_dimension <= 0: raise ValueError("input_dimension must be positive.") has_projections = ( self.input_dimension != effective_codebook_dimension ) self.input_projection = ( nn.Linear(self.input_dimension, effective_codebook_dimension) if has_projections else nn.Identity() ) if output_dimension is not None: self.output_projection = nn.Linear( effective_codebook_dimension, output_dimension ) else: self.output_projection = ( nn.Linear( effective_codebook_dimension, self.input_dimension ) if has_projections else nn.Identity() ) self.has_projections = has_projections self.codebook_size = int(self._levels.prod().item()) implicit_codebook = self.indices_to_codes( torch.arange(self.codebook_size), apply_output_projection=False, ) self.register_buffer("implicit_codebook", implicit_codebook, persistent=False) def bound_inputs(self, inputs: Tensor, epsilon: float = 1e-3) -> Tensor: """Bound inputs with shape (..., dimension).""" if self.training and self.jitter_spread: inputs = inputs + torch.randn_like(inputs) * self.jitter_spread half_width = (self._levels - 1) * (1 - epsilon) / 2 offset = torch.where(self._levels % 2 == 0, 0.5, 0.0) shift = (offset / half_width).tan() return (inputs + shift).tanh() * half_width - offset def quantize(self, inputs: Tensor) -> Tensor: """Quantize inputs and return values with the same shape.""" quantized = straight_through_round(self.bound_inputs(inputs)) half_width = self._levels // 2 # Renormalize to [-1, 1]. return quantized / half_width def _normalized_to_code_coordinates(self, normalized_codes: Tensor) -> Tensor: half_width = self._levels // 2 return (normalized_codes * half_width) + half_width def _code_coordinates_to_normalized(self, codes: Tensor) -> Tensor: half_width = self._levels // 2 return (codes - half_width) / half_width def codes_to_indices(self, normalized_codes: Tensor) -> Tensor: """Convert normalized scalar codes to integer codebook indices.""" if normalized_codes.shape[-1] != self.codebook_dimension: raise ValueError( f"Expected code dimension {self.codebook_dimension}, " f"received {normalized_codes.shape[-1]}." ) code_coordinates = self._normalized_to_code_coordinates( normalized_codes ) return (code_coordinates * self._basis).sum(dim=-1).to(int32) def indices_to_codes( self, indices: Tensor, apply_output_projection: bool = True, ) -> Tensor: """Inverse of `codes_to_indices`.""" has_spatial_dimensions = indices.ndim >= ( 3 + int(self.keep_codebook_dimension) ) indices = rearrange(indices, "... -> ... 1") code_coordinates = (indices // self._basis) % self._levels codes = self._code_coordinates_to_normalized(code_coordinates) if self.keep_codebook_dimension: codes = rearrange(codes, "... c d -> ... (c d)") if apply_output_projection: codes = self.output_projection(codes) if has_spatial_dimensions: codes = rearrange(codes, "b ... d -> b d ...") return codes def forward(self, inputs: Tensor) -> tuple[Tensor, Tensor]: """Return quantized outputs and their integer codebook indices.""" has_spatial_dimensions = inputs.ndim >= 4 # standardize image or video into (batch, seq, dimension) if has_spatial_dimensions: inputs = rearrange(inputs, "b d ... -> b ... d") inputs, packed_shape = _pack_single(inputs, "b * d") if inputs.shape[-1] != self.input_dimension: raise ValueError( f"Expected input dimension {self.input_dimension}, " f"received {inputs.shape[-1]}." ) projected_inputs = self.input_projection(inputs) projected_inputs = rearrange( projected_inputs, "b n (c d) -> b n c d", c=self.num_codebooks, ) codes = self.quantize(projected_inputs) indices = self.codes_to_indices(codes) codes = rearrange(codes, "b n c d -> b n (c d)") outputs = self.output_projection(codes) # reconstitute image or video dimensions if has_spatial_dimensions: outputs = _unpack_single(outputs, packed_shape, "b * d") outputs = rearrange(outputs, "b ... d -> b d ...") indices = _unpack_single(indices, packed_shape, "b * c") if not self.keep_codebook_dimension: indices = rearrange(indices, "... 1 -> ...") return outputs, indices