Randomized and Deterministic Attention Sparsification Algorithms for Over-parameterized Feature Dimension
Yichuan Deng, Sridhar Mahadevan, Zhao Song
Introduction
Attention mechanisms have become an essential tool in many natural language processing (NLP) applications, and large language models (LLMs). LLMs, for examples, Transformer , GPT-1 , BERT , GPT-2 , GPT-3 , PaLM and OPT have significantly advanced the state of the art in this field. These models rely on attention mechanisms to capture the dependencies and relations between different tokens in a sequence, making them ideal for tasks such as language modeling , machine translation , and sentiment analysis . A recent breakthrough, ChatGPT, which is a chatbot by OpenAI built with GPT-3, has shown the power of attention mechanism . Very recently, OpenAI released their technical report on GPT-4 , in many real-life tasks it has significant better performance than previous versions of GPT.
Inspired by previous works and applications of attention, we study the following question,
Is it possible to compute the attention matrix faster for large feature dimension?
To address this challenge, we seek for fast attention matrix computing algorithm. Here in this work, we focus on a specific form of attention matrix computation by assuming (ignore the effect of ):
where .
For this specific setting, we provide an affirmative response to the above question in this work.
Inspired by the approach of , formally, we prove:
Here . We use to denote the exponent of matrix multiplication. Currently .
Following the trend of derandomized algorithm for linear programming by Brand’s breakthrough result . In this work, we also consider how to design a deterministic sparsification algorithm for attention matrix. Inspired by another approach of deterministic sparsification method, we proposed another deterministic algorithm as follows.
Unlike the result of matching the linear programming algorithm by , here our deterministic algorithm is much slower than randomized algorithm. We are not sure if it’s possible to design a deterministic algorithm which can match the time of our randomized algorithm. We leave this as an open problem for the future work.
2 Related Work
There is a series of work focused on Attention computation . Studies have utilized locality sensitive hashing (LSH) techniques for attention approximation . With the observation that the kernel density estimation (KDE) problem can be used to simplify the denominator of the softmax function, the proposed the KDEformer which employed an efficient KDE solver so that it can approximate the attention in sub-quadratic time with provable spectral norm bounds. Furthermore, both static and dynamic versions of attention computation have been explored in recent work . The work of studied the regularized hyperbolic regression problems such exponential function, cosh functions and sinh functions.
Input Sparsity Time Algorithm
There has been a long history of the research in algorithms runs in input sparsity time . provides the first input sparsity time algorithm for linear regression and low-rank approximation. Given any input matrices , the product of a sparse embedding matrix and can be computed in time. Later work gives Oblivious Sparse Norm Approximating Projections (OSNAPs), improving the dimension in . Many works are generalizing the Frobenius norm low-rank approximation to other norms. Recent work proposed a faster algorithm for approximating the John Ellipsoid in input sparsity time. Another recent work gave an algorithm such that, for any real matrix, it can solve the discrepancy minimization problem in the input sparsity time of the matrix.
Sketching and Sampling
Sketching and sampling techniques are very powerful in numerical linear algebra. It has been applied to many fundamental tasks, Like linear programming (LP) , cutting plane method , empirical risk minimization , integral minimization , semi-definite programming , matrix sensing , federated learning , frank-wolfe method , matrix completion , dynamic sparsifier , differential privacy , clustering , tensor problems , training over-parameterized neural networks .
Roadmap.
We organize the paper as follows. In Section 2 we give the preliminary for our paper, including notations, algebraic facts and definition of sketching matrices. In Section 3 we provide the analysis of the correctness of our algorithms. In Section 4 we provide our sparsification tools. In Section 5 we give the analysis for the leverage score sampling tool we use.
Acknowledgements.
The authors would like to thank Josh Alman, Jan van den Brand, Beidi Chen, Yeqi Gao, Zhihang Li, Junze Yin, Lichen Zhang, and Tianyi Zhou for very helpful discussions.
Preliminary
In Section 2.1, we provide the notations to use acress the paper. In Section 2.2, we state some well-known facts. In Section 2.3 we introduce some sketching matrices.
For a matrix , we use to denote the number of nonzero in a matrix . For matrix , we use to denote its operator/spectral norm, i.e., . Note that means nothing, such notation should never exist in the draft. For a PSD matrix , let denote its SVD. We use to denote .
The Taylor expansion of is .
We use to denote the exponent of matrix multiplication, currently .
2 Basic Facts
For any ,
By Taylor expansion, for , we have
where the 1st step follows from the Taylor expansion of , the 2nd step follows from the triangle inequality, the 3rd step follows from lower bounding with since , the 4th step follows from summing up of the second term, and the last step follows from . ∎
Since for any the above equation is true, thus
which means .
3 Sketching Matrices
Here we formally define several kinds of sketching matrices.
Analysis
We consider a over-parameterized version attention matrix approximation problem. Suppose . Here in this section, we provide analysis of the two kinds of the algorithms. In Section 3.1 we prove some properties of PSD matrices to be used. In Section 3.2 we prove the perturbation of entry-wise exponentiating a matrix. In Section 3.3 we prove the perturbation of diagonal normalization matrix. In Section 3.4 we prove the perturbation of attention matrix. In Section 3.5 and Section 3.6 we provide our randomized algorithm and deterministic algorithm respectively.
Condition 1. ;
Condition 2. , for all .
Using Fact 2.3 and Condition 1, we know that
We note that the matrix is a PSD matrix, therefore,
where the 1st step follows from Lemma 2.2, the 2nd step follows (see Eq. (1)), and the 3rd step follows from the definition of from the lemma statement.
Combining the above two equations, we obtain the following range on :
We can do a symmetric argument using the fact that is a PSD matrix:
2 Perturbation of Entry-wise Exponentiating a Matrix
Condition 3. .
From Condition 2. and Condition 3., we have
where the 2nd step follows from .
where the 2nd step follows from Eq. (2), the last step follows from Condition 1 in the Lemma statement and Fact 2.1.
where the 2nd step follows from Eq. (2), and the last step follows from Condition 1 in the Lemma statement and Fact 2.1. ∎
3 Perturbation of Diagonal Normalization Matrix
Condition 1. for all
Condition 2. for all
where the 1st step follows from simple algebra, the 2nd step follows from triangle inequality, the 3rd step follows from Condition 1. in Lemma statement, and the last step follows from simple algebra.
where the 1st step follows from simple algebra, the 2nd step follows from triangle inequality, the 3rd step follows from Condition 2. in Lemma statement, and the last step follows from simple algebra.
4 Perturbation of Attention Matrix
Here we state the perturbation lemma for attention matrix as follows.
Let and denote two fixed constants. If the following conditions hold
Condition 2. For all ,
For each of the index pair , we have
where the first step follows from definition, the second steps follow from simple algebra, the third step follows from triangle inequality, the forth step follows from Condition 1 in the lemma statement, the fifth step follows from simple algebra, the sixth step follows from all the entries are positive, and the last step follows from definition of .
The second term.
For each of the index pair , we have
where the first step follows from simple algebra, the second step follows from triangle inequality, the third step follows from Condition 2 in the lemma statement, the forth step follows from simple algebra, and the last step follows from the definition of .
5 Randomized Symmetric Attention Approximation Algorithm
Here .
Proof of Running Time. The running time follows from Lemma 5.2.
Proof of Correctness. The correctness follows from Lemma 4.1, Lemma 3.1, Lemma 3.2, Lemma 3.3, and Lemma 3.4. ∎
6 Deterministic Symmetric Attention Approximation Algorithm
Here .
Proof of Running Time. The running time follows from Theorem 4.2.
Proof of Correctness. The correctness follows from Theorem 4.2, Lemma 3.1, Lemma 3.2, Lemma 3.3, and Lemma 3.4. ∎
Sparsification
In this section, we provide two kinds of sparsification tools. In Section 4.1 we provide our randomized sparsification algorithm and in Section 4.2 we provide the deterministic sparsification algorithm.
2 Deterministic Sparsification Algorithm
Leverage Score Sampling
Here in this section, we provide a sampling algorithm based on leverage score. In Section 5.1 we provide our fast leverage score computation algorithm. In Section 5.2 we provide the running time and correctness proof for our randomized sparsification algorithm. In Section 5.3 we provide our sampling process with its analysis.
In this section, we provide an algorithm which can fast approximate the leverage scores of an input matrix.
Previous work studied harder version of the leverage score sampling algorithm, which implies the following result directly.
with probability at least . The hides the factors.
2 Correctness and Running time
Here we present the correctness running time of our randomized sparsification algorithm based on the leverage score approximation algorithm.
Here hides the factors. Here we use to denote the exponent of matrix multiplication. For current FMM algorithm, we have .
We first approximately compute the leverage score, i.e., gives an -approximation to all leverage scores via Lemma 5.1. Then run we samples a number of rows according to the leverage scores, using Lemma 5.4, we can show the correctness. ∎
3 Sampling a Batch of Rank-111 Terms
In this section, we present the sampling process for sparsification of symmetric matrix in the shape of . First, we define the following sampling process,
Let . Let for some . We sample rows of matrix with respect to the probability with replacement. We use to represent the index of the row that is sampled during the -th trial. And the sampled matrix is defined as
where denotes the number of the trials.
For the above sampling process, previous work provided the following guarantee.
Let denote the precision parameter, and denote the failure probability. For the matrix sampled as Definition 5.3, it holds with probability at least that