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