HyperAttention: Long-context Attention in Near-Linear Time

Insu Han, Rajesh Jayaram, Amin Karbasi, Vahab Mirrokni, David P. Woodruff, Amir Zandieh

Introduction

Transformers have been successfully applied to a wide variety of learning tasks in areas such as natural language processing , computer vision , and time series forecasting . Despite their success, these models face serious scalability limitations because naïve exact computation of their attention layers incurs quadratic (in the sequence length) runtime and memory complexities. This presents a fundamental challenge for scaling transformer models to longer context lengths.

Various approaches have been explored to tackle the quadratic-time attention layer, with one notable direction focusing on approximating intermediate matrices in attention layers. Methods for doing this include approximations by sparse matrices , low-rank matrices , or a combination of both . These methods aim to provide faster approximation to various components of attention, but none of them provide end-to-end approximations of the full dot-product attention. Moreover, none of these works support the use of causal masking, which is a crucial part of modern transformer architectures. On the negative side, recent theoretical bounds suggest that entry-wise approximations to the attention matrix are impossible in sub-quadratic time in general .

The dot-product attention involves processing three input matrices: Q{\bm{Q}} (queries), K{\bm{K}} (keys), V{\bm{V}} (values), all of size n×dn\times d, where nn is the number of tokens in the input sequence and dd is the dimension of latent representations. This process outputs the following:

Here, matrix A:=exp⁡(QK⊤){\bm{A}}:=\exp\left({\bm{Q}}{\bm{K}}^{\top}\right) is defined as the element-wise exponential of QK⊤{\bm{Q}}{\bm{K}}^{\top}. Additionally, D{\bm{D}} is an n×nn\times n diagonal matrix derived from the sum of rows of A{\bm{A}}, Di,i=∥Ai,:∥1{\bm{D}}_{i,i}=\|{\bm{A}}_{i,:}\|_{1} for i∈[n]i\in[n]. In this context, matrix A{\bm{A}} is referred to as the “attention matrix”, and D−1A{\bm{D}}^{-1}{\bm{A}} is called the “softmax matrix”. It is important to note that calculating the attention matrix A{\bm{A}} directly requires Θ(n2d)\Theta(n^{2}d) operations, and storing it consumes Θ(n2)\Theta(n^{2}) memory. Consequently, a straightforward computation of Att\mathbf{Att} demands a runtime of Ω(n2d)\Omega(n^{2}d) and Ω(n2)\Omega(n^{2}) memory.

2 Our Contributions

We show that efficiently solving the matrix multiplication component of the attention approximation problem in ?? can be achieved by defining the sampling matrix S{\bm{S}} based on the row norms of V{\bm{V}}. The more challenging aspect lies in obtaining a reliable spectral approximation for the diagonal matrix D{\bm{D}}. In a recent result, Zandieh et al. effectively leverages fast KDE solvers to attain a high-quality approximation of D{\bm{D}}. However, we streamline the KDEformer procedure and demonstrate that uniform sampling is sufficient to achieve the desired spectral guarantee, eliminating the need for importance sampling based on kernel densities. This significant simplification allows us to develop a practical and provably linear time algorithm.

In contrast to prior work , our approach does not necessitate bounded entries or bounded stable rank. Furthermore, the fine-grained parameters we introduce to analyze the time complexity may remain small even when the entries in the attention matrix or the stable rank are large.

Prior work of Zandieh et al. used KDE to identify columns in the attention matrix with large norm and to perform approximate matrix product with the value matrix by sampling such columns. As mentioned, finding such columns requires at least O(n1.173)O(n^{1.173}) time. Instead, we observe that by doing a one-sided sampling from the squared row norms of VV, we can avoid the use of KDEs and achieve the same spectral norm guarantee in terms of the stable rank. Although our algorithm is simple and just samples by the row norms of the value matrix (or even samples uniformly in practice), the main technical challenge is that we do not know the row norms of the attention matrix needed in order to normalize it and produce a proper factorization of it. This is reminiscent of the quadratic time hard instance of where we may not be able to find a heavy entry in a row easily, and thus cannot normalize by its norm in the attention matrix. Our parameters (1) and (2) above allow us to argue that the heavy entries, if they exist, are not distributed in the worst possible way.

Empirically, HyperAttention demonstrates significant speed improvements, achieving over a 50×50\times acceleration in forward and backward propagation for sequence lengths of n=131n=131k. When dealing with causal masking, the method still delivers a substantial 5×5\times speedup. Moreover, when our approach is applied to pretrained LLMs, e.g., chatglm2\mathtt{chatglm2}-6b\mathtt{6b}-32k\mathtt{32k} and evaluated on long-context benchmark datasets, so-called LongBench , it maintains performance levels that closely match those of the original models, even without the need for fine-tuning. Furthermore, we investigate task-specific evaluations and discover summarization and code completion tasks are more robust to approximate attention layers than question answerings.

Preliminaries

Using this LSH function, as demonstrated by Zandieh et al. , we can sort keys and queries within an attention layer in such a way that large entries get shifted towards the diagonal of the attention matrix. Subsequently, these significant entries in the attention matrix can be captured by computing equal-sized blocks along the diagonal. This approach aligns with the block-memory access patterns of modern hardware and can be efficiently parallelized through batching across blocks.

Algorithm

Our procedure for approximating D{\bm{D}} consists of two steps. Initially, we identify the dominant entries within the attention matrix using an algorithm rooted in the Hamming sorted LSH, as defined in ??. The second step revolves around randomly selecting a small subset of keys K{\bm{K}}. We will demonstrate that under certain mild assumptions about matrices A{\bm{A}} and D{\bm{D}}, this simple approach allows us to establish spectral bounds on the estimated matrix. Our aim is to find a sufficiently precise approximate matrix D~\widetilde{{\bm{D}}} that satisfies:

The first step of our empirical algorithm involves identifying large entries of the attention matrix A{\bm{A}} through hashing keys and queries into uniformly-sized buckets using the Hamming sorted LSH, which we refer to as sortLSH. This process is detailed in Algorithm 1 and is visually illustrated in ??. Note that we also mention other was of identifying large patterns, such as checking for a known heavy hitter pattern, or using CountSketch which we describe more below.

Algorithm 1 returns a sparse mask designed to isolate the dominant entries of the attention matrix. Given this mask, we compute an approximation of the matrix D\mathbf{D} in Algorithm 2 that satisfies the spectral guarantee in ??. This algorithm accomplishes this by combining the attention values corresponding to the mask with a randomly chosen subset of columns from the attention matrix. The assumptions of Lemma 1 are used to ensure that the variance of the estimator is small, and the same complexity of the algorithm increases as a function of the parameters α,κ\alpha,\kappa. We remark that our algorithm is versatile and can function effectively with a predefined mask that specifies the positions of dominant entries within the attention matrix, mirroring the approach taken in . The main guarantee provided by this algorithm is given in ??.

Given a D~\widetilde{{\bm{D}}} that meets the spectral approximation conditions as in ??, we can achieve the spectral constraint in ??, by finding a sampling matrix that satisfies the following condition,

The above result is standard and for proof refer to .

Main Theorem.

Note that even if MH{\bm{M}}^{\mathcal{H}} is not given to us, but MH{\bm{M}}^{\mathcal{H}} can be found in d⋅n1+o(1)d\cdot n^{1+o(1)} time, the theorem holds. We also give examples when this is possible by using Hamming sorted LSH, which our experiments are based on, or using the ExpanderSketch of which is based on CountSketch but also gives a fast recovery time. In the supplementary we show:

Suppose all preconditions of ?? hold, where the mask matrix MH{\bm{M}}^{{\mathcal{H}}} is defined as follows. Suppose MH∈{0,1}n×n{\bm{M}}^{{\mathcal{H}}}\in\{0,1\}^{n\times n} is generated as in Algorithm 1 with block size b=no(1)b=n^{o(1)} and r=log⁡2nr=\log_{2}n in Definition 1. We further assume there are at most n1+o(1)n^{1+o(1)} pairs (i,j)(i,j) with θ(Qi,∗,Kj,∗)≤π2(1−o(1))\theta({\bm{Q}}_{i,*},{\bm{K}}_{j,*})\leq\frac{\pi}{2}(1-o(1)), where θ\theta is as in Definition 1. Then with probability 1−1/no(1)1-1/n^{o(1)}, the MH{\bm{M}}^{{\mathcal{H}}} we find in Algorithm 1 has at most n1+o(1)n^{1+o(1)} non-zero entries and with probability at least .98.98, the outputs S,D~{\bm{S}},\widetilde{{\bm{D}}} of Algorithm 3 satisfy ?? and the overall runtime is O(d⋅n1+o(1))O(d\cdot n^{1+o(1)}).

We note the assumption on the angles of the rows of Q{\bm{Q}} and K{\bm{K}} in Corollary 1 is satisfied if most rows are drawn uniformly at random from a dd-dimensional sphere, since in this case they will be nearly orthogonal, i.e., have angle at most π2(1−o(1))\frac{\pi}{2}(1-o(1)) with high probability. However, the corollarly also allows n1+o(1)n^{1+o(1)} pairs of rows to have arbitrary angle, which may be more realistic.

Suppose all preconditions of ?? hold, where the mask matrix MH{\bm{M}}^{{\mathcal{H}}} is defined as follows. Suppose MH∈{0,1}n×n{\bm{M}}^{{\mathcal{H}}}\in\{0,1\}^{n\times n} is defined such that there is a threshold τ=no(1)\tau=n^{o(1)} such that Mi,jH=1{\bm{M}}^{{\mathcal{H}}}_{i,j}=1 if and only if (QK⊤)i,j2≥∥QK⊤ej∥22τ({\bm{Q}}{\bm{K}}^{\top})_{i,j}^{2}\geq\frac{\|{\bm{Q}}{\bm{K}}^{\top}e_{j}\|_{2}^{2}}{\tau}. Then we can find MH{\bm{M}}^{{\mathcal{H}}} exactly with probability 1−O(1/n2)1-O(1/n^{2}), and with probability at least .98.98, the outputs S,D~{\bm{S}},\widetilde{{\bm{D}}} of Algorithm 3 satisfy ??. The runtime is O(d⋅n1+o(1))O(d\cdot n^{1+o(1)}).

The key idea behind the proof of Corollary 2 is to first sketch Q{\bm{Q}} by an ExpanderSketch TT, which is efficient since TT has a small number of rows. Then compute (T⋅Q)⋅K⊤(T\cdot{\bm{Q}})\cdot{\bm{K}}^{\top} which is again efficient since (T⋅Q)(T\cdot{\bm{Q}}) has a small number of rows. Thus, we never form the matrix Q⋅K⊤{\bm{Q}}\cdot{\bm{K}}^{\top}.

1 Causal Masking

Language models commonly employ causal masking. The causal mask is a lower triangular binary square matrix denoted as MC{\bm{M}}^{\mathcal{C}} where Mi,jC=1{i≥j}{\bm{M}}^{\mathcal{C}}_{i,j}=\bm{1}_{\{i\geq j\}}. The causal attention mechanism is defined as:

where A:=exp⁡(QK⊤){\bm{A}}:=\exp\left({\bm{Q}}{\bm{K}}^{\top}\right) is defined as before and DC{\bm{D}}_{\mathcal{C}} is an n×nn\times n diagonal matrix derived from the sum of rows of the masked attention MC⊙A{\bm{M}}^{\mathcal{C}}\odot{\bm{A}}, specifically [DC]i,i=⟨Mi,:C,Ai,:⟩[{\bm{D}}_{\mathcal{C}}]_{i,i}=\langle{\bm{M}}^{\mathcal{C}}_{i,:},{\bm{A}}_{i,:}\rangle for i∈[n]i\in[n]. To approximate causal attention with a spectral guarantee, we require two components. First, we need a spectral approximation for the diagonal matrix DC{\bm{D}}_{\mathcal{C}}. Second, we need to approximate the matrix product between DC−1(MC⊙A){\bm{D}}_{\mathcal{C}}^{-1}({\bm{M}}^{\mathcal{C}}\odot{\bm{A}}) and V{\bm{V}}, which can be achieved using the same sampling technique as described in Algorithm 3 and ??. The first component is more intricate, and we employ a recursive method to address it. So we focus on how to efficiently approximate the diagonal DC{\bm{D}}_{\mathcal{C}}.

Our approach is based on a key observation, as depicted in ??. The masked attention MC⊙A{\bm{M}}^{\mathcal{C}}\odot{\bm{A}} can be decomposed into three non-zero matrices, each of which has half the size of the original attention matrix. The block A21{\bm{A}}_{\bf 21}, located entirely below the diagonal is unmasked attention. Consequently, we can approximate its row sums using Algorithm 2. The two diagonal blocks M1C⊙A11{\bm{M}}^{\mathcal{C}}_{1}\odot{\bm{A}}_{\bf 11} and M2C⊙A22{\bm{M}}^{\mathcal{C}}_{2}\odot{\bm{A}}_{\bf 22} shown in ?? are causal attentions with half the original size. To handle these, we apply a recursive approach and further partition them into smaller blocks, and repeat this procedure. We present a pseudocode for this procedure in Algorithm 4.

Experiments

In this section, we benchmark our algorithms by scaling up existing large language models to handle long-range sequences. All experiments are performed on a single A100 GPU with 40 GB memory and we use FlashAttention 2 for the exact attention computation.

1 Monkey Patching Self-attention

We use LongBench , a collection of long context benchmark datasets, which contains 6 different tasks ranging from single and multiple-document question answering, summarization, few-shot learning, synthetic tasks, and code completion. We select a subset of dataset whose encoded sequence lengths are larger than 32,76832{,}768 and trim them if the length is over 32,76832{,}768 so that all data have sequence lengths of 32,76832{,}768. Then, we compute the perplexity (i.e., loss on next tokens prediction) of each model. To highlight the scalability on the long sequences, we calculate the total speedup on all attention layers whether performed by HyperAttention or FlashAttention.

The results are summarized in ??. Observe that chatglm2\mathtt{chatglm2}-6b\mathtt{6b}-32k\mathtt{32k} shows a reasonable perplexity even after monkey patched by HyperAttention, e.g., after replacing 20 layers the perplexity increases approximately by 1 and it slowly goes up until 24 layers. But it improves runtimes in attention layers about 50%50\%. If all the layers are replaced then the perplexity goes to up 12 but it runs about 2.3×2.3\times faster. For phi\mathtt{phi}-1.5\mathtt{1.5}, similar happens but the perplexities are linearly increasing as the number of HyperAttention grows.

In addition, we evaluate the performances of monkey patched chatglm2\mathtt{chatglm2}-6b\mathtt{6b}-32k\mathtt{32k} on LongBench datasets and compute task-specific evaluation scores on each task including single-document question answering, multiple-document question answering, summarization, few-shot learning, synthetic tasks and code completion. Results are provided in ??. While replacing HyperAttention generally leads to performance degradation, we observe that its role can vary depending on the task at hand. For example, summarization and code completion are more robust to other tasks. Notably, when half of all attention layers are patched (i.e., 14 layers), we verify that most of the tasks do not degrade more than 13%. In particular, the performance of the summarization task remained almost unchanged, suggesting that this task may be more robust to partial alterations in the attention mechanism. We recall that computations in attention layers can be 1.5×1.5\times faster when n=32n=32k.

2 Single Self Attention Layer

We further explore the speedup of HyperAttention with varying sequence lengths from 4,0964{,}096 to 131,072131{,}072. We measure wall-clock times of both forward and forward+backward operations when they are computed with FlashAttention or are accelerated by HyperAttention. We measure the times with and without causal masking. All inputs Q,K,V{\bm{Q}},{\bm{K}},{\bm{V}} have the same length and their dimensions are fixed to d=64d=64 and the number of attention heads is set by 1212. We chose the same parameters in HyperAttention as described in the previous section. In ??, we observe that HyperAttention runs to up 54×54\times faster without causal masking and 5.4×5.4\times when the causal masking applies. Although time complexities of both causal masking and non-masking are the same, a practical algorithm for causal masking (Algorithm 1) requires additional operations such as partitioning Q,K,V{\bm{Q}},{\bm{K}},{\bm{V}}, and merging attention outputs which result in an increase of practical runtime. However, those speedups will increase when the sequence length nn grows. We believe this opens the door to scale self-attention not only for inference but also for training or fine-tuning the LLMs to fit in significantly long sequences.

3 Empirical Verification of Assumption

To further investigate the dependence on nn, we utilize the chatglm2\mathtt{chatglm2}-6b\mathtt{6b}-32k\mathtt{32k} and LongBench narrative-qa dataset, changing the sequence length nn from 1k to 9k. We trim or pad the input context so that its length is strictly nn. Unlike the vision model, we notice that the first columns in D−1A{\bm{D}}^{-1}{\bm{A}} often contain heavy entries; hence we compute α\alpha as the largest squared norm excluding the first 3232 columns. We collect these values for all heads and layers and compute their average. ?? plots the value of αn\frac{\alpha}{n} with various sequence length nn. It is observed that the value of αn\frac{\alpha}{n} decreases as nn grows, supporting the claim that our assumption α=no(1)\alpha=n^{o(1)} holds in practice.

Conclusion

In this work, we propose a simple linear time attention approximation algorithm by simplifying the existing algorithm based on kernel density estimation (KDE). We introduce a more general parameterization for a spectral approximation guarantee based on the condition number, which does not require assumptions used in prior work. Our algorithm makes use of sortLSH to find large entries and we adopt fast matrix multiplication via row norm sampling. We additionally study how our algorithm is used for causal masking by recursive partitioning. Empirically, we illustrate that pre trained LLMs using our algorithm can enhance both inference and training speeds with only minimal performance degradation.

References

Appendix A Omitted proofs

Here we include the proofs that were omitted in the main body of the paper. First, we present the proof of ??.

Proof of ??: First, we show that τ\tau calculated in line 3 of Algorithm 2 is close to the maximum row sum of the matrix (1n−MH)⊙A(\bm{1}_{n}-{\bm{M}}^{{\mathcal{H}}})\odot{\bm{A}}. It is easy to check that τκ≤<1−Mi,:H,exp⁡(KQi,:⊤)>≤τκ\frac{\tau}{\kappa}\leq\left<1-{\bm{M}}^{{\mathcal{H}}}_{i,:},\exp({\bm{K}}{\bm{Q}}_{i,:}^{\top})\right>\leq\tau\kappa for all i∈[n]i\in[n] because of the definition of κ\kappa in the lemma statement. Furthermore, if we define the set:

Next, let us define the upper-capped version of matrix A{\bm{A}} where entries of ii-th row on positions where the mask MH{\bm{M}}^{{\mathcal{H}}} value is equal to zero are capped at value CiC_{i} (line 6 of the algorithm) as:

We proceed by bounding the total mass of large entries of matrix D−1A{\bm{D}}^{-1}{\bm{A}} lost through capping (i.e., entries of A{\bm{A}} that are larger than thresholds CiC_{i}). If we define constant C^:=ε2mκ2nlog⁡n\widehat{C}:=\frac{\varepsilon^{2}m}{\kappa^{2}n\log n}, we can write,

The inequality in ?? follows because, for every i∈[n]i\in[n], the cardinality of the set

Now to bound ∥Bt⋅v∥2\left\|B^{t}\cdot v\right\|_{2} we first find bounds on the number of 11’s in rows and columns of BtB^{t}. Using the definition of BtB^{t} in ?? and the fact that row sums in matrix D−1A{\bm{D}}^{-1}{\bm{A}} are equal to 11, we have:

Additionally, using the precondition of the lemma about α=n⋅max⁡i∈[n]∥D−1A⋅e(i)∥22\alpha=n\cdot\max_{i\in[n]}\left\|{\bm{D}}^{-1}{\bm{A}}\cdot e^{(i)}\right\|_{2}^{2}, we have:

Now we bound the norm ∥Bt⋅v∥2\left\|B^{t}\cdot v\right\|_{2} for an arbitrary integer t≥0t\geq 0 as follows:

where the inequality in second line above follows from ?? and the inequality in the last line follows from ??. The last equality follows from the assumption that ∥v∥2=1\left\|v\right\|_{2}=1. Therefore,

Now by plugging the above inequalities into ?? we find that:

Proof of Corollary 1: Because r=log⁡2nr=\log_{2}n, we have Pr⁡[H(Qi,∗)=H(Kj,∗)]≤1/n1−o(1)\Pr[{\mathcal{H}}({\bm{Q}}_{i,*})={\mathcal{H}}({\bm{K}}_{j,*})]\leq 1/n^{1-o(1)} whenever θ(Qi,∗,Kj,∗)≥π2(1−o(1))\theta({\bm{Q}}_{i,*},{\bm{K}}_{j,*})\geq\frac{\pi}{2}(1-o(1)). As there are at most n2n^{2} total pairs, the expected number of such pairs that collide under H{\mathcal{H}} is at most n1+o(1)n^{1+o(1)} and so by a Markov bound is at most n1+o(1)n^{1+o(1)} with failure probability 1/no(1)1/n^{o(1)}.

Since we also assume there are at most n1+o(1)n^{1+o(1)} pairs (i,j)(i,j) with θ(Qi,∗,Kj,∗)<π2(1−o(1))\theta({\bm{Q}}_{i,*},{\bm{K}}_{j,*})<\frac{\pi}{2}(1-o(1)), there can be at most n1+o(1)n^{1+o(1)} additional pairs that collide.

Thus, in total we have n1+o(1)n^{1+o(1)} collisions, and consequently the number of non-zero entries in MH{\bm{M}}^{{\mathcal{H}}} is at most n1+o(1)n^{1+o(1)} with failure probability 1/no(1)1/n^{o(1)}. The proof now follows by the assumptions of the corollary statement as well as Theorem 1. ∎

We compute T⋅Q{\bm{T}}\cdot{\bm{Q}}, followed by (T⋅Q)⋅K⊤({\bm{T}}\cdot{\bm{Q}})\cdot{\bm{K}}^{\top}. Note that the time for this computation is O(τnlog⁡n)O(\tau n\log n). Next, for each column jj of QK⊤{\bm{Q}}{\bm{K}}^{\top} this allows us to construct a set SjS_{j} with the property that if (Q⋅K⊤)i,j≥∥Q⋅K⊤ej∥22τ({\bm{Q}}\cdot{\bm{K}}^{\top})_{i,j}\geq\frac{\|{\bm{Q}}\cdot{\bm{K}}^{\top}e_{j}\|_{2}^{2}}{\tau}, then i∈Sji\in S_{j}. This holds simultaneously for all columns with probability at least 1−1n21-\frac{1}{n^{2}} by a union bound. The time for constructing all the sets SjS_{j} is n1+o(1)dn^{1+o(1)}d

Note that ∣Sj∣≤2τ|S_{j}|\leq 2\tau for all jj, and we can explicitly compute the exact value of (Q⋅K⊤)i,j({\bm{Q}}\cdot{\bm{K}}^{\top})_{i,j} for all i∈Sji\in S_{j} and all jj, in O(nτd)O(n\tau d) time. By the assumptions of the corollary, we have that SjS_{j} contains a superset of the support of the jj-th column of MH{\bm{M}}^{{\mathcal{H}}}, and since we can compute the values exactly, we can exactly construct the mask MH{\bm{M}}^{{\mathcal{H}}} matrix that the corollary requires, and in n1+o(1)dn^{1+o(1)}d time. The proof now follows by the assumptions of the statement as well as Theorem 1. ∎