| 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]) |
|
|