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 exp⁡\exp to denote the entry-wise exponential for matrices.

The straightforward algorithm for this problem computes the matrix AA and then performs the multiplications D−1AVD^{-1}AV, in time n2+o(1)n^{2+o(1)}. Since AA is an n×nn\times n matrix with n2n^{2} entries, it is impossible to improve on this much while explicitly computing the matrix AA. However, the input to the problem is not AA, but rather the three matrices Q,K,VQ,K,V which each have only n1+o(1)n^{1+o(1)} entries. An algorithm which only implicitly makes use of AA, 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 n1+o(1)n^{1+o(1)}?

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 EAttC\mathsf{EAttC} as follows.

Again, the straightforward algorithm for this problem runs in time O(n2d)≤n2+o(1)O(n^{2}d)\leq n^{2+o(1)}, but the input size is only O(nd)≤n1+o(1)O(nd)\leq n^{1+o(1)}. Our goal is to investigate when faster algorithms are possible in terms of the parameters d,Bd,B, and ϵa\epsilon_{a}.

We focus on the natural setting where d=O(log⁡n)d=O(\log n) (the setting where we model long sequences) and ϵa=1/poly⁡(n)\epsilon_{a}=1/\operatorname{poly}(n) (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 AAttC\mathsf{AAttC} critically depends on BB, the magnitudes of the entries in the input matrices.

We first show a lower bound, that when B≥Ω(log⁡n)B\geq\Omega(\sqrt{\log n}), it is impossible to design a truly subquadratic-time algorithm. Our lower bound makes use of the Strong Exponential Time Hypothesis (SETH\mathsf{SETH}) , a popular conjecture from the area of fine-grained complexity regarding the time required to solve kk-SAT. (See Section 4 below where we discuss SETH\mathsf{SETH} in more detail.)

Assuming SETH\mathsf{SETH}, for every q>0q>0, there are constants C,Ca,Cb>0C,C_{a},C_{b}>0 such that: there is no O(n2−q)O(n^{2-q}) time algorithm for the problem AAttC(n,d=Clog⁡n,B=Cblog⁡n,ϵa=n−Ca)\mathsf{AAttC}(n,d=C\log n,B=C_{b}\sqrt{\log n},\epsilon_{a}=n^{-C_{a}}).

Our second complementary result is a new algorithm, showing that when B<o(log⁡n)B<o(\sqrt{\log n}), the problem can be solved very efficiently, in almost linear time.

There is an algorithm (Algorithm 1) that solves AAttC(n,d=O(log⁡n),B=o(log⁡n),ϵa=1/poly⁡(n))\mathsf{AAttC}(n,d=O(\log n),B=o(\sqrt{\log n}),\epsilon_{a}=1/\operatorname{poly}(n)) in time n1+o(1)n^{1+o(1)}.

Our Theorems 1.3 and 1.4 show that the attention computation problem AAttC\mathsf{AAttC} exhibits a very tight transition at B=Θ(log⁡n)B=\Theta(\sqrt{\log n}) from almost linear time to trivial quadratic time. When B<o(log⁡n)B<o(\sqrt{\log n}) is smaller, the problem can be solved in almost linear time n1+o(1)n^{1+o(1)} in the input size, using our algorithm for Theorem 1.4. When B≥Ω(log⁡n)B\geq\Omega(\sqrt{\log n}) 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 SETH\mathsf{SETH}).

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 d=o(log⁡2n)d=o(\log^{2}n), they achieve a running time of roughly O(n1.17⋅d/ϵr2)O(n^{1.17}\cdot d/\epsilon_{r}^{2}), where ϵr\epsilon_{r} is a relative error parameter (which is similar, though not exactly the same, as our ϵa\epsilon_{a} from Definition 1.2). In particular, their algorithm applies for larger dd than ours (we require d=O(log⁡n)d=O(\log n)), but we achieve almost linear time n1+o(1)n^{1+o(1)} (whereas their running time is bounded below by Ω(n1.17)\Omega(n^{1.17})), and our algorithm can handle any polynomial error ϵa=1/poly⁡(n)\epsilon_{a}=1/\operatorname{poly}(n) (whereas they require ϵr≥1/no(1)\epsilon_{r}\geq 1/n^{o(1)} 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 SETH\mathsf{SETH}. They prove, among other results, that AAttC\mathsf{AAttC} cannot be solved in truly subquadratic time in the case when d=ω(log⁡n)d=\omega(\log n). Our Theorem 1.3 improves their result to also hold for d=Θ(log⁡n)d=\Theta(\log n), and to show how the complexity changes with the magnitude of entries BB (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 d,Bd,B 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 AAttC\mathsf{AAttC}, 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 AA. 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 AAttC\mathsf{AAttC} 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 ANN\mathsf{ANN}. In ANN\mathsf{ANN}, one is given as input nn vectors of dimension dd, and an error parameter ϵ>0\epsilon>0, and the goal is to find a pair of vectors whose distance is at most (1+ϵ)(1+\epsilon) times the minimum distance between any pair of the vectors. The straightforward algorithm for ANN\mathsf{ANN} runs in quadratic time, and it is known that it is impossible to solve ANN\mathsf{ANN} in truly subquadratic time assuming SETH\mathsf{SETH} .

In order to prove our lower bound, we show that AAttC\mathsf{AAttC} can be used to solve ANN\mathsf{ANN}. The key idea is that, if the matrices QQ and KK from AAttC\mathsf{AAttC} are formed by concatenating the input vectors to the ANN\mathsf{ANN} problem, then the nearest neighbor vectors correspond to the largest entries of the attention matrix AA. It is not immediately clear that AAttC\mathsf{AAttC} can be used to detect large entries of AA, since the output is rescaled by the matrix D−1D^{-1}, but we show that this can be overcome with some modifications to the input vectors which approximately balance the rows of AA. 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 [n][n] to denote set {1,2,⋯ ,n}\{1,2,\cdots,n\}.

We use 1n{\bf 1}_{n} to denote a length-nn vector whose entries are all 11s. We use 0n{\bf 0}_{n} to denote a length-nn 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, PP can be computed efficiently: its coefficients are rational numbers with poly⁡(g)\operatorname{poly}(g)-bit integer numerators and denominators which can be computed in poly⁡(g)\operatorname{poly}(g) 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

∣A~i,j−Ai,j∣≤ϵ⋅Ai,j|\widetilde{A}_{i,j}-A_{i,j}|\leq\epsilon\cdot A_{i,j} for all (i,j)∈[n]2(i,j)\in[n]^{2}.

2 From Low Degree Polynomials to Low Rank Matrices

Let P(x)P(x) denote the degree-gg polynomial. Expand it in terms of its coefficients as

K\mathsf{K} is a degree-2g2g polynomial in the 2d2d entries u1,⋯ud,v1,⋯ ,vdu_{1},\cdots u_{d},v_{1},\cdots,v_{d} of the vectors u,vu,v. Define the set VV of its variables,

Let F{\cal F} denote the set of functions

Each entry of these matrices can be constructed by multiplying together at most gg variables, so these n×rn\times r matrices can be constructed in time O(nrg)O(nrg) as desired. ∎

4 Key Lemma

Our key lemma shows that, even though the attention matrix AA may have full rank, it has a low-rank approximation that is easy to compute:

and a positive integer rr bounded above by

Let M:=QK⊤/dM:=QK^{\top}/d. From Lemma 3.3, we know that ∥M∥∞≤B2\|M\|_{\infty}\leq B^{2}. Thus, applying Corollary 2.2 (with bound B2B^{2} on its entries), there is a degree-gg polynomial PP such that the matrix A~=P(M)\widetilde{A}=P(M) is an (ϵ,r)(\epsilon,r)-approximation to AA (See the definition of (ϵ,r)(\epsilon,r)-approximation in Definition 3.1.) We can then compute U1,U2U_{1},U_{2} 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 ∣(Di,i−D~i,i)/Di,i∣≤ϵD|(D_{i,i}-\widetilde{D}_{i,i})/D_{i,i}|\leq\epsilon_{D}, the fifth step follows from D~i−1>0\widetilde{D}_{i}^{-1}>0 and A~i,l>0\widetilde{A}_{i,l}>0, and the last step follows from our assumption on VV.

where the second step follows from triangle inequality, the third step follows from Di,i−1>0D_{i,i}^{-1}>0, the forth step follows from ∣A~i,l−Ai,l∣≤ϵA⋅Ai,l|\widetilde{A}_{i,l}-A_{i,l}|\leq\epsilon_{A}\cdot A_{i,l} and the last step follows from definition of Di,iD_{i,i}.

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 ϵ=1/poly⁡(n)\epsilon=1/\operatorname{poly}(n) and the third step follows from B=o(log⁡n)B=o(\sqrt{\log n}).

Since g=o(log⁡n)g=o(\log n), let us write g=(log⁡n)/fg=(\log n)/f for some f=ω(1)f=\omega(1). We thus have that

The second step follows from the generic bound (ab)≤(ea/b)b\binom{a}{b}\leq(ea/b)^{b} for 1≤b≤a1\leq b\leq a, and the third step uses that d=O(log⁡n)d=O(\log n).

Since d,r,gd,r,g are all bounded by no(1)n^{o(1)}, our final running time is n1+o(1)n^{1+o(1)} 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 (SETH\mathsf{SETH}), the main hardness assumption we will use. In Section 4.2, we define the approximate nearest neighbor search problem, and its known hardness assuming SETH\mathsf{SETH}. 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 P≠NP\mathsf{P}\neq\mathsf{NP} conjecture, which asserts that our current best SAT\mathsf{SAT} algorithms are roughly optimal:

For every ϵ>0\epsilon>0 there is a positive integer k≥3k\geq 3 such that kk-SAT\mathsf{SAT} on formulas with nn variables cannot be solved in O(2(1−ϵ)n)O(2^{(1-\epsilon)n}) 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 SETH\mathsf{SETH} and approximate nearest neighbor search.

For a parameter ϵ>0\epsilon>0, in the (1+ϵ)(1+\epsilon)-Approximate Hamming Nearest Neighbor Search problem for nn vectors of dimension dd, we are given as input two sets A,B⊂{0,1}dA,B\subset\{0,1\}^{d} with ∣A∣=∣B∣=n|A|=|B|=n, and our goal is to find an a∗∈Aa^{*}\in A and b∗∈Bb^{*}\in B satisfying ∥a∗−b∗∥0≤(1+ϵ)⋅min⁡a∈A,b∈B∥a−b∥0\|a^{*}-b^{*}\|_{0}\leq(1+\epsilon)\cdot\min_{a\in A,b\in B}\|a-b\|_{0}.

(This is sometimes called the ‘bichromatic’ ANN\mathsf{ANN} 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 ANN\mathsf{ANN} assuming SETH\mathsf{SETH}:

Assuming SETH, for every q>0q>0, there are ϵ∈(0,1)\epsilon\in(0,1) and C>0C>0 such that (1+ϵ)(1+\epsilon)-Approximate Hamming Nearest Neighbor Search in dimension d=Clog⁡nd=C\log n requires Ω(n2−q)\Omega(n^{2-q}) time.

We may assume that 4.3 holds even in the special case where each input vector from AA and BB has half its entries equal to and half equal to 11. Indeed, for any vector a∈{0,1}da\in\{0,1\}^{d}, we can construct a new vector a~∈{0,1}2d\widetilde{a}\in\{0,1\}^{2d} given by a~=[a⊤a‾⊤]⊤\widetilde{a}=\begin{bmatrix}a^{\top}&\overline{a}^{\top}\end{bmatrix}^{\top}. Here a‾∈{0,1}d\overline{a}\in\{0,1\}^{d} is the binary complement of vector aa, i.e., a‾i=1−ai\overline{a}_{i}=1-a_{i} for all i∈[d]i\in[d]. Thus, ∥a~∥0=d\|\widetilde{a}\|_{0}=d. We can similarly construct a new vector b~∈{0,1}2d\widetilde{b}\in\{0,1\}^{2d} for each b∈Bb\in B. After this transformation, for any a∈Aa\in A and b∈Bb\in B, we have ∥a~−b~∥0=2⋅∥a−b∥0\|\widetilde{a}-\widetilde{b}\|_{0}=2\cdot\|a-b\|_{0}, 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 Gap−ANN(n,d,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d,t,\epsilon).

Let n,dn,d denote two positive integers. Let t>0t>0 denote a threshold parameter. Let ϵ\epsilon denote a accuracy parameter. Given two sets of points A={a1,⋯ ,an}⊂{0,1}dA=\{a_{1},\cdots,a_{n}\}\subset\{0,1\}^{d} and B={b1,⋯ ,an}⊂{0,1}dB=\{b_{1},\cdots,a_{n}\}\subset\{0,1\}^{d}: For each i∈[n]i\in[n], we need to distinguish the following two cases

Case 1. There exists a j∈[n]j\in[n] such that ∥ai−bj∥0<t\|a_{i}-b_{j}\|_{0}<t.

Case 2. For all j∈[n]j\in[n] we have ∥ai−bj∥22≥(1+ϵ)⋅t\|a_{i}-b_{j}\|_{2}^{2}\geq(1+\epsilon)\cdot t.

An algorithm for Gap−ANN(n,d,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d,t,\epsilon) can be called log⁡(nd)\log(nd) times to binary search for the answer to ANN\mathsf{ANN}, so Lemma 4.3 holds as well for Gap−ANN(n,d,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d,t,\epsilon).

3 Hardness Result

In the remainder of this section, we prove our lower bound for attention computation:

Assuming SETH, for every sufficiently small q>0q>0, there are constants C>0C>0 and Cα>0C_{\alpha}>0 and Cβ>1C_{\beta}>1 such that Approximate Attention Computation AAttC\mathsf{AAttC} (Definition 1.2) for parameters (n,d=Clog⁡n,B=Cβlog⁡n,ϵa=n−Cα)(n,d=C\log n,B=C_{\beta}\sqrt{\log n},\epsilon_{a}=n^{-C_{\alpha}}) requires Ω(n2−q)\Omega(n^{2-q}) 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 Cγ∈(0,0.1)C_{\gamma}\in(0,0.1): For every ϵ>0\epsilon>0 and C>0C>0, there exist constants Ca>0C_{a}>0 and Cb>0C_{b}>0 and such that, if AAttC\mathsf{AAttC} (Definition 1.2) for parameters (2n,d=2Clog⁡n,B=Cblog⁡n,ϵa=n−Ca)(2n,d=2C\log n,B=C_{b}\sqrt{\log n},\epsilon_{a}=n^{-C_{a}}) can be solved in time TT, then Gap−ANN(n,d=Clog⁡n,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d=C\log n,t,\epsilon) (Definition 4.5) can be solved in time O(T+n2−Cγ)O(T+n^{2-C_{\gamma}}).

We give an algorithm with the stated running time for Gap−ANN(n,d=Clog⁡n,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d=C\log n,t,\epsilon). Let c>0c>0 be a parameter we will choose later (it will be a function of CC and CγC_{\gamma}). Our algorithm will proceed to one of two cases depending on the value of tt. If t<clog⁡nt<c\log n, then we will use one algorithm which runs in time O(n2−Cγ)O(n^{2-C_{\gamma}}). Otherwise, if t≥clog⁡nt\geq c\log n, we will use another algorithm which runs in time O(T)O(T).

Let a1,⋯ ,an,b1,⋯ ,bn∈{0,1}da_{1},\cdots,a_{n},b_{1},\cdots,b_{n}\in\{0,1\}^{d} be the input vectors to Gap−ANN{\sf Gap}{\rm-}{\sf ANN}, and let t∈[0,d]t\in[0,d] denote the target distance. Recall that d=Clog⁡nd=C\log n.

In this t<clog⁡nt<c\log n case, we will simply brute-force for the answer in the following way: We first store the vectors b1,⋯ ,bnb_{1},\cdots,b_{n} in a lookup table, then for each i∈[n]i\in[n], we iterate over all vectors b′∈{0,1}db^{\prime}\in\{0,1\}^{d} which have Hamming distance at most tt from aia_{i} and check whether b′b^{\prime} is in the lookup table. This determines whether there is a b∈Bb\in B at distance at most tt from aia_{i}, as desired.

For each i∈[n]i\in[n], we need to iterate over (dt){d\choose t} choices for the vector b′b^{\prime}, so the total running time will be O(n⋅(dt))O(n\cdot{d\choose t}). By standard bounds on binomial coefficients, we know that

We can thus pick a sufficiently small constant c>0c>0, depending only on CγC_{\gamma} and CC such that f(C,c)<1−Cγf(C,c)<1-C_{\gamma} and this entire brute-force takes O(n2−Cγ)O(n^{2-C_{\gamma}}) time.

Let a1,⋯ ,an,b1,⋯ ,bn∈{0,1}da_{1},\cdots,a_{n},b_{1},\cdots,b_{n}\in\{0,1\}^{d} denote the input of Gap−ANN(n,d,t,ϵ){\sf Gap}{\rm-}{\sf ANN}(n,d,t,\epsilon) (Definition 4.5), and recall from Remark 4.4 that we may assume each has half its entries and half its entries 11. We will explain how to construct an Attention matrix using this instance.

Let β>0\beta>0 and d~≥d\widetilde{d}\geq d denote parameters we will choose later (see Eq. (9) and Eq. (6), respectively). Define τ>0\tau>0 by

Intuitively, our goal in picking these parameters is that τ\tau 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 AAttC(n~,d~,B,ϵa)\mathsf{AAttC}(\widetilde{n},\widetilde{d},B,\epsilon_{a}) problem, for the following parameters:

Since each entry of QQ and KK is either β\sqrt{\beta} or , it follows that

(Note that we do not explicitly compute all the entries of AA in our algorithm; we will make use of it only through calling our algorithm for the Attention problem.)

For each (i,j)∈[n]×[n](i,j)\in[n]\times[n], we know that

On the other hand, we know that that for each (i,j)∈[n]×[n](i,j)\in[n]\times[n],

since it is an exponential of an entry of QK⊤/d~QK^{\top}/\widetilde{d}.

Using Eq. (4.3) and Eq. (11), combined with our expression for AA, it thus follows that

Since Di,i=(A1n~)iD_{i,i}=(A{\bf 1}_{\widetilde{n}})_{i}, thus we know that

We can show that t~≥ϵa\widetilde{t}\geq\epsilon_{a} as follows:

where the second step follows from simple algebra, the third step follows from t=C0log⁡nt=C_{0}\log n (Eq. (4)) and d=Clog⁡nd=C\log n (assumption in Lemma statement), the second step follows from choice of β\beta (Eq. (7)), and the sixth step follows from choice of CaC_{a} (Eq. (8)), and the last step follows from Eq. (8).

Therefore, for any (i,j)∈[n]×[n](i,j)\in[n]\times[n],

where the second step follows from ∥ai∥22=∥bj∥22=d/2\|a_{i}\|_{2}^{2}=\|b_{j}\|_{2}^{2}=d/2, and the last step follows from simple algebra.

Recall that our goal is to determine, for each i∈[n]i\in[n], whether there is a j∈[n]j\in[n] such that ∥ai−bj∥22≤t\|a_{i}-b_{j}\|_{2}^{2}\leq t, or whether ∥ai−bj∥22≥(1+ϵa)t\|a_{i}-b_{j}\|_{2}^{2}\geq(1+\epsilon_{a})t for all j∈[n]j\in[n]. We will show next that we can distinguish these two cases by seeing whether uiu_{i} is greater than or less than the value t~0:=2t~\widetilde{t}_{0}:=2\widetilde{t}.

If there exists an (i,j)∈[n]×[n](i,j)\in[n]\times[n] such that ∥ai−bj∥22≤t\|a_{i}-b_{j}\|_{2}^{2}\leq t, then

where the first step follows from 2d=d~2d=\widetilde{d} (see Eq. (6)).

where the second step follows from the definition of t~\widetilde{t} (see Eq. (12)), and the last step follows from the definition of t~0\widetilde{t}_{0}.

If for all (i,j)∈[n]×[n](i,j)\in[n]\times[n], we have ∥ai−bj∥22>t(1+ϵ)\|a_{i}-b_{j}\|_{2}^{2}>t(1+\epsilon), this implies

where the third step follows from definition of t~\widetilde{t} (see Eq. (12)), the forth step follows from the calculation in Eq. (4.3) below, and the last step follows from t~=t~0/2\widetilde{t}=\widetilde{t}_{0}/2.

Finally, b our choice of β\beta and tt, we can see that

where the first step follows t=C0log⁡nt=C_{0}\log n (Eq. (4)), the second step follows from d=Clog⁡nd=C\log n, and the third step follows from β=B2\beta=B^{2} (Eq. (9)) and the choice of BB (Eq. (7)). ∎

Conclusion

In this work, we showed that how quickly one can perform attention computation depends critically on the magnitude, BB, 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.

References