Fast Gauss Sums via Flash Attention
arXiv:2609.04910v1 [cs.LG] 4 Sep 2026
Nicolaj Rux Sebastian Neumayer Faculty of Mathematics Chemnitz University of Technology Reichenhainer Str. 39, 09126 Chemnitz, Germany {nicolaj.rux, sebastian.neumayer}@math.tu-chemnitz.de
Abstract Gaussian kernel sums are the computational core of maximum mean discrepancies (MMDs), kernel gradient flows, Stein variational gradient descent (SVGD), and many other kernel methods. At the same time, softmax attention has received an extraordinary amount of hardware-aware code engineering, culminating in flash attention. We show that Gauss kernel sums with arbitrary, signed weights can be evaluated via flash attention: two small input augmentations turn the normalized softmax reduction into the unnormalized Gauss sum, without writing a single line of custom GPU code. For feature dimension D > 8 in fp16, this approach beats compiled PyTorch code as well as PyKeOps kernels (often significantly) in speed, memory-overhead and accuracy. Indeed, its memory scaling remains linear.
1
Introduction τ
2
Let Φτ (q, k) = e− 2 ∥q−k∥2 denote the rescaled Gauss kernel. Sums of the form PN sm := n=1 Φτ (qm , kn ) vn ∈ RC , m = 1, . . . , M, D
(1)
C
with points k1 , . . . , kN , q1 , . . . , qM ∈ R , values v1 , . . . , vN ∈ R and bandwidth τ > 0, appear in essentially every Gaussian kernel method. They are the bottleneck of MMD gradient flows [Arbel et al., 2019, Hertrich et al., 2024] and of Stein variational gradient descent [Liu and Wang, 2016]. Indeed, evaluating (1) naively requires O(M N (C+D)) operations with an O(M N ) memory footprint. Classical fast summation methods trade exactness for better asymptotics in restricted regimes [Greengard and Strain, 1991, Beatson and Newsam, 1992, Yang et al., 2004, Rahimi and Recht, 2007, Hertrich, 2024], while PyKeOps [Charlier et al., 2021] performs the exact reduction based on a fused, memory-efficient GPU kernel. Meanwhile, a different community has devoted substantial code engineering effort to a single, specific reduction: softmax attention [Vaswani et al., 2017] was implemented as flash attention [Dao et al., 2022, Dao, 2024]. In this note, we show that the Gaussian kernel sum (1) reduces to a flash evaluation after two small augmentations, without any custom CUDA or Triton code. Our contributions are: (i) two reductions of (1) to attention calls, one fully differentiable through the public PyTorch API and one that additionally reads the logits returned by the attention backends; (ii) an analysis of the fp16 pitfalls together with simple safeguards; and (iii) benchmarks against compiled PyTorch and PyKeOps implementations across different input shapes, including gradients.
2
Gaussian kernel sums from flash attention
Softmax attention is denoted by Attτ (q, k, v) and given for m = 1, . . . , M via PN ⊤ N τ qm kn X ⊤ vn C n=1 e := Attτ (q, k, v)m := PN ∈ R , lse (q, k) log eτ qm kn ∈ R. τ m ⊤k τ q n m n=1 n=1 e Preprint.
(2)
Algorithm 1 (prescale) Gaussian kernel summation via flash attention # 1 2 3 4 5 6
Step Input k ∈ RN ×D , q ∈ RM ×D , v ∈ RN ×C , τ > 0 PN Output sm = n=1 Φτ (qm , kn )vn , m = 1, . . . , M ṽn := vn exp(− τ2 ∥kn ∥2 ), n = 1, . . . , N Compute Attτ (q, k, ṽ) ∈ RM ×C , lseτ (q, k) ∈ RM sm := Attτ (q, k, ṽ)m exp(lseτ (q, k)m ) exp(− τ2 ∥qm ∥2 ) return s
Work – – N (C + D) M N (D+C) MC
Call
mult Flash mult
Flash attention evaluates (2) with O(M N (D+C)) operations and O((M +N )(D+C)) memory. The logits lseτ are computed as a by-product in fp32, but no gradient is implemented. To connect (2) with Gaussian kernel sums, we first remove the normalization by multiplying with exp(lseτ (q, k)m ) to get PN ⊤ τ qm kn vn = Attτ (q, k, v)m · elseτ (q,k)m . (3) n=1 e τ
2
Then, we introduce ṽ = (vn e− 2 ∥kn ∥ )N n=1 to obtain N N X X ⊤ Attτ (q, k, ṽ)m elseτ (q,k)m τ qm kn − τ2 ∥kn ∥2 − τ2 ∥qm ∥2 = e v e e = Φτ (qm , kn )vn . τ n 2 e 2 ∥qm ∥ n=1 n=1
(4)
Algorithm 1 summarizes this method. Unfortunately, this implementation cannot handle gradients. Interestingly, we can proceed without using the (non-differentiable) logits by adding two additional entries to the queries and keys, see Algorithm 2. Then, it holds indeed for [α, β] := Attτ (q̃, k̃, ṽ) that PN PN ⊤ 2 2 2 ⊤ τ τ τ κ n=1 eτ qm kn − 2 ∥kn ∥ vn e 2 ∥qm ∥ + n=1 eτ qm kn − 2 ∥kn ∥ καm = sm . = τ P τ 2 τ 2 ⊤ 2 N βm e 2 ∥qm ∥ κ e 2 ∥qm ∥ + n=1 eτ qm kn − 2 ∥kn ∥ Proposition 1 (Stability). In Algorithm 2 (reweight) every m = 1, . . . , M satisfies κ(N + 1)−1 ≤ βm ≤ κ, ∥αm ∥∞ ≤ max ∥vn ∥∞ . n=1,...,N √ −1 With κ = N + 1, this gives βm ∈ [κ , κ].
(5)
(6)
Usually the weights are bounded. Thus, the critical condition is to ensure that β does not underflow. The dynamic range of fp16 is [2−14 , 216 ], which restricts us to κ ≤ 214 or N ≤ 268 435 455. Algorithm 2 (reweight) Gaussian kernel summation via flash attention # 1 2 3 4 5 6 7 8
Step √ Input k ∈ RN ×D , q ∈ RM ×D , v ∈ RN ×C , τ > 0, κ = N +1 PN Output sm = n=1 Φτ (qm , kn )vn , m = 1, . . . , M q̃m := [qm , 1, 21 ∥qm ∥22 ] ∈ RD+2 , m = 1, . . . , M k̃n := [kn , − 12 ∥kn ∥22 , 0], n = 1, . . . , N , k̃0 := eD+2 ṽn := [vn , 0], n = 1, . . . , N , ṽ0 := κ eC+1 [α, β] := Attτ (q̃, k̃, ṽ) with α ∈ RM ×C and β ∈ RM sm := καm /βm , m = 1, . . . , M return s
Work – – MD ND NC M N (D+C) MC
Call
pad pad pad Flash mult
Remark 2 (Practical implementation). Algorithm 2 is implemented entirely through the public scaled_dot_product_attention (SDPA) interface: automatic differentiation applies, and the backward pass again runs flash kernels. Algorithm 1 is algebraically leaner (no auxiliary key, channel or division) but requires the logits, which the PyTorch backends flash, cudnn and memory_efficient expose. In practice, we shift q and k by their common mean to avoid numerical overflow. Moreover, flash requires half precision (fp16 or bf16) and a multiple of 8 for the dimension of queries, keys and values. Flash is optimized for D = C ∈ {32, 64, 128} and only works up to D, C ≤ 256. Both implementations therefore run at dimension Dreweight = ⌈max(D+2, C+1)⌉8 and Dprescale = ⌈max(D, C)⌉8 , respectively, where ⌈·⌉8 rounds up to a multiple of 8. 2
3
Applications
Maximum mean discrepancy. For discrete measures µ and ν, write their signed difference as PN σ := µ − ν = n=1 vn δkn , with v ∈ RN . Then, their squared maximum mean discrepancy is ZZ
2
MMDτ (µ, ν) =
Φτ (q, k)dσ(q)dσ(k) =
N X
vm sm ,
sm =
m=1
N X
Φτ (qm , kn ) vn .
(7)
n=1
Thus, MMD amounts to a kernel sum with q = k and C = 1, followed by the inner product v T s. MMD flows and Stein variational gradient descent. with respect to the query points, namely ∇q m
N X
Φτ (qm , kn )vn = τ
Particle flows require the gradient of (1)
P
PN N n=1 Φτ (qm , kn )vn kn − qm n=1 Φτ (qm , kn )vn
,
(8)
n=1
which is again a single kernel sum with values (vn kn , vn ) ∈ RD+1 . This is the regime C = D + 1 and it covers both terms of the SVGD update [Liu and Wang, 2016, Eq. (8)], whose kernel part has the same form. Alternatively, plain autograd through Algorithm 2 applies.
4
Numerical results
Here, we benchmark different methods for computing (1), namely a naive torch.compile [Ansel et al., 2024] version (PyTorch), PyKeOps [Charlier et al., 2021] (PyKeOps) and the two proposed flash-based variants reweight and prescale from Section 2. The inputs are B independent standard Gaussian point clouds scaled by D−1/2 with M = N and τ = 1, which are evaluated in parallel. For reproducibility, we fix the random seed and always perform 3 warm-up runs. Then, we call each method 8 times and compute the mean and standard deviation of the elapsed time, the memory footprint (excluding input tensors), and the relative L2 error against a fp64 reference summation. All timings were measured on an NVIDIA GeForce RTX 5090 with 32 GB using PyKeOps 2.3, PyTorch 2.13 with CUDA 13.0 and cuDNN 9.20.0. Since flash relies on Tensor cores running in half precision, we report results for fp16. Appendix 6.2 briefly discusses the fp32 regime.
102 101 100 10 1 10 2
29
211
213
215 N
217
219
109
107 106
time [ms]
103
29
PyTorch PyKeOps reweight prescale
101 100 10 1 29
211
213
215 N
217
(d) Elapsed time.
211
213
215 N
217
219
221
29
(b) Memory overhead. 1010
102
10 2
10 3
105
221
memory overhead [bytes]
104
PyTorch PyKeOps reweight prescale
108
(a) Elapsed time. 105
PyTorch PyKeOps reweight prescale
relative L2 error
1010
PyTorch PyKeOps reweight prescale
219
221
109
211
213 N
215
217
215
217
(c) Accuracy.
PyTorch PyKeOps reweight prescale
PyTorch PyKeOps reweight prescale
relative L2 error
time [ms]
103
memory overhead [bytes]
104
108 107
10 3
106 29
211
213
215 N
217
219
(e) Memory overhead.
221
29
211
213 N
(f) Accuracy.
Figure 1: Sweep over N computing the forward of the Gauss sum (1) for fixed B = 4 in fp16 with backend flash. The upper row is D = 3, C = 1 and the lower row is D = 32, C = 32. 3
In Figure 1, we sweep over N with fixed D. As PyTorch allocates an entire N × N matrix, this method quickly runs out of memory. For D = 3, PyKeOps is marginally faster than reweight and prescale for large N . Further, it uses consistently 3 times less memory than the flash variants. However, it has a worse error, especially for large N . For D = 32, the flash variants outperform PyTorch and PyKeOps in both speed and memory overhead, while maintaining a healthy relative accuracy. In Figure 2, we additionally include the gradient computations for the Gauss sum (1). Recall that prescale is not applicable here. The performance comparison turns out similar to the forward call (see also Figure 1 bottom row).
103 102 101 100 10 1
29
211
213
215
N
217
219
PyTorch PyKeOps reweight
109
10 3 relative L2 error
memory overhead [bytes]
104 time [ms]
1010
PyTorch PyKeOps reweight
105
108 107 106
4 × 10 4 3 × 10 4
221
29
(a) Elapsed time.
PyTorch PyKeOps reweight
6 × 10 4
211
213
215
N
217
219
221
28
(b) Memory overhead.
29
210
211 N
212
213
214
(c) Accuracy.
Figure 2: Sweep over N computing the forward and backward of the Gauss sum (1) for fixed B = 4, D = 32 and C = 1 in fp16 with backend flash.
PyKeOps reweight prescale
21
22
23
D
24
25
(a) Elapsed time.
26
27
relative L2 error
memory overhead [bytes]
101 20
PyKeOps reweight prescale
109
time [ms]
102
108
107
20
21
22
23
D
24
25
(b) Memory overhead.
26
27
PyKeOps reweight prescale
10 3
20
21
22
23
D
24
25
26
27
(c) Accuracy.
Figure 3: Sweep over D computing the forward of the Gauss sum (1) for fixed B = 64, C = 1 and N = 16384 in fp16 with backend flash. In Figure 3, the performance across different dimensions D is compared. PyKeOps consistently has the worst accuracy with an error around 2.7 × 10−3 , while all other methods lie an order of magnitude below around 3.5 × 10−4 . Its speed degenerates with increasing dimension, while the memory scales linearly and even beats both flash variants for D ≤ 8. Both speed and memory wise PyKeOps is best in the low dimensional regime (D ≤ 8). The flash variants are the fastest for D > 8. If D is a power of 2 and C ≤ D, reweight is worse compared to prescale, because reweight needs to pad the dimension to ⌈max(D+2, C+1)⌉8 = D + 8 while prescale only pads to ⌈max(D, C)⌉8 = D. This effect becomes visible for D ≥ 16 in Figure 3a. For D ≥ 8, the memory overhead of the flash variants is comparable to the one of PyKeOps with a roughly linear incline. The accuracy of the flash variants lies around 4.0 × 10−4 and slightly above PyTorch, which sits at 3.3 × 10−4 .
5
Conclusion
Two small input augmentations make flash attention a fast and memory efficient drop-in method for evaluating unnormalized Gaussian kernel sums with signed weights. It excels at fp16, which often suits sampling, flows and testing, but not ill-conditioned solvers. For D ≥ 16, the flashbased Gaussian kernel sums run 2−21 times faster than PyKeOps forward and 3−10 times faster with gradients, at up to 7 times lower error and similar linear memory scaling. Despite the code optimization, both variants have quadratic complexity. Thus, combining the constant-factor gains shown here with subquadratic approximations [Hertrich, 2024, Rux et al., 2025] is a promising future direction. 4
References J. Ansel, E. Yang, H. He, N. Gimelshein, A. Jain, M. Voznesensky, B. Bao, P. Bell, D. Berard, E. Burovski, G. Chauhan, A. Chourdia, W. Constable, A. Desmaison, Z. DeVito, E. Ellison, W. Feng, J. Gong, M. Gschwind, B. Hirsh, S. Huang, K. Kalambarkar, L. Kirsch, M. Lazos, M. Lezcano, Y. Liang, J. Liang, Y. Lu, C. K. Luk, B. Maher, Y. Pan, C. Puhrsch, M. Reso, M. Saroufim, M. Y. Siraichi, H. Suk, S. Zhang, M. Suo, P. Tillet, X. Zhao, E. Wang, K. Zhou, R. Zou, X. Wang, A. Mathews, W. Wen, G. Chanan, P. Wu, and S. Chintala. PyTorch 2: Faster machine learning through dynamic Python bytecode transformation and graph compilation. In Proceedings of the 29th ACM International Conference on Architectural Support for Programming Languages and Operating Systems, volume 2, pages 929–947, 2024. doi: 10.1145/3620665.3640366. M. Arbel, A. Korba, A. Salim, and A. Gretton. Maximum mean discrepancy gradient flow. In Advances in Neural Information Processing Systems, volume 32. Curran Associates, Inc., 2019. URL https://proceedings.neurips.cc/paper_files/paper/2019/file/ 944a5ae3483ed5c1e10bbccb7942a279-Paper.pdf. R. K. Beatson and G. N. Newsam. Fast evaluation of radial basis functions: I. Comput. Math. Appl., 24(12):7–19, 1992. doi: 10.1016/0898-1221(92)90167-G. B. Charlier, J. Feydy, J. A. Glaunès, F.-D. Collin, and G. Durif. Kernel operations on the GPU, with autodiff, without memory overflows. J. Mach. Learn. Res., 22(74):1–6, 2021. URL https: //jmlr.org/papers/v22/20-275.html. T. Dao. FlashAttention-2: Faster attention with better parallelism and work partitioning. In International Conference on Learning Representations, pages 35549– 35562, 2024. URL https://proceedings.iclr.cc/paper_files/paper/2024/file/ 98ed250b203d1ac6b24bbcf263e3d4a7-Paper-Conference.pdf. T. Dao, D. Y. Fu, S. Ermon, A. Rudra, and C. Ré. FlashAttention: Fast and memory-efficient exact attention with IO-awareness. In Advances in Neural Information Processing Systems, volume 35. Curran Associates, Inc., 2022. URL https://proceedings.neurips.cc/paper_files/paper/ 2022/file/67d57c32e20fd0a7a302cb81d36e40d5-Paper-Conference.pdf. L. Greengard and J. Strain. The fast Gauss transform. SIAM J. Sci. Stat. Comput., 12(1):79–94, 1991. doi: 10.1137/0912004. J. Hertrich. Fast kernel summation in high dimensions via slicing and Fourier transforms. SIAM J. Math. Data Sci., 6(4):1109–1137, 2024. doi: 10.1137/24M1632085. J. Hertrich, C. Wald, F. Altekrüger, and P. Hagemann. Generative sliced MMD flows with Riesz kernels. In International Conference on Learning Representations, pages 20923– 20949, 2024. URL https://proceedings.iclr.cc/paper_files/paper/2024/file/ 5b288823575bb29654b0953a251e933b-Paper-Conference.pdf. Q. Liu and D. Wang. Stein variational gradient descent: A general purpose Bayesian inference algorithm. In Advances in Neural Information Processing Systems, volume 29. Curran Associates, Inc., 2016. URL https://proceedings.neurips.cc/paper_files/paper/2016/ file/b3ba8f1bee1238a2f37603d90b58898d-Paper.pdf. A. Rahimi and B. Recht. Random features for large-scale kernel machines. In Advances in Neural Information Processing Systems, volume 20. Curran Associates, Inc., 2007. URL https://proceedings.neurips.cc/paper_files/paper/2007/file/ 013a006f03dbc5392effeb8f18fda755-Paper.pdf. N. Rux, J. Hertrich, and S. Neumayer. Numerical methods for kernel slicing. arXiv preprint, 2025. URL https://arxiv.org/abs/2510.11478. A. Vaswani, N. Shazeer, N. Parmar, J. Uszkoreit, L. Jones, A. N. Gomez, Ł. Kaiser, and I. Polosukhin. Attention is all you need. In Advances in Neural Information Processing Systems, volume 30. Curran Associates, Inc., 2017. URL https://proceedings.neurips.cc/paper_files/paper/ 2017/file/3f5ee243547dee91fbd053c1c4a845aa-Paper.pdf. 5
C. Yang, R. Duraiswami, and L. S. Davis. Efficient kernel machines using the improved fast Gauss transform. In Advances in Neural Information Processing Systems, volume 17. MIT Press, 2004. URL https://proceedings.neurips.cc/paper_files/paper/2004/file/ 85353d3b2f39b9c9b5ee3576578c04b7-Paper.pdf.
6
Appendix
6.1
Proof of Proposition 1
Proof. Since 0 < Φτ (qm , kn ) ≤ 1, the estimate κ(N + 1)−1 ≤ βm ≤ κ follows from 2
τ
βm =
κe 2 ∥qm ∥ κ = . P PN τ 2 ⊤ k − τ ∥k ∥2 N ∥q ∥ τ q e 2 m + n=1 e m n 2 n 1 + n=1 Φτ (qm , kn )
(9)
The bound on α follows directly as each αm is a convex combination of v and ∥ · ∥∞ is convex. 6.2
Single precision
For fp32, flash is no longer available as it relies on highly optimized tensor cores. Instead, the default backend for softmax attention is memory_efficient. Figure 3 and Figure 4 use the same setup (B = 64, C = 1, N = 16384) but in fp16 with flash and fp32 with memory_efficient. For D ≤ 32, memory_efficient in fp32 is around 10 times slower than its flash-fp16 counterpart. In contrast, PyKeOps is optimized for fp32. Remarkably, its memory overhead remains tiny across dimensions. While the runtime of PyKeOps does not change much between fp16 and fp32, the softmax based variants become competitive only around D ≥ 32. While all methods have a healthy accuracy of around 5 × 10−7 , PyKeOps consistently has the lowest relative L2 error. Thus, for fp32, PyKeOps remains the state-of-the-art implementation.
time [ms]
102
101
20
21
22
23
D
24
25
26
PyKeOps reweight prescale
109
relative L2 error
memory overhead [bytes]
PyKeOps reweight prescale
108
107
(a) Elapsed time.
PyKeOps reweight prescale
4 × 10 7 3 × 10 7
20
27
6 × 10 7
21
22
23
D
24
26
25
20
27
(b) Memory overhead.
21
22
23
D
24
25
26
27
(c) Accuracy.
Figure 4: Sweep over D computing the forward of the Gauss sum (1) for fixed B = 64, C = 1 and N = 16384 in fp32 with backend memory_efficient.
PyTorch PyKeOps reweight prescale
1010
time [ms]
103 102 101 100 10 1 29
211
213
215 N
217
(a) Elapsed time.
219
221
109
PyTorch PyKeOps reweight prescale
PyTorch PyKeOps reweight prescale
relative L2 error
104
memory overhead [bytes]
105
108 107
10 6
106 105
29
211
213
215 N
217
219
(b) Memory overhead.
221
29
211
213 N
215
217
(c) Accuracy.
Figure 5: Sweep over N computing the forward of the Gauss sum (1) for fixed B = 4, D = 32 and C = 32 in fp32 with backend memory_efficient.
6