State Space Models: Mamba и альтернатива Transformer
opensourceaimambastate-spacessmsequence-modelingit
Введение: State Space Models
Что такое SSM?
State Space Models (SSM) = нейросети из систем управления
Классическая теория систем:
x'(t) = A·x(t) + B·u(t) # состояние
y(t) = C·x(t) + D·u(t) # выход
Нейросетевой SSM:
x'(t) = f(A, B, u(t)) # hidden state
y(t) = f(C, D, u(t)) # output
Преимущества:
- Линейная сложность O(n)
- Постоянная память O(1)
- Бесконечный контекст
Факт: Mamba (2024) показал результаты, сравнимые с Transformer, но в 2-4x быстрее при инференсе.
S4: Structured State Space
Базовая модель S4
S4 Model:
1. Continuous-time SSM:
x'(t) = A·x(t) + B·u(t)
y(t) = C·x(t)
2. Discretization (Zero-Order Hold):
x[k+1] = A̅·x[k] + B̅·u[k]
y[k] = C·x[k]
где:
A̅ = exp(Δ·A)
B̅ = (Δ·A)⁻¹ · (exp(Δ·A) - I) · Δ·B
3. Convolution form:
y = K * u
где K = [C·B, C·A̅·B, C·A̅²·B, ...]
class S4Layer(torch.nn.Module):
"""Structured S4 layer"""
def __init__(self, d_model, d_state=16, d_rank=16, L=128):
super().__init__()
self.d_model = d_model
self.d_state = d_state
self.d_rank = d_rank
# Projection
self.B = torch.nn.Linear(d_rank, d_state)
self.C = torch.nn.Linear(d_rank, d_state)
# Diagonal parameterization (HippSZ)
self.real_A = torch.nn.Parameter(torch.randn(1, d_model, d_rank))
self.imag_A = torch.nn.Parameter(torch.randn(1, d_model, d_rank))
# Delta (step size)
self.Delta = torch.nn.Parameter(torch.ones(1, d_model, 1))
# Output projection
self.out_proj = torch.nn.Linear(d_model, d_model)
# S
self.S = torch.nn.Parameter(torch.randn(d_model, d_state))
def discretize(self):
"""Discretize continuous parameters"""
Delta = torch.nn.functional.softplus(self.Delta) # Δ > 0
# A̅ = exp(Δ·A)
real_A = self.real_A # (1, d_model, d_rank)
imag_A = self.imag_A
# Complex eigenvalues: λ = Δ · (real_A + i·imag_A)
Lambda = Delta * (real_A + 1j * imag_A) # (1, d_model, d_rank)
# exp(Λ)
exp_Lambda = torch.exp(Lambda) # (1, d_model, d_rank)
# B̅ = Δ · B / (1 - exp(Λ)) (diagonal approx)
B_cont = self.B.weight # (d_state, d_rank)
B_disc = Delta * B_cont / (1 - exp_Lambda.real + 1e-3)
return exp_Lambda, B_disc
def forward(self, x):
"""Forward pass through S4 layer"""
B, T, D = x.size()
# Discretize
exp_Lambda, B_disc = self.discretize()
# Compute kernel: K[t] = C · A^t · B
C = self.C.weight # (d_state, d_rank)
S = self.S # (d_model, d_state)
# S4 convolution via FFT
K = self._compute_kernel(exp_Lambda, B_disc, C)
# FFT-based convolution
x_fft = torch.fft.rfft(x, dim=1) # (B, T, D)
K_fft = torch.fft.rfft(K, dim=1, signal_dim=(1,))
y_fft = x_fft * K_fft
y = torch.fft.irfft(y_fft, n=T, dim=1)
# Add skip connection and output projection
y = y + x # skip
y = self.out_proj(y)
return y
def _compute_kernel(self, exp_Lambda, B_disc, C, kernel_len=256):
"""Compute S4 kernel"""
d_model = exp_Lambda.size(1)
d_rank = exp_Lambda.size(2)
t = torch.arange(kernel_len, device=exp_Lambda.device)
t_exp = t.unsqueeze(0).unsqueeze(0) # (1, 1, L)
K = (exp_Lambda ** t_exp) # (1, d_model, d_rank, L)
K = torch.einsum('bdrk,sk->bdrs', K, B_disc) # (1, d_model, d_state, L)
K = torch.einsum('bdrs,sk->bdrk', K, C) # (1, d_model, d_rank, L)
K = K.sum(dim=2) # (1, d_model, L)
return K
Mamba: Selective State Space
Ключевое улучшение: Selective Mechanism
Mamba vs S4:
S4:
A, B — фиксированные параметры
Один и тот же для всех входов
Mamba:
A, B, C — зависят от входа!
A(u) = f_A(u), B(u) = f_B(u), C(u) = f_C(u)
→ Модель "выбирает" что запомнить, что забыть
Selective SSM:
x'(t) = A(u(t))·x(t) + B(u(t))·u(t)
y(t) = C(u(t))·x(t)
class MambaBlock(torch.nn.Module):
"""Mamba selective SSM block"""
def __init__(self, d_model=512, d_state=16, d_conv=4, expand=2):
super().__init__()
self.d_model = d_model
self.d_inner = d_model * expand
self.d_state = d_state
self.d_conv = d_conv
# Input projection
self.in_proj = torch.nn.Linear(d_model, self.d_inner * 2)
# Convolution (local context)
self.conv1d = torch.nn.Conv1d(
in_channels=self.d_inner,
out_channels=self.d_inner,
kernel_size=d_conv,
padding=d_conv - 1,
groups=self.d_inner
)
# Activation
self.act = torch.nn.SiLU()
# SSM parameters (input-dependent!)
self.x_proj = torch.nn.Linear(self.d_inner, d_state * 5)
# Discrete SSM parameters
self.dt_proj = torch.nn.Linear(self.d_inner, self.d_inner)
# A parameter: log(a) ∈ [-log(1/2), log(1/2)]
self.A_log = torch.nn.Parameter(torch.log(torch.ones(self.d_inner, d_state)))
self.D = torch.nn.Parameter(torch.ones(self.d_inner))
# Output projection
self.out_proj = torch.nn.Linear(self.d_inner, d_model)
def forward(self, x):
"""Forward pass through Mamba block"""
B, L, D = x.size()
# 1. Project input
x_and_res = self.in_proj(x) # (B, L, 2·d_inner)
x, res = x_and_res.split([self.d_inner, self.d_inner], dim=-1)
# 2. Convolution (local context)
x = x.transpose(1, 2) # (B, d_inner, L)
x = self.conv1d(x)[:, :, :L] # truncate
x = x.transpose(1, 2) # (B, L, d_inner)
x = self.act(x)
# 3. SSM encoding
y = self.ssm_encode(x)
# 4. Apply residual and output
y = y * self.act(res)
y = self.out_proj(y)
return y
def ssm_encode(self, x):
"""Selective SSM encoding"""
B, L, D = x.size()
# Compute input-dependent parameters
dt, B, C = self.x_proj(x.split(
[self.d_inner, self.d_state, self.d_state], dim=-1
))
# Delta (step size) — input-dependent!
dt = torch.nn.functional.softplus(self.dt_proj(dt)) # (B, L, D)
# Discretize
A = torch.exp(self.A_log) # (D, d_state)
dA = torch.einsum('bl,dn->bldn', dt, A) # (B, L, D, d_state)
exp_dA = torch.exp(dA) # (B, L, D, d_state)
# B, C are input-dependent
# Scan through sequence
y = self._scan(x, B, C, exp_dA)
return y
def _scan(self, x, B, C, exp_dA):
"""Sequential scan through SSM"""
B, L, D, d_state = exp_dA.size()
y = torch.zeros(B, L, D, device=x.device)
h = torch.zeros(B, D, d_state, device=x.device)
for l in range(L):
# h[t] = A_t · h[t-1] + B_t · x[t]
h = exp_dA[:, l] * h + B[:, l].unsqueeze(-1) * x[:, l].unsqueeze(-1)
y[:, l] = (C[:, l].unsqueeze(-1) * h).sum(dim=-1) + self.D * x[:, l]
return y
SSM Scan Algorithm
SSM Scan: последовательное вычисление
Рекуррентная форма:
h[t] = A̅_t · h[t-1] + B̅_t · x[t]
y[t] = C_t · h[t] + D · x[t]
Параллельная форма (PGM):
h[t] = (Π_{i=1}^{t} A̅_i) · h[0] + Σ_{j=1}^{t} (Π_{i=j+1}^{t} A̅_i) · B̅_j · x[j]
Prefix sums позволяют параллелизовать!
Mamba Architecture
Полная архитектура Mamba
Mamba Model:
Input → Embedding
↓
[Stack of Mamba Blocks]
├─ Selective SSM
├─ Feed Forward (Gated)
└─ LayerNorm
↓
[Residual Connection]
↓
Output Head
class MambaLayer(torch.nn.Module):
"""Single Mamba layer with residual and norm"""
def __init__(self, d_model, d_state, d_conv, expand):
super().__init__()
self.norm = torch.nn.LayerNorm(d_model)
self.mamba = MambaBlock(d_model, d_state, d_conv, expand)
def forward(self, x):
return x + self.mamba(self.norm(x))
class Mamba(torch.nn.Module):
"""Full Mamba model"""
def __init__(self, vocab_size=50000, d_model=1024, d_state=16,
d_conv=4, expand=2, num_layers=24):
super().__init__()
# Embedding
self.embedding = torch.nn.Embedding(vocab_size, d_model)
# Layers
self.layers = torch.nn.ModuleList([
MambaLayer(d_model, d_state, d_conv, expand)
for _ in range(num_layers)
])
# Output head
self.ln_final = torch.nn.LayerNorm(d_model)
self.lm_head = torch.nn.Linear(d_model, vocab_size, bias=False)
# Weight tying
self.embedding.weight = self.lm_head.weight
def forward(self, x):
B, T = x.size()
# Embed
x = self.embedding(x) # (B, T, d_model)
# Pass through layers
for layer in self.layers:
x = layer(x)
# Final norm and head
x = self.ln_final(x)
logits = self.lm_head(x)
return logits
def generate(self, x, max_tokens=100, temperature=0.8):
"""Text generation"""
for _ in range(max_tokens):
# Forward
logits = self(x)
logits = logits[:, -1, :] / temperature
# Softmax
probs = torch.softmax(logits, dim=-1)
# Sample
x_next = torch.multinomial(probs, 1)
# Append
x = torch.cat([x, x_next], dim=1)
return x
Сравнение: Transformer vs Mamba
┌─────────────────┬──────────────────┬──────────────────┐
│ Параметр │ Transformer │ Mamba │
├─────────────────┼──────────────────┼──────────────────┤
│ Complexity │ O(n²) │ O(n) │
│ Memory │ O(n²) │ O(n) │
│ Context length │ Ограничен │ Практически бес- │
│ │ │ конечен │
│ Parallel train │ Отлично │ Хорошо (scan) │
│ Sequential infer│ Хорошо │ Отлично (RNN) │
│ Hardware util. │ Высокая │ Средняя │
│ Implementation │ Зрелая │ Новая │
└─────────────────┴──────────────────┴──────────────────┘
SSM vs RNN vs Transformer
Архитектурное сравнение:
RNN (LSTM/GRU):
h[t] = f(h[t-1], x[t])
- O(n) complexity
- Последовательный инференс
- Проблема затухания градиента
SSM (Mamba):
h[t] = A_t · h[t-1] + B_t · x[t]
- O(n) complexity
- Selective mechanism
- Структурированные параметры
Transformer (Attention):
Attention(Q,K,V) = softmax(QK^T/√d) · V
- O(n²) complexity
- Полный контекст за один шаг
- Параллельный тренинг
Практическое применение Mamba
Установка и использование
# Установка Mamba
pip install mamba-ssm
# Компиляция CUDA kernels
cd mamba/ssm
python setup.py install
import torch
from mamba_ssm import Mamba
# Создание модели Mamba
model = Mamba(
d_model=1024, # размерность модели
d_state=16, # размерность состояния
d_conv=4, # размерность свёртки
expand=2, # коэффициент расширения
num_layers=24 # количество слоёв
)
# Forward pass
x = torch.randint(0, 50000, (2, 512)) # (batch, seq_len)
output = model(x) # (2, 512, 1024)
Mamba в гибридных архитектурах
Hybrid Transformer-Mamba:
[Transformer Blocks] ← хороший параллельный тренинг
↓
[Mamba Blocks] ← эффективный инференс
↓
[Transformer Blocks] ← внимание к ключевым токенам
↓
[Mamba Blocks] ← линейная сложность
Преимущества:
- Лучшая производительность
- Эффективный инференс
- Гибкость архитектуры
Заключение
State Space Models, особенно Mamba, представляют собой перспективную альтернативу Transformer для обработки последовательностей. Ключевые преимущества:
- Линейная сложность O(n) вместо O(n²)
- Selective mechanism для управления памятью
- Бесконечный контекст без деградации качества
- Быстрый инференс в режиме RNN
Mamba уже показывает результаты, сравнимые с Transformer на многих задачах, и продолжает развиваться.