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-nn input is given to the attention unit (like a sentence or paragraph of nn words), we embed it into three matrices Q,K,VQ,K,V (the query, key, and value token matrices) where each has nn rows and dd columns. Here dd is the feature dimension; one has d≪nd\ll n in the long sentence regime. Mathematically, the attention unit computes D−1exp⁡(QK⊤)VD^{-1}\exp(QK^{\top})V, where D=\diag(exp⁡(QK⊤)1n)D=\diag(\exp(QK^{\top}){\bf 1}_{n}) is a diagonal matrix, 1n{\bf 1}_{n} denotes the length-nn vector with all entries equal to 11, and exp⁡\exp 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 QK⊤QK^{\top}. 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 n×dn\times d matrices Q,K1,K2,V1,V2Q,K_{1},K_{2},V_{1},V_{2}, the goal is to construct another n×dn\times d matrix

1n2{\bf 1}_{n^{2}} here denotes a length-n2n^{2} vector whose entries are all ones.

One may naturally view AA as an n×n×nn\times n\times n 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 AA 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.

∥Q∥∞≤B\|Q\|_{\infty}\leq B, ∥K1∥∞≤B\|K_{1}\|_{\infty}\leq B, ∥K2∥∞≤B\|K_{2}\|_{\infty}\leq B, ∥V1∥∞≤B\|V_{1}\|_{\infty}\leq B, ∥V2∥∞≤B\|V_{2}\|_{\infty}\leq B

the other matrices are defined as in Definition 1.1 above.

We focus here on the natural setting with d=O(log⁡n)d=O(\log n) (so that we are modeling long sequences) and ϵa=1/\poly(n)\epsilon_{a}=1/\poly(n) (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 O(n2)O(n^{2}) time algorithm is essentially optimal for exact computation, assuming the Strong Exponential Time Hypothesis (\SETH\SETH). SETH\mathsf{SETH} is a popular conjecture from fine-grained complexity which posits that one cannot substantially improve our current best algorithms for kk-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 Q,KQ,K: If they are smaller than o(log⁡n)o(\sqrt{\log n}), then there is a fast algorithm running in time n1+o(1)n^{1+o(1)}; this near-linear time algorithm is essentially as fast as one could hope for. On the other hand, if they are at least Ω(log⁡n)\Omega(\sqrt{\log n}), then there is no algorithm substantially faster than the trivial O(n2)O(n^{2}) assuming \SETH\SETH. 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 AAttC\mathsf{AAttC} depends on the parameter BB, the magnitudes of the entries in the query, key, and value 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 subcubic-time algorithm (assuming \SETH\SETH). 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 Ω(log⁡n)\Omega(\sqrt{\log n}).

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 algorithm running in time O(n3−q)O(n^{3-q}) for the problem ATAttC(n,d=Clog⁡n,B=Cblog⁡n,ϵa=n−Ca)\mathsf{ATAttC}(n,d=C\log n,B=C_{b}\sqrt{\log n},\epsilon_{a}=n^{-C_{a}}).

Our second result is a new algorithm, showing that when B<o(log⁡n)B<o(\sqrt{\log n}), then there is an almost linear time algorithm for solving the problem.

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

Our Theorems 1.3 and 1.4 together show that the complexity of ATAttC\mathsf{ATAttC} has a very tight transition at B=Θ(log⁡n)B=\Theta(\sqrt{\log n}). When B<o(log⁡n)B<o(\sqrt{\log n}) is smaller than the threshold, the problem can be solved essentially as quickly as one could hope for, in time n1+o(1)n^{1+o(1)}. Meanwhile, when B≥Ω(log⁡n)B\geq\Omega(\sqrt{\log n}) is greater than the threshold, it is impossible to achieve a subcubic running time, no matter what algorithmic techniques are used (assuming SETH\mathsf{SETH}).

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 B<log⁡nB<\sqrt{\log n}, whereas for tensor-based attention, we show it is necessary and sufficient that B<log⁡nB<\sqrt{\log n}.

More generally, for any positive integer k≥2k\geq 2, we study a higher-order tensor generalization of attention which can detect kk-wise correlations. (Regular attention corresponds to k=2k=2 and ATAttC\mathsf{ATAttC} corresponds to k=3k=3.) For this problem, we further generalize our results to show that there is a near-linear time algorithm when the entries satisfy B<log⁡nkB<\sqrt[k]{\log n}, and that the trivial O(nk)O(n^{k}) 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 BB be the largest magnitude of an entry, then we select the largest kk for which B<log⁡nkB<\sqrt[k]{\log n}, and we can quickly perform kk-th order attention computations for our tokens, but not higher-order attention.

Suppose we are given n×dn\times d matrices Q,K1,K2,⋯ ,Kk−1Q,K_{1},K_{2},\cdots,K_{k-1} and V1,V2,⋯ ,Vk−1V_{1},V_{2},\cdots,V_{k-1}, our target is to construct another n×dn\times d matrix

1nk−1{\bf 1}_{n^{k-1}} is the length-nk−1n^{k-1} 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 TT can be written in the form A⊙B⊙CA\odot B\odot C like this if and only if its tensor rank is at most dd.

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 ⊘\oslash (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 D=\diag(exp⁡(Q(K1⊘K2)⊤/d)1n2).D=\diag(\exp(Q(K_{1}\oslash K_{2})^{\top}/d){\bf 1}_{n^{2}}). Notice that the intermediate matrix exp⁡(Q(K1⊘K2)⊤/d)\exp(Q(K_{1}\oslash K_{2})^{\top}/d) has n3n^{3} 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 U1,U2,U3U_{1},U_{2},U_{3} (which we discuss in more detail shortly) and formulate D~=\diag(U1(U2⊘U3)⊤1n2)\widetilde{D}=\diag(U_{1}(U_{2}\oslash U_{3})^{\top}{\bf 1}_{n^{2}}) such that D~≈D\widetilde{D}\approx D. Given the matrices U1,U2,U3U_{1},U_{2},U_{3}, and using the above tool for ⊘\oslash, we can compute D~\widetilde{D} quickly in O(nd)O(nd) time.

Approximating AA

Finding approximating matrices U1,U2,U3U_{1},U_{2},U_{3}

Thus, it remains to find matrices U1,U2,U3U_{1},U_{2},U_{3} which appropriately approximate DD and AA as above. We show how to efficiently find such matrices as long as the inputs Q,K1,K1,V1,V2Q,K_{1},K_{1},V_{1},V_{2} 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 MM is a low-rank matrix, and pp is a low-degree polynomial, then p(M)p(M) (where pp is applied entry-wise) also has relatively low rank. Furthermore, its low-rank decomposition can be found efficiently given the decomposition of MM. By applying this method where pp is an appropriate polynomial approximation of the exp⁡\exp function (see ), we get a low-rank approximation of exp⁡(M)\exp(M).

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 AA interchangeably as both an n×n×nn\times n\times n tensor and an n×n2n\times n^{2} 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} (Definition 4.6). In this problem, one is given as input 3n3n vectors a1,…,an,b1,…,bn,c1,…,cn∈{0,1}da_{1},\ldots,a_{n},b_{1},\ldots,b_{n},c_{1},\ldots,c_{n}\in\{0,1\}^{d} as well as a threshold tt, and the goal is to distinguish between the cases

⟨ai,bj,ck⟩≤t\langle a_{i},b_{j},c_{k}\rangle\leq t for all i,j,k∈[n]i,j,k\in[n], or

⟨ai,bj,ck⟩≥2t\langle a_{i},b_{j},c_{k}\rangle\geq 2t for some i,j,k∈[n]i,j,k\in[n].

We first prove that Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} cannot be solved in truly subcubic time assuming \SETH\SETH. We then show that a truly subcubic time algorithm for our generalized ATAttC\mathsf{ATAttC} (Definition 1.2) problem with large entries would yield one for Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} 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 SAT\mathsf{SAT} to Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP}. A standard reduction shows that SAT\mathsf{SAT} reduces to the 3OV\mathsf{3OV} 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} (Definition 4.6) problem to our ATAttC\mathsf{ATAttC} (Definition 1.2) problem. The key idea is that, by defining the matrices Q,K1,K2,V1,V2Q,K_{1},K_{2},V_{1},V_{2} of generalized attention in terms of the inputs to Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP}, we can make large entries of the attention matrix AA correspond to the triples with largest inner product. Some manipulation similar to prior work allows us to detect large entries from the output of ATAttC\mathsf{ATAttC}. 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 ϵ>0\epsilon>0 there exists an integer k≥3k\geq 3 such that CNF−SAT\mathsf{CNF}-\mathsf{SAT} on formulas with clauses size at most kk (the so called kk-SAT\mathsf{SAT} problem) and nn variables cannot be solved in O(2(1−ϵ)n)O(2^{(1-\epsilon)n}) time even by a randomized algorithm.

Given three sets A,B,C⊂{0,1}dA,B,C\subset\{0,1\}^{d} where ∣A∣=∣B∣=∣C∣=n|A|=|B|=|C|=n, the goal is to find a tuple (i1,i2,i3)∈[n]×[n]×[n](i_{1},i_{2},i_{3})\in[n]\times[n]\times[n] such that ⟨ai1,bi2,ci3⟩=0\langle a_{i_{1}},b_{i_{2}},c_{i_{3}}\rangle=0.

For every ϵ>0\epsilon>0, there is a c≥1c\geq 1 such that 3OV\mathsf{3OV} cannot be solved in n3−ϵn^{3-\epsilon} time on instances with d=clog⁡nd=c\log n.

It is known that \SETH\SETH implies 3OVC\mathsf{3OVC}; 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. C{\cal C} and C′{\cal C}^{\prime} are linear codes. For each w1,w2,w3∈Cw_{1},w_{2},w_{3}\in{\cal C}, there exists w′∈C′w^{\prime}\in{\cal C}^{\prime} such that for each i∈Rni\in{\cal R}_{n}, w′(i)=w1(i)⋅w2(i)⋅w3(i)w^{\prime}(i)=w_{1}(i)\cdot w_{2}(i)\cdot w_{3}(i)

Efficiency. Both codes can be encoded in \poly(n)\poly(n) time and checked in \poly(n)\poly(n) time.

Parameters. Both codes have relative rate at least 0.010.01 and relative distance at least 0.010.01.

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 T∈[2,m]T\in[2,m]. There is a MA\mathsf{MA}-communication protocol for Set Disjointness over universe [m][m]. This protocol is computationally efficient.

In particular, the details of protocol are

Merlin sends Alice O(mlog⁡TT)O(\frac{m\log T}{T}) bits

Alice, Bob, Charlie toss O(log⁡m)O(\log m) coins

If the three sets do not have any element in common, Alice always accepts. Otherwise, she accepts with probability at most 1/21/2.

We assume that TT divides mm, i.e., there is some positive integer rr such that m=Trm=Tr. Otherwise, increase mm to the nest multiple of TT; this at most doubles mm. We partition the universe into TT disjoint sets of size rr: [m]=U1∪⋯∪UT.[m]=U^{1}\cup\cdots\cup U^{T}.

Let α,β,γ⊆[m]\alpha,\beta,\gamma\subseteq[m] denote the inputs of Alice, Bob, and Charlie. Our goal is to determine whether there is an element in the intersection α∩β∩γ\alpha\cap\beta\cap\gamma.

For each t∈[T]t\in[T], we define the tt-th parts of the three sets:

Let ρC,δC\rho_{C},\delta_{C} be the rate and distance of the code; recall these are at least a positive constant. Let nC=mT⋅ρC=O(m/T)n_{C}=\frac{m}{T\cdot\rho_{C}}=O(m/T) be the length of the codewords of CC.

For each t∈[T]t\in[T], we write C(αt)C(\alpha^{t}), C(βt)C(\beta^{t}), C(γt)C(\gamma^{t}) to denote the encodings of αt\alpha^{t}, βt\beta^{t} and γt\gamma^{t}. Thus, their entry-wise product μt\mu^{t} ( i.e., μit:=C(αt)i⋅C(βt)i⋅C(γt)i\mu_{i}^{t}:=C(\alpha^{t})_{i}\cdot C(\beta^{t})_{i}\cdot C(\gamma^{t})_{i} ) is a codeword in the second code C′C^{\prime}. Furthermore, since C′C^{\prime} is a linear code, the entry-wise sum of the μt\mu^{t}’s (μi=∑t=1Tμit\mu_{i}=\sum_{t=1}^{T}\mu_{i}^{t}) is also a codeword of CC’.

CC is a systematic code, so we may assume that for each i∈[n/T]i\in[n/T], the entries C(αt)i,C(βt)i,C(γt)iC(\alpha^{t})_{i},C(\beta^{t})_{i},C(\gamma^{t})_{i} are from {0,1}\{0,1\} and represent membership in the set. Similarly, μit∈{0,1}\mu_{i}^{t}\in\{0,1\}, and the sets are disjoint if and only if μit=0\mu_{i}^{t}=0 for all i∈[m/T]i\in[m/T] and t∈[T]t\in[T], or equivalently, μi=0\mu_{i}=0 for all i∈[m/T]i\in[m/T].

Step 1. Merlin sends Alice μ^\widehat{\mu}, which is supposed to be the encoding of μ\mu

Step 2. Charlie, Bob and Alice pick a random i∗∈[nC]i^{*}\in[n_{C}]

Step 3. Charlie sends Alice C(γt)i∗C(\gamma^{t})_{i^{*}} for all t∈[T]t\in[T]

Step 4. Bob sends Alice C(βt)i∗C(\beta^{t})_{i^{*}} for all t∈[T]t\in[T]

Step 5. Alice accepts iff all of the following hold:

μ^\widehat{\mu} is a codeword in C′C^{\prime}

μ^i∗=∑t=1TC(αt)i∗⋅C(βt)i∗⋅C(γt)i∗\widehat{\mu}_{i^{*}}=\sum_{t=1}^{T}C(\alpha^{t})_{i^{*}}\cdot C(\beta^{t})_{i^{*}}\cdot C(\gamma^{t})_{i^{*}}

μ^i=0\widehat{\mu}_{i}=0 for all i∈[m/T]i\in[m/T]

First, we observe that Merlin’s message length is nc⋅log⁡T=O((log⁡T)⋅m/T)n_{c}\cdot\log T=O((\log T)\cdot m/T) , and both Bob and Charlie’s message lengths are T⋅O(log⁡T)T\cdot O(\log T), as desired. To see correctness, note that if Alice ever accepts given Merlin’s message μ^\widehat{\mu}, then μ^\widehat{\mu} must in particular be a codeword of C′C^{\prime}. If Alice accepts with probability greater than 1−δC′1-\delta_{C^{\prime}} (where δC′\delta_{C^{\prime}} is a positive constant) then μ^\widehat{\mu} is also equal to the true μ\mu by definition of δC′\delta_{C^{\prime}}. This means μi=0,∀i∈[m/T],\mu_{i}=0,\forall i\in[m/T], so the sets are disjoint.

3 Showing 3-𝖬𝖠𝖷\mathsf{MAX}-𝖨𝖯\mathsf{IP} is hard

We now define the appropriate gap 33-MAX\mathsf{MAX}-IP\mathsf{IP} problem, which we use as our intermediate hard problem.

We use t>0t>0 to represent a threshold parameter.

We use ϵ\epsilon to represent an accuracy parameter.

Suppose n,dn,d denote two positive integers.

A={a1,⋯ ,an}⊂{0,1}dA=\{a_{1},\cdots,a_{n}\}\subset\{0,1\}^{d}

B={b1,⋯ ,bn}⊂{0,1}dB=\{b_{1},\cdots,b_{n}\}\subset\{0,1\}^{d}

C={c1,⋯ ,cn}⊂{0,1}dC=\{c_{1},\cdots,c_{n}\}\subset\{0,1\}^{d}

For every index i∈[n]i\in[n], we need to distinguish the following two cases

Case 1. There exists a pair (j1,j2)∈[n]×[n](j_{1},j_{2})\in[n]\times[n] such that ⟨ai,bj1,cj2⟩≥t\langle a_{i},b_{j_{1}},c_{j_{2}}\rangle\geq t.

Case 2. For all pairs (j1,j2)∈[n]×[n](j_{1},j_{2})\in[n]\times[n] we have ⟨ai,bj1,cj2⟩≤(1−ϵ)⋅t\langle a_{i},b_{j_{1}},c_{j_{2}}\rangle\leq(1-\epsilon)\cdot t.

Implicit in previous work () is a proof that the analogue of Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP} with two sets of points is hard. Here we generalize this to three sets.

Unless SETH\mathsf{SETH} and OVC\mathsf{OVC} are false, the following holds: for every δ>0\delta>0 there are constants α1>α2>0\alpha_{1}>\alpha_{2}>0 such that for integer nn, solving Gap−MaxIP(n,d=α1log⁡n,t=α2log⁡n,ϵ=1/2){\sf Gap}{\rm-}{\sf MaxIP}(n,d=\alpha_{1}\log n,t=\alpha_{2}\log n,\epsilon=1/2) requires time Ω(n3−δ)\Omega(n^{3-\delta}).

We reduce from 3\OV3\OV to Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP}. Let δ\OV=δ/2\delta_{\OV}=\delta/2. Our reduction takes as input an instance (A\OV,B\OV,C\OV)(A_{\OV},B_{\OV},C_{\OV}) of orthogonal vectors over {0,1}m\{0,1\}^{m}. These sets have sizes ∣A\OV∣=∣B\OV∣=∣C\OV∣=2m/c|A_{\OV}|=|B_{\OV}|=|C_{\OV}|=2^{m/c} for a constant cc depending on δ\OV\delta_{\OV} from Definition 4.2 and Conjecture 4.3, and 3OVC\mathsf{3OVC} posits there is no algorithm solving this problem in time O((2m/c)3−δ\OV)O((2^{m/c})^{3-\delta_{\OV}}).

For a constant k>0k>0 to be determined, pick ϵ>0\epsilon>0 to be a constant such that

We use the protocol of Theorem 4.5, instantiated with parameter

Suppose that T′=2O((log⁡T)⋅T)T^{\prime}=2^{O((\log T)\cdot T)} is representing the number of different possible messages sent by Bob and Charlie in the protocol.

Let us choose TT so that T′=O(1/ϵ)T^{\prime}=O(1/\epsilon).

For each vector γ∈C\OV\gamma\in C_{\OV}, we construct a new vector c~γ∈{0,1}(T′)2×m\widetilde{c}^{\gamma}\in\{0,1\}^{(T^{\prime})^{2}\times m} by setting c~iB,iC,jγ:=1\widetilde{c}^{\gamma}_{i_{B},i_{C},j}:=1 iff Charlie send message iC∈[T′]i_{C}\in[T^{\prime}] on input γ′\gamma^{\prime} and randomness j∈[m]j\in[m]. (The value is independent of iBi_{B}.)

For each vector β∈B\OV\beta\in B_{\OV}, we construct a new vector b~β∈{0,1}(T′)2×m\widetilde{b}^{\beta}\in\{0,1\}^{(T^{\prime})^{2}\times m} by setting b~iB,iC,jβ:=1\widetilde{b}^{\beta}_{i_{B},i_{C},j}:=1 iff Bob sends message iB∈[T′]i_{B}\in[T^{\prime}] on input β′\beta^{\prime} and randomness j∈[m]j\in[m]. (The value is independent of iCi_{C}.)

For each Merlin-message μ∈{0,1}O((log⁡T)⋅m/T)\mu\in\{0,1\}^{O((\log T)\cdot m/T)} and vector α∈A\OV\alpha\in A_{\OV}, we construct a new vector a~μ,α∈{0,1}(T′)2×m\widetilde{a}^{\mu,\alpha}\in\{0,1\}^{(T^{\prime})^{2}\times m} as follows: a~iB,iC,jμ,α:=1\widetilde{a}^{\mu,\alpha}_{i_{B},i_{C},j}:=1 iff Alice accepts on

message iBi_{B} from Bob, message iCi_{C} from Charlie, and randomness jj.

Notice also that the inner product of three vectors

is exactly proportional to the probability that Alice, Bob and Charlie accept on inputs α,β\alpha,\beta, γ\gamma and message μ\mu from Merlin.

In particular, if α\alpha, β\beta and γ\gamma are orthogonal (i.e., ⟨α,β,γ⟩=0\langle\alpha,\beta,\gamma\rangle=0), the inner product is at most

Otherwise, there exists a μ\mu such that

In particular, these can be distinguished by an algorithm for

which must therefore be as hard as solving the original instance of 3OV3\mathsf{OV}. By 3OVC\mathsf{3OVC}, this means it requires time

where the last step follows from choosing kk large enough in the definition of ϵ\epsilon.

At the end, we notice that the vectors we construct have dimension 2(T′)2⋅m=O(m)=O(log⁡n)2(T^{\prime})^{2}\cdot m=O(m)=O(\log n) 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP}.

Appendix A Preliminary

For any positive integer nn, we write [n][n] to denote {1,2,⋯ ,n}\{1,2,\cdots,n\}.

We use 1n{\bf 1}_{n} to denote a length-nn vector whose entries are all ones.

Given three sets A,B,C⊆{0,1}dA,B,C\subseteq\{0,1\}^{d} of vectors where ∣A∣=∣B∣=∣C∣=n|A|=|B|=|C|=n, 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 Gap−MaxIP{\sf Gap}{\rm-}{\sf MaxIP}. 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 Cγ∈(0,0.1)C_{\gamma}\in(0,0.1), every ϵ>0\epsilon>0, and every C>C0>0C>C_{0}>0, there exist constants Ca>0C_{a}>0 and Cb>0C_{b}>0 and such that, if ATAttC\mathsf{ATAttC} (Definition D.1) 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−MaxIP(n,d=Clog⁡n,t=C0log⁡n,ϵ){\sf Gap}{\rm-}{\sf MaxIP}(n,d=C\log n,t=C_{0}\log n,\epsilon) (Definition 4.6) can be solved in time O(T+n3−Cγ)O(T+n^{3-C_{\gamma}}).

We give an algorithm for Gap−MaxIP(n,d=Clog⁡n,t=C0log⁡n,ϵ){\sf Gap}{\rm-}{\sf MaxIP}(n,d=C\log n,t=C_{0}\log n,\epsilon) (Definition 4.6). Let a1,⋯ ,an,b1,⋯ ,bn,c1,⋯ ,cn∈{0,1}da_{1},\cdots,a_{n},b_{1},\cdots,b_{n},c_{1},\cdots,c_{n}\in\{0,1\}^{d} denote the inputs to this problem. Using them, we will construct appropriate inputs to the ATAttC\mathsf{ATAttC} problem so that its output will help us to detect triples with large inner product.

Let β>0\beta>0 and d~≥d\widetilde{d}\geq d be parameters to be determined (in Eq. (5) and Eq. (2) below). Define τ>0\tau>0 by

We pick these parameters so that τ\tau will be an upper bound on entries of the attention matrix, namely

We will use an algorithm for the ATAttC(n~,d~,B,ϵa)\mathsf{ATAttC}(\widetilde{n},\widetilde{d},B,\epsilon_{a}) problem with parameters:

Since each entry of QQ and K1K_{1}, K2K_{2} is either β\sqrt{\beta} or 00, it follows that

is naturally partitioned into eight submatrices

where each AiA_{i} (∀i∈\forall i\in) is a matrix of size n×n2n\times n^{2}, defined as follows. For each j0∈[n],j1∈[n],j2∈[n]j_{0}\in[n],j_{1}\in[n],j_{2}\in[n],

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A1A_{1} is

exp⁡(β(⟨aj0,bj1,cj2⟩+⟨1d,0d,0d⟩)/d~)=exp⁡(β⟨aj0,bj1,cj2⟩/d~)\exp(\beta(\langle a_{j_{0}},b_{j_{1}},c_{j_{2}}\rangle+\langle{\bf 1}_{d},{\bf 0}_{d},{\bf 0}_{d}\rangle)/\widetilde{d})=\exp(\beta\langle a_{j_{0}},b_{j_{1}},c_{j_{2}}\rangle/\widetilde{d})

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A2A_{2} is

exp⁡(β(⟨aj0,bj1,0d⟩+⟨aj0,0d,1d⟩)/d~)=exp⁡(0)=1\exp(\beta(\langle a_{j_{0}},b_{j_{1}},{\bf 0}_{d}\rangle+\langle a_{j_{0}},{\bf 0}_{d},{\bf 1}_{d}\rangle)/\widetilde{d})=\exp(0)=1

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A3A_{3} is

exp⁡(β(⟨aj0,0d,cj2⟩+⟨aj0,1d,0d⟩)/d~)=exp⁡(0)=1\exp(\beta(\langle a_{j_{0}},{\bf 0}_{d},c_{j_{2}}\rangle+\langle a_{j_{0}},{\bf 1}_{d},{\bf 0}_{d}\rangle)/\widetilde{d})=\exp(0)=1

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A4A_{4} is

exp⁡(β(⟨aj0,0d,0d⟩+⟨1d,1d,1d⟩)/d~)=exp⁡(βd/d~)=τ\exp(\beta(\langle a_{j_{0}},{\bf 0}_{d},{\bf 0}_{d}\rangle+\langle{\bf 1}_{d},{\bf 1}_{d},{\bf 1}_{d}\rangle)/\widetilde{d})=\exp(\beta d/\widetilde{d})=\tau

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A5A_{5} is 11

exp⁡(β(⟨0d,bj1,cj2⟩+⟨1d,0d,0d⟩)/d~)=exp⁡(0)=1\exp(\beta(\langle{\bf 0}_{d},b_{j_{1}},c_{j_{2}}\rangle+\langle{\bf 1}_{d},{\bf 0}_{d},{\bf 0}_{d}\rangle)/\widetilde{d})=\exp(0)=1

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A6A_{6} is 11

exp⁡(β(⟨0d,bj1,0d⟩+⟨aj0,0d,1d⟩)/d~)=exp⁡(0)=1\exp(\beta(\langle{\bf 0}_{d},b_{j_{1}},{\bf 0}_{d}\rangle+\langle a_{j_{0}},{\bf 0}_{d},{\bf 1}_{d}\rangle)/\widetilde{d})=\exp(0)=1

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A7A_{7} is 11

exp⁡(β(⟨0d,0d,cj2⟩+⟨aj0,1d,0d⟩)/d~)=exp⁡(0)=1\exp(\beta(\langle{\bf 0}_{d},{\bf 0}_{d},c_{j_{2}}\rangle+\langle a_{j_{0}},{\bf 1}_{d},{\bf 0}_{d}\rangle)/\widetilde{d})=\exp(0)=1

The (j0,j1+(j2−1)n)(j_{0},j_{1}+(j_{2}-1)n)-th entry of A8A_{8} is

exp⁡(β(⟨0d,0d,0d⟩+⟨1d,1d,1d⟩)/d~)=exp⁡(βd/d~)=τ\exp(\beta(\langle{\bf 0}_{d},{\bf 0}_{d},{\bf 0}_{d}\rangle+\langle{\bf 1}_{d},{\bf 1}_{d},{\bf 1}_{d}\rangle)/\widetilde{d})=\exp(\beta d/\widetilde{d})=\tau

For each (i,j1,j2)∈[n]×[n]×[n](i,j_{1},j_{2})\in[n]\times[n]\times[n], we know that

Here we used the fact that d<d~d<\widetilde{d} (see Eq. (2)), and the last step uses the definition of τ\tau (see Eq. (1)).

We also know that for each (i,j1,j2)∈[n]×[n]×[n](i,j_{1},j_{2})\in[n]\times[n]\times[n],

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

By combining our expression for AA with Eq. (B.1) and Eq. (7), we see that

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

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

Here, the last two steps follow from Eq. (4).

Recall that for each i∈[n]i\in[n] we need to distinguish between two cases: either there is a pair (j1,j2)∈[n](j_{1},j_{2})\in[n] such that ⟨ai,bj1,cj2⟩≥t\langle a_{i},b_{j_{1}},c_{j_{2}}\rangle\geq t, or else for all pairs (j1,j2)∈[n](j_{1},j_{2})\in[n] the inner product ⟨ai,bj1,cj2⟩≤(1−ϵa)t\langle a_{i},b_{j_{1}},c_{j_{2}}\rangle\leq(1-\epsilon_{a})t. We will distinguish between these cases by checking whether uiu_{i} is greater than a threshold value t~0:=2t~\widetilde{t}_{0}:=2\widetilde{t}. We next consider the two cases to see why this is.

For a given i∈[n]i\in[n], if there are (j1,j2)∈[n]×[n](j_{1},j_{2})\in[n]\times[n] such that ⟨ai,bj1,bj2⟩≥t\langle a_{i},b_{j_{1}},b_{j_{2}}\rangle\geq t, then

where the 1st step follows from 2d=d~2d=\widetilde{d} (see Eq. (2)). This means as desired that

For a given i∈[n]i\in[n], if for all (j1,j2)∈[n]×[n](j_{1},j_{2})\in[n]\times[n] we have ⟨ai,bj1,cj2⟩<t(1−ϵ)\langle a_{i},b_{j_{1}},c_{j_{2}}\rangle<t(1-\epsilon), then

Here, the 4th step follows because, by our choice of β\beta and tt, we have

where we used that t=C0log⁡nt=C_{0}\log n (by Lemma statement), that d=Clog⁡nd=C\log n, that β=B3\beta=B^{3} (Eq. (5)) and the choice of BB (Eq. (3)). ∎

B.2 Main Hardness Result

We can finally conclude our main lower bound.

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 algorithm running in time O(n3−q)O(n^{3-q}) 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}}).

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 n×dn\times d matrices, there are two standard variants on their Kronecker product one may consider: The standard Kronecker product (denoted ⊗\otimes) is a new n2×d2n^{2}\times d^{2} matrix, whereas the column-wise Kronecker product (denoted ⊘\oslash) is a new n2×dn^{2}\times d 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 ⊘\oslash be defined as Definition 2.4.

We define C1:=A1⊤B1,C2:=A2⊤B2C_{1}:=A_{1}^{\top}B_{1},C_{2}:=A_{2}^{\top}B_{2}

For each i∈[n]i\in[n], let b1,i⊤b_{1,i}^{\top} denote the ii-th row of B1B_{1}.

For each i∈[n]i\in[n], let b2,i⊤b_{2,i}^{\top} denote the ii-th row of B2B_{2}.

From the above, we can calculate that the entry of CC in location k1,k2k_{1},k_{2} 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 ii and the summation over jj, and the last step follows from definition of matrices C1C_{1} and C2C_{2}.

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 ⊘\oslash be defined as Definition 2.4.

Let ⊙\odot be defined as Definition 2.2.

Let (⋅,⋅,⋅)(\cdot,\cdot,\cdot) operator be defined as Definition C.4.

Let ∘\circ be defined as Definition 2.1.

Part 1. Ai,j1+(j2−1)n=Ai,j1,j2A_{i,j_{1}+(j_{2}-1)n}=\mathsf{A}_{i,j_{1},j_{2}} for i∈[n],j1∈[n],j2∈[n]i\in[n],j_{1}\in[n],j_{2}\in[n] (This means A\mathsf{A} can be viewed as the tensor version of AA)

Part 2. A1n2=A(I,1n,1n)=Q⊙(1n⊤K1)⊙(1n⊤K2)A{\bf 1}_{n^{2}}=\mathsf{A}(I,{\bf 1}_{n},{\bf 1}_{n})=Q\odot({\bf 1}_{n}^{\top}K_{1})\odot({\bf 1}_{n}^{\top}K_{2})

Part 3. A(V1⊘V2)=Q(K1⊘K2)⊤(V1⊘V2)=Q((K1⊤V1)∘(K2⊤V2))A(V_{1}\oslash V_{2})=Q(K_{1}\oslash K_{2})^{\top}(V_{1}\oslash V_{2})=Q((K_{1}^{\top}V_{1})\circ(K_{2}^{\top}V_{2}))

Directly follows from definition of AA and A\mathsf{A}.

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 AA to DD, then in Section D.6, we analyze the error propagation from AA and DD to the attention matrix.

∥Q∥∞≤B\|Q\|_{\infty}\leq B, ∥K1∥∞≤B\|K_{1}\|_{\infty}\leq B, ∥K2∥∞≤B\|K_{2}\|_{\infty}\leq B, ∥V1∥∞≤B\|V_{1}\|_{\infty}\leq B, ∥V2∥∞≤B\|V_{2}\|_{\infty}\leq B

Notice that the straightforward algorithm for this problem will spend at least Ω(n3)\Omega(n^{3}) time to write the matrix AA (we can also think of AA as an tensor that has size n×n×nn\times n\times n).

D.2 An Error Control Tool From Previous Work

Let g:=Θ(max⁡{log⁡(1/ϵ)log⁡(log⁡(1/ϵ)/B),B})g:=\Theta(\max\{\frac{\log(1/\epsilon)}{\log(\log(1/\epsilon)/B)},B\}).

D.3 Tensor Q⊙K1⊙K2Q\odot K_{1}\odot K_{2} Has Bounded Entries

Let ⊙\odot operation be defined as Definition 2.2.

For every index triple (i,j1,j2)∈[n]×[n]×[n](i,j_{1},j_{2})\in[n]\times[n]\times[n], we are able to prove

D.4 Tensor Low-Rank Approximation

In the following definition, we view the n×n2n\times n^{2} size matrix as an n×n×nn\times n\times n size attention matrix.

We use r≥1r\geq 1 to denote a positive integer.

We use ϵ∈(0,0.1)\epsilon\in(0,0.1) to represent an accuracy parameter.

∣A~i,j1,j2−Ai,j1,j2∣≤ϵ⋅Ai,j1,j2|\widetilde{\mathsf{A}}_{i,j_{1},j_{2}}-A_{i,j_{1},j_{2}}|\leq\epsilon\cdot\mathsf{A}_{i,j_{1},j_{2}} for all (i,j1,j2)∈[n]×[n]×[n](i,j_{1},j_{2})\in[n]\times[n]\times[n].

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 i∈[n]i\in[n], 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 ϵA,ϵD∈(0,0.1)\epsilon_{A},\epsilon_{D}\in(0,0.1)

We bound each of these two terms to get our desired result.

First of all, for every index pair (i,j)∈[n]×[d](i,j)\in[n]\times[d],

Second, for every (i,j)∈[n]×[d](i,j)\in[n]\times[d],

Here, again, the 2nd step uses the triangle inequality, the 3rd step follows because Di,i−1D_{i,i}^{-1} is positive, the 4th step follows from the assumption that ∣A~i,l−Ai,l∣≤ϵA⋅Ai,l|\widetilde{A}_{i,l}-A_{i,l}|\leq\epsilon_{A}\cdot A_{i,l} and the final step follows by definition of Di,iD_{i,i}.

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 P(x)P(x) denote a degree-gg single-variable polynomial. We apply P(M)P(M) entry-wisely, i.e, P(M)i,j,l=P(Mi,j,l)P(M)_{i,j,l}=P(M_{i,j,l}).

Let rr be rank parameter that is r:=(3(g+d)3g).r:={3(g+d)\choose 3g}.

We define set VV and provide names for variables in set VV in the following sense,

Thus, function K\mathsf{K} can be viewed a degree-3g3g polynomial in the 3d3d entries in VV of the vectors u,v,wu,v,w.

Here, we can view K\mathsf{K} function as

Therefore, we should construct three matrices U1,U2U_{1},U_{2} and U3U_{3} as the following way, for each i∈[n]i\in[n]

These n×rn\times r matrices can be constructed in time O(nrg)O(nrg) in the straightforward way, since each entry depends on gg variables. ∎

E.2 Time for Constructing U1,U2,U3U_{1},U_{2},U_{3}

∥K1∥∞≤B\|K_{1}\|_{\infty}\leq B, ∥K2∥∞≤B\|K_{2}\|_{\infty}\leq B,

∥V1∥∞≤B\|V_{1}\|_{\infty}\leq B, and ∥V2∥∞≤B\|V_{2}\|_{\infty}\leq B.

For bounded number B>0B>0 and accuracy parameter ϵ∈(0,1)\epsilon\in(0,1), there are positive integers gg and rr

it takes O(nr)O(nr) time to construct U1,U2U_{1},U_{2} and U3U_{3}.

Recall that the definition of (ϵ,r)(\epsilon,r)-approximation can be found in Definition D.4.

We can then compute U1,U2U_{1},U_{2}, and U3U_{3} 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 ATAttC(n,d=O(log⁡n),B=o(log⁡n),ϵa=1/\poly(n))\mathsf{ATAttC}(n,d=O(\log n),B=o(\sqrt{\log n}),\epsilon_{a}=1/\poly(n)) in time n1+o(1)n^{1+o(1)}.

Using Lemma E.2, we know that Step 1 (in Algorithm 1) can be implemented in O(nrg)O(nrg) time

Using Lemma C.3, we know that Step 2 (in Algorithm 1) can be implemented in O(nr)O(nr) time

Step 3 can implemented in O(n)O(n) 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 n×dn\times d matrix

We combine Corollary D.2, Lemma D.5, Lemma D.6, and simple algebra. ∎

References