from __future__ import annotations import torch from torch import nn class LinearCodec(nn.Module): def __init__(self) -> None: super().__init__() self.encoder = nn.Linear(2, 2) self.decoder = nn.Linear(2, 2) def encode(self, observations: torch.Tensor) -> torch.Tensor: return self.encoder(observations) def forward(self, observations: torch.Tensor) -> torch.Tensor: return self.decoder(self.encode(observations)) class CoordinatePredictor(nn.Module): def __init__(self) -> None: super().__init__() self.network = nn.Sequential( nn.Linear(1, 24), nn.Tanh(), nn.Linear(24, 24), nn.Tanh(), nn.Linear(24, 1), ) def forward(self, coordinate: torch.Tensor) -> torch.Tensor: return self.network(coordinate) def parameter_count(module: nn.Module) -> int: return sum(parameter.numel() for parameter in module.parameters())