mmacosha's picture
add mnist_vae/ (checkpoint + code)
5e3f03f verified
Raw
History Blame Contribute Delete
4.66 kB
from typing import Any
import jax
import jax.numpy as jnp
from flax import linen as nn
Array = jax.Array
class ResBlock(nn.Module):
features: int
groups: int = 8
@nn.compact
def __call__(self, x: Array) -> Array:
residual = x
if x.shape[-1] != self.features:
residual = nn.Conv(self.features, kernel_size=(1, 1))(residual)
h = nn.GroupNorm(num_groups=min(self.groups, self.features))(x)
h = nn.swish(h)
h = nn.Conv(self.features, kernel_size=(3, 3))(h)
h = nn.GroupNorm(num_groups=min(self.groups, self.features))(h)
h = nn.swish(h)
h = nn.Conv(self.features, kernel_size=(3, 3))(h)
return h + residual
class Downsample(nn.Module):
features: int
@nn.compact
def __call__(self, x: Array) -> Array:
return nn.Conv(self.features, kernel_size=(3, 3), strides=(2, 2))(x)
class Upsample(nn.Module):
features: int
@nn.compact
def __call__(self, x: Array) -> Array:
b, h, w, _ = x.shape
x = jax.image.resize(x, (b, h * 2, w * 2, x.shape[-1]), method="nearest")
return nn.Conv(self.features, kernel_size=(3, 3))(x)
class Encoder(nn.Module):
latent_dim: int
base_channels: int = 64
channel_mults: tuple[int, ...] = (1, 2, 4)
@nn.compact
def __call__(self, x: Array) -> tuple[Array, Array]:
h = nn.Conv(self.base_channels, kernel_size=(3, 3))(x)
for i, mult in enumerate(self.channel_mults):
features = self.base_channels * mult
h = ResBlock(features)(h)
h = ResBlock(features)(h)
if i < len(self.channel_mults) - 1:
h = Downsample(features)(h)
h = Downsample(self.base_channels * self.channel_mults[-1])(h)
h = nn.GroupNorm(num_groups=8)(h)
h = nn.swish(h)
h = h.reshape(h.shape[0], -1)
mu = nn.Dense(self.latent_dim)(h)
log_var = nn.Dense(self.latent_dim)(h)
log_var = jnp.clip(log_var, -10.0, 10.0)
return mu, log_var
class Decoder(nn.Module):
out_channels: int
base_channels: int = 64
channel_mults: tuple[int, ...] = (1, 2, 4)
image_size: int = 32
@nn.compact
def __call__(self, z: Array) -> Array:
n_down = len(self.channel_mults)
start = self.image_size // (2**n_down)
features_top = self.base_channels * self.channel_mults[-1]
h = nn.Dense(start * start * features_top)(z)
h = h.reshape(z.shape[0], start, start, features_top)
h = Upsample(features_top)(h)
for i, mult in enumerate(reversed(self.channel_mults)):
features = self.base_channels * mult
h = ResBlock(features)(h)
h = ResBlock(features)(h)
if i < len(self.channel_mults) - 1:
next_features = self.base_channels * list(reversed(self.channel_mults))[i + 1]
h = Upsample(next_features)(h)
h = nn.GroupNorm(num_groups=8)(h)
h = nn.swish(h)
h = nn.Conv(self.out_channels, kernel_size=(3, 3))(h)
return nn.sigmoid(h)
class VAE(nn.Module):
latent_dim: int = 128
in_channels: int = 3
base_channels: int = 64
channel_mults: tuple[int, ...] = (1, 2, 4)
image_size: int = 32
obs_sigma: float = 0.1
def setup(self):
self.encoder = Encoder(
latent_dim=self.latent_dim,
base_channels=self.base_channels,
channel_mults=self.channel_mults,
)
self.decoder = Decoder(
out_channels=self.in_channels,
base_channels=self.base_channels,
channel_mults=self.channel_mults,
image_size=self.image_size,
)
def encode(self, x: Array) -> tuple[Array, Array]:
return self.encoder(x)
def decode(self, z: Array) -> Array:
return self.decoder(z)
def reparameterize(self, rng: Array, mu: Array, log_var: Array) -> Array:
eps = jax.random.normal(rng, mu.shape)
return mu + jnp.exp(0.5 * log_var) * eps
def __call__(self, x: Array, rng: Array) -> dict[str, Array]:
mu, log_var = self.encode(x)
z = self.reparameterize(rng, mu, log_var)
x_hat = self.decode(z)
return {"x_hat": x_hat, "mu": mu, "log_var": log_var, "z": z}
def sample_prior(self, rng: Array, n: int) -> Array:
z = jax.random.normal(rng, (n, self.latent_dim))
return self.decode(z)
def latent_log_prior(z: Array) -> Array:
"""Standard normal prior log density, summed over latent dimensions."""
d = z.shape[-1]
return -0.5 * (jnp.sum(z * z, axis=-1) + d * jnp.log(2 * jnp.pi))