Spaces:
Running
Running
| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| class SelectiveSSM(nn.Module): | |
| def __init__( | |
| self, | |
| vocab_size: int = 14, | |
| d_model: int = 32, | |
| state_dim: int = 4, | |
| selective: bool = True, | |
| ) -> None: | |
| super().__init__() | |
| self.d_model = d_model | |
| self.state_dim = state_dim | |
| self.selective = selective | |
| self.embedding = nn.Embedding(vocab_size, d_model) | |
| self.marker_projection = nn.Linear(1, d_model, bias=False) | |
| self.input_projection = nn.Linear(d_model, 2 * d_model) | |
| self.convolution = nn.Conv1d( | |
| d_model, | |
| d_model, | |
| kernel_size=3, | |
| groups=d_model, | |
| padding=2, | |
| ) | |
| self.a_log = nn.Parameter(torch.zeros(d_model, state_dim)) | |
| self.d_skip = nn.Parameter(torch.ones(d_model)) | |
| if selective: | |
| self.delta_projection = nn.Linear(d_model, d_model) | |
| self.b_projection = nn.Linear(d_model, state_dim) | |
| self.c_projection = nn.Linear(d_model, state_dim) | |
| else: | |
| self.delta = nn.Parameter(torch.zeros(d_model)) | |
| self.b = nn.Parameter(torch.randn(state_dim) * 0.05) | |
| self.c = nn.Parameter(torch.randn(state_dim) * 0.05) | |
| self.normalization = nn.LayerNorm(d_model) | |
| self.classifier = nn.Linear(d_model, 10) | |
| def scan( | |
| self, values: torch.Tensor | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| batch, length, _ = values.shape | |
| state = values.new_zeros(batch, self.d_model, self.state_dim) | |
| outputs = [] | |
| update_strength = [] | |
| stable_a = -torch.exp(self.a_log) | |
| for step in range(length): | |
| item = values[:, step] | |
| if self.selective: | |
| delta = F.softplus(self.delta_projection(item)) | |
| b = self.b_projection(item) | |
| c = self.c_projection(item) | |
| else: | |
| delta = F.softplus(self.delta).expand(batch, -1) | |
| b = self.b.expand(batch, -1) | |
| c = self.c.expand(batch, -1) | |
| decay = torch.exp(delta.unsqueeze(-1) * stable_a) | |
| state = ( | |
| decay * state | |
| + delta.unsqueeze(-1) | |
| * item.unsqueeze(-1) | |
| * b.unsqueeze(1) | |
| ) | |
| output = (state * c.unsqueeze(1)).sum(-1) + self.d_skip * item | |
| outputs.append(output) | |
| update_strength.append(delta.mean(dim=1)) | |
| return torch.stack(outputs, dim=1), torch.stack(update_strength, dim=1) | |
| def forward( | |
| self, | |
| tokens: torch.Tensor, | |
| markers: torch.Tensor, | |
| *, | |
| return_trace: bool = False, | |
| ): | |
| embedded = self.embedding(tokens) + self.marker_projection( | |
| markers.unsqueeze(-1) | |
| ) | |
| projected, gate = self.input_projection(embedded).chunk(2, dim=-1) | |
| convolved = self.convolution(projected.transpose(1, 2))[ | |
| :, :, : tokens.shape[1] | |
| ].transpose(1, 2) | |
| values = F.silu(convolved) | |
| scanned, trace = self.scan(values) | |
| hidden = self.normalization(embedded + scanned * torch.sigmoid(gate)) | |
| logits = self.classifier(hidden[:, -1]) | |
| if return_trace: | |
| return logits, trace | |
| return logits | |
| class GRUControl(nn.Module): | |
| def __init__(self, vocab_size: int = 14, d_model: int = 32) -> None: | |
| super().__init__() | |
| self.embedding = nn.Embedding(vocab_size, d_model) | |
| self.marker_projection = nn.Linear(1, d_model, bias=False) | |
| self.gru = nn.GRU(d_model, d_model, batch_first=True) | |
| self.classifier = nn.Linear(d_model, 10) | |
| def forward(self, tokens: torch.Tensor, markers: torch.Tensor) -> torch.Tensor: | |
| embedded = self.embedding(tokens) + self.marker_projection( | |
| markers.unsqueeze(-1) | |
| ) | |
| hidden, _ = self.gru(embedded) | |
| return self.classifier(hidden[:, -1]) | |
| def parameter_count(module: nn.Module) -> int: | |
| return sum(parameter.numel() for parameter in module.parameters()) | |