from __future__ import annotations import torch from torch import nn from torch.nn import functional as F class PocketMoE(nn.Module): def __init__(self, experts: int = 4, top_k: int = 2) -> None: super().__init__() self.expert_count = experts self.top_k = top_k self.encoder = nn.Sequential( nn.Linear(64, 32), nn.GELU(), ) self.router = nn.Linear(32, experts) self.experts = nn.ModuleList( [ nn.Sequential( nn.Linear(32, 16), nn.GELU(), nn.Linear(16, 10), ) for _ in range(experts) ] ) def forward( self, pixels: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: hidden = self.encoder(pixels) router_probabilities = F.softmax(self.router(hidden), dim=1) top_probabilities, top_indices = router_probabilities.topk( self.top_k, dim=1, ) sparse_weights = torch.zeros_like(router_probabilities).scatter( 1, top_indices, top_probabilities, ) sparse_weights = sparse_weights / sparse_weights.sum(dim=1, keepdim=True) expert_logits = torch.stack( [expert(hidden) for expert in self.experts], dim=1, ) logits = (expert_logits * sparse_weights.unsqueeze(-1)).sum(dim=1) return logits, router_probabilities, sparse_weights class DenseControl(nn.Module): def __init__(self) -> None: super().__init__() self.network = nn.Sequential( nn.Linear(64, 48), nn.GELU(), nn.Linear(48, 40), nn.GELU(), nn.Linear(40, 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()) def active_parameter_count(model: PocketMoE) -> int: shared = sum(parameter.numel() for parameter in model.encoder.parameters()) router = sum(parameter.numel() for parameter in model.router.parameters()) experts = sorted( [ sum(parameter.numel() for parameter in expert.parameters()) for expert in model.experts ], reverse=True, ) return shared + router + sum(experts[: model.top_k])