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: (queries), (keys), (values), all of size , where is the number of tokens in the input sequence and is the dimension of latent representations. This process outputs the following:
Here, matrix is defined as the element-wise exponential of . Additionally, is an diagonal matrix derived from the sum of rows of , for . In this context, matrix is referred to as the “attention matrix”, and is called the “softmax matrix”. It is important to note that calculating the attention matrix directly requires operations, and storing it consumes memory. Consequently, a straightforward computation of demands a runtime of and 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 based on the row norms of . The more challenging aspect lies in obtaining a reliable spectral approximation for the diagonal matrix . In a recent result, Zandieh et al. effectively leverages fast KDE solvers to attain a high-quality approximation of . 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 time. Instead, we observe that by doing a one-sided sampling from the squared row norms of , 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 acceleration in forward and backward propagation for sequence lengths of k. When dealing with causal masking, the method still delivers a substantial speedup. Moreover, when our approach is applied to pretrained LLMs, e.g., -- 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 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 . We will demonstrate that under certain mild assumptions about matrices and , this simple approach allows us to establish spectral bounds on the estimated matrix. Our aim is to find a sufficiently precise approximate matrix that satisfies:
The first step of our empirical algorithm involves identifying large entries of the attention matrix 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 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 . 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 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 is not given to us, but can be found in 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 is defined as follows. Suppose is generated as in Algorithm 1 with block size and in Definition 1. We further assume there are at most pairs with , where is as in Definition 1. Then with probability , the we find in Algorithm 1 has at most non-zero entries and with probability at least , the outputs of Algorithm 3 satisfy ?? and the overall runtime is .
We note the assumption on the angles of the rows of and in Corollary 1 is satisfied if most rows are drawn uniformly at random from a -dimensional sphere, since in this case they will be nearly orthogonal, i.e., have angle at most with high probability. However, the corollarly also allows pairs of rows to have arbitrary angle, which may be more realistic.
Suppose all preconditions of ?? hold, where the mask matrix is defined as follows. Suppose is defined such that there is a threshold such that if and only if . Then we can find exactly with probability , and with probability at least , the outputs of Algorithm 3 satisfy ??. The runtime is .
The key idea behind the proof of Corollary 2 is to first sketch by an ExpanderSketch , which is efficient since has a small number of rows. Then compute which is again efficient since has a small number of rows. Thus, we never form the matrix .
1 Causal Masking
Language models commonly employ causal masking. The causal mask is a lower triangular binary square matrix denoted as where . The causal attention mechanism is defined as:
where is defined as before and is an diagonal matrix derived from the sum of rows of the masked attention , specifically for . To approximate causal attention with a spectral guarantee, we require two components. First, we need a spectral approximation for the diagonal matrix . Second, we need to approximate the matrix product between and , 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 .
Our approach is based on a key observation, as depicted in ??. The masked attention can be decomposed into three non-zero matrices, each of which has half the size of the original attention matrix. The block , located entirely below the diagonal is unmasked attention. Consequently, we can approximate its row sums using Algorithm 2. The two diagonal blocks and 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 and trim them if the length is over so that all data have sequence lengths of . 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 -- 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 . If all the layers are replaced then the perplexity goes to up 12 but it runs about faster. For -, similar happens but the perplexities are linearly increasing as the number of HyperAttention grows.
In addition, we evaluate the performances of monkey patched -- 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 faster when k.
2 Single Self Attention Layer
We further explore the speedup of HyperAttention with varying sequence lengths from to . 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 have the same length and their dimensions are fixed to and the number of attention heads is set by . We chose the same parameters in HyperAttention as described in the previous section. In ??, we observe that HyperAttention runs to up faster without causal masking and 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 , and merging attention outputs which result in an increase of practical runtime. However, those speedups will increase when the sequence length 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 , we utilize the -- and LongBench narrative-qa dataset, changing the sequence length from 1k to 9k. We trim or pad the input context so that its length is strictly . Unlike the vision model, we notice that the first columns in often contain heavy entries; hence we compute as the largest squared norm excluding the first columns. We collect these values for all heads and layers and compute their average. ?? plots the value of with various sequence length . It is observed that the value of decreases as grows, supporting the claim that our assumption 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 calculated in line 3 of Algorithm 2 is close to the maximum row sum of the matrix . It is easy to check that for all because of the definition of in the lemma statement. Furthermore, if we define the set:
Next, let us define the upper-capped version of matrix where entries of -th row on positions where the mask value is equal to zero are capped at value (line 6 of the algorithm) as:
We proceed by bounding the total mass of large entries of matrix lost through capping (i.e., entries of that are larger than thresholds ). If we define constant , we can write,
The inequality in ?? follows because, for every , the cardinality of the set
Now to bound we first find bounds on the number of ’s in rows and columns of . Using the definition of in ?? and the fact that row sums in matrix are equal to , we have:
Additionally, using the precondition of the lemma about , we have:
Now we bound the norm for an arbitrary integer 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 . Therefore,
Now by plugging the above inequalities into ?? we find that:
Proof of Corollary 1: Because , we have whenever . As there are at most total pairs, the expected number of such pairs that collide under is at most and so by a Markov bound is at most with failure probability .
Since we also assume there are at most pairs with , there can be at most additional pairs that collide.
Thus, in total we have collisions, and consequently the number of non-zero entries in is at most with failure probability . The proof now follows by the assumptions of the corollary statement as well as Theorem 1. ∎
We compute , followed by . Note that the time for this computation is . Next, for each column of this allows us to construct a set with the property that if , then . This holds simultaneously for all columns with probability at least by a union bound. The time for constructing all the sets is
Note that for all , and we can explicitly compute the exact value of for all and all , in time. By the assumptions of the corollary, we have that contains a superset of the support of the -th column of , and since we can compute the values exactly, we can exactly construct the mask matrix that the corollary requires, and in time. The proof now follows by the assumptions of the statement as well as Theorem 1. ∎