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 input 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}. $$

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 layer wrapping SLiCE with a post-activation stage (GLU or tanh)
  • 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).

Use SLiCELayer as a residual sequence layer

import torch
from slices import SLiCELayer

x = torch.randn(4, 256, 64)

layer = SLiCELayer(
    input_dim=64,
    block_size=4,
    diagonal_dense=True,
    dropout_rate=0.01,
    use_glu=True,
)

y = layer(x)  # (4, 256, 64)

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,
    use_glu=True,
)

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, data_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)

Benchmarking

To compare recurrent vs parallel throughput across sequence lengths and hidden dimensions:

python examples/benchmark_parallel_vs_recurrent.py

This script:

  • benchmarks all four SLiCE matrix modes (diagonal, block_diagonal, diagonal_dense, dense)
  • 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.2.0.tar.gz (12.8 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.2.0-py3-none-any.whl (9.5 kB view details)

Uploaded Python 3

File details

Details for the file torch_slices-0.2.0.tar.gz.

File metadata

  • Download URL: torch_slices-0.2.0.tar.gz
  • Upload date:
  • Size: 12.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.14.3

File hashes

Hashes for torch_slices-0.2.0.tar.gz
Algorithm Hash digest
SHA256 7f0c673eb3811ad468328c121b2f50d956db626c1286d8353ef86ebbcc052a53
MD5 d68711113b122162fdad99ca0a67e85f
BLAKE2b-256 015e7f85c822dedc8fc1ff5fb97eb08442f271fbd2cbcf26ef04c4ac6bb91aa2

See more details on using hashes here.

File details

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

File metadata

  • Download URL: torch_slices-0.2.0-py3-none-any.whl
  • Upload date:
  • Size: 9.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.14.3

File hashes

Hashes for torch_slices-0.2.0-py3-none-any.whl
Algorithm Hash digest
SHA256 00ec784f5b0ce5a3db30a50700ad29d4b6d71fc7d6a7736328dbfa31814d4d96
MD5 7080d99c33424c9cfcb8bf5c48a1a304
BLAKE2b-256 4fc8addce871b5710f8b32a0e95c467285b5c55258614973bda5efe6bf4d93e2

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