How to Capture Higher-order Correlations? Generalizing Matrix Softmax Attention to Kronecker Computation
Josh Alman, Zhao Song
Introduction
Large language models, such as Transformer , BERT , GPT-1 , GPT-2 , GPT-3 , PaLM , OPT , GPT-3.5, Bard, GPT-4 , Llama , Llama 2 and its successors, have gained immense importance and found a wide range of applications due to their ability to understand and generate human-like text. These models are trained on massive amounts of text data, enabling them to learn patterns, structures, and nuances of human language. They have applications in many areas, including understanding natural language, content generation, improved human-computer interaction, translation and multilingual communication, and rapid prototyping.
The fundamental computational structure at the core of LLMs is called an attention unit. When a length- input is given to the attention unit (like a sentence or paragraph of words), we embed it into three matrices (the query, key, and value token matrices) where each has rows and columns. Here is the feature dimension; one has in the long sentence regime. Mathematically, the attention unit computes , where is a diagonal matrix, denotes the length- vector with all entries equal to , and is applied entry-wise.
Intuitively, the attention unit is finding pairwise correlations between tokens in the input since it computes inner products between pairs of tokens when computing . However, if the input data has correlated triples of tokens, it is not clear an attention unit can detect this.
A recent and exciting work formalized this intuition. They defined a simple task about learning correlations between triples of words, and showed that attention units are unable to solve it. By contrast, they are able to solve the analogous problem of learning correlations between pairs of words. Toward resolving this, proposed a generalization of attention computation:
Given as input matrices , the goal is to construct another matrix
here denotes a length- vector whose entries are all ones.
One may naturally view as an tensor, which is why we call this a ‘tensor generalization’; this view will be important in our proofs below.
In this generalization, entries of the matrix now correspond to triples of tokens, so one may hope that this generalization can detect triple-wise correlations. And indeed, show that this is the case: the tensor generalization gets around their expressivity barrier and is able to detect correlations among triples of tokens.
A fundamental question arises naturally: how quickly can generalized attention computations be performed? The running time of attention computations is critically important, since it forms the time bottleneck of LLM training and inference. By generalizing attention to make it more expressive, have we also made it intractably slow?
To answer this question, we focus on an approximate version of the tensor attention computation problem. In practical applications, it is sufficient to approximately perform these computations , and this often helps lead to faster algorithms.
, , , ,
the other matrices are defined as in Definition 1.1 above.
We focus here on the natural setting with (so that we are modeling long sequences) and (so that one can combine the errors from attention computations over an entire network).
In the case of (non-tensor) attention, the computational complexity of exact and approximate attention computation is very well-understood. showed that the trivial time algorithm is essentially optimal for exact computation, assuming the Strong Exponential Time Hypothesis (). is a popular conjecture from fine-grained complexity which posits that one cannot substantially improve our current best algorithms for -SAT; see the survey for more details.
studied the approximate (non-tensor) attention problem and showed that its complexity depends on the magnitude of the entries of the matrices : If they are smaller than , then there is a fast algorithm running in time ; this near-linear time algorithm is essentially as fast as one could hope for. On the other hand, if they are at least , then there is no algorithm substantially faster than the trivial assuming . This theoretical result mirrors practical observations that bounded entries are essential for fast attention .
Our main results tightly resolve the computational complexity of the tensor generalization of attention. Generalizing the situation for (non-tensor) attention, we show that whether or not there is a fast algorithm for depends on the parameter , the magnitudes of the entries in the query, key, and value matrices.
We first show a lower bound, that when , it is impossible to design a truly subcubic-time algorithm (assuming ). Note that the straigtforward algorithm for this problem runs in cubic time, so our result shows that one cannot substantially improve on the straightforward algorithm when the entries have magnitude at least .
Assuming , for every , there are constants such that: there is no algorithm running in time for the problem .
Our second result is a new algorithm, showing that when , then there is an almost linear time algorithm for solving the problem.
There is an algorithm (Algorithm 1) that solves in time .
Our Theorems 1.3 and 1.4 together show that the complexity of has a very tight transition at . When is smaller than the threshold, the problem can be solved essentially as quickly as one could hope for, in time . Meanwhile, when is greater than the threshold, it is impossible to achieve a subcubic running time, no matter what algorithmic techniques are used (assuming ).
It is exciting that, even for the more expressive tensor generalization of attention, there is a near-linear time algorithm in the bounded entry regime. Interestingly, though, the bound must be smaller than for regular attention: for regular attention to have a near-linear time algorithm, it is necessary and sufficient that , whereas for tensor-based attention, we show it is necessary and sufficient that .
More generally, for any positive integer , we study a higher-order tensor generalization of attention which can detect -wise correlations. (Regular attention corresponds to and corresponds to .) For this problem, we further generalize our results to show that there is a near-linear time algorithm when the entries satisfy , and that the trivial time essentially cannot be beaten otherwise. This suggests an intriguing tradeoff between the boundedness of the entries, and the expressiveness of attention we can perform quickly: Given vectors corresponding to tokens for LLM training or inference, we let be the largest magnitude of an entry, then we select the largest for which , and we can quickly perform -th order attention computations for our tokens, but not higher-order attention.
Suppose we are given matrices and , our target is to construct another matrix
is the length- vector whose entries are all ones.
In Section 2, we provide a number of basic notations and definitions. In Section 3, we give a technique overview, summarizing our proofs for both our upper bound result and our lower bound result. In Section 4, we prove the key intermediate results for our lower bound result. Our upper bound result, and the remainder of our lower bound result, are proved in the Appendix.
Preliminary
Tensor Operations
Many of our proofs will involve manipulating tensors. Here we introduce three different tensor operations we will frequently use.
We note that a tensor can be written in the form like this if and only if its tensor rank is at most .
In this work, we will primarily use the following column-wise version of the Kronecker product.
Matrix Multiplication
Technique Overview
Generalizing prior work on the computational complexity of the attention problem to our tensor generalization requires overcoming a number of technical challenges. Here we summarize our approach, with an emphasis on differences with the prior work on (non-tensor) attention that we build on .
We begin by introducing a basic tool for manipulating computations involving the column-wise Kronecker product (see details in Lemma C.3 below). Define the following matrices.
Approximating DD
In order to perform generalized attention, we aim to compute the matrix Notice that the intermediate matrix has entries. We thus cannot compute it in subcubic time. We instead aim to use an implicit representation of an approximation of this matrix which can be quickly manipulated.
Toward this goal, we find appropriate matrices (which we discuss in more detail shortly) and formulate such that . Given the matrices , and using the above tool for , we can compute quickly in time.
Approximating AA
Finding approximating matrices U1,U2,U3U_{1},U_{2},U_{3}
Thus, it remains to find matrices which appropriately approximate and as above. We show how to efficiently find such matrices as long as the inputs have bounded entries. The key idea is to use the polynomial method, a key tool from prior work which allows one to find low-rank representations of matrices.
The method generally says that if is a low-rank matrix, and is a low-degree polynomial, then (where is applied entry-wise) also has relatively low rank. Furthermore, its low-rank decomposition can be found efficiently given the decomposition of . By applying this method where is an appropriate polynomial approximation of the function (see ), we get a low-rank approximation of .
This polynomial method approach was also taken in the prior work on (non-tensor) attention . Here we generalize it, showing that the same line of attack can be applied to low-rank tensors. Viewing interchangeably as both an tensor and an matrix allows us to take advantage of this low-rank tensor approximation as well as the aforementioned matrix multiplication algorithms. See details in Lemma E.1.
2 Hardness
Our hardness proof proceeds by introducing and considering a new intermediate problem we call (Definition 4.6). In this problem, one is given as input vectors as well as a threshold , and the goal is to distinguish between the cases
for all , or
for some .
We first prove that cannot be solved in truly subcubic time assuming . We then show that a truly subcubic time algorithm for our generalized (Definition 1.2) problem with large entries would yield one for as well.
Previous work on (non-tensor) attention used as its intermediate problem the approximate Hamming Nearest Neighbor problem. However, it is not obvious how to directly generalize this to the tensor setting, since there is no way to define a ‘distance’ function for triples of vectors which satisfies the needed properties to generalize the original proof. We instead investigate the problem, which can itself be seen as a generalization of an intermediate step in the proof of hardness for approximate Hamming Nearest Neighbor .
Hardness of 𝖦𝖺𝗉−𝖬𝖺𝗑𝖨𝖯{\sf Gap}{\rm-}{\sf MaxIP}
Fine-grained complexity results for approximation problems like have previously been shown using a distributed probabilistically checkable proof framework , which we also use here.
We begin by generalizing the approach of using Merlin-Arthur (MA) communication protocols (). We construct a four party communication protocol for the disjointness problem: Alice, Bob and Charlie are each given subsets of a universe, and want to determine whether there is an element in all three of their sets. In an MA protocol, Merlin first sends an advice string to the three players to convince them their sets are disjoint. Alice, Bob and Charlie may then flip private random coins and communicate to come to an answer. (See details in Theorem 4.5).
Generalizing known three-party protocols for disjointness , our protocol is algebraic in nature, and critically makes use of algebraic geometry codes from coding theory .
We then use this protocol to reduce from to . A standard reduction shows that reduces to the problem, which is a computational version of the three player disjointness problem. We can convert inputs to this problem into vectors by corresponding entries of the vectors to possible transcripts of the communication protocol. The gap in inner products will arise naturally from the correctness guarantees of the protocol. See reduction details in Theorem 4.7 and its proofs.
Reducing from 𝖦𝖺𝗉−𝖬𝖺𝗑𝖨𝖯{\sf Gap}{\rm-}{\sf MaxIP} to 𝖠𝖳𝖠𝗍𝗍𝖢\mathsf{ATAttC}
Finally, we reduce the (Definition 4.6) problem to our (Definition 1.2) problem. The key idea is that, by defining the matrices of generalized attention in terms of the inputs to , we can make large entries of the attention matrix correspond to the triples with largest inner product. Some manipulation similar to prior work allows us to detect large entries from the output of . This approach has been used for the fine-grained hardness of many attention and kernel density estimation problems . See details in Lemma B.1 and its proofs.
Hardness
In this section, we begin the formal proof of our hardness result. We begin by introducing the fine-grained hypotheses we will use.
For every there exists an integer such that on formulas with clauses size at most (the so called - problem) and variables cannot be solved in time even by a randomized algorithm.
Given three sets where , the goal is to find a tuple such that .
For every , there is a such that cannot be solved in time on instances with .
It is known that implies ; see, e.g., .
We state a important tool from the field of algebraic geometry codes. For more background on algebraic geometry codes, we refer the reader to .
3-way Polynomial Closure. and are linear codes. For each , there exists such that for each ,
Efficiency. Both codes can be encoded in time and checked in time.
Parameters. Both codes have relative rate at least and relative distance at least .
2 A Four Party MA Communication Protocol
Prior work () constructed a protocol for three party communication, which includes Merlin, Alice and Bob. Here we modify this protocol for four parties.
For any . There is a -communication protocol for Set Disjointness over universe . This protocol is computationally efficient.
In particular, the details of protocol are
Merlin sends Alice bits
Alice, Bob, Charlie toss coins
If the three sets do not have any element in common, Alice always accepts. Otherwise, she accepts with probability at most .
We assume that divides , i.e., there is some positive integer such that . Otherwise, increase to the nest multiple of ; this at most doubles . We partition the universe into disjoint sets of size :
Let denote the inputs of Alice, Bob, and Charlie. Our goal is to determine whether there is an element in the intersection .
For each , we define the -th parts of the three sets:
Let be the rate and distance of the code; recall these are at least a positive constant. Let be the length of the codewords of .
For each , we write , , to denote the encodings of , and . Thus, their entry-wise product ( i.e., ) is a codeword in the second code . Furthermore, since is a linear code, the entry-wise sum of the ’s () is also a codeword of ’.
is a systematic code, so we may assume that for each , the entries are from and represent membership in the set. Similarly, , and the sets are disjoint if and only if for all and , or equivalently, for all .
Step 1. Merlin sends Alice , which is supposed to be the encoding of
Step 2. Charlie, Bob and Alice pick a random
Step 3. Charlie sends Alice for all
Step 4. Bob sends Alice for all
Step 5. Alice accepts iff all of the following hold:
is a codeword in
for all
First, we observe that Merlin’s message length is , and both Bob and Charlie’s message lengths are , as desired. To see correctness, note that if Alice ever accepts given Merlin’s message , then must in particular be a codeword of . If Alice accepts with probability greater than (where is a positive constant) then is also equal to the true by definition of . This means so the sets are disjoint.
3 Showing 3-𝖬𝖠𝖷\mathsf{MAX}-𝖨𝖯\mathsf{IP} is hard
We now define the appropriate gap -- problem, which we use as our intermediate hard problem.
We use to represent a threshold parameter.
We use to represent an accuracy parameter.
Suppose denote two positive integers.
For every index , we need to distinguish the following two cases
Case 1. There exists a pair such that .
Case 2. For all pairs we have .
Implicit in previous work () is a proof that the analogue of with two sets of points is hard. Here we generalize this to three sets.
Unless and are false, the following holds: for every there are constants such that for integer , solving requires time .
We reduce from to . Let . Our reduction takes as input an instance of orthogonal vectors over . These sets have sizes for a constant depending on from Definition 4.2 and Conjecture 4.3, and posits there is no algorithm solving this problem in time .
For a constant to be determined, pick to be a constant such that
We use the protocol of Theorem 4.5, instantiated with parameter
Suppose that is representing the number of different possible messages sent by Bob and Charlie in the protocol.
Let us choose so that .
For each vector , we construct a new vector by setting iff Charlie send message on input and randomness . (The value is independent of .)
For each vector , we construct a new vector by setting iff Bob sends message on input and randomness . (The value is independent of .)
For each Merlin-message and vector , we construct a new vector as follows: iff Alice accepts on
message from Bob, message from Charlie, and randomness .
Notice also that the inner product of three vectors
is exactly proportional to the probability that Alice, Bob and Charlie accept on inputs , and message from Merlin.
In particular, if , and are orthogonal (i.e., ), the inner product is at most
Otherwise, there exists a such that
In particular, these can be distinguished by an algorithm for
which must therefore be as hard as solving the original instance of . By , this means it requires time
where the last step follows from choosing large enough in the definition of .
At the end, we notice that the vectors we construct have dimension as desired. ∎
Acknowledgements
The authors would like to thank Yichuan Deng, Yeqi Gao, Junze Yin, Lichen Zhang, Ruizhe Zhang, Tianyi Zhou for helpful discussions of attention literature.
Appendix
In Section A, we provide the definitions of several notations. In Section C, we provide the running time proofs for our upper bound result. In Section D, we provide the error analysis for our upper bound result. In Section E, we combine everything together, and also present our algorithm. In Section B. we show how to reduce our problem to .
Appendix A Preliminary
For any positive integer , we write to denote .
We use to denote a length- vector whose entries are all ones.
Given three sets of vectors where , the goal is to compute
Appendix B Hardness: From 𝖬𝖺𝗑𝖨𝖯\mathsf{MaxIP} to Our Problem
In Section B.1, we show how to reduce our problem to . In Section B.2, we present our main lower bound (hardness) result.
We now generalize the hardness proof of to the tensor attention case.
For every constant , every , and every , there exist constants and and such that, if (Definition D.1) for parameters can be solved in time , then (Definition 4.6) can be solved in time .
We give an algorithm for (Definition 4.6). Let denote the inputs to this problem. Using them, we will construct appropriate inputs to the problem so that its output will help us to detect triples with large inner product.
Let and be parameters to be determined (in Eq. (5) and Eq. (2) below). Define by
We pick these parameters so that will be an upper bound on entries of the attention matrix, namely
We will use an algorithm for the problem with parameters:
Since each entry of and , is either or , it follows that
is naturally partitioned into eight submatrices
where each () is a matrix of size , defined as follows. For each ,
The -th entry of is
The -th entry of is
The -th entry of is
The -th entry of is
The -th entry of is
The -th entry of is
The -th entry of is
The -th entry of is
For each , we know that
Here we used the fact that (see Eq. (2)), and the last step uses the definition of (see Eq. (1)).
We also know that for each ,
since it is the exponential of an entry of .
By combining our expression for with Eq. (B.1) and Eq. (7), we see that
Since , it follows that
We can see that as follows:
Here, the last two steps follow from Eq. (4).
Recall that for each we need to distinguish between two cases: either there is a pair such that , or else for all pairs the inner product . We will distinguish between these cases by checking whether is greater than a threshold value . We next consider the two cases to see why this is.
For a given , if there are such that , then
where the 1st step follows from (see Eq. (2)). This means as desired that
For a given , if for all we have , then
Here, the 4th step follows because, by our choice of and , we have
where we used that (by Lemma statement), that , that (Eq. (5)) and the choice of (Eq. (3)). ∎
B.2 Main Hardness Result
We can finally conclude our main lower bound.
Assuming , for every , there are constants such that: there is no algorithm running in time for the problem .
Follows from combining Theorem 4.5, Theorem 4.7, and Lemma B.1. ∎
Appendix C Upper Bound: Running Time
In Section C.1, we review the standard “matrix” attention computation problem. In Section C.2, we define the “tensor” attention computation problem. In Section C.3, we provide an efficient tool for implementing tensor related computations. In Section C.4, we provide several tools for rearranging tensor computations that we will use in our algorithm.
We first review the attention computation definition in ,
C.2 Tensor Attention Computation
Given two matrices, there are two standard variants on their Kronecker product one may consider: The standard Kronecker product (denoted ) is a new matrix, whereas the column-wise Kronecker product (denoted ) is a new matrix. For more literature on tensor computations and their applications in learning algorithms, we refer the readers to .
Next, we generalize the matrix attention computation (in ) into tensor attention computation as follows:
C.3 Efficient Column-wise Kronecker Computation
We prove an important tool which will be used in analyze the running time of our algorithm.
Let be defined as Definition 2.4.
We define
For each , let denote the -th row of .
For each , let denote the -th row of .
From the above, we can calculate that the entry of in location is
where the first step follows from Eq. (C.3), the second step follows from simple algebra, the third step follows from separating the summation over and the summation over , and the last step follows from definition of matrices and .
C.4 Simple Equivalent Tools for Tensor Notations
We define a standard tensor notation, for example see .
Next, we present several equivalence results for tensors.
Let be defined as Definition 2.4.
Let be defined as Definition 2.2.
Let operator be defined as Definition C.4.
Let be defined as Definition 2.1.
Part 1. for (This means can be viewed as the tensor version of )
Part 2.
Part 3.
Directly follows from definition of and .
Follows from tensor notations in Definition 2.2 and Definition C.4.
Directly follows from applying Part 1 of Lemma C.3 here. ∎
Appendix D Upper Bound: Error Analysis
In Section D.1, we provide the definition of approximate tensor attention computation. In Section D.2, we state a polynomial approximation tool from previous work. In Section D.3, we show a bound on the entries of the attention matrix. In Section D.4, we provide a low-rank decomposition for the tensor version of the attention matrix. Finally, in Section D.5, we compute the error propagation from to , then in Section D.6, we analyze the error propagation from and to the attention matrix.
, , , ,
Notice that the straightforward algorithm for this problem will spend at least time to write the matrix (we can also think of as an tensor that has size ).
D.2 An Error Control Tool From Previous Work
Let .
D.3 Tensor Q⊙K1⊙K2Q\odot K_{1}\odot K_{2} Has Bounded Entries
Let operation be defined as Definition 2.2.
For every index triple , we are able to prove
D.4 Tensor Low-Rank Approximation
In the following definition, we view the size matrix as an size attention matrix.
We use to denote a positive integer.
We use to represent an accuracy parameter.
for all .
D.5 From AA to DD
In this section and the next, we generalize the proof of for error propagation from the matrix setting to the tensor setting. The proofs are nearly identical.
Then, for every index , the following bound holds
where the second step follows from triangle inequality.
D.6 From A,DA,D to Tensor Attention
The goal of this section is to prove Lemma D.6.
Suppose the following conditions are true
Let
We bound each of these two terms to get our desired result.
First of all, for every index pair ,
Second, for every ,
Here, again, the 2nd step uses the triangle inequality, the 3rd step follows because is positive, the 4th step follows from the assumption that and the final step follows by definition of .
The Lemma conclusion then becomes true by substituting Eq. (D.6) and Eq. (D.6) into Eq. (11). ∎
Appendix E Upper Bound: Putting It All Together
In Section E.1, we provide a low-rank decomposition for approximating the original attention tensor. In Section E.2, we calculate the running time of constructing that low-rank decomposition. In Section E.3, we put everything together, and prove our main upper bound theorem.
The goal of this section is to prove Lemma E.1.
Let denote a degree- single-variable polynomial. We apply entry-wisely, i.e, .
Let be rank parameter that is
We define set and provide names for variables in set in the following sense,
Thus, function can be viewed a degree- polynomial in the entries in of the vectors .
Here, we can view function as
Therefore, we should construct three matrices and as the following way, for each
These matrices can be constructed in time in the straightforward way, since each entry depends on variables. ∎
E.2 Time for Constructing U1,U2,U3U_{1},U_{2},U_{3}
, ,
, and .
For bounded number and accuracy parameter , there are positive integers and
it takes time to construct and .
Recall that the definition of -approximation can be found in Definition D.4.
We can then compute , and using Lemma E.1, which gives the bound
E.3 Main Algorithmic Result
We present our main algorithmic result as follows:
There is an algorithm (Algorithm 1) that solves in time .
Using Lemma E.2, we know that Step 1 (in Algorithm 1) can be implemented in time
Using Lemma C.3, we know that Step 2 (in Algorithm 1) can be implemented in time
Step 3 can implemented in in a straightforward way.
To compute Step 4 efficiently, we need to use Lemma C.3 again.
Computing Step 5 is just standard matrix multiplication
Step 6 is just rescaling the matrix
We combine Corollary D.2, Lemma D.5, Lemma D.6, and simple algebra. ∎