from __future__ import annotations import torch from torch import nn class FederatedMLP(nn.Module): def __init__(self) -> None: super().__init__() self.network = nn.Sequential( nn.Linear(64, 32), nn.GELU(), nn.Linear(32, 10), ) def forward(self, pixels: torch.Tensor) -> torch.Tensor: return self.network(pixels) def parameter_count(model: nn.Module) -> int: return sum(parameter.numel() for parameter in model.parameters())