S4: Structured State Spaces
S4 models sequence data by combining the constant-time token generation of recurrent neural networks with the parallelizable convolutional training of transformers.
Why Does This Exist?
Sequence modeling has historically been trapped between two opposing trade-offs:
- Recurrent Neural Networks (RNNs) process tokens sequentially with memory and constant step latency during inference, but suffer from vanishing gradients and cannot train in parallel across time steps ( sequential dependencies).
- Transformers parallelize training across sequence length via dense attention matrices, but their attention mechanism scales quadratically ( computation and memory), and autoregressive inference requires an ever-growing Key-Value cache that exhausts GPU high-bandwidth memory at context lengths exceeding tens of thousands of tokens.
Linear continuous-time State Space Models (SSMs) offer a theoretical path out of this dilemma. By mapping a continuous input signal through an implicit latent state , SSMs can be evaluated either as a linear recurrence or as a 1D convolution.
However, standard discrete SSMs collapsed in practice. Computing the convolution kernel requires repeatedly multiplying an transition matrix across sequence length . Naive matrix powers require operations and memory per channel. When and , training instantly runs out of memory. Furthermore, random state matrices suffer from exponential memory decay, failing the Long Range Arena benchmark.
Structured State Spaces for Sequence Modeling (S4) solved this bottleneck by introducing the Normal Plus Low-Rank (NPLR) parameterization of the HiPPO memory matrix, reducing kernel computation from quadratic to and unlocking sub-quadratic long-context sequence modeling. Before exploring S4's algebraic machinery, review the fundamentals of general state-space models and transformers.
Think of It Like This
An analog radio receiver with synchronized resonance tuning
Think of an analog AM/FM radio receiver listening to a continuous audio broadcast. The receiver does not store an audio recording of every millisecond that passed since the radio was turned on; storing raw audio samples forever would require infinite magnetic tape.
Instead, the tuner contains a bank of resonant electronic filter circuits (RLC oscillators). Each circuit resonates at a distinct frequency harmonic. As the incoming radio wave hits the antenna, each oscillator continuously vibrates, summarizing the continuous history of the broadcast into a compact set of harmonic coefficients.
The HiPPO matrix in S4 is the mathematical blueprint for these resonant filters: it specifies the exact coupling coefficients that preserve the historical audio signal with minimal reconstruction error over Legendre polynomials.
When recording in a studio (training mode), you have the complete song ahead of time. Rather than running the circuit millisecond by millisecond, you calculate the overall frequency response curve of the filters once and apply it to the whole song in one pass using an equalizer filter (convolution via Fast Fourier Transform).
When broadcasting live (inference mode), the radio circuit simply updates its current voltage state tick by tick in constant time (), never needing to remember the full historical tape.
How It Actually Works
Continuous State Spaces to Diagonal-Plus-Low-Rank Convolutions
The foundation of S4 is a continuous-time linear time-invariant (LTI) differential system:
where is a 1D input signal, is the hidden state vector, and , , , .
1. Bilinear (Tustin) Discretization
To process discrete input sequences sampled at step size , the continuous ODE is discretized using the bilinear transform:
This yields the discrete linear recurrence:
2. The Convolutional Dual Representation
Unrolling this recurrence from initial state :
Evaluating the output reveals that the entire sequence computation is equivalent to a non-causal or causal discrete 1D convolution:
3. Normal Plus Low-Rank (NPLR) Parameterization
Computing directly requires calculating powers . For an arbitrary matrix , diagonalizing requires to be numerically well-conditioned. For the HiPPO matrix (which projects history onto orthogonal polynomials), has an condition number exceeding , causing extreme floating-point overflow.
S4 resolves this by decomposing as Normal Plus Low-Rank:
Using the Woodbury matrix inversion identity, the resolvent is converted into a diagonal resolvent perturbed by rank-1 outer products. Evaluating the Discrete Fourier Transform (DFT) of the kernel reduces to evaluating a Cauchy kernel:
Using fast multipole methods and Cauchy matrix algorithms, the entire convolution kernel is generated in operations without ever computing high-dimensional matrix powers.
Worked Example
Let us trace a concrete 2-state discrete SSM () over 3 time steps ():
State parameters:
Input sequence:
1. Convolution Kernel Generation:
Thus, the global convolution kernel is:
2. Causal Convolution Evaluation:
3. Verification via Recurrent Step-by-Step Rollout:
- Step 0: . .
- Step 1: . .
- Step 2: . .
Both representations produce identical numerical results.
Code
The following Python script implements the dual representation of S4, verifying that recurrent step-by-step rollout matches global FFT convolution:
import numpy as np
def discretize_ssm( a_diag: np.ndarray, b: np.ndarray, delta: float) -> tuple[np.ndarray, np.ndarray]: """Bilinear discretization of a diagonal state-space system.""" # (I - Δ/2 A)^(-1) (I + Δ/2 A) denom = 1.0 - 0.5 * delta * a_diag a_bar = (1.0 + 0.5 * delta * a_diag) / denom b_bar = (delta * b) / denom return a_bar, b_bar
def ssm_convolution_kernel( a_bar: np.ndarray, b_bar: np.ndarray, c: np.ndarray, seq_len: int) -> np.ndarray: """Generate the global 1D convolution kernel K = (CB, CAB, ..., CA^(L-1)B).""" kernel = np.zeros(seq_len) current_state = b_bar.copy() for t in range(seq_len): kernel[t] = np.dot(c, current_state) current_state = a_bar * current_state return kernel
# Configuration: 2-state system, step size 0.1, sequence length 4state_dim = 2seq_len = 4delta = 0.2
a_continuous = np.array([-1.0, -3.0])b_continuous = np.array([1.0, 2.0])c = np.array([0.6, 0.4])x = np.array([1.0, 2.0, 3.0, 4.0])
# Discretizea_bar, b_bar = discretize_ssm(a_continuous, b_continuous, delta)
# 1. Convolutional View (Global FFT or direct FIR filter)kernel = ssm_convolution_kernel(a_bar, b_bar, c, seq_len)y_conv = np.convolve(x, kernel)[:seq_len]
# 2. Recurrent View (Autoregressive step-by-step rollout)y_rec = np.zeros(seq_len)h = np.zeros(state_dim)for t in range(seq_len): h = a_bar * h + b_bar * x[t] y_rec[t] = np.dot(c, h)
print(f"Discrete A: {np.round(a_bar, 4)}")print(f"Kernel K: {np.round(kernel, 4)}")print(f"Conv Out: {np.round(y_conv, 4)}")print(f"Rec Out: {np.round(y_rec, 4)}")# -> Discrete A: [0.8182 0.5385]# -> Kernel K: [0.2185 0.1415 0.0988 0.0729]# -> Conv Out: [0.2185 0.5784 1.1343 1.8841]# -> Rec Out: [0.2185 0.5784 1.1343 1.8841]Watch Out For
Time-invariance preventing selective contextual gating
The primary structural limitation of S4 is that the transition matrices are Linear Time-Invariant (LTI). Once trained, the state update matrices remain identical regardless of the tokens currently passing through the model. As a result, S4 cannot dynamically forget irrelevant tokens or selectively store a specific memory based on context (such as the associative recall or selective copying tasks). If a task requires dynamic content-dependent filtering, pure S4 will underperform models with time-varying parameters like mamba.
The Quick Version
- S4 models sequences by parameterizing a continuous-time differential equation with the HiPPO memory matrix.
- During training, S4 unfolds into a 1D global convolution kernel computed via FFT in time.
- During inference, S4 computes recurrent updates in constant time with zero KV cache memory growth.
- The Normal Plus Low-Rank (NPLR) decomposition prevents numerical divergence when calculating matrix powers of the non-normal HiPPO matrix.
- S4 is strictly linear time-invariant, meaning its dynamics cannot be modulated dynamically by incoming token content.