Jet#

class stable_pretraining.backbone.Jet(image_size: int | tuple[int, int] = 64, patch_size: int = 4, in_channels: int = 3, coupling_layers: int = 8, hidden_dim: int = 256, depth: int = 2, num_heads: int = 8, coupling_types: tuple[str, ...] = ('channel', 'spatial'), scale_parameterization: str = 'exp_floor', scale_eps: float = 0.0001, checkpoint_conditioner: bool = False, capture_stats: bool = False)[source]#

Bases: Module

Invertible transformer image encoder with per-sample exact logdet.

Parameters:
  • image_size – Image height/width, or a square side length.

  • patch_size – Square patch side; must divide both image dimensions.

  • in_channels – Image channels. patch_size**2 * in_channels must be even.

  • coupling_layers – Number of affine coupling layers.

  • hidden_dim – Transformer conditioner width, divisible by num_heads.

  • depth – Transformer blocks in each conditioner.

  • num_heads – Attention heads per conditioner.

  • coupling_types – Repeated sequence of channel and/or spatial. Spatial coupling requires an even number of patches.

  • scale_parameterization – exp_floor (default) or jet_sigmoid.

  • scale_eps – Lower floor for exp_floor, strictly between zero and one.

  • checkpoint_conditioner – Recompute conditioner activations during backward.

  • capture_stats – Populate detached diagnostics after each forward pass.

Note

Forward returns tokens and logdet, not a pooled embedding. Pooling tokens is lossy: the reported determinant belongs to the full token output only. Images and affine arithmetic use float32 minimum even under autocast. Reconstruction is accurate to floating-point precision, not bitwise guaranteed after training. Use the same precision context for inverse. Scale semantics are stored in state_dict extra state and checked on load. This port does not load upstream or deepstats state_dicts directly.

forward(images: Tensor) → tuple[Tensor, Tensor][source]#

Encode images and compute the full-output log absolute determinant.

Parameters:

images – Floating-point BCHW images.

Returns:

Tokens (B, N, D) and per-sample logdet (B,).

get_extra_state() → dict[str, Any][source]#

Return checkpoint metadata protecting affine scale semantics.

Returns:

Scale configuration plus image geometry and coupling layout.

inverse(tokens: Tensor) → tuple[Tensor, Tensor][source]#

Decode tokens with the identical affine scale and opposite logdet.

Parameters:

tokens – Full encoder output (B, N, D), without pooling or projection.

Returns:

Reconstructed BCHW images and per-sample inverse logdet (B,).

patchify(images: Tensor) → Tensor[source]#

Rearrange BCHW images into tokens without changing volume.

Parameters:

images – Batch matching the configured image geometry.

Returns:

Tensor shaped (batch, patches, patch_dim).

set_extra_state(state: dict[str, Any]) → None[source]#

Validate checkpoint metadata before accepting its scale semantics.

Parameters:

state – Metadata produced by get_extra_state.

Raises:

RuntimeError – The checkpoint describes a different transform.

unpatchify(tokens: Tensor) → Tensor[source]#

Invert patchify into BCHW images.

Parameters:

tokens – Tensor shaped (batch, patches, patch_dim).

Returns:

Image tensor with the configured geometry.