Skip to main content

subquadratic_ops_torch

Introduction

subquadratic_ops_torch provides CUDA kernels for subquadratic operations, e.g. long and short convolution. It contains PyTorch bindings to optimized kernels.

Installation

Please install using pip install subquadratic-ops-torch-cu[12,13]

Documentation

For detailed usage information of the kernels, please refer to the docstrings in their respective functions.

Usage

You can import the library from python:

import subquadratic_ops_torch as subq

Kernels are primarily exposed as function calls underlying torch.ops, which also provide a lower-level interface as torch.library operators. This allows you to export models using these operations via torch.export and run inference on them using TensorRT.

Support and Feedback

Please contact the developers for any issues you might encounter.

Requirements

  • CUDA-compatible NVIDIA GPU (Ampere+)
  • CUDA Toolkit 12.0 or higher
  • Python 3.11-3.13

Modules

B2B CausalConv1d

Operation

Back-to-back causal conv1d for the Striped Hyena 2 architecture used in the Evo2 model. The operation is performed in a causal manner, meaning each position only attends to previous positions in the sequence. In code terms,

in_dim = 8192

width_proj = 8
width_mixer = 128

dtype = torch.float32

class Conv1DModel(nn.Module):
    def __init__(self, in_dim, width, dtype, skip_bias=False):
        super(Conv1DModel, self).__init__()

        self.conv = nn.Conv1d(
            in_dim,
            in_dim,
            width,
            padding=width - 1,
            groups=in_dim,
            bias=False,
            dtype=dtype,
            device="cuda:0",
        )
        self.width = width
        self.weight = self.conv.weight.reshape(-1, width)
        if skip_bias:
            self.skip_bias = nn.Parameter(torch.zeros(in_dim, dtype=dtype, device="cuda:0").reshape(1, -1, 1))
        else:
            self.skip_bias = None

    def forward(self, x):
        seqlen = x.shape[-1]
        out = self.conv(x)
        return out[..., :seqlen]

def model(x, conv1d_proj, conv1d_mixer):
    xv = conv1d_proj(x)
    z = xv[:,1::3, :] * xv[:, 2::3, :]
    y = conv1d_mixer(z) + conv1d_mixer.skip_bias * z
    return y * xv[:, ::3, :]
x = torch.randn(batch_size, 3*in_dim, seq_dim)
conv1d_proj = Conv1DModel(3*in_dim, width_proj)
conv1d_mixer = Conv1DModel(in_dim, width_mixer, True)

y = model(x, conv1d_proj, conv1d_mixer)

is equivalent to,

weight_proj = torch.randn(3*in_dim, width_proj).to(dtype)
weight_mixer = torch.randn(in_dim, width_mixer).to(dtype)
skip_bias = torch.randn(in_dim).to(dtype)

b2b_causal_conv1d(x, weight_proj, weight_mixer, skip_bias)

Supported Kernel Sizes for B2B Causal Conv1d

Kernel Type Supported Sizes
Projection 2, 3, 4, 8, 16, 32
Mixer 2, 3, 4, 5, 6, 7, 8, 16, 32, 64, 128, 256

CausalConv1d

Causal conv1d: the convolution operation is performed in a causal manner, meaning each position only attends to previous positions in the sequence. In code terms,

in_dim = 8192
width = 8
dtype = torch.float32
model = nn.Conv1d(
            in_dim,
            in_dim,
            width,
            padding=width - 1,
            groups=in_dim,
            bias=False,
            dtype=dtype,
            device="cuda:0",
        )
y = model(x)

is equivalent to,

weight = torch.randn((in_dim, width))
causal_conv1d(x, weight)

Supported Kernel Sizes for Causal Conv1d

Kernel Type Supported Sizes Channel Last
CausalConv1d <= 256 False
CausalConv1d <= 128 (64 fp64) True

FFT Conv1d

Non-causal 1D convolution using real FFT. Supports sequences up to FFT size 8192.

from subquadratic_ops_torch.fft_conv1d import fft_conv1d

batch_size, dim, seq_len, filter_dim = 64, 128, 512, 1024
x = torch.randn(batch_size, dim, seq_len, device="cuda")
weight = torch.randn(dim, filter_dim, device="cuda")
y = fft_conv1d(x, weight)  # shape: (64, 128, 512)

FFT CausalConv1d

FFT Causal Conv1d: the convolution operation is performed in a causal manner, meaning each position only attends to previous positions in the sequence. It uses real FFT and IFFT instead of direct summation for convolution. In code terms,

in_dim = 8192
width = 1024
dtype = torch.float32
weight = torch.randn(1, in_dim, width)
def model(x, w):
    fft_size = x.shape[-1] * 2
    xf = torch.fft.rfft(x, n=fft_size, dim=-1)
    wf = torch.fft.rfft(w, n=fft_size, dim=-1)
    return torch.fft.irfft(xf*wf, n=fft_size, dim=-1)[..., :x.shape[-1]]

y = model(x, weight)

is equivalent to,

weight = torch.randn((in_dim, width))
y = fft_causal_conv1d(x, weight)

Release files for subquadratic-ops-torch-cu12 0.2.1

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Built distributions (wheels)

Table of built distributions (wheels) for subquadratic-ops-torch-cu12 0.2.1
File
subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_34_aarch64.whl CPython 3.14 CPython 3.14 Linux glibc 2.34+ ARM64 Details
subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl CPython 3.14 CPython 3.14 Linux glibc 2.28+ x86-64, Linux glibc 2.24+ x86-64 Details
subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_34_aarch64.whl CPython 3.13 CPython 3.13 Linux glibc 2.34+ ARM64 Details
subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl CPython 3.13 CPython 3.13 Linux glibc 2.28+ x86-64, Linux glibc 2.24+ x86-64 Details
subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_34_aarch64.whl CPython 3.12 CPython 3.12 Linux glibc 2.34+ ARM64 Details
subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl CPython 3.12 CPython 3.12 Linux glibc 2.28+ x86-64, Linux glibc 2.24+ x86-64 Details
subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_34_aarch64.whl CPython 3.11 CPython 3.11 Linux glibc 2.34+ ARM64 Details
subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl CPython 3.11 CPython 3.11 Linux glibc 2.28+ x86-64, Linux glibc 2.24+ x86-64 Details

Total release size: 3.3 GB

Release files / subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_34_aarch64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_34_aarch64.whl
Size 402.5 MB
Tags CPython 3.14 Linux glibc 2.34+ ARM64
SHA-256 checksum
How to use checksums
fd00c987026c832be084c62a1ffd11be89c8750abc8c2007b5be847b485b1a94
BLAKE2b-256 checksum
How to use checksums
2b042400141169f249b1b5ab151790b3c93bf6e621048cf54c63cd38618e4358
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Size 419.4 MB
Tags CPython 3.14 Linux glibc 2.24+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
019ca9dcf935217c6f39814e68c5565f66a271c12789761687eee236f201bd9b
BLAKE2b-256 checksum
How to use checksums
01be7d671f11de953409371b2909ff4d56cee1690dba5d538a26d38c9eaa02f3
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_34_aarch64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_34_aarch64.whl
Size 402.5 MB
Tags CPython 3.13 Linux glibc 2.34+ ARM64
SHA-256 checksum
How to use checksums
319b0fc410e3faa966d8ab05608aa4a9486fccda1d8ab4e83346421fa53e8bad
BLAKE2b-256 checksum
How to use checksums
372089e9ef5b1e28904a864e17409e0f8df7fd28b73acae8a94b88e301f9c4a3
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Size 419.3 MB
Tags CPython 3.13 Linux glibc 2.24+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
ee90a930835b2eee5403532731e54d63659579ca3b0bf09fbd47dc9a6ba76d2b
BLAKE2b-256 checksum
How to use checksums
293f5f2927773316c09d45d2771725b6414a5e0bdc24db902ad51b7c9a097dc1
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_34_aarch64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_34_aarch64.whl
Size 402.5 MB
Tags CPython 3.12 Linux glibc 2.34+ ARM64
SHA-256 checksum
How to use checksums
9d9f68cd797dd19221dd5bab081dc73b6b025d1827f3405d033841c5710ee6cd
BLAKE2b-256 checksum
How to use checksums
793bd4ae0b84b0ad6a59f413d4258344fc2319d6b4607b69a69a428888a258c4
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Size 419.3 MB
Tags CPython 3.12 Linux glibc 2.24+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
df475180c89b6706fefa65d60395daee7e9f09aa7115cb6e9e161908b39efedd
BLAKE2b-256 checksum
How to use checksums
7a13a36682b43c615981788aab27f8ec1a1f1c7fff9ab360cb521ec7bc65db01
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_34_aarch64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_34_aarch64.whl
Size 402.5 MB
Tags CPython 3.11 Linux glibc 2.34+ ARM64
SHA-256 checksum
How to use checksums
b0b1d2ee705e1aea5882f9474f3f4f74d72823f4190187b9c73cdca05eff330c
BLAKE2b-256 checksum
How to use checksums
7ab4af801be2dbef9100bb895d08f5b6769c44e55c2027b010c04be7eaa3f162
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release files / subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl

Download URL subquadratic_ops_torch_cu12-0.2.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Size 419.3 MB
Tags CPython 3.11 Linux glibc 2.24+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
d2d0eed323d27879dda702159fd37a6ac42f2669e7c6caf5c11188072bf32ebc
BLAKE2b-256 checksum
How to use checksums
5c91d05287435d516f8235153a0131c1ee71a3bcd77c2ae514ca12e064559073
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.14.6

Release history Release notifications | RSS feed

This release

0.2.1 This release

8 release files

0.2.0

1 release file

0.1.1

6 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page