from __future__ import annotations import torch from torch import nn class TeacherCNN(nn.Module): def __init__(self) -> None: super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 16, kernel_size=3, padding=1), nn.GELU(), nn.MaxPool2d(2), nn.Conv2d(16, 32, kernel_size=3, padding=1), nn.GELU(), nn.MaxPool2d(2), ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(32 * 2 * 2, 64), nn.GELU(), nn.Dropout(0.1), nn.Linear(64, 10), ) def forward(self, pixels: torch.Tensor) -> torch.Tensor: return self.classifier(self.features(pixels)) class TinyStudentCNN(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())