File size: 5,094 Bytes
4ca4e4c | 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 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 | import torch
import torch.nn as nn
import torch.nn.functional as F
class SimpleSSM(nn.Module):
"""
Minimal discrete-time linear SSM:
x_{t+1} = A x_t + B u_t
y_t = C x_t
"""
def __init__(self, d_model, d_state):
super().__init__()
self.d_model = d_model
self.d_state = d_state
# Learnable parameters
self.A = nn.Parameter(torch.randn(d_model, d_state, d_state) * 0.01)
self.B = nn.Parameter(torch.randn(d_model, d_state) * 0.01)
self.C = nn.Parameter(torch.randn(d_model, d_state) * 0.01)
self.Cb = nn.Parameter(torch.zeros(d_model, d_state))
self.Delta = nn.Parameter(torch.randn(d_model) * 0.01)
self.Wo = nn.Linear(d_model, d_model, bias=False)
self.h0 = nn.Parameter(torch.zeros(d_state, d_model))
def forward(self, x):
"""
x: (batch, seq_len, d_model)
returns: (batch, seq_len, d_model)
"""
B, T, _ = x.shape
h = self.h0.unsqueeze(0).expand(B, -1, -1) # (batch, d_state, d_model)
outputs = []
for t in range(T):
xt = x[:, t] # (batch, d_model)
Delta = xt @ self.Delta
# edA = torch.exp(-torch.einsum('b,aij->baij', Delta, self.A))
edA = torch.einsum('b,aij->baij', 1-Delta, torch.eye(self.d_state, device=x.device).unsqueeze(0)) \
+ torch.einsum('b,aij->baij', Delta, self.A) # (batch, d_model, d_state, d_state)
B = xt @ self.B
# B = torch.einsum('ba,ak->bk', xt, self.B) # (batch, d_model, d_state)
C = xt @ self.C
# print(edA.shape, h.shape, xt.shape, Delta.shape, B.shape)
h = torch.einsum('bja,bakj->bka', h, edA) + torch.einsum('b,ba,bk->bka', Delta, xt, B) # state update
# h = h @ self.A.T + xt @ self.B.T # state update
# print(h.shape, self.Cb.shape)
yt = torch.einsum('bka,bk->ba', h, C) + torch.einsum('bka,ak->ba', h, self.Cb) # output
# yt = h @ self.C.T # output
outputs.append(yt)
# if Delta == 1:
# print("Next is number!")
# print(torch.max(torch.abs(xt - h)))
# print(h[0,0])
# print(torch.max(edA), torch.max(torch.einsum('baj,akj->bak', h, edA)), torch.max(torch.einsum('b,bk,ba->bak', Delta, xt, B)), torch.max(yt), torch.max(torch.abs(h)))
# print(yt[0, 25:30])
# print(torch.stack(outputs, dim=1)[0, -1, 25:30])
return self.Wo(torch.stack(outputs, dim=1))
class CausalSelfAttention(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_head = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
self.out = nn.Linear(d_model, d_model, bias=False)
def forward(self, x):
"""
x: (batch, seq_len, d_model)
"""
B, T, D = x.shape
qkv = self.qkv(x)
q, k, v = qkv.chunk(3, dim=-1)
# reshape for heads
q = q.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
k = k.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
v = v.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
# scaled dot-product attention
scores = (q @ k.transpose(-2, -1)) / (self.d_head ** 0.5)
# causal mask
mask = torch.tril(torch.ones(T, T, device=x.device))
# mask = torch.tril(torch.ones(T, T, device=x.device)).T
scores = scores.masked_fill(mask == 0, float('-inf'))
attn = F.softmax(scores, dim=-1)
out = attn @ v
out = out.transpose(1, 2).contiguous().view(B, T, D)
return self.out(out), attn
class SSMTransformerBlock(nn.Module):
def __init__(self, d_model, d_state, n_heads, d_ff, layer_type):
super().__init__()
self.layer_type = layer_type
if layer_type == 'SSM':
self.layer = SimpleSSM(d_model, d_state)
elif layer_type == 'TF':
self.layer = CausalSelfAttention(d_model, n_heads)
def forward(self, x):
out = self.layer(x)
if self.layer_type == 'SSM':
x = x + out
elif self.layer_type == 'TF':
x = x + out[0]
return x
class SSMTransformer(nn.Module):
def __init__(self, num_vocab, d_model, d_state, n_heads, d_ff, layers):
super().__init__()
self.embedding = nn.Embedding(num_vocab, d_model)
self.pos_emb = nn.Parameter(torch.randn(1, 100, d_model))
self.layers = nn.ModuleList([
SSMTransformerBlock(d_model, d_state, n_heads, d_ff, layer)
for layer in layers
])
self.lm_head = nn.Linear(d_model, num_vocab)
def forward(self, x):
x = self.embedding(x.long()) + self.pos_emb
for layer in self.layers:
x = layer(x)
return self.lm_head(x)
|