about
Fast transformers via sketches for polynomial kernels (arxiv.org)
3 points by TaurenHunter on Jan 18, 2024 | hide | past | pdf | 1 comment on HN

In plain words: Softmax attention slows as text grows; this model replaces it with a polynomial formula, then compresses it to run in linear time without dropping attention links. For 32,000-word contexts, it trained 2.5 to 4 times faster than standard fast attention with no quality loss.

Abstract · PolySketchFormer: Fast Transformers via Sketching Polynomial Kernels

The quadratic time and memory complexity inherent to self-attention mechanisms, with respect to sequence length, presents a critical computational bottleneck in the training and deployment of large-scale Transformer-based language models. Recent theoretical results indicate the intractability of sub-quadratic softmax attention approximation under reasonable complexity assumptions. This paper addresses this challenge by first demonstrating that polynomial attention with high degree can effectively replace softmax without sacrificing model quality. Next, we develop polynomial sketching techniques from numerical linear algebra to achieve linear-time polynomial attention with approximation guarantees. Crucially, our approach achieves this speedup without requiring the sparsification of attention matrices. We also present a block-based algorithm to apply causal masking efficiently. Combining these techniques, we provide \emph{PolySketchFormer}, a practical linear-time Transformer architecture for language modeling that offers provable guarantees. We validate PolySketchFormer empirically by training language models capable of handling long contexts. These experiments utilize both synthetic and real-world datasets (PG19, Wikipedia and C4) on Google Cloud TPUs. For context lengths of 32k and GPT-2 style models, our model achieves a 2.5-4x speedup in training compared to FlashAttention, with no observed degradation in quality across our experiments.

Praneeth Kacham, Vahab Mirrokni, Peilin Zhong
arXiv:2310.01655 · cs.LG · submitted Oct 2, 2023 · updated Mar 17, 2024
abstract · pdf · html · Added results of more experiments. Added a link to our JAX implementation of models

add comment on HN
Also discussed: Oct 2023 (2 points, 0 comments)

"... we show that sketches for Polynomial Kernel from the randomized numerical linear algebra literature can be used to approximate the polynomial attention which leads to a significantly faster attention mechanism without assuming any sparse structure for the attention matrix that has been done in many previous works. ... we propose an efficient block-based algorithm that lets us apply the causal mask to the attention matrix without explicitly realizing the n×n attention matrix and compute the output of the polynomial attention mechanism in time linear in the context length. ... "