State Space Models: Mamba и альтернатива Transformer

opensourceaimambastate-spacessmsequence-modelingit
← Back to Blog

Введение: 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 на многих задачах, и продолжает развиваться.


См. также