Quasi Linear Kernel Attention with Infinite Capacity
Authors: Nicolaj Rux, Johannes Hertrich, Sebastian Neumayer
Organizations: Faculty of Mathematics, Chemnitz University of Technology, 09126 Chemnitz, Germany · Institute of Computer Science, University of Göttingen, 37073 Göttingen, Germany
The evaluation cost of transformers with softmax attention scales quadratically with sequence length. Kernel attention addresses this by replacing softmax with a more general kernel function. In this paper, we aim to identify kernels that retain the expressivity of attention while enabling quasi linear computation. To quantify expressivity, we introduce a capacity for each kernel, measuring the maximum sequence length for which the attention matrix can approximate the identity. A higher capacity thus indicates greater expressivity. We show that expressive kernels like softmax, Gauss, and Laplace have infinite capacity. In contrast, common quasi linear kernels, such as those derived from finite dimensional feature maps, exhibit finite capacity. As a solution, we propose additive kernels constructed from univariate spline and polynomial exponential kernels. We prove that these maintain infinite capacity while allowing quasi linear computation via sorting. Finally, we implement additive sorting kernels efficiently and benchmark them against modern softmax backends, demonstrating advantages for long sequences.
Figures & tables
#
Step
Work
Memory
1
Input s∈RM , t∈RN and v∈RN×C
-
(C+1)N+M
2
Output zm:=∑n=1N∣sm−tn∣vn∈RCm=1,…,M
-
MC
3
σ:=argsort(t) s.t. tσ(1)≤…≤tσ(N)
Nlog2N
N
4
pm:=max({n:tσ(n)≤sm}∪{0}) for m=1,…,M
Mlog2N
M
5
pref_tn:=∑j=1ntσ(j)vσ(j)∈RC for n=0,…,N
(N+1)C
(N+1)C
6
pref_vn:=∑j=1nvσ(j)∈RC for n=0,…,N
(N+1)C
(N+1)C
Algorithm 1 Weighted absolute value sum
Name
Formula
dim
stationary
additive
capacity
complexity
softmax
eq⊤k
∞
no
no
∞
quadratic
gauss
e−21∥q−k∥22
∞
yes
no
∞
quadratic
laplace
e−∥q−k∥1
∞
yes
no
∞
quasi linear
riesz
∥q∥2+∥k∥2−∥q−k∥2+ε
∞
partially
no
−
quadratic
add_riesz
∣s∣+∣t∣−∣s−t∣+ε
∞
partially
yes
≥2D
quasi linear
add_laplace
e−∣s−t∣
∞
yes
yes
∞
quasi linear
Table 1: For additive kernels, we state the univariate ϕ(s,t) . dpfp is described in Appendix A and the remaining kernels are defined for q,k∈RD . All kernels except tri are nonnegative. riesz and add_riesz use ε=10−3 and are partially stationary, i.e., stationary up to terms depending only on q or only on k . The last column is the best known kernel sum complexity in N+M .
add_laplace fp32
add_riesz fp32
softmax
N
GLOBAL
GLOBAL
FUSED
KEOPS
MEM_EFF32
MEM_EFF16
FLASH16
CUDNN16
128
0.442
0.455
0.305
0.381
0.029
0.021
0.021
0.023
256
0.575
0.539
0.353
0.425
0.058
0.031
0.022
0.030
512
0.824
0.694
0.478
0.688
0.144
0.058
0.041
0.043
1024
1.305
1.014
0.719
1.648
0.431
0.152
0.093
0.119
2048
2.260
1.690
1.278
4.754
1.645
0.537
0.266
0.265
Table 2: Forward-pass of kernel attention ( 1 ) with runtime in milliseconds, mean over 10 runs after 5 warm-up iterations. Relative standard deviations are below 18% for N≤512 , below 4% at N=1024 and below 3% for N≥2048 . Bold marks the fastest method overall; underline marks the fastest among the five fp32 methods. Shape: B=4 , H=12 , D=64 , C=64 , M=N .
MTEB ↑
LoCo (nDCG@10) ↑
Kernel
scale τ
≤512
2048
4096
8192
softmax (teacher)
2.828
64.41
0.873
0.882
0.883
gauss
2.828
64.41
0.875
0.884
0.886
laplace
6.0
64.38
0.885
0.891
0.887
riesz
1.0
61.69
0.805
0.839
0.858
add_riesz
1.0
61.72
0.769
0.813
0.832
Table 3: Results after full two-stage distillation for MTEB and LoCo. Best performance among quasi linear kernels is underlined.
MTEB ↑
LoCo (nDCG@10) ↑
Kernel
scale τ
≤512
2048
4096
8192
softmax (teacher)
2.828
64.41
0.873
0.882
0.883
gauss
2.828
64.41
0.875
0.884
0.886
laplace
6.0
64.38
0.885
0.891
0.887
riesz
1.0
61.69
0.805
0.839
0.858
add_riesz
1.0
61.72
0.769
0.813
0.832
Table 3: Results after full two-stage distillation for MTEB and LoCo. Best performance among quasi linear kernels is underlined.
Appendix figures & tables5 assets
Supplementary material from the paper’s appendix.
Appendix
Paper
φ
J
fj
Katharopoulos et al. (2020)
φelu
1
f1(x)=elu(x)+1
Peng et al. (2021)
φtri
2
f1(x)=sin(x),f2(x)=cos(x)
Choromanski et al. (2021)
φReLU
1
f1(x)=ReLU(x)
Appendix
Table 4: Examples of FFMs of the form ( 19 ).
add_riesz fp32
softmax
N
GLOBAL
FUSED
KEOPS
MEM_EFF32
MEM_EFF16
FLASH16
CUDNN16
128
1.066
0.913
1.598
0.098
0.071
0.059
0.057
256
1.329
1.147
2.424
0.233
0.104
0.067
0.069
512
1.832
1.620
4.087
0.609
0.231
0.132
0.143
1024
2.950
2.664
11.96
1.923
0.634
0.324
0.363
2048
5.283
4.927
40.33
7.360
2.262
0.978
1.050
Appendix
Table 5: Forward+backward of kernel attention ( 1 ) with runtime in milliseconds, mean over 10 runs after 5 warm-up iterations, with forward and backward timed together in a single region. Relative standard deviations are below 10% for N≤512 , below 1.2% at N=1024 and below 0.7% for N≥2048 . Bold marks the fastest method overall; underline marks the fastest among the four fp32 methods. Shape: B=4 , H=12 , D=64 , C=64 , M=N .
1
Input s∈RM , t∈RN and v∈RN×C
2
Output (zm)m=1M=CausalSorting(s,t,v) given by zm:=∑n=1mϕ(sm,tn)vn∈RC
3
Denote by Sorting(s,t,v) the application of Algorithm 1