| from __future__ import annotations |
|
|
| import torch |
| from schema import ACTIONS, VOCABULARY_SIZE |
| from torch import nn |
|
|
|
|
| class KernelMindPolicy(nn.Module): |
| def __init__(self, dimensions: int = 32) -> None: |
| super().__init__() |
| self.token_embedding = nn.Embedding(VOCABULARY_SIZE, dimensions) |
| self.class_token = nn.Parameter(torch.zeros(1, 1, dimensions)) |
| self.position_embedding = nn.Parameter(torch.randn(1, 8, dimensions) * 0.02) |
| layer = nn.TransformerEncoderLayer( |
| d_model=dimensions, |
| nhead=4, |
| dim_feedforward=64, |
| dropout=0.0, |
| batch_first=True, |
| activation="gelu", |
| norm_first=True, |
| ) |
| self.encoder = nn.TransformerEncoder(layer, num_layers=2) |
| self.action_heads = nn.Linear(dimensions, 3 * len(ACTIONS)) |
|
|
| def forward(self, tokens: torch.Tensor) -> torch.Tensor: |
| embedded = self.token_embedding(tokens) |
| class_token = self.class_token.expand(len(tokens), -1, -1) |
| sequence = torch.cat([class_token, embedded], dim=1) |
| hidden = self.encoder(sequence + self.position_embedding) |
| return self.action_heads(hidden[:, 0]).reshape(len(tokens), 3, len(ACTIONS)) |
|
|
|
|
| def parameter_count(model: nn.Module) -> int: |
| return sum(parameter.numel() for parameter in model.parameters()) |
|
|