File size: 5,922 Bytes
8399861
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
e1170d9
 
8399861
 
e1170d9
8399861
 
 
e1170d9
8399861
 
 
 
e1170d9
8399861
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
 
e1170d9
8399861
e1170d9
8399861
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e1170d9
8399861
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
---
library_name: diffusers
base_model: microsoft/Mage-Flow
base_model_relation: adapter
tags:
- vae
- autoencoder
- flux.2
- image-generation
---

# MageFlow VAE (diffusers)

A 🧨 diffusers-native `AutoencoderKLMage` port of the VAE from
[microsoft/Mage-Flow](https://huggingface.co/microsoft/Mage-Flow)
(DConvEncoder + DConvDenoiser/CoD decoder), usable as a drop-in adapter for
pipelines that expect a Flux.2-style VAE.

- **Encode:** one-step diffusion encoder (t=0) → `latent_dist` over `(B, 32, H/8, W/8)`
- **Decode:** DConvDenoiser + CoD decoder (t=0) → image `(B, 3, H, W)` in `[-1, 1]`
- **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](https://arxiv.org/abs/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](#comparison-with-the-original-flux2-vae) 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:

```python
import sys

import torch
from huggingface_hub import snapshot_download

repo_dir = snapshot_download("<your-username>/MageFlow-VAE-diffusers")
sys.path.insert(0, repo_dir)

from autoencoder_kl_mage import AutoencoderKLMage

vae = AutoencoderKLMage.from_pretrained(repo_dir, torch_dtype=torch.bfloat16).to("cuda")

image = torch.rand(1, 3, 1024, 1024, device="cuda", dtype=torch.bfloat16) * 2 - 1
latents = vae.encode(image).latent_dist.sample()   # (1, 32, 128, 128), Flux.2 latent space
recon = vae.decode(latents, return_dict=False)[0]  # (1, 3, 1024, 1024) in [-1, 1]
```

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:

```python
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:

```bash
python autoencoder_kl_mage.py   # convert_mage_ckpt: MageFlow/vae -> MageFlow/vae_diffusers
```

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.