from __future__ import annotations import torch from torch import nn class JEncoder(nn.Module): def __init__(self, latent_dim: int = 32) -> None: super().__init__() self.network = nn.Sequential( nn.Linear(64, 96), nn.LayerNorm(96), nn.GELU(), nn.Linear(96, 64), nn.LayerNorm(64), nn.GELU(), nn.Linear(64, latent_dim), ) def forward(self, images: torch.Tensor) -> torch.Tensor: return self.network(images.flatten(1)) class JPredictor(nn.Module): def __init__(self, latent_dim: int = 32) -> None: super().__init__() self.network = nn.Sequential( nn.Linear(latent_dim, 64), nn.GELU(), nn.Linear(64, latent_dim), ) def forward(self, embeddings: torch.Tensor) -> torch.Tensor: return self.network(embeddings) def parameter_count(module: nn.Module) -> int: return sum(parameter.numel() for parameter in module.parameters())