Image Autoencoder (8x, 32 channels, flow-matching decoder)

Inference-only package for an image autoencoder intended as a latent space for diffusion models.

  • Compression: 8x spatially, 32 latent channels โ€” an H x W RGB image becomes a 32 x H/8 x W/8 latent.
  • Encoder: deterministic transformer with local window attention (2D RoPE), works at any resolution whose sides are multiples of 8.
  • Decoder: conditional flow-matching model (x0-prediction). It starts from Gaussian noise and is integrated in N steps (default 4) conditioned on the latent.
  • Normalization built in: encode() returns latents already normalized per channel ((z - mean) / std, roughly zero mean / unit variance); decode() undoes it. Use encode_raw() / decode_raw() for unnormalized latents.

Install

pip install torch safetensors huggingface_hub numpy pillow
# optional: triton (CUDA) for the fused window-attention kernel; otherwise PyTorch SDPA is used

The model code lives in the ae/ folder of this repo:

huggingface-cli download Muinez/f8c32-ae --local-dir f8c32-ae
export PYTHONPATH=$PWD/f8c32-ae:$PYTHONPATH

Usage

import torch
from ae import AutoEncoder

model = AutoEncoder.from_pretrained("Muinez/f8c32-ae", device="cuda")   # bf16 compute on CUDA

x = ...                                   # (B, 3, H, W) float in [-1, 1], H and W multiples of 8
z = model.encode(x)                       # (B, 32, H/8, W/8), normalized latents
y = model.decode(z, steps=4)              # (B, 3, H, W), approximately in [-1, 1]
y = y.clamp(-1, 1)

# reproducible decoding
g = torch.Generator("cuda").manual_seed(0)
y = model.decode(z, steps=4, generator=g)

# raw (unnormalized) latents
z_raw = model.encode_raw(x)
y = model.decode_raw(z_raw, steps=4)

Batch of different resolutions / aspect ratios

encode / decode (and the _raw variants) also take a list of tensors of different sizes and return a list. The whole list goes through the model in one forward pass: each image keeps its own size (no resizing or cropping), is cut into its own grid of 8x8 patch tokens, and the grids are packed into one sequence batch padded only up to the largest image area. Attention never crosses into other images or padding, so results match encoding each image on its own. Only the decoder's small convolutional upsampling head runs per group of equal-shape images.

from PIL import Image
import numpy as np

def load(path, max_side=1024):
    im = Image.open(path).convert("RGB")
    s = min(1.0, max_side / max(im.size))
    w, h = (int(im.width * s) // 8 * 8, int(im.height * s) // 8 * 8)   # sides: multiples of 8
    im = im.resize((w, h), Image.LANCZOS)
    return torch.from_numpy(np.asarray(im)).permute(2, 0, 1)            # uint8 (3, H, W)

images = [load("portrait.jpg"), load("landscape.png"), load("square.webp")]   # any sizes
latents = model.encode(images)              # list of (32, H_i/8, W_i/8), normalized
recons = model.decode(latents, steps=4)     # list of (3, H_i, W_i)

Images may be uint8 [0, 255] or float [-1, 1].

See example.py for a complete round trip (python example.py input.png recon.png).

Notes

  • steps trades speed for detail: 1 step is fastest and slightly softer, 4 is the default. Because the decoder samples from noise, different seeds give slightly different fine detail.
  • dtype in from_pretrained sets the compute dtype (default: bfloat16 on CUDA, float32 on CPU). Encoder weights are cast to it; the decoder keeps float32 weights and runs under autocast.
  • On CUDA with Triton installed, a fused window-attention kernel is used for any image sizes; otherwise (CPU, no Triton) an equivalent masked SDPA path is used. Force SDPA with import ae.window_attention as wa; wa.BACKEND = "sdpa".

Files

  • config.json โ€” architecture and latent normalization statistics
  • model.safetensors โ€” encoder and decoder weights (float32)
  • ae/ โ€” model code (AutoEncoder, AEConfig)
  • example.py โ€” round-trip example
  • LICENSE โ€” Apache License 2.0

License

Apache License 2.0 โ€” see LICENSE.

Downloads last month
18
Safetensors
Model size
92.3M params
Tensor type
F32
ยท
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support