| from __future__ import annotations |
|
|
| import torch |
| from torch import nn |
|
|
|
|
| class TuringSurrogate(nn.Module): |
| def __init__(self) -> None: |
| super().__init__() |
| self.network = nn.Sequential( |
| nn.Conv2d(4, 32, kernel_size=3, padding=1, padding_mode="circular"), |
| nn.SiLU(), |
| nn.Conv2d(32, 32, kernel_size=3, padding=1, padding_mode="circular"), |
| nn.SiLU(), |
| nn.Conv2d(32, 16, kernel_size=3, padding=1, padding_mode="circular"), |
| nn.SiLU(), |
| nn.Conv2d(16, 2, kernel_size=1), |
| ) |
|
|
| def forward( |
| self, |
| state: torch.Tensor, |
| feed: torch.Tensor, |
| kill: torch.Tensor, |
| ) -> torch.Tensor: |
| batch, _, height, width = state.shape |
| parameters = torch.stack([feed, kill], dim=1) |
| parameter_fields = parameters[:, :, None, None].expand( |
| batch, |
| 2, |
| height, |
| width, |
| ) |
| delta = self.network(torch.cat([state, parameter_fields], dim=1)) |
| return torch.clamp(state + delta, 0, 1) |
|
|
|
|
| def parameter_count(model: nn.Module) -> int: |
| return sum(parameter.numel() for parameter in model.parameters()) |
|
|