| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| class AffinityRegressor(nn.Module): | |
| def __init__( | |
| self, | |
| input_dim: int, | |
| hidden_dims: list[int] | tuple[int, ...] = (256, 128, 64), | |
| dropout: float = 0.2, | |
| ) -> None: | |
| super().__init__() | |
| layers: list[nn.Module] = [] | |
| previous = input_dim | |
| for hidden in hidden_dims: | |
| layers.extend( | |
| [ | |
| nn.Linear(previous, hidden), | |
| nn.LayerNorm(hidden), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| ] | |
| ) | |
| previous = hidden | |
| layers.append(nn.Linear(previous, 1)) | |
| self.network = nn.Sequential(*layers) | |
| def forward(self, features: torch.Tensor) -> torch.Tensor: | |
| return self.network(features).squeeze(-1) | |