| name | distribution-matching-vae |
| title | Distribution Matching Variational AutoEncoder |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2512.07778 |
| keywords | ["variational autoencoders","distribution matching","latent space","generative models","image synthesis"] |
| description | Align latent distributions with arbitrary reference distributions via explicit matching constraints rather than fixed priors. DMVAE achieves gFID 3.2 on ImageNet with 64 epochs—when you need flexibility in latent representation design for image generation. |
Overview
DMVAE generalizes beyond conventional Gaussian priors by explicitly aligning the encoder's latent distribution with an arbitrary reference distribution. The framework enables matching with distributions derived from self-supervised learning, diffusion noise, or other priors—moving beyond rigid architectural constraints.
When to Use
- Image generation models where latent distribution affects quality
- Scenarios exploring optimal latent structures for specific tasks
- Applications where self-supervised learning distributions outperform Gaussians
- Models needing flexibility in distribution choice
- Seeking efficient training (64 epochs on ImageNet)
When NOT to Use
- Models already achieving satisfactory results with Gaussian priors
- Fixed, rigid latent assumptions are acceptable
- Scenarios where distribution flexibility adds complexity without benefit
Core Technique
Explicit distribution matching for flexible latent space design:
class DistributionMatchingVAE:
def __init__(self, encoder, decoder, reference_distribution=None):
self.encoder = encoder
self.decoder = decoder
self.reference_dist = reference_distribution or self.default_gaussian()
def forward(self, x):
"""
Encode to latent space and reconstruct.
Explicitly match latent distribution to reference.
"""
latent, encoder_params = self.encoder(x)
recon = self.decoder(latent)
return recon, latent, encoder_params
def compute_distribution_matching_loss(self, latent_samples):
matching_loss = .compute_distribution_distance(
latent_samples,
.reference_dist
)
matching_loss
():
x = batch
recon, latent, encoder_params = .forward(x)
recon_loss = torch.nn.functional.mse_loss(recon, x)
dist_matching_loss = .compute_distribution_matching_loss(latent)
total_loss = recon_loss + .beta * dist_matching_loss
total_loss
():
ssl_features = []
batch dataset:
features = ssl_model.encode(batch)
ssl_features.append(features)
ssl_features = torch.cat(ssl_features, dim=)
.reference_dist = .fit_distribution_to_features(
ssl_features
)
.reference_dist
():
noise_samples = diffusion_model.sample_noise_distribution()
.reference_dist = .fit_distribution_to_features(
noise_samples
)
.reference_dist
():
.metric == :
distance = .wasserstein_distance(samples, reference)
.metric == :
distance = .maximum_mean_discrepancy(samples, reference)
.metric == :
distance = .ks_distance(samples, reference)
:
ValueError()
distance