ARotting's picture
Publish 9.5K parameter conditional variational autoencoder
36072d3 verified
Raw
History Blame Contribute Delete
2.52 kB
from __future__ import annotations
import torch
from torch import nn
class ConditionalVAE(nn.Module):
def __init__(self, latent_dimensions: int = 8) -> None:
super().__init__()
self.latent_dimensions = latent_dimensions
self.label_embedding = nn.Embedding(10, 8)
self.encoder = nn.Sequential(
nn.Linear(64, 64),
nn.GELU(),
nn.Linear(64, 32),
nn.GELU(),
)
self.mean = nn.Linear(32, latent_dimensions)
self.log_variance = nn.Linear(32, latent_dimensions)
self.decoder = nn.Sequential(
nn.Linear(latent_dimensions + 8, 32),
nn.GELU(),
nn.Linear(32, 64),
nn.Sigmoid(),
)
def encode(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
hidden = self.encoder(pixels)
return self.mean(hidden), self.log_variance(hidden)
def reparameterize(
self,
mean: torch.Tensor,
log_variance: torch.Tensor,
) -> torch.Tensor:
if not self.training:
return mean
standard_deviation = torch.exp(0.5 * log_variance)
return mean + torch.randn_like(standard_deviation) * standard_deviation
def decode(self, latent: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
condition = self.label_embedding(labels)
return self.decoder(torch.cat([latent, condition], dim=1))
def forward(
self,
pixels: torch.Tensor,
labels: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
mean, log_variance = self.encode(pixels)
latent = self.reparameterize(mean, log_variance)
return self.decode(latent, labels), mean, log_variance
class TinyVisionJudge(nn.Module):
def __init__(self) -> None:
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(1, 8, kernel_size=3, padding=1),
nn.GELU(),
nn.Conv2d(8, 8, kernel_size=3, padding=1, groups=8),
nn.GELU(),
nn.Conv2d(8, 12, kernel_size=1),
nn.GELU(),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.Flatten(),
nn.Linear(12 * 4 * 4, 10),
)
def forward(self, pixels: torch.Tensor) -> torch.Tensor:
return self.classifier(self.features(pixels))
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())