micro-mamba-scan / model.py
ARotting's picture
Publish Interactive selective state-update explorer
2bc5689 verified
Raw
History Blame Contribute Delete
4.14 kB
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())