subquadratic-ops-torch-cu13 0.3.0


pip install subquadratic-ops-torch-cu13

  Latest version

Released: Aug 28, 2026


Meta
Author: Alireza Moradzadeh

Classifiers

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_conv1d and channel-first causal_conv1d only: cuDNN 9.24+ and nvidia-cudnn-frontend 1.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:

  • float64 is not supported and raises ValueError. Earlier releases accepted it; the cuDNN path has no double-precision implementation. Channel-last causal_conv1d still 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 ValueError before a kernel runs. Channel-first causal_conv1d serves 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.

Extras: None
Dependencies:
scikit-build-core (>=0.10)
apache-tvm-ffi (<0.2,>=0.1.12)
warp-lang (>=1.8.0)
nvidia-ml-py
nvidia-cudnn-cu13 (>=9.24.0.43)
nvidia-cudnn-frontend (>=1.27.0)
packaging (>=23.2)