File size: 4,144 Bytes
231a212
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
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())