Skip to main content

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 recurrence
  • SLiCELayer: a residual SLiCE layer with RMSNorm + GELU MLP by default
  • StackedSLiCE: stacks multiple SLiCELayers 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=False
  • block_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=False
  • block_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=True
  • block_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=1
  • ff_activation="glu" or ff_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 StackedSLiCE end-to-end
  • evaluates validation accuracy every --eval-every training 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

torch_slices-0.3.1.tar.gz (14.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

torch_slices-0.3.1-py3-none-any.whl (11.1 kB view details)

Uploaded Python 3

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

Hashes for torch_slices-0.3.1.tar.gz
Algorithm Hash digest
SHA256 7eccfdcf02b02c2895ddbe1e26c65ec2a70f4c6fd36fc715efa70e7987b0aaa0
MD5 58c74e0faa21ff3d0b20d825726529a9
BLAKE2b-256 b562979664760ac406d8b47a1f5ae13de46f030bce3d1387b46b201d63bc4c35

See more details on using hashes here.

File details

Details for the file torch_slices-0.3.1-py3-none-any.whl.

File metadata

File hashes

Hashes for torch_slices-0.3.1-py3-none-any.whl
Algorithm Hash digest
SHA256 0b2f964fee3f57482539fd3da854db38147a214beb5cc741a533796d4d60d115
MD5 9c58208c4d7e05131d297d0abb26f0cc
BLAKE2b-256 f9d7fa1e540ed9795edb662b5c23408f96cdcdd0af7b1451e5b9ae92dbad6d49

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page