subquadratic_ops_torch
CUDA kernels and PyTorch bindings for subquadratic sequence operations — causal and non-causal convolutions, direct and FFT-based, in short, long and variable-length (packed) forms.
Full API reference, including every argument and shape: documentation.
Install
pip install subquadratic-ops-torch-cu13
Quick start
import torch
from subquadratic_ops_torch.causal_conv1d import causal_conv1d
x = torch.randn(2, 8192, 4096, device="cuda") # (batch, dim, seq_len)
weight = torch.randn(8192, 8, device="cuda") # (dim, kernel_size)
y = causal_conv1d(x, weight) # (2, 8192, 4096)
Every operator is depthwise and causal unless stated otherwise, takes (batch, dim, seq_len)
input and (dim, kernel_size) weights, and is differentiable. Most register torch.library
operators under torch.ops.subquadratic_ops_torch, so models using them can be traced with
torch.export and run under TensorRT; b2b_causal_conv1d and channel-first causal_conv1d are
the exceptions (see below).
Operators
| Operator | Purpose | Envelope |
|---|---|---|
causal_conv1d |
direct depthwise causal conv, optional bias and activation | kernel 2–256 (channel-first, cuDNN-served, fp32/fp16/bf16); kernel ≤ 128 channel-last (64 in fp64) |
b2b_causal_conv1d |
fused back-to-back projection → gate → mixer → gate, for Striped Hyena 2 / Evo2 | projection 2–32, mixer 2–256, both contiguous; fp32/fp16/bf16 only |
fft_causal_conv1d |
causal conv via rFFT, for long filters | dispatches single-shot, blocked or decomposed internally; FFT size ≤ 8192 (4096 in fp64) |
fft_causal_conv1d_packed |
causal FFT conv over a variable-length batch, cu_seqlens-addressed |
sequence ≤ 16384; the 32768 transform needs SM90+; fp32/fp16/bf16 |
fft_conv1d |
non-causal 1D conv via rFFT | FFT size ≤ 8192 |
fft_conv2d, fused_fft_conv2d |
2D FFT convolution, unfused and fused-filter variants | see documentation |
rearrange |
(b,h,l) ↔ (l,b,h) layout change |
— |
implicit_filter |
implicit long-filter generation | — |
b2b_causal_conv1d(x, weight_proj, weight_mixer, skip_bias) computes, with xv = proj(x):
z = xv[:, 1::3] * xv[:, 2::3]
y = mixer(z) + skip_bias * z
out = y * xv[:, ::3]
Channels are interleaved, and both convolutions correlate (taps are not flipped) — matching
torch.nn.Conv1d with padding=kernel_size-1 and the tail truncated.
Requirements
- NVIDIA GPU, Ampere or newer
- CUDA 13
- Python 3.11–3.13
b2b_causal_conv1dand channel-firstcausal_conv1donly: cuDNN 9.24+ andnvidia-cudnn-frontend1.27.0+, both declared as dependencies
b2b_causal_conv1d and channel-first causal_conv1d are served by cuDNN
These dispatch to the cuDNN frontend rather than to kernels in this wheel. causal_conv1d with
channel_last=True (NWH) keeps this package's native kernels and none of the following applies
to it. Three consequences:
float64is not supported and raisesValueError. Earlier releases accepted it; the cuDNN path has no double-precision implementation. Channel-lastcausal_conv1dstill serves fp64.- Kernel widths are contiguous ranges, where earlier releases dispatched a discrete allowlist
(projection 2/3/4/8/16/32, mixer 2–8/16/32/64/128/256, with fp32 backward additionally refusing a
mixer width of 256). Everything that worked before still works — widths such as 5, 9 and 31 now
work too. Anything outside the range raises
ValueErrorbefore a kernel runs. Channel-firstcausal_conv1dserves 2–256, matching its previous JIT envelope exactly. - There is no fallback. If cuDNN or the frontend is missing or too old, the call raises a diagnostic naming the GPU, its compute capability, and the frontend and backend versions actually resolved.
pip check reports a conflict between this package's cuDNN 9.24 floor and the exact pin PyTorch
carries. That is expected and benign: PyTorch is bit-identical on 9.24.
Support
Please contact the developers with any issues.