Fast Attention Requires Bounded Entries
Josh Alman, Zhao Song
Introduction
Large language models (LLMs) such as Transformer , BERT , GPT-3 , PaLM , and OPT can process natural language more effectively than smaller models or traditional algorithms. This means that they can understand and generate more complex and nuanced language, which can be useful for a variety of tasks such as language translation, question answering, and sentiment analysis. LLMs can also be adapted to multiple purposes without needing to be retained from scratch. Their power is particularly exemplified by the recent success of ChatGPT, a chat software by OpenAI built on top of GPT-3 .
The key technical backbone of LLMs is the attention matrix . An attention matrix is a square matrix whose rows and columns correspond to words or “tokens”, and whose entries correspond to the correlations between these tokens in natural text. The attention matrix is then used to calculate the importance of each input token in a sequence when producing an output. In an attention mechanism, each input token is given a weight or score, which reflects its importance or relevance to the current output being generated. These scores are calculated based on a comparison between the current output state and the input states, using a similarity function.
We formally define Attention computation as follows. Throughout this paper, we write to denote the entry-wise exponential for matrices.
The straightforward algorithm for this problem computes the matrix and then performs the multiplications , in time . Since is an matrix with entries, it is impossible to improve on this much while explicitly computing the matrix . However, the input to the problem is not , but rather the three matrices which each have only entries. An algorithm which only implicitly makes use of , without explicitly computing all its entries, could hope to run in almost linear time!
In this paper, we investigate the possibility of accelerating attention computations in this way. The two main questions we address are:
Q1. When can we perform attention computations in almost linear time ?
Q2. When can we prove that subquadratic-time algorithms for attention computations are impossible?
In most LLMs, it suffices to approximately perform attention computations throughout the inference process as long as there are reasonable precision guarantees . We therefore focus here on approximate attention computation, which can potentially be performed even faster than exact computation. Mathematically, we define the approximate version of as follows.
Again, the straightforward algorithm for this problem runs in time , but the input size is only . Our goal is to investigate when faster algorithms are possible in terms of the parameters , and .
We focus on the natural setting where (the setting where we model long sequences) and (low enough error so that attention computations over an entire network can be combined). Our main results show that whether or not there is a fast algorithm for critically depends on , the magnitudes of the entries in the input matrices.
We first show a lower bound, that when , it is impossible to design a truly subquadratic-time algorithm. Our lower bound makes use of the Strong Exponential Time Hypothesis () , a popular conjecture from the area of fine-grained complexity regarding the time required to solve -SAT. (See Section 4 below where we discuss in more detail.)
Assuming , for every , there are constants such that: there is no time algorithm for the problem .
Our second complementary result is a new algorithm, showing that when , the problem can be solved very efficiently, in almost linear time.
There is an algorithm (Algorithm 1) that solves in time .
Our Theorems 1.3 and 1.4 show that the attention computation problem exhibits a very tight transition at from almost linear time to trivial quadratic time. When is smaller, the problem can be solved in almost linear time in the input size, using our algorithm for Theorem 1.4. When is greater, our algorithm from Theorem 1.4 no longer applies, and furthermore our lower bound from Theorem 1.3 shows that it is impossible to solve the problem in truly subquadratic time, no matter what algorithmic techniques one uses (assuming ).
It has been observed in LLM implementations in practice that computations are much faster when one assumes that the matrix entries are bounded or can be well-approximated using a small number of bits (see, e.g., [30, Section 2] and [19, Section 3.2.1]). Our work can be viewed as giving a theoretical explanation for this phenomenon, and helping to explain why techniques like quantization and low-degree polynomial approximation have been so effective in practice.
A recent work by Zandieh, Han, Daliri, and Karbasi was the first to give an algorithm with provable guarantees for attention approximation. Their algorithm makes use of locality sensitive hashing (LSH) techniques which, as we will discuss next, is quite different from our algorithm for Theorem 1.4 which uses the polynomial method .
In the case when , they achieve a running time of roughly , where is a relative error parameter (which is similar, though not exactly the same, as our from Definition 1.2). In particular, their algorithm applies for larger than ours (we require ), but we achieve almost linear time (whereas their running time is bounded below by ), and our algorithm can handle any polynomial error (whereas they require to not increase the running time by a polynomial factor).
It is natural to wonder whether further improvements are possible by combining our techniques with those of . However, our lower bound of Theorem 1.3 shows that our algorithm of Theorem 1.4 is already essentially tight and cannot be substantially improved.
Another recent work by Keles, Wijewardena, and Hedge was the first to prove a lower bound for attention computation assuming . They prove, among other results, that cannot be solved in truly subquadratic time in the case when . Our Theorem 1.3 improves their result to also hold for , and to show how the complexity changes with the magnitude of entries (which is not studied by ). As we discuss more shortly, both our lower bound proof and use the high-level technique of , although our more fine-grained analysis of the parameters requires a more intricate analysis and the use of other techniques from fine-grained complexity related to approximate nearest neighbor search and the polynomial method .
2 Technique Overview
Our high-level approach is to make use of similarities between attention computation and other computational problems related to Kernel Density Estimation (KDE). Such a relationship was investigated by recent work . In particular, was inspired to apply LSH techniques to attention computation because of the prevalence of LSH in KDE algorithms . The main conceptual idea behind our results is that different techniques from the KDE literature, other than LSH, can be modified to apply in this setting and yield tight algoriths and lower bounds.
To use this to solve , we make use of a recent result which bounds the degree required to approximate the exponential function by a polynomial in order to find a low-rank approximation of the attention matrix . Prior work applied these polynomials in a similar way to solve the Gaussian KDE problem; our main observation is that by an appropriate rescaling, this approach can be modified to apply to as well.
The proof of our lower bound Theorem 1.3 builds off of another line of work on the fine-grained complexity of KDE problems . The main idea is to give a fine-grained reduction from the well-studied problem of Approximate Nearest Neighbor search . In , one is given as input vectors of dimension , and an error parameter , and the goal is to find a pair of vectors whose distance is at most times the minimum distance between any pair of the vectors. The straightforward algorithm for runs in quadratic time, and it is known that it is impossible to solve in truly subquadratic time assuming .
In order to prove our lower bound, we show that can be used to solve . The key idea is that, if the matrices and from are formed by concatenating the input vectors to the problem, then the nearest neighbor vectors correspond to the largest entries of the attention matrix . It is not immediately clear that can be used to detect large entries of , since the output is rescaled by the matrix , but we show that this can be overcome with some modifications to the input vectors which approximately balance the rows of . Prior work used a very similar approach to give lower bounds for KDE problems, although KDE doesn’t involve any rescaling factors.
In Section 2, we introduce relevant notation and tools from prior work. In Section 3, we present and analyze our attention algorithm. In Section 4, we prove our fine-grained attention lower bound. In Section 5, we provide a conclusion for this paper.
Preliminaries
We work in the standard real-RAM model and assume arithmetic operations on real numbers can be performed in constant time in our algorithms.
For any positive integer, we use to denote set .
We use to denote a length- vector whose entries are all s. We use to denote a length- vector whose entries are all s.
Our algorithm for attention computation will critically make use of a polynomial approximation for the exponential function. In particular, we use the following tight construction from previous work .
Moreover, can be computed efficiently: its coefficients are rational numbers with -bit integer numerators and denominators which can be computed in time.
2 From Additive Error to Relative Error
We note that in our setting, Lemma 2.1 can be used to give a relative error approximation as well:
Attention Algorithm
for all .
2 From Low Degree Polynomials to Low Rank Matrices
Let denote the degree- polynomial. Expand it in terms of its coefficients as
is a degree- polynomial in the entries of the vectors . Define the set of its variables,
Let denote the set of functions
Each entry of these matrices can be constructed by multiplying together at most variables, so these matrices can be constructed in time as desired. ∎
4 Key Lemma
Our key lemma shows that, even though the attention matrix may have full rank, it has a low-rank approximation that is easy to compute:
and a positive integer bounded above by
Let . From Lemma 3.3, we know that . Thus, applying Corollary 2.2 (with bound on its entries), there is a degree- polynomial such that the matrix is an -approximation to (See the definition of -approximation in Definition 3.1.) We can then compute using Lemma 3.2, which gives the bound
5 From A𝐴A to D𝐷D
6 From A𝐴A and D𝐷D to Attention Matrix
We now bound each of these two terms separately.
where the second step follows from the triangle inequality, the forth step follows from , the fifth step follows from and , and the last step follows from our assumption on .
where the second step follows from triangle inequality, the third step follows from , the forth step follows from and the last step follows from definition of .
The result follows by combining Eq. (1), and two inequalities (Eq. (3.6) and Eq. (3.6)). ∎
7 Main Upper Bound
The running time of each step is shown in Algorithm 1; its running time follows from Lemma 3.4. Its correctness follows from Lemma 3.5 and Lemma 3.6. ∎
8 Proof of Theorem 1.4
where the second step follows from and the third step follows from .
Since , let us write for some . We thus have that
The second step follows from the generic bound for , and the third step uses that .
Since are all bounded by , our final running time is as desired. ∎
Hardness
In this section, we prove our fine-grained lower bound for attention computation. In Section 4.1, we state the Strong Exponential Time Hypothesis (), the main hardness assumption we will use. In Section 4.2, we define the approximate nearest neighbor search problem, and its known hardness assuming . Finally, in Section 4.3, we give a reduction from approximate nearest neighbor search to attention computation, which implies our hardness result.
The Strong Exponential Time Hypothesis (SETH) was introduced by Impagliazzo and Paturi over 20 years ago. It is a strengthening of the conjecture, which asserts that our current best algorithms are roughly optimal:
For every there is a positive integer such that - on formulas with variables cannot be solved in time, even by a randomized algorithm.
SETH is a popular conjecture which has been used to prove fine-grained lower bounds for a wide variety algorithmic problems. See, for instance, the survey .
2 Nearest Neighbor Search
We will make use of a known relationship between and approximate nearest neighbor search.
For a parameter , in the -Approximate Hamming Nearest Neighbor Search problem for vectors of dimension , we are given as input two sets with , and our goal is to find an and satisfying .
(This is sometimes called the ‘bichromatic’ problem, and a monochromatic version has also been studied; see, for instance, .) Rubinstein showed that for certain parameters, it is impossible to substantially improve on the straightforward quadratic-time algorithm for assuming :
Assuming SETH, for every , there are and such that -Approximate Hamming Nearest Neighbor Search in dimension requires time.
We may assume that 4.3 holds even in the special case where each input vector from and has half its entries equal to and half equal to . Indeed, for any vector , we can construct a new vector given by . Here is the binary complement of vector , i.e., for all . Thus, . We can similarly construct a new vector for each . After this transformation, for any and , we have , so it suffices to find an approximate nearest neighbor among these transformed vectors.
For convenience in our the analysis, we define a gap version of approximate nearest neighbor search problem .
Let denote two positive integers. Let denote a threshold parameter. Let denote a accuracy parameter. Given two sets of points and : For each , we need to distinguish the following two cases
Case 1. There exists a such that .
Case 2. For all we have .
An algorithm for can be called times to binary search for the answer to , so Lemma 4.3 holds as well for .
3 Hardness Result
In the remainder of this section, we prove our lower bound for attention computation:
Assuming SETH, for every sufficiently small , there are constants and and such that Approximate Attention Computation (Definition 1.2) for parameters requires time.
This follows from combining Lemma 4.3 (hardness for approximation nearest neighbor search) and Lemma 4.7 (a reduction from approximate nearest neighbor search to approximate attention computation) which we prove below. ∎
For any constant : For every and , there exist constants and and such that, if (Definition 1.2) for parameters can be solved in time , then (Definition 4.5) can be solved in time .
We give an algorithm with the stated running time for . Let be a parameter we will choose later (it will be a function of and ). Our algorithm will proceed to one of two cases depending on the value of . If , then we will use one algorithm which runs in time . Otherwise, if , we will use another algorithm which runs in time .
Let be the input vectors to , and let denote the target distance. Recall that .
In this case, we will simply brute-force for the answer in the following way: We first store the vectors in a lookup table, then for each , we iterate over all vectors which have Hamming distance at most from and check whether is in the lookup table. This determines whether there is a at distance at most from , as desired.
For each , we need to iterate over choices for the vector , so the total running time will be . By standard bounds on binomial coefficients, we know that
We can thus pick a sufficiently small constant , depending only on and such that and this entire brute-force takes time.
Let denote the input of (Definition 4.5), and recall from Remark 4.4 that we may assume each has half its entries and half its entries . We will explain how to construct an Attention matrix using this instance.
Let and denote parameters we will choose later (see Eq. (9) and Eq. (6), respectively). Define by
Intuitively, our goal in picking these parameters is that will be an upper bound on entries of the attention matrix, i.e., we will have:
We will make use of an algorithm for the problem, for the following parameters:
Since each entry of and is either or , it follows that
(Note that we do not explicitly compute all the entries of in our algorithm; we will make use of it only through calling our algorithm for the Attention problem.)
For each , we know that
On the other hand, we know that that for each ,
since it is an exponential of an entry of .
Using Eq. (4.3) and Eq. (11), combined with our expression for , it thus follows that
Since , thus we know that
We can show that as follows:
where the second step follows from simple algebra, the third step follows from (Eq. (4)) and (assumption in Lemma statement), the second step follows from choice of (Eq. (7)), and the sixth step follows from choice of (Eq. (8)), and the last step follows from Eq. (8).
Therefore, for any ,
where the second step follows from , and the last step follows from simple algebra.
Recall that our goal is to determine, for each , whether there is a such that , or whether for all . We will show next that we can distinguish these two cases by seeing whether is greater than or less than the value .
If there exists an such that , then
where the first step follows from (see Eq. (6)).
where the second step follows from the definition of (see Eq. (12)), and the last step follows from the definition of .
If for all , we have , this implies
where the third step follows from definition of (see Eq. (12)), the forth step follows from the calculation in Eq. (4.3) below, and the last step follows from .
Finally, b our choice of and , we can see that
where the first step follows (Eq. (4)), the second step follows from , and the third step follows from (Eq. (9)) and the choice of (Eq. (7)). ∎
Conclusion
In this work, we showed that how quickly one can perform attention computation depends critically on the magnitude, , of the entries of the input matrices. Our main idea was to make use of similarities between attention computation and KDE, and to show how many known techniques for KDE can also be used in this setting. Since KDE is a very well-studied problem, it would be exciting to see what other techniquues can be applied to attention computation as well.
The authors would like to thank Beidi Chen for helpful discussions related to LLMs, and Feyza Duman, Chinmay Hegde, and Piotr Indyk for helpful comments on an earlier draft.