from __future__ import annotations import torch from torch import nn class PocketDenoiser(nn.Module): def __init__(self, diffusion_steps: int = 50) -> None: super().__init__() self.diffusion_steps = diffusion_steps self.time_embedding = nn.Embedding(diffusion_steps, 24) self.label_embedding = nn.Embedding(11, 24) self.network = nn.Sequential( nn.Linear(64 + 24 + 24, 160), nn.GELU(), nn.Linear(160, 160), nn.GELU(), nn.Linear(160, 64), ) def forward( self, noisy_pixels: torch.Tensor, timesteps: torch.Tensor, labels: torch.Tensor, ) -> torch.Tensor: features = torch.cat( [ noisy_pixels, self.time_embedding(timesteps), self.label_embedding(labels), ], dim=1, ) return self.network(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())