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:
ModuleInvertible 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
channeland/orspatial. Spatial coupling requires an even number of patches.scale_parameterization –
exp_floor(default) orjet_sigmoid.scale_eps – Lower floor for exp_floor, strictly between zero and one.
checkpoint_conditioner – Recompute conditioner activations during backward.
capture_stats – Populate detached
diagnosticsafter 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).