State Space Models (SSMs) Technical Notes¶
Quick Reference¶
- Definition: State Space Models (SSMs) are dynamic system models that represent sequences through evolving hidden states, offering efficient alternatives to transformers for long-range dependencies in machine learning tasks.
- Key Use Cases: Long-context language modeling, time series forecasting with irregular sampling, and audio processing where efficiency over long sequences is crucial.
- Prerequisites: Proficiency in linear algebra, basic differential equations, experience with PyTorch or similar ML frameworks, and understanding of sequence models like RNNs.
Table of Contents¶
- Introduction
- Core Concepts
- Implementation Details
- Real-World Applications
- Tools & Resources
- References
- Appendix
Introduction¶
What¶
State Space Models (SSMs) generalize linear dynamical systems to handle sequential data, parameterizing state transitions for efficient computation, particularly in modern variants like S4 or Mamba.
Why¶
SSMs provide O(N) time complexity for sequence modeling via parallelizable scans, enabling handling of million-length contexts where transformers scale quadratically.
Where¶
SSMs are used in natural language processing for long-document tasks, genomics for sequence analysis, and control systems with continuous-time dynamics.
Core Concepts¶
Fundamental Understanding¶
- Basic Principles: SSMs model sequences as continuous-time systems discretized for computation, with state evolution dx/dt = Ax + Bu and output y = Cx + Du.
- Key Components:
- Continuous Parameters: Matrices A, B, C for state dynamics, input coupling, and output mapping.
- Discretization: Conversion to discrete steps using methods like bilinear transform for stability.
- Selective Scan: Parallel computation of states with selective updates for context-awareness.
- Common Misconceptions:
- SSMs are linear only: Modern SSMs incorporate non-linearities via gating.
- Inefficient for training: Structured kernels enable fast FFT-based computation.
- Limited to time series: Excel in general sequence tasks like NLP.
Visual Architecture¶
graph TD
A["Continuous Input u(t)"] -->|Discretize| B[Discrete Input u_k]
B -->|B_bar| C[State x_k]
C -->|A_bar| C
C -->|C_bar| D[Output y_k]
D -->|Non-linear Activation| E["Next Layer/Input"]
- System Overview: Continuous signals are discretized, states evolve with barred matrices, outputs generated, often with non-linearities for deep models.
- Component Relationships: Discretization enables discrete computation, states integrate history, outputs feed downstream processing.
Implementation Details¶
Intermediate Patterns¶
import torch
import torch.nn as nn
class SimpleSSM(nn.Module):
def __init__(self, state_dim, input_dim, output_dim):
super().__init__()
self.A = nn.Parameter(torch.randn(state_dim, state_dim))
self.B = nn.Parameter(torch.randn(state_dim, input_dim))
self.C = nn.Parameter(torch.randn(output_dim, state_dim))
def discretize(self, dt=1.0):
# Bilinear discretization
I = torch.eye(self.A.size(0))
A_inv = torch.inverse(I - dt/2 * self.A)
A_bar = A_inv @ (I + dt/2 * self.A)
B_bar = A_inv @ (dt * self.B)
return A_bar, B_bar
def forward(self, u):
# u: (batch, seq_len, input_dim)
A_bar, B_bar = self.discretize()
states = []
x = torch.zeros(u.size(0), A_bar.size(0), device=u.device)
for t in range(u.size(1)):
x = A_bar @ x + B_bar @ u[:, t]
states.append(x)
states = torch.stack(states, dim=1) # (batch, seq_len, state_dim)
y = states @ self.C.mT # Assuming output_dim=1 for simplicity
return y.squeeze(-1)
# Usage
model = SimpleSSM(state_dim=4, input_dim=1, output_dim=1)
inputs = torch.randn(2, 10, 1) # batch=2, seq=10
outputs = model(inputs)
print(outputs.shape) # (2, 10)
Real-World Applications¶
Industry Examples¶
- Use Case: Long-context code generation in LLMs.
- Implementation Pattern: Use SSM layers instead of attention for efficiency.
- Success Metrics: Handle 1M+ tokens with constant memory.
Hands-On Project¶
- Project Goals: Implement an SSM for time series forecasting.
- Implementation Steps:
- Load a dataset (e.g., stock prices).
- Define and discretize SSM parameters.
- Train end-to-end with PyTorch.
- Evaluate on test sequences.
- Validation Methods: Compare MSE with RNN baselines.
Tools & Resources¶
Essential Tools¶
- Development Environment: PyTorch 2.0+.
- Key Frameworks: triton for custom kernels.
- Testing Tools: WandB for logging.
Learning Resources¶
- Documentation: PyTorch SSM examples.
- Tutorials: "Mamba: Linear-Time Sequence Modeling".
- Community Resources: Hugging Face forums.
References¶
- S4 Paper: "Efficiently Modeling Long Sequences".
- Mamba Paper: "Mamba: Linear-Time Sequence Modeling".
- Kalman Filter: Original 1960 paper.
Appendix¶
Glossary¶
- Discretization: Converting continuous to discrete-time.
- Scan: Cumulative operation over sequences.
- Selective SSM: Context-aware state updates.
Setup Guides¶
- Install PyTorch:
pip install torch.