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 WRGB image becomes a32 x H/8 x W/8latent. - 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
Nsteps (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. Useencode_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
stepstrades 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.dtypeinfrom_pretrainedsets 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 statisticsmodel.safetensorsโ encoder and decoder weights (float32)ae/โ model code (AutoEncoder,AEConfig)example.pyโ round-trip exampleLICENSEโ Apache License 2.0
License
Apache License 2.0 โ see LICENSE.
- Downloads last month
- -