Structured Linear CDE (SLiCE) layers for sequence modelling in PyTorch
Project description
SLiCEs
slices is a small PyTorch package providing Structured Linear CDE (SLiCE) recurrences.
Mathematical form
Given an input sequence $x_i \in \mathbb{R}^D$ for $i=1,\dots,T$, a SLiCE computes hidden states $y_i \in \mathbb{R}^H$ via
$$ y_i = y_{i-1} + A(X_i)y_{i-1} + B(X_i), $$
where $A(\cdot): \mathbb{R}^D \rightarrow \mathbb{R}^{H \times H}$ and $B(\cdot): \mathbb{R}^D \rightarrow \mathbb{R}^H$ are learned linear maps, the initial state $y_0$ is either a function of $X_0$ or a learnt vector, and the driving path is augmented with an extra channel:
inc= a constant “increment” channel (all ones)
such that
$$ X_i = [inc_i, x_i] \in \mathbb{R}^{D+1}. $$
By default, SLiCE treats the provided sequence as path values and internally
computes first differences with torch.diff(..., prepend=zeros). Set
path_mode="increments" to treat the provided sequence as increments instead.
Installation
pip install torch-slices
Or install from source:
pip install git+https://github.com/datasig-ac-uk/slices.git
What's included
SLiCE: the core structured recurrenceSLiCELayer: a residual SLiCE layer with RMSNorm + GELU MLP by defaultStackedSLiCE: stacks multipleSLiCELayers with an embedding + output projection (supports tokens or continuous inputs)
SLiCE supports both:
- Recurrent execution (step-by-step update)
- Parallel chunked scan execution using
torch.associative_scan
Structured transition matrices
SLiCE supports different $A(X_i)$ structures:
1) Diagonal (elementwise update)
Set:
diagonal_dense=Falseblock_size=1
Then $A(X_i)$ is diagonal, which aligns with the approach used by Mamba (see here for more details).
2) Block-diagonal
Set:
diagonal_dense=Falseblock_size > 1
Then $A(X_i)$ is block-diagonal with blocks of shape (block_size × block_size).
3) Diagonal + dense tail block
Set:
diagonal_dense=Trueblock_size > 1
Then the first (hidden_dim - block_size) dimensions are diagonal, and the final block_size dimensions are updated via a dense (block_size × block_size) matrix.
Quickstart
Use SLiCE directly
import torch
from slices import SLiCE
x = torch.randn(8, 128, 32) # (batch, seq, input_dim)
layer = SLiCE(
input_dim=32,
hidden_dim=64,
block_size=4,
diagonal_dense=False,
bias=True,
use_parallel=True,
chunk_size=256,
)
y = layer(x) # (8, 128, 64)
print(y.shape)
Execution mode is configured via constructor arguments (use_parallel, chunk_size).
path_mode determines how SLiCE treats the sequence you pass in.
Use SLiCELayer as a residual sequence layer
import torch
from slices import SLiCELayer
x = torch.randn(4, 256, 64) # (batch, seq, input_dim)
layer = SLiCELayer(
input_dim=64,
block_size=4,
diagonal_dense=True,
)
y = layer(x) # (4, 256, 64)
SLiCELayer defaults to this structure:
- RMSNorm -> SLiCE -> residual
- RMSNorm -> Linear -> GELU -> Linear -> residual
Optional toggle for the post-norm wrapper:
prenorm=False
Optional toggles for the LayerNorm + single-stage wrapper include:
norm_type="layernorm"ff_style="single"ff_mult=1ff_activation="glu"orff_activation="tanh"dropout_position="output"
Stack layers for a full model
Token sequence mode (tokens=True)
Uses an nn.Embedding(data_dim, hidden_dim) front-end.
import torch
from slices import StackedSLiCE
batch, seq_len = 2, 128
vocab_size = 5000
x = torch.randint(0, vocab_size, (batch, seq_len))
model = StackedSLiCE(
num_layers=4,
data_dim=vocab_size,
hidden_dim=256,
label_dim=vocab_size,
tokens=True,
block_size=4,
diagonal_dense=False,
)
logits = model(x) # (batch, seq_len, vocab_size)
Continuous time-series mode (tokens=False)
Uses an nn.Linear(data_dim, hidden_dim) front-end.
import torch
from slices import StackedSLiCE
x = torch.randn(16, 100, 12) # (batch, seq, input_dim)
model = StackedSLiCE(
num_layers=3,
data_dim=12,
hidden_dim=64,
label_dim=10,
tokens=False,
block_size=4,
diagonal_dense=True,
)
y = model(x) # (16, 100, 10)
Training Example
examples/language_disambiguation.py is a simple example of training a
StackedSLiCE model end-to-end on a real dataset.
This example:
- uses a real dataset (wikimedia/wikipedia, English/French subset) for character-level language disambiguation
- trains a compact token-mode
StackedSLiCEend-to-end - evaluates validation accuracy every
--eval-everytraining steps - prints sample predictions so you can inspect model behaviour quickly
To run it, first install the example dependencies:
uv sync --group examples
and then execute the script:
uv run python examples/language_disambiguation.py
Useful flags:
--device cpu|cuda--train-steps 3000--eval-every 300--train-size 12000 --val-size 3000--max-seq-len 192
Benchmarking
To compare recurrent vs parallel throughput across sequence lengths and hidden dimensions:
uv run python examples/benchmark_parallel_vs_recurrent.py
This script:
- benchmarks all four SLiCE matrix modes (
diagonal,block_diagonal,diagonal_dense,dense) - uses the default value-path semantics unless
path_mode="increments"is set in code - prints timing/speedup tables
- saves a combined 3D plot to
examples/images/parallel_vs_recurrent_speedup_3d_all_modes.png
For plotting in development, install development dependencies (includes matplotlib):
uv sync --dev
Requirements
- Python ≥ 3.11
- PyTorch ≥ 2.8.0
- NumPy ≥ 2.4.1
License
MIT License. See LICENSE.
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file torch_slices-0.3.1.tar.gz.
File metadata
- Download URL: torch_slices-0.3.1.tar.gz
- Upload date:
- Size: 14.9 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: uv/0.7.8
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
7eccfdcf02b02c2895ddbe1e26c65ec2a70f4c6fd36fc715efa70e7987b0aaa0
|
|
| MD5 |
58c74e0faa21ff3d0b20d825726529a9
|
|
| BLAKE2b-256 |
b562979664760ac406d8b47a1f5ae13de46f030bce3d1387b46b201d63bc4c35
|
File details
Details for the file torch_slices-0.3.1-py3-none-any.whl.
File metadata
- Download URL: torch_slices-0.3.1-py3-none-any.whl
- Upload date:
- Size: 11.1 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: uv/0.7.8
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
0b2f964fee3f57482539fd3da854db38147a214beb5cc741a533796d4d60d115
|
|
| MD5 |
9c58208c4d7e05131d297d0abb26f0cc
|
|
| BLAKE2b-256 |
f9d7fa1e540ed9795edb662b5c23408f96cdcdd0af7b1451e5b9ae92dbad6d49
|