Source code for bayesflow.experimental.autoencoder.variational_autoencoder

import keras

from bayesflow.metrics.functional import maximum_mean_discrepancy
from bayesflow.types import Tensor
from bayesflow.utils import resolve_seed, non_batch_axis, weighted_mean
from bayesflow.utils.serialization import serializable, serialize
from .autoencoder import AutoEncoder


[docs] @serializable("bayesflow.experimental") class VariationalAutoEncoder(AutoEncoder): """Information-Maximizing Variational Autoencoder according to [1]. The loss is computed as loss = reconstruction_loss + w_kl * KL[q(z | x) || p(z)] + w_mmd * MMD[q(z), p(z)] with w_kl = 1 - alpha w_mmd = alpha + beta - 1 Useful settings are: Vanilla VAE (default): w_kl=1, w_mmd=0 -> alpha=0, beta=1 beta-VAE: w_kl=a, w_mmd=0 -> alpha=1-a, beta=a MMD/InfoVAE: w_kl=0, w_mmd=b -> alpha=1, beta=b Mixed objective: w_kl=a, w_mmd=b -> alpha=1-a, beta=a+b [1] Zhao, S., Song, J., & Ermon, S. (2019). InfoVAE: Balancing learning and inference in variational autoencoders. In Proceedings of the AAAI Conference on Artificial Intelligence (Vol. 33, No. 01, pp. 5885-5892). Parameters ---------- latent_dim Dimensionality of the latent variable. encoder_network Network mapping inputs to an encoder representation. decoder_network Network mapping latent samples to a decoder representation. alpha InfoVAE information parameter. Controls the weight of the conditional encoder KL through ``1 - alpha``. beta InfoVAE marginal distribution matching parameter. Together with alpha, controls the MMD weight through ``alpha + beta - 1``. In [1], this parameter is named lambda. mmd_kwargs Optional keyword arguments forwarded to ``maximum_mean_discrepancy``. """ def __init__( self, latent_dim: int, encoder_network: keras.Layer, decoder_network: keras.Layer, alpha: float = 0.0, beta: float = 1.0, mmd_kwargs: dict | None = None, **kwargs, ): super().__init__( latent_dim=latent_dim, encoder_network=encoder_network, decoder_network=decoder_network, **kwargs, ) self.encoder_projector.units = 2 * latent_dim self.alpha = alpha self.beta = beta self.mmd_kwargs = mmd_kwargs or {} self.seed_generator = keras.random.SeedGenerator() @property def kl_weight(self) -> float: return 1.0 - self.alpha @property def mmd_weight(self) -> float: return self.alpha + self.beta - 1.0
[docs] def get_config(self): base_config = super().get_config() config = {"alpha": self.alpha, "lambd": self.beta, "mmd_kwargs": self.mmd_kwargs} return base_config | serialize(config)
def _encode( self, x: Tensor, training: bool = False, seed: int | keras.random.SeedGenerator | None = None, **kwargs, ): seed = resolve_seed(seed) z = super()._forward(x, training=training, **kwargs) mean, log_var = keras.ops.split(z, 2, axis=-1) epsilon = keras.random.normal( shape=keras.ops.shape(mean), seed=seed, dtype=mean.dtype, ) sample = mean + keras.ops.exp(0.5 * log_var) * epsilon return z, mean, log_var, epsilon, sample def _forward( self, x: Tensor, training: bool = False, seed: int | keras.random.SeedGenerator | None = None, **kwargs, ): *_, sample = self._encode( x, training=training, seed=seed, **kwargs, ) return sample def _conditional_kl(self, mean: Tensor, log_var: Tensor) -> Tensor: """Per-example KL[q(z | x) || p(z)] for diagonal Gaussian q and standard normal p.""" return 0.5 * keras.ops.sum( keras.ops.square(mean) + keras.ops.exp(log_var) - 1.0 - log_var, axis=non_batch_axis(mean), ) def _marginal_mmd(self, sample: Tensor, seed: int | keras.random.SeedGenerator | None = None) -> Tensor: """MMD[q(z), p(z)] using samples from the aggregate posterior and prior.""" targets = keras.random.normal( shape=keras.ops.shape(sample), seed=seed, dtype=sample.dtype, ) return maximum_mean_discrepancy(sample, targets, **self.mmd_kwargs) def _reconstruction_loss(self, x: Tensor, reconstruction: Tensor) -> Tensor: """Per-example mean squared reconstruction error.""" return keras.ops.mean( keras.ops.square(x - reconstruction), axis=non_batch_axis(x), )
[docs] def compute_metrics( self, x: Tensor, sample_weight: Tensor = None, stage: str = "training", seed: int | keras.random.SeedGenerator | None = None, **kwargs, ) -> dict[str, Tensor]: training = stage == "training" seed = resolve_seed(seed) _, mean, log_var, _, sample = self._encode( x, training=training, seed=seed, **kwargs, ) reconstruction = self( sample, training=training, inverse=True, **kwargs, ) recon_loss = weighted_mean(self._reconstruction_loss(x, reconstruction), sample_weight) kl_loss = keras.ops.mean(self._conditional_kl(mean, log_var)) loss = recon_loss + self.kl_weight * kl_loss if self.mmd_weight != 0.0: mmd_loss = self._marginal_mmd(sample, seed=seed) loss = loss + self.mmd_weight * mmd_loss else: mmd_loss = keras.ops.zeros((), dtype=sample.dtype) return {"loss": loss, "recon_loss": recon_loss, "kl_loss": kl_loss, "mmd_loss": mmd_loss, "z": sample}