Jet: invertible images, tokens, and exact logdet

Jet: invertible images, tokens, and exact logdet#

Jet is a backbone, exposed as spt.backbone.Jet. It does not choose a training objective or create its own training loop. This is an unreleased addition; install the current source as described in the quickstart.

import torch
import stable_pretraining as spt

encoder = spt.backbone.Jet(
    image_size=16, patch_size=2, in_channels=3,
    coupling_layers=2, hidden_dim=16, depth=1, num_heads=2,
    scale_parameterization="exp_floor", scale_eps=1e-4,
    capture_stats=True,
)
images = torch.randn(2, 3, 16, 16)
tokens, logdet = encoder(images)
reconstructed, inverse_logdet = encoder.inverse(tokens)
torch.testing.assert_close(reconstructed, images)
torch.testing.assert_close(logdet + inverse_logdet, torch.zeros_like(logdet), atol=1e-6, rtol=0)
assert tokens.shape == (2, 64, 12)
assert logdet.shape == (2,)

Inputs are BCHW. Patch dimensions must be even; spatial coupling also needs an even number of patches. The conditioner uses transformers with zero-initialized output projections, so initialization is identity in patch coordinates. The port adapts btrude/jet-pytorch under Apache-2.0; see the repository’s THIRD_PARTY_NOTICES.md and packaged license. It has its own state-dict layout; upstream/deepstats checkpoints need an explicit conversion.

The default scale is eps + (1-eps)*exp(raw). Its logarithm uses logaddexp. The scale itself is evaluated by adding the floor directly, since exp(log(eps)) can round below that floor. There is no upper clipping and no log(scale+eps) determinant approximation. jet_sigmoid selects the original 2*sigmoid(raw) bounded ablation. Choose one parameterization for the entire encoder; settings are included in its state dict and checked on reload.

The original pre-scale bias convention is preserved: (x + bias)*scale, with inverse y/scale - bias. This equals x*scale + shift for shift=bias*scale. Coupling determinants sum log_scale over transformed coordinates. Scale, affine arithmetic, and logdet use float32 minimum; float64 is preserved. Use the same autocast context for forward and inverse. Floating-point rounding still limits reconstruction accuracy, particularly after large expansion.

With capture_stats=True, encoder.diagnostics contains detached tensors:

  • flow/layer_k/log_scale_{mean,std,min,max}

  • flow/layer_k/scale_{mean,std,min,max}

  • flow/logdet_per_dim_{mean,std,min,max}

Pass them to Module.log_dict in the forward callable. Values with abs(log_scale)>10 warn; nonfinite scales, shifts, outputs, and determinants raise. No warning changes the optimization or clamps a value. Checks synchronize with the device, so this correctness-first implementation has a runtime cost. checkpoint_conditioner=True optionally trades recomputation for activation memory.

MSE plus entropy demonstration#

python -m stable_pretraining.quickstart --method jet-entropy --cache-dir ./jet-demo
spt web ./jet-demo/runs

This smoke example predicts mean-pooled tokens across two augmented views, using MSE minus 0.01 times mean full-token logdet per input dimension. It logs semantic mean-token energy, residual-token energy, and flow diagnostics. It reuses spt.Module, spt.OnlineProbe, spt.data.DataModule, and spt.Manager. It is not a claim that the objective has a finite optimum or learns semantic representations: residual directions can expand without being penalized by a mean-token MSE. Inspect that behavior in a real experiment.

The exact determinant belongs to all tokens, not their mean or any reduced projection. A floor on coupling scales also does not bound every singular value of the full Jacobian. For a continuous input distribution, expected full-output logdet gives its differential-entropy change under the invertible transform; it does not by itself supply an absolute entropy estimate for discrete images.