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 X=Q=KX=Q=K (ignore the effect of VV):

where D⁡(X):=diag⁡(exp⁡(XX⊤)1n)\operatorname{\mathsf{D}}(X):=\operatorname{diag}(\exp(XX^{\top})\mathbf{1}_{n}).

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 D⁡(X):=diag⁡(exp⁡(XX⊤)1n)\operatorname{\mathsf{D}}(X):=\operatorname{diag}(\exp(XX^{\top}){\bf 1}_{n}). We use ω\omega to denote the exponent of matrix multiplication. Currently ω≈2.373\omega\approx 2.373.

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 AA, the product of a sparse embedding matrix SS and AA can be computed in nnz⁡(A)\operatorname{nnz}(A) 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 AA, we use nnz⁡(A)\operatorname{nnz}(A) to denote the number of nonzero in a matrix AA. For matrix AA, we use ∥A∥\|A\| to denote its operator/spectral norm, i.e., ∥A∥:=max⁡x∥Ax∥2/∥x∥2\|A\|:=\max_{x}\|Ax\|_{2}/\|x\|_{2}. Note that ∥A∥2\|A\|_{2} means nothing, such notation should never exist in the draft. For a PSD matrix AA, let UΣU⊤U\Sigma U^{\top} denote its SVD. We use A1/2A^{1/2} to denote UΣ1/2U⊤U\Sigma^{1/2}U^{\top}.

The Taylor expansion of exp⁡(x)\exp(x) is ∑i=0∞xii!=1+x+x22!+⋯\sum_{i=0}^{\infty}\frac{x^{i}}{i!}=1+x+\frac{x^{2}}{2!}+\cdots.

We use ω\omega to denote the exponent of matrix multiplication, currently ω≈2.373\omega\approx 2.373 .

2 Basic Facts

For any x∈[−0.1,0.1]x\in[-0.1,0.1], ∣1−exp⁡(x)∣≤2∣x∣|1-\exp(x)|\leq 2|x|

By Taylor expansion, for x∈[−0.1,0.1]x\in[-0.1,0.1], we have

where the 1st step follows from the Taylor expansion of exp⁡(x)\exp(x), the 2nd step follows from the triangle inequality, the 3rd step follows from lower bounding i!i! with 2i2^{i} since i>2i>2, the 4th step follows from summing up of the second term, and the last step follows from ∣x∣≤0.1|x|\leq 0.1. ∎

Since for any a,ba,b the above equation is true, thus

which means Bi,iBj,j≥Bi,j2B_{i,i}B_{j,j}\geq B_{i,j}^{2}.

(1−ϵ)A⪯B⪯(1+ϵ)A(1-\epsilon)A\preceq B\preceq(1+\epsilon)A

3 Sketching Matrices

Here we formally define several kinds of sketching matrices.

Analysis

We consider a over-parameterized version attention matrix approximation problem. Suppose d≫nd\gg n. 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. (1−ϵ)B⪯A⪯(1+ϵ)B(1-\epsilon)B\preceq A\preceq(1+\epsilon)B;

Condition 2. −r≤Ai,j≤r-r\leq A_{i,j}\leq r, for all i,j∈[n]×[n]i,j\in[n]\times[n].

Using Fact 2.3 and Condition 1, we know that

We note that the matrix B−(1−ϵ)⋅AB-(1-\epsilon)\cdot A is a PSD matrix, therefore,

where the 1st step follows from Lemma 2.2, the 2nd step follows Bi,i≤(1+ϵ)Ai,iB_{i,i}\leq(1+\epsilon)A_{i,i} (see Eq. (1)), and the 3rd step follows from the definition of Ai,jA_{i,j} from the lemma statement.

Combining the above two equations, we obtain the following range on Bi,jB_{i,j}:

We can do a symmetric argument using the fact that (1+ϵ)⋅A−B(1+\epsilon)\cdot A-B is a PSD matrix:

2 Perturbation of Entry-wise Exponentiating a Matrix

Condition 3. Bi,j∈[−(1+ϵ)r,(1+ϵ)r]B_{i,j}\in[-(1+\epsilon)r,(1+\epsilon)r].

From Condition 2. and Condition 3., we have

where the 2nd step follows from ϵ∈(0,1)\epsilon\in(0,1).

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. ∣exp⁡(Ai,j)−exp⁡(Bi,j)∣≤exp⁡(Ai,j)⋅6r|\exp(A_{i,j})-\exp(B_{i,j})|\leq\exp(A_{i,j})\cdot 6r for all i∈[n],j∈[n]i\in[n],j\in[n]

Condition 2. ∣exp⁡(Ai,j)−exp⁡(Bi,j)∣≤exp⁡(Bi,j)⋅6r|\exp(A_{i,j})-\exp(B_{i,j})|\leq\exp(B_{i,j})\cdot 6r for all i∈[n],j∈[n]i\in[n],j\in[n]

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 c1>0c_{1}>0 and c2>0c_{2}>0 denote two fixed constants. If the following conditions hold

Condition 2. For all i∈[n]i\in[n], j∈[n]j\in[n]

For each of the index pair (i,j)∈[n]×[n](i,j)\in[n]\times[n], 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 D⁡\operatorname{\mathsf{D}}.

The second term.

For each of the index pair (i,j)∈[n]×[n](i,j)\in[n]\times[n], 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 D⁡\operatorname{\mathsf{D}}.

5 Randomized Symmetric Attention Approximation Algorithm

Here D⁡(X)=diag⁡(exp⁡(XX⊤)1n)\operatorname{\mathsf{D}}(X)=\operatorname{diag}(\exp(XX^{\top}){\bf 1}_{n}).

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 D⁡(X)=diag⁡(exp⁡(XX⊤)1n)\operatorname{\mathsf{D}}(X)=\operatorname{diag}(\exp(XX^{\top}){\bf 1}_{n}).

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 1−δσ1-\delta_{\sigma}. The O~\widetilde{O} hides the log⁡(d/δσ)\log(d/\delta_{\sigma}) 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 O~\widetilde{O} hides the log⁡(d/δ)\log(d/\delta) factors. Here we use ω\omega to denote the exponent of matrix multiplication. For current FMM algorithm, we have ω≈2.373\omega\approx 2.373.

We first approximately compute the leverage score, i.e., gives an O(1)O(1)-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 A⊤AA^{\top}A. First, we define the following sampling process,

Let H=A⊤AH=A^{\top}A. Let pj≥β⋅σj(A)/np_{j}\geq{\beta\cdot\sigma_{j}(A)}/{n} for some β≥1\beta\geq 1. We sample TT rows of matrix AA with respect to the probability pjp_{j} with replacement. We use jtj_{t} to represent the index of the row that is sampled during the tt-th trial. And the sampled matrix is defined as

where TT denotes the number of the trials.

For the above sampling process, previous work provided the following guarantee.

Let ϵ0∈(0,1)\epsilon_{0}\in(0,1) denote the precision parameter, and δ0∈(0,0.1)\delta_{0}\in(0,0.1) denote the failure probability. For the matrix H~\widetilde{H} sampled as Definition 5.3, it holds with probability at least 1−δ01-\delta_{0} that

References