from __future__ import annotations import math import torch from torch import nn class ConditionalGenerator(nn.Module): def __init__(self, noise_dimensions: int = 32, embedding_dimensions: int = 16) -> None: super().__init__() self.noise_dimensions = noise_dimensions self.label_embedding = nn.Embedding(10, embedding_dimensions) self.network = nn.Sequential( nn.Linear(noise_dimensions + embedding_dimensions, 128), nn.LayerNorm(128), nn.SiLU(), nn.Linear(128, 128), nn.LayerNorm(128), nn.SiLU(), nn.Linear(128, 64), nn.Sigmoid(), ) def forward(self, noise: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: condition = self.label_embedding(labels) return self.network(torch.cat([noise, condition], dim=1)) @torch.inference_mode() def generate( self, labels: torch.Tensor, *, seed: int, temperature: float = 1.0, ) -> torch.Tensor: generator = torch.Generator(device=labels.device).manual_seed(seed) noise = torch.randn( len(labels), self.noise_dimensions, generator=generator, device=labels.device, ) return self(noise * temperature, labels) class ProjectionCritic(nn.Module): def __init__(self, feature_dimensions: int = 64) -> None: super().__init__() self.features = nn.Sequential( nn.Linear(64, 128), nn.LeakyReLU(0.2), nn.Linear(128, feature_dimensions), nn.LeakyReLU(0.2), ) self.score = nn.Linear(feature_dimensions, 1) self.label_projection = nn.Embedding(10, feature_dimensions) self.classifier = nn.Linear(feature_dimensions, 10) def forward( self, pixels: torch.Tensor, labels: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: features = self.features(pixels) projection = (features * self.label_projection(labels)).sum(dim=1) projection = projection / math.sqrt(features.shape[1]) score = self.score(features).squeeze(1) + projection return score, self.classifier(features) 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())