A 🧨 diffusers-native AutoencoderKLMage port of the VAE from
microsoft/Mage-Flow
(DConvEncoder + DConvDenoiser/CoD decoder), usable as a drop-in adapter for
pipelines that expect a Flux.2-style VAE.
Latent space: shaped and valued like the raw Flux.2 VAE latent — the native
128-channel H/16 code is 2×2-unpatchified and denormalized with the Flux.2 BN
latent stats stored in config.json (anchor-latent regularization,
arXiv:2607.19064). No dependency on the
Flux.2 VAE at runtime (0.972 mean latent correlation over a 50-image test
set; see the comparison below).
Checkpoint: bf16, with the t=0 adaLN MLPs constant-folded at conversion time
(folded: true in the config), so it is ready to run on load.
[!NOTE]
The latents exposed by this model are not the native MageVAE code. The
native (B, 128, H/16, W/16) code is 2×2-unpatchified to
(B, 32, H/8, W/8) and denormalized with the Flux.2 BN latent statistics
(stored in config.json) so that encode/decode operate directly in the
original Flux.2 VAE latent space. Latents from this model can be decoded by
the Flux.2 VAE and vice versa.
Comparison with the original Flux.2 VAE
Measured against black-forest-labs/FLUX.2-dev (subfolder vae,
AutoencoderKLFlux2) on an NVIDIA L4, bf16, batch 1 (torch 2.8.0,
diffusers 0.37.1). Speed/memory at 1024×1024, means over 5 runs after warmup,
memory is peak CUDA allocation during the op. Quality is the mean over a
50-image test set (06JPEG, ~2040×1524 photos, center-cropped to a multiple
of 16, evaluated at native resolution with deterministic .mode() latents):
MageFlow VAE (this repo)
Flux.2 VAE
Parameters
100.8M
84.0M
Encode time
23 ms (~10× faster)
240 ms
Decode time
71 ms (~6× faster)
441 ms
Encode peak memory
0.46 GiB
1.90 GiB
Decode peak memory
0.98 GiB
2.78 GiB
Roundtrip PSNR (50 images)
34.1 dB
34.7 dB
Quality is on par with the original: same-VAE roundtrip reconstruction
averages within ~0.5 dB of the Flux.2 VAE over the 50-image set. The latent
spaces are interchangeable — mean latent correlation with raw Flux.2 latents
is 0.972 (stable per image, ~0.97 on every one of the 50), and
cross-decoding (flux.decode(mage.encode(x)) at 34.0 dB,
mage.decode(flux.encode(x)) at 33.8 dB) stays within ~0.7 dB of the
same-VAE roundtrip.
Repeated roundtrips hold up slightly better than the original: MageFlow
starts ~0.5 dB below the Flux.2 VAE at one roundtrip but degrades more
slowly, matching it by the second cycle and leading by ~1.1 dB after five —
so it is well suited to iterative editing workflows. Alternating the two
VAEs each cycle (MageFlow encode → Flux.2 decode) tracks in between
(50-image means):
PSNR vs original (dB)
k=1
k=2
k=3
k=4
k=5
MageFlow VAE ×k
34.1
32.1
30.3
28.9
27.7
Flux.2 VAE ×k
34.7
32.1
29.9
28.1
26.5
Alternating (Mage enc → Flux.2 dec)
34.0
31.5
29.5
27.8
26.4
Speed comes from the one-step architecture: both directions run a single
t=0 forward pass of depthwise-conv (DiCo) blocks with the adaLN modulation
constant-folded into the checkpoint, instead of the Flux.2 VAE's deep
ResNet/attention encoder-decoder.
Memory stays flat at high resolution: the decoder's per-patch MLP tail is
chunked (decode_chunk_size, default 4096 = one 1024×1024 image worth of
16×16 patches), so decode peak memory is roughly constant beyond 1024×1024
instead of growing with image area.
Files
File
Purpose
autoencoder_kl_mage.py
Self-contained: AutoencoderKLMage (ModelMixin/ConfigMixin), all network blocks, and the one-time checkpoint converter
config.json
Model config, incl. Flux.2 BN latent stats
diffusion_pytorch_model.safetensors
Folded bf16 weights
Usage
The model class lives in this repo, so grab the code files alongside the weights:
Because the latents are shape- and value-compatible with the Flux.2 VAE, you can
swap this in as the vae of a Flux.2 pipeline:
pipe.vae = vae
Input H/W must be multiples of 16. from_pretrained accepts
fold_adaln=False if you need the unfolded adaLN MLPs (only relevant for
unfolded checkpoints).
Requirements
torch
diffusers>=0.37
safetensors
loguru
Conversion
The checkpoint here was produced from the original CoD checkpoint layout with:
which remaps student.dconv_encoder.* -> encoder.* and pipeline.* -> decoder.*,
carries the Flux.2 BN stats into the config, constant-folds the t=0 adaLN MLPs
(~74 MB smaller), and saves in bf16.