from __future__ import annotations import torch from torch import nn class ConditionalEnergyNetwork(nn.Module): def __init__(self) -> None: super().__init__() self.features = nn.Sequential( nn.Linear(64, 128), nn.SiLU(), nn.Linear(128, 64), nn.SiLU(), ) self.energy_heads = nn.Linear(64, 10) def all_energies(self, pixels: torch.Tensor) -> torch.Tensor: return self.energy_heads(self.features(pixels)) def forward(self, pixels: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: energies = self.all_energies(pixels) return energies.gather(1, labels[:, None]).squeeze(1) def langevin_sample( model: ConditionalEnergyNetwork, pixels: torch.Tensor, labels: torch.Tensor, *, steps: int, step_size: float, noise_scale: float, generator: torch.Generator | None = None, ) -> torch.Tensor: was_training = model.training model.eval() current = pixels.detach().clone() for _ in range(steps): current.requires_grad_(True) energy = model(current, labels).sum() gradient = torch.autograd.grad(energy, current)[0] with torch.no_grad(): noise = torch.randn( current.shape, generator=generator, device=current.device, ) current = current - step_size * gradient + noise_scale * noise current.clamp_(0, 1) model.train(was_training) return current.detach() 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())