Conv-Basis: A New Paradigm for Efficient Attention Inference and Gradient Computation in Transformers

Yingyu Liang, Heshan Liu, Zhenmei Shi, Zhao Song, Zhuoyan Xu, Junze Yin

Introduction

Numerous notable large language models (LLMs) in natural language processing (NLP) have emerged in these two years, such as BERT , OPT , PaLM , Mistral , Gemini , Gemma , Claude3 , GPT-3 , GPT-4 , Llama , Llama2 , Llama3 and so on. These models have profoundly changed the world and have been widely used in human activities, such as education , law , finance , bio-informatics , coding , and even creative writing such as top AI conference reviews . The key component of the generative LLMs success is the decoder-only transformer architecture introduced by . The transformer uses the self-attention mechanism, allowing the model to capture long-range dependencies in the input sequence. Self-attention computes a weighted sum of the input tokens, where the weights are determined by the similarity between each pair of tokens. This enables the model to attend to relevant information from different parts of the sequence when generating the output. However, the computational complexity of the self-attention in transformers grows quadratically O(n2)O(n^{2}) with the input length nn, limiting their applicability to long context, e.g., 128k, 200k, 1000k input tokens for GPT4 , Claude3 , Gemma respectively.

To overcome the complexity obstacle of softmax(QK⊤)\mathsf{softmax}(QK^{\top}), many works study more efficient attention computation methods that can scale almost linearly with the sequence length while maintaining the model’s performance. show that if all entry of QK⊤QK^{\top} is bounded and d=O(log⁡n)d=O(\log n), softmax(QK⊤)\mathsf{softmax}(QK^{\top}) will be “close” to a low-rank matrix. Then, they present an algorithm that can approximate attention computation in almost linear time. Similarly, by uniform softmax column norms assumption and sparse assumption, solve attention computation in almost linear time, where they identify large entries in the attention matrix and only focus on them.

On the other hand, many works find that the attention pattern has conv\mathsf{conv}-like (or “diagonalized”) structure (see Figure 1 (b)), mathematically, Ai,j≈Ai′,j′A_{i,j}\approx A_{i^{\prime},j^{\prime}} when i−j=i′−j′i-j=i^{\prime}-j^{\prime}, where we can see i−ji-j as the position distance between two tokens. It is relevant to the bag-of-words or nn-gram concept, i.e., nn adjacent symbols or words in NLP. Furthermore, the conv\mathsf{conv}-like structure can be connected to convolution recurrent models and structured state space models (SSMs) as Mamba . More specifically, we can use multiple convolution matrices to approximate an attention matrix, whose intuition is similar to the low-rank approximation in the sense of computation acceleration. Note that the matrix product of a convolution matrix and a vector can be computed by Fast Fourier Transform (FFT) with time complexity O(nlog⁡(n))O(n\log(n)), while the naive way takes O(n2)O(n^{2}) time (see details in Figure 1 (a)). Therefore, it is natural to ask:

Can we use the convolution matrix to accelerate the attention computation?

Thus, our algorithm can achieve attention inference in O(kndlog⁡(n))O(knd\log(n)), without needing any parameter updates e.g., re-train or finetune. Our theorems can also applied to accelerate attention training . In detail, our methods take O(dknlog⁡n+nd2)O(dkn\log n+nd^{2}) time for forward computation and O(d2knlog⁡n)O(d^{2}kn\log n) time for backward gradient computation (Theorem 5.6). Furthermore, applying our technique to low-rank approximation of attention matrices , we can extend their results to more general settings (Theorem 6.5). In detail, only works on attention approximation without an attention mask, while ours can be applied to different kinds of attention masks including the most popular causal attention mask (Definition 3.2). It shows the broad applicability of our analysis. We summarize our contributions below.

Our results are beyond or comparable to the two brilliant previous works in the following sense. (1) To guarantee a small approximation error, for the attention matrix, needs bounded entries assumption and d=O(log⁡n)d=O(\log n) assumption, while needs uniform softmax column norms assumption and sparse assumption. However, without all these assumptions, our algorithm can still guarantee a small approximation error (Corollary 4.5), i.e., our algorithm can apply to any Q,KQ,K including unbounded matrices, dense matrices, and any hidden dimension dd. (2) To guarantee a truly subquadratic running time, needs to assume d=O(log⁡n)d=O(\log n) to get n1+o(1)n^{1+o(1)} time complexity. However, for our algorithm, as long as d=no(1)d=n^{o(1)} and k=no(1)k=n^{o(1)}, we achieve running time n1+o(1)n^{1+o(1)}. This has much less restriction on dd. Moreover, our time complexity covers from n1+o(1)n^{1+o(1)} to n2−Ω(1)n^{2-\Omega(1)} with different dd, while can only handle d=O(log⁡n)d=O(\log n). (3) To guarantee a truly subquadratic running time, needs to assume dm=n2−Ω(1)dm=n^{2-\Omega(1)}, as they get O(dn1+o(1)+dm)O(dn^{1+o(1)}+dm) time complexity where mm is the number of large entries in attention matrices. Our work gets O(kndlog⁡(n))O(knd\log(n)) time complexity and we need kd=n1−Ω(1)kd=n^{1-\Omega(1)} to get truly subquadratic running time. For the situation m=n1+o(1),d=no(1)m=n^{1+o(1)},d=n^{o(1)} and k=no(1)k=n^{o(1)}, both our algorithm and run in n1+o(1)n^{1+o(1)} time. For the situation m=n1+Ω(1)m=n^{1+\Omega(1)}, d=no(1)d=n^{o(1)} and k=no(1)k=n^{o(1)}, running time in will be truly super-linear n1+Ω(1)n^{1+\Omega(1)} while our algorithm remains almost n1+o(1)n^{1+o(1)} linear timeConsidering the case where attention matrix is all 11 lower triangular matrix, we have k=1k=1 and m=n(n+1)/2m=n(n+1)/2..

Our contributions.

Our algorithm can quickly decompose any lower triangular matrix into its kk convolution basis (Algorithm 2), so, via FFT, we can solve Exact Attention Computation task in O(kndlog⁡(n))O(knd\log(n)) (Algorithm 1 and Corollary 4.5). When kd=no(1)kd=n^{o(1)}, our method takes almost linear time n1+o(1)n^{1+o(1)}. Our results are beyond or comparable to previous works (see comparison above).

During attention inference, our algorithm takes O(kndlog⁡(n))O(knd\log(n)), without needing any parameter updates e.g., re-train or fine-tune (Theorem 4.4). It may enable further improvement and scalability of LLMs in the longer context.

During attention training, our methods take O(dknlog⁡n+nd2)O(dkn\log n+nd^{2}) time for forward computation and O(d2knlog⁡n)O(d^{2}kn\log n) time for backward gradient computation (Theorem 5.6). It may save time, resources, and energy for nowadays LLMs training.

Our broad applicable technique can be applied to the low-rank approximation of attention matrices and extend existing results to more general settings (Theorem 6.5).

Roadmap.

In Section 2, we introduce the related work. In Section 3, we present the background for this paper. In Section 4, we present our algorithms and mathematical analysis for the convolution approximation for attention computation. In Section 5, we present the conv\mathsf{conv} approximation during training forward and backward gradient. In Section 6, we explain how to apply our technique to the low-rank approximation. In Section 7, we discuss two case studies: Longlora and RoPE.

Related Work

Very recent works study the conv\mathsf{conv}-like attention matrix. find that in-context learning is driven by the formation of “induction heads”–attention heads that copy patterns from earlier in the input sequence. This is reflected in the attention matrix becoming more diagonal, with tokens attending primarily to preceding tokens that match the current token. In Figure 6, they show a similar conv\mathsf{conv}-like attention pattern for other important attention circuits. Figure 3 of shows that in a minimal classification task, the abrupt emergence of in-context learning coincides with the formation of an induction head, characterized by a diagonal attention pattern. proves that for a simplified task, gradient descent causes a transformer to encode the causal graph structure of the task in the attention matrix. This results in tokens attending primarily to their causal parents reflected in a sparse diagonal structure (Figure 2).

Fast attention computation and long context LLM.

The development of efficient attention computation has been an active area of research in recent years. The standard self-attention mechanism, introduced in the transformer architecture , has a quadratic complexity with respect to the sequence length, which limits its applicability to long sequences. To address this limitation, various approaches have been proposed to improve the efficiency of attention computation. One line of research focuses on patterns of sparse attention that reduce the number of computations . Another approach is to use low-rank approximations or random features for the attention matrix , which reduces the computational complexity to linear in the sequence length. In addition, using linear attention as a proxy of softmax attention is a rich line of work . These developments in efficient attention computation have enabled transformer-based models to process longer sequences and have opened up new possibilities for their application in various domains .

Convolution in language model and FFT.

There are many subquadratic-time architectures are proposed to address Transformers’ computational inefficiency on long sequences, gated convolution recurrent models , and structured state space models (SSMs) . They can use global or local convolution operations to replace attention while keeping a comparable performance. The convolution operation can be computed by fast Fourier transform (FFT) efficiently . Moreover, the development of efficient convolution algorithms like Winograd and FFT-based convolutions has further optimized the computation, reducing the memory footprint and improving the overall speed. There are many other works studying Fourier transform .

Preliminary

In Section 3.1, we introduce the basic definitions and mathematical properties. In Section 3.2, we give the formal definition of the sub-convolution matrix and present it basic properties.

1 Basic Definitions and Facts about Attention and 𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv}

Now, we present basic definitions. We start by introducing the input and weight matrix.

It is straightforward to see QK⊤=XWQWK⊤X⊤QK^{\top}=XW_{Q}W_{K}^{\top}X^{\top}. In generative LLMs, there is a causal attention mask MM to guarantee the later tokens cannot see the previous tokens during generation. It is defined as follows:

We define the causal attention mask as M∈{0,1}n×nM\in\{0,1\}^{n\times n}, where Mi,j=1M_{i,j}=1 if i≥ji\geq j and Mi,j=0M_{i,j}=0 otherwise. We define MjM_{j} be the jj-th column of MM.

Now, we are ready to introduce the mathematical definition of the exact attention computation with a mask.

In Definition 3.3, we divide the softmax\mathsf{softmax} operation into an element-wise exp⁡\exp operation and a diagonal normalization matrix DD to obtain a clear formulation.

Efficiently computing the attention needs to exploit structured matrices that enable fast multiplication algorithms. Here, we define the convolution matrix, which is a structured matrix where each row vector is rotated one element to the right relative to the preceding row vector, which is defined as follows:

By the following fact, we know that the rank of a convolution matrix can be an arbitrary number. Thus, our conv\mathsf{conv}-basis is totally different from the rank basis.

The proof is trivial by Definition 3.5. ∎

Efficient computation of the convolution operation is crucial for many applications. The convolution theorem states that the circular convolution of two vectors can be computed efficiently using the Fast Fourier Transform (FFT). This leads to the following claim (see proof in Appendix A.1):

One property of convolution matrices is that they are additive with respect to the input vectors. In other words, the convolution of the sum of two vectors is equal to the sum of the convolutions of the individual vectors. This is stated formally in the following claim:

This is trivial by Definition 3.5 and matrix product operation is additive. ∎

2 Sub-convolution Matrix: Definitions and Properties

If we would like to use conv\mathsf{conv} as a basis system, we need to introduce some new concepts. Recall that, in general, the sum of two rank-1 matrices is a rank-2 two matrix. Due to conv\mathsf{conv} being additive, the sum of two convolution matrices is another convolution matrix, which does not hold the above property. Thus, we need to introduce sub-convolution matrices to be the basis.

Similarly, sub-convolution can be computed in O(nlog⁡n)O(n\log n) time via FFT.

This is trivial by considering the calculation between the truncated matrix of conv(a,m)\mathsf{conv}(a,m) and the truncated vector of xx with Claim 3.7. ∎

Here, we present the definition of the matrix with kk-conv\mathsf{conv} basis which is non-reducible.

The following lemma establishes that any non-zero lower triangular matrix can be represented as a matrix with a kk-conv\mathsf{conv} basis for some unique kk between 11 and nn. The proof is Appendix D.1.

𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv} Approximation during Inference

In Section 4.1, we introduce the basic definitions to support our algorithmic analysis in this section. In Section 4.2, we present the binary search and recover kk-conv\mathsf{conv} algorithms and present their theoretical guarantees. In Section 4.3, we provide the formal version of our main result.

Recall that any non-zero lower triangular matrix can be represented as a matrix with a kk-conv\mathsf{conv} basis for some unique kk between 11 and nn (Lemma 3.12). However, exactly getting kk is hard and the definition is too strict for the algorithm design. Thus, for more flexibility, we introduce a more general definition of non-degenerate kk-conv\mathsf{conv} basis as below, which is a proxy notion to relax the conditions required in algorithm design.

Here (T,δ)(T,\delta)-non-degenerate kk-conv\mathsf{conv} basis means that each conv\mathsf{conv} basis cannot be “covered” by the other basis easily.

The following theorem establishes that any non-zero lower triangular matrix can be represented as an ϵ\epsilon-close (T,δ)(T,\delta)-non-degenerate kk-conv\mathsf{conv} basis matrix. There may be many different choices of (k,T,δ,ϵ)(k,T,\delta,\epsilon), which provide flexibility for our Algorithm 1.

By Lemma 3.12, we have GG is a matrix with kk-conv\mathsf{conv} basis for some k∈[n]k\in[n]. We finish the proof by setting T=1T=1 and δ=ϵ=0\delta=\epsilon=0. ∎

2 Algorithms and Their Properties

Now, we present our main Algorithm 1. We also present Algorithm 2 and Algorithm 3 being used.

3 Main Theoretical Result

In this section, we present our main result.

whose time complexity is O(kndlog⁡(n))O(knd\log(n)) given M,Q,K,VM,Q,K,V.

The proof idea is that using binary search to recover all non-degenerate conv\mathsf{conv} basis (Lemma A.15), which takes O(nkdlog⁡(n))O(nkd\log(n)) time and has upto 2(exp⁡(2ϵ)−1)∥V∥∞2(\exp(2\epsilon)-1)\|V\|_{\infty} error (Lemma A.16). Then, via FFT (Claim 3.10), we finish the proof. ∎

Furthermore, we can exactly get YY, i.e., ϵ=0\epsilon=0, through Algorithm 1 with time complexity O(n2dlog⁡(n))O(n^{2}d\log(n)) in the worst case.

We set k=nk=n, T=1T=1, δ=0\delta=0 and ϵ=0\epsilon=0 as the input of Algorithm 1. Then, the proof follows Theorem 4.3 and Theorem 4.4 . ∎

By Theorem 4.4, when ϵ=O(1)\epsilon=O(1), we directly get the attention inference time complexity is O(kndlog⁡(n))O(knd\log(n)) with error upto O(ϵ)O(\epsilon) as claimed in Section 1. It may enable further improvement and scalability of LLMs in the longer context.

𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv} Approximation during Training Forward and Backward Gradient

We can apply our algorithm to accelerate attention training including forward and back propagation. We first define the attention training task, which is also used in .

Recall that during inference, we have the n×nn\times n size matrix QK⊤QK^{\top}. Similarly, in gradient calculation, we have an n×nn\times n size matrix, and we denote it as u(x)u(x).

Then, we are ready to present our main results for attention training.

During backward computation, we can convey the properties of low-rank and convolution at the same time (Lemma B.13 and Lemma B.15). Then, by tensor trick, we can compute the attention gradient based on attention inference (Lemma B.9). We finish the proof by Theorem 4.4. ∎

Note that only needs to convey the low-rank property, while our proof needs to convey the properties of low-rank and convolution simultaneously, which is a more general analysis.

Our Theorem 5.6 shows that our algorithm can accelerate Transformer training as well. It may save time, resources, and energy for nowadays LLMs training.

Low Rank Approximation

We can apply our analysis technique to low-rank approximation setting in . In detail, only works on attention approximation without an attention mask. Equipped with our mask analysis trick, we can generalize their results with different kinds of attention masks including the most popular causal attention mask. We first introduce some practical attention masks.

We define the continuous row mask as W∈{0,1}n×nW\in\{0,1\}^{n\times n}, where for each i∈[n]i\in[n], we are given si,ti∈[n]s_{i},t_{i}\in[n] such that Wi,j=1W_{i,j}=1 if si≤j≤tis_{i}\leq j\leq t_{i} and Wi,j=0W_{i,j}=0 otherwise.

Then, we have the following main results for the low-rank setting. The proof is in Appendix C.2.

The time complexity to get Y~\widetilde{Y} is

O(knd)O(knd) when WW is a causal mask defined in Definition 3.2.

O(kd∑j=1nBj)O(kd\sum_{j=1}^{n}B_{j}) when WW is a row change mask defined in Definition 6.1.

O(kndlog⁡(n))O(knd\log(n)) when WW is a continuous row mask defined in Definition 6.2.

O(rnd)O(rnd) when WW is a distinct rr columns / rows mask defined in Definition 6.3 / Definition 6.4.

Our Theorem 6.5 has the same error guarantee as . For the normal mask, e.g., casual attention mask (Definition 3.2), Theorem 6.5 shares the same time complexity as theirs.

Discussion

In this section, we discuss how to implement our methods in some popular long-context LLMs, e.g., LongLora and the RoPE model family.

Our conv\mathsf{conv} and low-rank approximation can be applied to LongLora , whose mask is shown in the left of Figure 3. They use this kind of sparse mask to extend the context sizes of pre-trained large language models, with limited computation cost, e.g., extending Llama2 70B from 4k context to 32k on a single 8×8\times A100 machine. As the “diagonalized” mask structure, we can directly apply our Algorithm 1 by replacing the causal attention mask (Definition 3.2) with their sparse mask for the conv\mathsf{conv} approximation with time complexity O(kndlog⁡(n))O(knd\log(n)). Similarly, for the low-rank approximation, we directly use the second statement in Theorem 6.5 by considering row change by amortized constant mask defined in Definition 6.1 with time complexity O(knd)O(knd), where Bj=O(1)B_{j}=O(1) for any j∈[n]j\in[n].

RoPE.

Acknowledgement

Research is partially supported by the National Science Foundation (NSF) Grants 2023239-DMS, CCF-2046710, and Air Force Grant FA9550-18-1-0166.

In Section A, we present additional details and proofs related to the convolution approximation approach. In Section B, we introduce the conv\mathsf{conv} approximation in gradient. In Section C, we include supplementary material for the low-rank approximation. In Section D, we present a collection of useful tools and lemmas that are referenced throughout the main text and the appendix. In Section E, we provide more related work.

Appendix A Technical Details About 𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv} Approximation

In Section A.1, we present the background of Toeplitz, circulant, and convolution matrices. In Section A.2, we develop more mathematical tools for studying the conv\mathsf{conv} approximation. In Section A.3, we give the key lemmas we used. In Section A.4, we use these tools and lemmas to prove our main theorem for the conv\mathsf{conv} approximation. In Section A.5, we analyze our case study.

The integer ii may have different ranges. We will specify these ranges in later text, corresponding to different contexts.

The Toeplitz matrix is one such structured matrix that has constant values along its diagonals. We define it as follows:

Furthermore, we define the circulant matrix, which is a structured matrix where each row vector is rotated one element to the right relative to the preceding row vector, which is defined as follows:

Finally, we present a basic fact about the Hadamard product.

Below, we explore the properties of conv\mathsf{conv}, Toep\mathsf{Toep}, Resi\mathsf{Resi}, and Circ\mathsf{Circ}.

The proof directly follows the Definition 3.2, Definition A.2, and Definition 3.5. ∎

where the first step follows Claim A.6, i.e., conv(a)=Toep([0n−1a])\mathsf{conv}(a)=\mathsf{Toep}(\begin{bmatrix}{\bf 0}_{n-1}\\ a\end{bmatrix}), the second step follows Fact A.7 and the last step follows Fact A.8. We finish the proof by O(nlog⁡n)O(n\log n) for FFT. ∎

A.2 Mathematical Tools Development for k𝑘k-𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv} Basis

When a lower triangular matrix HH is expressed as the sum of kk convolution matrices, it is useful to understand the structure of the entries in HH. The following claim provides an explicit formula for the entries of HH in terms of the basis vectors of the convolution matrices.

For any i<j∈[n]i<j\in[n], we have Hi,j=0H_{i,j}=0.

This is trivial by following H=∑i∈[k]conv(bi,mi)H=\sum_{i\in[k]}\mathsf{conv}(b_{i},m_{i}), the Definition 3.5 and Definition 3.9. ∎

We present the property of H~=M∘(QK⊤)\widetilde{H}=M\circ(QK^{\top}) as follows:

with time complexity O(nd)O(nd), where (K⊤)j(K^{\top})_{j} denotes the jj-th row of KK.

where the first step follows from the definition of H~\widetilde{H} (see Definition A.10), the second step follows from simple algebra, the third step follows from the fact that the jj-th column of K⊤K^{\top} is equal to the jj-th row of KK.

For any vector vv, we need O(n)O(n) time to get Mj∘vM_{j}\circ v.

Thus, in total, the time complexity is O(nd)O(nd). ∎

The key idea behind our approach is to express the matrix exponential of a matrix with kk-conv\mathsf{conv} basis as the sum of kk sub-convolution matrices involving the basis vectors. This allows us to efficiently approximate the exponential of the attention matrix. We show how to compute the new basis vectors of the convolution matrices from the original basis vectors below.

and M∘exp⁡(H)=∑r∈[k]conv(b~r,mr)M\circ\exp(H)=\sum_{r\in[k]}\mathsf{conv}(\widetilde{b}_{r},m_{r}) with time complexity O(nk)O(nk).

As exp⁡\exp is an element-wise function, when i≥ji\geq j we have (M∘exp⁡(H))i,j=exp⁡(H)i,j(M\circ\exp(H))_{i,j}=\exp(H)_{i,j} and

When i<ji<j we have (M∘exp⁡(H))i,j=0=∑r=1kconv(b~r,mr)i,j(M\circ\exp(H))_{i,j}=0=\sum_{r=1}^{k}\mathsf{conv}(\widetilde{b}_{r},m_{r})_{i,j}.

Thus, we have M∘exp⁡(H)=∑r∈[k]conv(b~r,mr)M\circ\exp(H)=\sum_{r\in[k]}\mathsf{conv}(\widetilde{b}_{r},m_{r}).

We need O(nk)O(nk) time to get ∑l∈[r]bl\sum_{l\in[r]}b_{l} for any r∈[k]r\in[k]. Then, we need O(1)O(1) time for element-wise exp⁡\exp and minus operation for O(nk)O(nk) terms. Thus, in total, we need O(nk)O(nk) time complexity. ∎

where the first step follows the lemma statement, the second step follows the property of Hadamard product and the last step follows the lemma statement. ∎

A.3 Lemma Used in Main Theorem Proof

Part 1: v=∑r∈[i](br′)1:Tv=\sum_{r\in[i]}(b^{\prime}_{r})_{1:T} and u=∑r∈[i]br′u=\sum_{r\in[i]}b^{\prime}_{r}

Part 3: ∥∑r∈[i](br′)1:T−∑r∈[i](br)1:T∥1≤Tϵ\|\sum_{r\in[i]}(b^{\prime}_{r})_{1:T}-\sum_{r\in[i]}(b_{r})_{1:T}\|_{1}\leq T\epsilon

Part 4: ∣∑r∈[i](br′)l−∑r∈[i](br)l∣≤ϵ|\sum_{r\in[i]}(b^{\prime}_{r})_{l}-\sum_{r\in[i]}(b_{r})_{l}|\leq\epsilon for any l∈[n]l\in[n].

We use the math induction to prove the correctness.

Part 1: v=∑r∈[i](br′)1:Tv=\sum_{r\in[i]}(b^{\prime}_{r})_{1:T} and u=∑r∈[i]br′u=\sum_{r\in[i]}b^{\prime}_{r}

Part 2: s=n−mi+1s=n-m_{i}+1 (Denote s=0s=0, after the -th loop.)

Part 3: ∥∑r∈[i](br′)1:T−∑r∈[i](br)1:T∥1≤Tϵ\|\sum_{r\in[i]}(b^{\prime}_{r})_{1:T}-\sum_{r\in[i]}(b_{r})_{1:T}\|_{1}\leq T\epsilon

Part 4: ∣∑r∈[i](br′)l−∑r∈[i](br)l∣≤ϵ|\sum_{r\in[i]}(b^{\prime}_{r})_{l}-\sum_{r\in[i]}(b_{r})_{l}|\leq\epsilon for any l∈[n]l\in[n]

We have v=∑r∈[i+1](br′)1:Tv=\sum_{r\in[i+1]}(b^{\prime}_{r})_{1:T} and u=∑r∈[i+1]br′u=\sum_{r\in[i+1]}b^{\prime}_{r} by the line 9 and line 10 in Algorithm 2.

We denote the output of Search(Q,K,k,T,δ,ϵ,∑r∈[i](br′)1:T,mi,n−T+1Q,K,k,T,\delta,\epsilon,\sum_{r\in[i]}(b^{\prime}_{r})_{1:T},m_{i},n-T+1) as yy. Now, we prove y=n−mi+1+1y=n-m_{i+1}+1.

It is clear that n−mi+1≤y≤n−T+1n-m_{i}+1\leq y\leq n-T+1. For any j∈{n−mi+1,…,n−T+1}j\in\{n-m_{i}+1,\dots,n-T+1\}, we have line 7 in Algorithm 3 as

where the first step follows from Definition A.10 (H~=H+R\widetilde{H}=H+R), the second step follows from Part 1, and the last step follows from Definition 4.2 (H=∑r∈[k]conv(br,mr)H=\sum_{r\in[k]}\mathsf{conv}(b_{r},m_{r})).

where the first step follows from the triangle inequality, the second step follows from Definition 4.2 (∥R∥∞≤ϵ\|R\|_{\infty}\leq\epsilon), the third step follows from j<n−mi+1+1j<n-m_{i+1}+1, the fourth step follows from Definition 3.9, the fifth step follows from Part 3, and the last step follows from Definition 4.2 (ϵ≤δ5T<δ4T\epsilon\leq\frac{\delta}{5T}<\frac{\delta}{4T}).

Similarly, when j≥n−mi+1+1j\geq n-m_{i+1}+1, we have Eq. (2) as

where the first step follows from the triangle inequality, the second step follows from Definition 4.2 (∥R∥∞≤ϵ\|R\|_{\infty}\leq\epsilon), the third step follows from simple algebra, the fourth step follows from the triangle inequality, the fifth step follows from Part 3, and the last step follows from Definition 4.1.

Thus, we can claim, when α<δ−2Tϵ\alpha<\delta-2T\epsilon, we have j<n−mi+1+1j<n-m_{i+1}+1, and we have j≥n−mi+1+1j\geq n-m_{i+1}+1 otherwise. Therefore, by binary search, we can get s=y=n−mi+1+1s=y=n-m_{i+1}+1.

We have s=n−mi+1+1s=n-m_{i+1}+1 and u=∑r∈[i]br′u=\sum_{r\in[i]}b^{\prime}_{r} at line 8 in Algorithm 2. Thus, we have

where the first step follows from simple algebra, the second step follows from Algorithm 2 (line 8), the third step follows from u=∑r∈[i]br′u=\sum_{r\in[i]}b^{\prime}_{r}, the fourth step follows from Definition A.10 (H~=H+R\widetilde{H}=H+R), the fifth step follows from Definition 4.2 (H=∑r∈[k]conv(br,mr)H=\sum_{r\in[k]}\mathsf{conv}(b_{r},m_{r})), the sixth step follows from s=n−mi+1+1s=n-m_{i+1}+1, the seventh step follows from Definition 3.9, the eighth step follows from simple algebra, and the last step follows from Definition 4.2 (∥R∥∞≤ϵ\|R\|_{\infty}\leq\epsilon).

We can get ∣∑r∈[i+1](br′)l−∑r∈[i](br)l∣≤ϵ|\sum_{r\in[i+1]}(b^{\prime}_{r})_{l}-\sum_{r\in[i]}(b_{r})_{l}|\leq\epsilon for any l∈[n]l\in[n] similarly as Proof of Part 3.

We can check the initial conditions hold. Thus, we finish the whole proof by math induction. ∎

Building upon Lemma A.15, we now analyze the overall error of our approach for approximating the attention computation. Recall that our goal is to efficiently approximate the matrix Y=D−1AVY=D^{-1}AV, where A=M∘exp⁡(QK⊤)A=M\circ\exp(QK^{\top}) and D=diag⁡(A1n)D=\operatorname{diag}(A{\bf 1}_{n}). We will show that by using the approximate basis vectors recovered by Algorithm 2, we can construct matrices A~\widetilde{A} and D~\widetilde{D} such that the approximation error ∥Y−D~−1A~V∥∞\|Y-\widetilde{D}^{-1}\widetilde{A}V\|_{\infty} is bounded. The following lemma provides this error analysis:

where the first step follows from triangle inequality, the second step follows from H~=H+R\widetilde{H}=H+R and the last step follows from ∥R∥∞≤ϵ\|R\|_{\infty}\leq\epsilon and Eq. (3).

by Lemma A.13 and line 12 in Algorithm 2.

In each loop, we call O(log⁡(n))O(\log(n)) times of binary search function. In each binary search function, we take O(nd)O(nd) time for line 6 in Algorithm 3 by Lemma A.12. Thus, we take O(ndlog⁡(n))O(nd\log(n)) in total for the search (Algorithm 3) in each loop.

In each loop, we take O(nd)O(nd) time for line 7 in Algorithm 2 by Lemma A.12.

Thus, we take total O(k(nd+ndlog⁡(n)))=O(kndlog⁡(n))O(k(nd+nd\log(n)))=O(knd\log(n)) for the whole loop.

We take O(nk)O(nk) time for the line 12 in Algorithm 2 by Lemma A.13.

In total, we take O(nk+kndlog⁡(n))=O(kndlog⁡(n))O(nk+knd\log(n))=O(knd\log(n)) time. ∎

We are now ready to prove our main result for the conv\mathsf{conv} approximation approach. Theorem A.17 brings together the key components we have developed: the existence of a kk-conv\mathsf{conv} basis for the attention matrix (Definition 4.2), the ability to efficiently recover an approximate kk-conv\mathsf{conv} basis (Algorithm 2 and Lemma A.15), and the bounded approximation error when using this approximate basis (Lemma A.16). The theorem statement is a formal version of our main conv\mathsf{conv} result, Theorem 4.4 and Algorithm 1, which was presented in the main text. It specifies the input properties, the approximation guarantees, and the time complexity of our approach.

A.4 Proof of Main Theorem

whose time complexity is O(kndlog⁡(n))O(knd\log(n)) given M,Q,K,VM,Q,K,V.

Denote A~:=∑r∈[k]conv(b~r,mr)\widetilde{A}:=\sum_{r\in[k]}\mathsf{conv}(\widetilde{b}_{r},m_{r}). By Claim 3.10, we take O(kndlog⁡(n))O(knd\log(n)) time to get A~V\widetilde{A}V via FFT as kk-conv\mathsf{conv} basis and dd columns in VV. Similarly, by Claim 3.10, we take O(knlog⁡(n))O(kn\log(n)) time for D~=diag⁡(A~1n)\widetilde{D}=\operatorname{diag}(\widetilde{A}{\bf 1}_{n}) via FFT as kk-conv\mathsf{conv} basis. Finally, we take O(nd)O(nd) time to get D~−1A~V\widetilde{D}^{-1}\widetilde{A}V as D~−1\widetilde{D}^{-1} is a diagonal matrix.

Thus, in total, we take O(kndlog⁡(n)+kndlog⁡(n)+knlog⁡(n)+nd)=O(kndlog⁡(n))O(knd\log(n)+knd\log(n)+kn\log(n)+nd)=O(knd\log(n)) time complexity. ∎

A.5 Construction for Case Study

For each i∈[n]i\in[n], let xi,1=eiiθx_{i,1}=e^{\mathbf{i}i\theta} and ei,l=0e_{i,l}=0 for all l≠1l\neq 1

Then we have for all i∈[n]i\in[n], for all j∈[n]j\in[n], ∥xi−xj∥22=f(i−j)\|x_{i}-x_{j}\|_{2}^{2}=f(i-j) for some function ff.

where the first step follows from the assumption that for each i∈[n]i\in[n] and l≠1l\neq 1, xi,1=eiiθx_{i,1}=e^{\mathbf{i}i\theta} and ei,l=0e_{i,l}=0, the second step follows from simple algebra, the third step follows from the ∣eijθ∣=1|e^{\mathbf{i}j\theta}|=1, and the last step follows from the definition of the function ff.

xi,1=cos⁡(iθ)x_{i,1}=\cos(i\theta) and xi,2=sin⁡(iθ)x_{i,2}=\sin(i\theta). For all l∉{1,2}l\notin\{1,2\}, we have xi,l=0x_{i,l}=0.

Then we have for all i∈[n]i\in[n], for all j∈[n]j\in[n], ∥xi−xj∥22=f(i−j)\|x_{i}-x_{j}\|_{2}^{2}=f(i-j) for some function ff.

where the first step follows from construction condition, the second step follows from simple algebra, and the last step follows from the trigonometric properties.

Let (s1,s2,…,sd)(s_{1},s_{2},\dots,s_{d}) be a permutation of (1,2,…,d)(1,2,\dots,d).

Then we have for all i∈[n]i\in[n], for all j∈[n]j\in[n], ∥xi−xj∥22=f(i−j)\|x_{i}-x_{j}\|_{2}^{2}=f(i-j) for some function ff.

When dd is odd, we can show similar results by the same way. Thus, we complete the proof. ∎

(QK⊤)i,j=bi−j+1(QK^{\top})_{i,j}=b_{i-j+1} if i≥ji\geq j

Then, there is a vector a=exp⁡(b)a=\exp(b) such that

where the second step follows from the fact that exp⁡(⋅)\exp(\cdot) is applied entry-wisely to a vector.

By the assumption from the Lemma statement that (QK⊤)i,j=bi−j+1(QK^{\top})_{i,j}=b_{i-j+1} if i≥ji\geq j and (QK⊤)i,j=bi−j+n+1(QK^{\top})_{i,j}=b_{i-j+n+1} if i<ji<j, we get

which is exactly equal to Circ(b)\mathsf{Circ}(b) (see Definition A.3).

Therefore, combining with Eq. (A.5), we have

For each i,j∈[n]i,j\in[n], (QK⊤)i,j=bi−j(QK^{\top})_{i,j}=b_{i-j}.

Then, there is a vector a=exp⁡(b)a=\exp(b) such that

Let z1,…,znz_{1},\dots,z_{n} defined in Definition A.24 satisfy the properties in Lemma A.20.

Then, there is a vector a=exp⁡(b)a=\exp(b) such that

By Lemma A.20, we have for all i∈[n]i\in[n], for all j∈[n]j\in[n],

where the first two steps from Definition A.24, and the last step from Lemma A.20. We finish the proof by denote bi−jb_{i-j} as g(i−j)g(i-j) in Lemma A.22. ∎

Appendix B 𝖼𝗈𝗇𝗏𝖼𝗈𝗇𝗏\mathsf{conv} Approximation in Gradient

In Section B.1, we present the basic definitions. In Section B.2, we combine all these definitions to form the loss function. In Section B.3, we analyze the running time. In Section B.4, we present the proof of the main theorem of conv\mathsf{conv} approximation in gradient.

B.2 Loss Functions

Now, we start the construction of the loss function.

For every j0∈[n]j_{0}\in[n], for every i0∈[d]i_{0}\in[d], we define L(x)j0,i0L(x)_{j_{0},i_{0}} to be :=0.5c(x)j0,i02:=0.5c(x)_{j_{0},i_{0}}^{2}.

The proof is trivial by element-wise multiplication. ∎

From the Lemma statement, by Lemma B.8, we have

where the first step is from the chain rule and the second step follows from Mj0,∗∘f(x)j0=f(x)j0M_{j_{0},*}\circ f(x)_{j_{0}}=f(x)_{j_{0}}.

where the last step is by simple algebra.

Let q(x)j0q(x)_{j_{0}} be defined as in Definition B.6:

Let p(x)j0p(x)_{j_{0}} be define as in Definition B.7:

where the 1st step is because of Definition 5.1, the second step follows from Eq. (B.2), the third step follows from Eq. (8), the fourth step follows from Eq. (9), and the fifth step follows from Fact D.9. ∎

B.3 Running Time

In this section, we analyze the running time of the conv\mathsf{conv} approximation approach for computing the training forward pass and backward gradient. We build upon the key definitions and loss functions introduced in the previous sections to derive the running time of the algorithm.

Suppose u(x)u(x) is a kk-conv\mathsf{conv} matrix defined in Definition 3.11 with known basis.

which can be done in O(knlog⁡n)O(kn\log n) time by Fact A.5.

The second part is trivial by Definition B.3. ∎

Suppose f(x)wf(x)w takes O(knlog⁡n)O(kn\log n) time.

c(x)c(x) can be expiciltiy computed in O(dknlog⁡n)O(dkn\log n) time.

Firstly we can compute f(x)h(y)f(x)h(y), this can be done in O(dknlog⁡n)O(dkn\log n), since we run f(x)f(x) times a vector oracle (Lemma B.10) for dd times.

q(x)q(x)’s rank-dd factorization can be explilcitly computed in O(nd)O(nd) time.

Note that q(x)=c(x)h(y)⊤q(x)=c(x)h(y)^{\top}. Since both c(x)c(x) and h(y)h(y) are known. Thus, the result is trivial. ∎

Let q(x)q(x) denote a rank-τ\tau matrix with known low-rank factorizations.

Using a standard linear algebra trick, we can show that

Let r(x)j0:=⟨f(x)j0,q(x)j0⟩r(x)_{j_{0}}:=\langle f(x)_{j_{0}},q(x)_{j_{0}}\rangle.

Let q(x)q(x) denote a rank-τ\tau matrix with known low-rank factorizations.

Let q(x)=UaUb⊤q(x)=U_{a}U_{b}^{\top}. It is easy to see that f(x)q(x)⊤f(x)q(x)^{\top} can be written as f(x)UbUa⊤f(x)U_{b}U_{a}^{\top}.

We firstly compute f(x)Ubf(x)U_{b}, since UbU_{b} has τ\tau columns, each column will take O(knlog⁡n)O(kn\log n) time, so in total it takes O(τknlog⁡n)O(\tau kn\log n) time.

Then, we know that r(x)j0=⟨(f(x)Ub)j0,∗,(Ua)j0,∗⟩r(x)_{j_{0}}=\langle(f(x)U_{b})_{j_{0},*},(U_{a})_{j_{0},*}\rangle which takes O(τ)O(\tau) time per j0j_{0}. There are nn different j0j_{0}, so it takes O(nτ)O(n\tau) time.

Overall it takes O(τknlog⁡n)O(\tau kn\log n) time.

Let p2(x)=diag⁡(r(x))f(x)p_{2}(x)=\operatorname{diag}(r(x))f(x) (This is obvious from definition of r(x)r(x))

For any vector ww, we firstly compute f(x)wf(x)w, then we compute diag⁡(r(x))(f(x)w)\operatorname{diag}(r(x))(f(x)w).

Firstly, we can compute p1(x)A2p_{1}(x)A_{2}, this takes dTp1d\mathcal{T}_{p_{1}} time.

Second, we can compute p2(x)A2p_{2}(x)A_{2}, this takes dTp2d\mathcal{T}_{p_{2}} time.

Putting it all together we complete the proof. ∎

B.4 Proof of Main Theorem

In this section, we present the formal proof of our main theorem regarding the conv\mathsf{conv} approximation approach for efficiently computing the training forward pass and backward gradient of the attention mechanism.

Suppose u(x)u(x) is a kk-conv\mathsf{conv} matrix defined in Definition 3.11 with known basis. Then there is an algorithm that runs in time O(d2knlog⁡n)O(d^{2}kn\log n) time to compute the gradient of attention loss defined in Definition 5.1.

We need to choose τ=d\tau=d, thus total running time is

by putting everything together from Lemma B.9, Lemma B.10, Lemma B.11, Lemma B.12, Lemma B.13, Lemma B.14, Lemma B.15, Lemma B.16. ∎

For the forward, we directly get the correctness by Theorem 4.4. For the backward, we directly run error propagation analysis which is similar to and proof of Lemma A.16.

Appendix C Incorporating Weighted Low Rank Approximation

In Section C.1, we introduce the preliminary for this section. In Section C.2, we present the proof of our main result for the low-rank approximation. In Section C.3, we present the algorithm and its mathematical properties for causal attention mask. In Section C.4, we analyze the algorithm and its mathematical properties for row change by amortized constant mask. In Section C.5, we study the algorithm and its mathematical properties for continuous row mask. In Section C.6, we analyze the property of the mask matrix with rr distinct columns or rr distinct rows.

∣H~i,j−Hi,j∣≤ϵ⋅Hi,j|\widetilde{H}_{i,j}-H_{i,j}|\leq\epsilon\cdot H_{i,j} with any arbitrary (i,j)∈[n]×[n](i,j)\in[n]\times[n].

In the following lemma, we prove the validity of the statement that if there exists an algorithm whose output is Y′=(W∘(U1U2⊤))vY^{\prime}=(W\circ(U_{1}U_{2}^{\top}))v in O(t)O(t) time, then there exists an algorithm outputs Y=D−1(W∘(U1U2⊤))vY=D^{-1}(W\circ(U_{1}U_{2}^{\top}))v in O(t+n)O(t+n) time. We will combine everything together and show the soundness of this statement later in the proof of Theorem C.4.

which takes O(t)O(t) time, then, there exists an algorithm promise that

Suppose there exists an algorithm whose output is Y′Y^{\prime} satisfying Y′=(W∘(U1U2⊤))vY^{\prime}=(W\circ(U_{1}U_{2}^{\top}))v and takes O(t)O(t) time. We denote this algorithm as Alg.

Let Y′=\textscAlg(U1,U2,v)Y^{\prime}=\textsc{Alg}(U_{1},U_{2},v). Let Y~=\textscAlg(U1,U2,1n)\widetilde{Y}=\textsc{Alg}(U_{1},U_{2},{\bf 1}_{n}). Then, Y=diag⁡(Y~)−1Y′Y=\operatorname{diag}(\widetilde{Y})^{-1}Y^{\prime}.

Computing Y′Y^{\prime} and Y~\widetilde{Y} takes O(t)O(t) time. Computing Y=diag⁡(Y~)−1Y′Y=\operatorname{diag}(\widetilde{Y})^{-1}Y^{\prime} takes O(n)O(n) time. Therefore, it takes O(t+n)O(t+n) time in total. ∎

C.2 Proof of Main Results

The time complexity to get Y~\widetilde{Y} is

O(knd)O(knd) when WW is a causal mask defined in Definition 3.2.

O(kd∑j=1nBj)O(kd\sum_{j=1}^{n}B_{j}) when WW is a row change mask defined in Definition 6.1.

O(kndlog⁡(n))O(knd\log(n)) when WW is a continuous row mask defined in Definition 6.2.

O(rnd)O(rnd) when WW is a distinct rr columns / rows mask defined in Definition 6.3 / Definition 6.4.

where the first step follows A~=W∘U1U2⊤\widetilde{A}=W\circ U_{1}U_{2}^{\top} and A=W∘HA=W\circ H, the second step follows mask is element-wise operation, the third step follows Definition C.1, and the last step follows A=W∘HA=W\circ H.

By Lemma C.2, the matrices U1U_{1} and U2U_{2} defining H~\widetilde{H} can be computed in O(nk)O(nk) time.

By Lemma C.3, if we can compute Y′=(W∘(U1U2⊤))VY^{\prime}=(W\circ(U_{1}U_{2}^{\top}))V in O(td)O(td) time, we can compute Y~\widetilde{Y} in O(td+nd)O(td+nd) time.

Finally, we finish the proof by following Lemma C.6 for the causal mask, Lemma C.8 for row change by amortized constant mask, Lemma C.9 for continuous row mask, and Lemma C.12 for distinct rr columns mask or distinct rr rows mask. ∎

C.3 Causal Attention Mask

In this section, we present the causal attention mask.

Let (U2⊤)j(U_{2}^{\top})_{j} denote the jj-th row of U2U_{2}.

Let SjS_{j} be the support set defined in Lemma C.5. Note that for the causal attention mask, we have Sj=[j]S_{j}=[j] for any j∈[n]j\in[n]. Thus, by Lemma C.5, we have

Computing (U2⊤)jvj(U_{2}^{\top})_{j}v_{j}, for all j∈[n]j\in[n] takes O(nk)O(nk) time.

Note that by the definition of inner product

Therefore, it also takes O(nk)O(nk) to compute (U1⊤)j⊤cj(U_{1}^{\top})_{j}^{\top}c_{j} for all j∈[n]j\in[n].

Therefore, it takes O(nk)O(nk) times in total. ∎

C.4 Row Change by Amortized Constant Mask

In this section, we analyze the row change by amortized constant mask.

Let W∈{0,1}n×nW\in\{0,1\}^{n\times n} be the causal attention mask defined in Definition 3.2. Then we have WW is a row change by amortized constant mask defined in Definition 6.1, where Bj=1B_{j}=1, ∀j∈[n]\forall j\in[n].

The proof directly follows the two Definitions. ∎

which takes O(k∑j=1nBj)O(k\sum_{j=1}^{n}B_{j}) time.

We will prove it by induction. It is obvious that base case Y1Y_{1} is correct, because S0=∅S_{0}=\emptyset.

For a fixed jj, we suppose YjY_{j} has the correct answer. This means cjc_{j} is correct for that jj, i.e., cj=∑l∈Sjbl=∑l∈Sj(U2⊤)lvlc_{j}=\sum_{l\in S_{j}}b_{l}=\sum_{l\in S_{j}}(U_{2}^{\top})_{l}v_{l}.

Now we use Qj+1+Q_{j+1}^{+} and Qj+1−Q_{j+1}^{-} to generate cj+1c_{j+1} by adding terms in Qj+1+Q_{j+1}^{+} and deleting terms in Qj+1−Q_{j+1}^{-},

where the first step follows Algorithm 5 line 10 and line 12, the second step follows Sj=(Sj∩Sj+1)∪(Sj∖Sj+1)S_{j}=(S_{j}\cap S_{j+1})\cup(S_{j}\setminus S_{j+1}), (Sj∩Sj+1)(S_{j}\cap S_{j+1}) and (Sj∖Sj+1)(S_{j}\setminus S_{j+1}) are disjoint, the third step follows simple algebra, and the last step follows the as the second step.

Therefore, we have cj+1c_{j+1} is correct, i.e., cj+1=∑l∈Sj+1bl=∑l∈Sj+1(U2⊤)lvlc_{j+1}=\sum_{l\in S_{j+1}}b_{l}=\sum_{l\in S_{j+1}}(U_{2}^{\top})_{l}v_{l}. Thus, Yj+1Y_{j+1} is also correct by Lemma C.5. Finally, we finish proving the correctness by math induction.

Note that there are two for-loops in this algorithm. Inside the inner for-loops, it takes O(k)O(k) time to compute

The inner for-loop has ∣Qj+∪Qj−∣=Bj|Q_{j}^{+}\cup Q_{j}^{-}|=B_{j} iterations, and the outer for-loop has nn iterations.

Therefore, it takes O(k∑j=1nBj)O(k\sum_{j=1}^{n}B_{j}) time in total. ∎

C.5 Continuous Row Mask

In this section, we study the continuous row mask.

The correctness is trivially from the construction of the segment tree.

The running time is dominated by O(nklog⁡n)O(nk\log n). This time comes from two parts, where the first is from building the segment tree by O(nk)O(nk), and the second part is from for-loop by O(nklog⁡n)O(nk\log n). ∎

C.6 Distinct r𝑟r Columns or Rows

Now, we analyze the mask matrix with rr distinct columns.

Let WW be the distinct rr columns mask defined in Definition 6.3. Let S1,⋯ ,Sr⊆[n]S_{1},\cdots,S_{r}\subseteq[n] denote rr disjoint subsets and ∪j∈[r]Sj=[n]\cup_{j\in[r]}S_{j}=[n] be defined in Definition 6.3. Let h:[r]→[n]h:[r]\rightarrow[n] denote that h(j)∈Sjh(j)\in S_{j} and h(j)h(j) is the smallest index in SjS_{j}.

Now, we analyze the mask matrix with rr distinct rows.

Let WW be the distinct rr rows mask defined in Definition 6.4. Let S1,⋯ ,Sr⊆[n]S_{1},\cdots,S_{r}\subseteq[n] denote rr disjoint subsets and ∪j∈[r]Sj=[n]\cup_{j\in[r]}S_{j}=[n] be defined in Definition 6.4. Let h:[r]→[n]h:[r]\rightarrow[n] denote that h(j)∈Sjh(j)\in S_{j} and h(j)h(j) is the smallest index in SjS_{j}.

Therefore, we have shown Eq. (10), which completes the proof. ∎

The correctness and running time is directly follows Lemma C.10 for the column case and Lemma C.11 for the row case. ∎

Appendix D Supporting Lemmas and Technical Results

In Section D.1, we present the matrix and vector properties. In Section D.2, we analyze and develop the tools for error analysis.

As H≠0n×nH\neq{\bf 0}_{n\times n}, it must have at least 11 conv\mathsf{conv} basis, and we proved the first part.

Now, we prove the second part by math induction.

where the first step follows from the fact that GG is a lower triangular matrix and Definition 3.9, the second step follows from simple algebra, and the last step follows from simple algebra.

As GG and G′G^{\prime} are lower triangular matrices, we have that G−conv(G~i+1,n−i)G-\mathsf{conv}(\widetilde{G}_{i+1},n-i) is a lower triangular matrix. Thus, we proved the following statement.

D.2 Tools for Error Analysis

It is trivial by exp⁡(a+b)=exp⁡(a)exp⁡(b)\exp(a+b)=\exp(a)\exp(b). ∎

where the first step follows simple algebra, and the last step follows triangle inequality.

For the first part, for any i∈[n],j∈[n]i\in[n],j\in[n], we have

where the first step follows simple algebra, the second step follows triangle inequality, the third step follows simple algebra, the fourth step follows D=diag⁡(A1n)D=\operatorname{diag}(A{\bf 1}_{n}), D~=diag⁡(A~1n)\widetilde{D}=\operatorname{diag}(\widetilde{A}{\bf 1}_{n}), A=exp⁡(H)A=\exp(H), A~=exp⁡(H~)\widetilde{A}=\exp(\widetilde{H}), the fifth steps follows triangle inequality, the sixth step follows Lemma D.3 and the last step follows D~i,i=∑k=1nexp⁡(H~i,k)\widetilde{D}_{i,i}=\sum_{k=1}^{n}\exp(\widetilde{H}_{i,k}) and Di,i=∑l=1nAi,lD_{i,i}=\sum_{l=1}^{n}A_{i,l}.

For the second part, for any i∈[n],j∈[n]i\in[n],j\in[n], we have

where the first step follows simple algebra, the second step follows triangle inequality, the third step follows A=exp⁡(H)A=\exp(H), A~=exp⁡(H~)\widetilde{A}=\exp(\widetilde{H}), the fourth step follows Lemma D.3, and the last step follows D~i,i=∑l=1nexp⁡(H~i,l)\widetilde{D}_{i,i}=\sum_{l=1}^{n}\exp(\widetilde{H}_{i,l}).

Let a,b≥0a,b\geq 0 and ϵ∈(0,0.1)\epsilon\in(0,0.1). If ∣a−b∣≤ϵa|a-b|\leq\epsilon a, then ∣a−b∣≤2ϵmin⁡{a,b}.|a-b|\leq 2\epsilon\min\{a,b\}.

It is trivial by considering two cases when b≥ab\geq a and b<ab<a. ∎

where the first step follows simple algebra, and the last step follows triangle inequality.

For the first part, for any i∈[n],j∈[n]i\in[n],j\in[n], we have

where the first step follows simple algebra, the second step follows triangle inequality, the third step follows simple algebra, the fourth step follows D=diag⁡(A1n)D=\operatorname{diag}(A{\bf 1}_{n}), D~=diag⁡(A~1n)\widetilde{D}=\operatorname{diag}(\widetilde{A}{\bf 1}_{n}), the fifth step follows triangle inequality, the sixth step follows Lemma D.5 and the last step follows D~i,i=∑k=1nA~i,k\widetilde{D}_{i,i}=\sum_{k=1}^{n}\widetilde{A}_{i,k} and Di,i=∑l=1nAi,lD_{i,i}=\sum_{l=1}^{n}A_{i,l}.

For the second part, for any i∈[n],j∈[n]i\in[n],j\in[n], we have

where the first step follows simple algebra, the second step follows triangle inequality, the third step follows Lemma D.5, and the last step follows D~i,i=∑l=1nA~i,l\widetilde{D}_{i,i}=\sum_{l=1}^{n}\widetilde{A}_{i,l}.

D.3 Tensor Tools for Gradient Computation

where the first step follows from the definition of outer product, the second step follows from the definition of vectorization operator vec⁡(⋅)\operatorname{vec}(\cdot) which stacks rows of a matrix into a column vector, and the last step follows from Definition 5.4. ∎

where the first step follows from that matrix can be written as a summation of vectors, the second step follows from Fact D.8, the third step follows from that matrix can be written as a summation of vectors, and the last step follows from the definition of vectorization operator vec⁡(⋅)\operatorname{vec}(\cdot). ∎

Appendix E More Related Work

introduced sparse factorizations of the attention matrix, which scale linearly with the sequence length. proposed a combination of local and global attention, where local attention captures short-range dependencies, and global attention captures long-range dependencies. introduced a random attention pattern that scales linearly with the sequence length while maintaining the ability to capture global dependencies. introduced a method called Performer, which uses random feature maps to approximate the attention matrix, resulting in linear complexity.

Other works have focused on studying the attention regression problem. The work by demonstrates that the forward computation of attention can be performed in sub-quadratic time complexity without explicitly constructing the full n×nn\times n attention matrix. answer the question of how efficiently the gradient can be computed in the field of fast attention acceleration. There are many works simplifying the attention regression problem. For example, studies the softmax regression problem, proposes and analyzes the attention kernel regression problem, studies the rescaled hyperbolic functions regression, all of which fill the need of using iterative methods analyzing the attention regression related problems. After that, provides an un-simplified version of the attention regression problem, and construct and study two-layer regression problems related to attention.

Moreover, applies the attention computation to design a decentralized LLM. shows that attention layers in transformers learn 2-dimensional cosine functions, similar to how single-hidden layer neural networks learn Fourier-based circuits when solving modular arithmetic tasks. analyzes the data recovery by using attention weight. designs the quantum algorithm for solving the attention computation problem. analyzes the sparsification of the attention problem. applies the Zero-th Order method to approximate the gradient of the attention regression problem. analyzes the techniques for differentially private approximating the attention matrix. extends the work of to study the multiple softmax regression via the tensor trick, which is the same as our Fact D.9. analyzes the hardness for dynamic attention maintenance in large language models. studies how to balance the trade-off between creativity and reality using attention computation. formulates the exp⁡\exp, cosh⁡\cosh and sinh⁡\sinh regression problems from the original attention computation. replace the softmax unit by the polynomial unit. introduces an outlier-efficient modern Hopfield model and corresponding Hopfield layers that provide an alternative to the standard attention mechanism in transformer-based models.

Convolution computation.

Convolution computation has been a focal point of research in the field of deep learning and computer vision. Convolutional Neural Networks (CNNs) have demonstrated remarkable performance in various tasks such as image classification , object detection , and semantic segmentation . The efficiency and effectiveness of convolution computation play a crucial role in the success of CNNs.

One notable development in convolution computation is the introduction of depthwise separable convolutions, proposed by . This technique separates the convolution operation into two steps: depthwise convolution and pointwise convolution, which significantly reduces the computational cost and model size while maintaining comparable accuracy. This advancement has enabled the deployment of CNNs on resource-constrained devices such as mobile phones and embedded systems. Another important contribution is the development of strided convolutions, which allow for downsampling the spatial dimensions of feature maps without the need for additional pooling layers. This technique has been widely adopted in modern CNN architectures like ResNet and Inception , leading to more efficient and compact models.

To further improve the computational efficiency of convolutions, several works have explored the use of group convolutions. ResNeXt and ShuffleNet have demonstrated that by dividing the input channels into groups and performing convolutions within each group, the computational cost can be significantly reduced while maintaining or even improving the model’s performance. In addition to architectural innovations, there have been advancements in the implementation of convolution computation on hardware. The introduction of specialized hardware such as GPUs and TPUs has greatly accelerated the training and inference of CNNs . Recent research has also focused on reducing the computational cost of convolutions through network pruning and quantization techniques. proposed a pruning method that removes unimportant connections in the network, resulting in sparse convolutions that require less computation. Quantization techniques, such as those introduced by , reduce the precision of weights and activations, allowing for faster convolutions with minimal loss in accuracy.

References