| from __future__ import annotations |
|
|
| import torch |
| from torch import nn |
|
|
|
|
| def squash(vectors: torch.Tensor) -> torch.Tensor: |
| squared_norm = vectors.square().sum(dim=-1, keepdim=True) |
| scale = squared_norm / (1 + squared_norm) |
| return scale * vectors / torch.sqrt(squared_norm + 1e-8) |
|
|
|
|
| class DynamicRoutingCapsuleNet(nn.Module): |
| def __init__(self, routing_iterations: int = 3) -> None: |
| super().__init__() |
| self.routing_iterations = routing_iterations |
| self.primary = nn.Linear(64, 28) |
| self.transforms = nn.Parameter(torch.randn(7, 10, 4, 8) * 0.08) |
|
|
| def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| primary = squash(torch.tanh(self.primary(pixels)).reshape(-1, 7, 4)) |
| votes = torch.einsum("bpd,pcde->bpce", primary, self.transforms) |
| routing_logits = torch.zeros( |
| len(pixels), |
| 7, |
| 10, |
| device=pixels.device, |
| ) |
| digit_capsules = None |
| for iteration in range(self.routing_iterations): |
| coupling = torch.softmax(routing_logits, dim=2) |
| digit_capsules = squash((coupling[..., None] * votes).sum(dim=1)) |
| if iteration + 1 < self.routing_iterations: |
| agreement = (votes * digit_capsules[:, None]).sum(dim=-1) |
| routing_logits = routing_logits + agreement |
| assert digit_capsules is not None |
| return digit_capsules, digit_capsules.norm(dim=-1) |
|
|
|
|
| class MatchedMLP(nn.Module): |
| def __init__(self) -> None: |
| super().__init__() |
| self.network = nn.Sequential( |
| nn.Linear(64, 54), |
| nn.GELU(), |
| nn.Linear(54, 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()) |
|
|
|
|