Beyond Linear Approximations: A Novel Pruning Approach for Attention Matrix

Yingyu Liang, Jiangxuan Long, Zhenmei Shi, Zhao Song, Yufa Zhou

Introduction

Large Language Models (LLMs) based on the transformer architecture , including GPT-4o , Claude , and OpenAI’s recent o1 , have shown immense potential to enhance our daily lives. They revolutionize fields like conversational AI , AI agents , search AI , and AI assistants . With their growing capabilities, LLMs are powerful tools shaping the future of technology. However, the current state-of-the-art LLM weights number is extremely large. For instance, the smallest version of Llama 3.1 needs 8 billion parameters, which takes more than 16GB GPU memory with half float precision and requires significant inference time. Due to large memory and high computational cost, deploying such models on edge devices such as smartphones becomes challenging.

In the following, we introduce some key backgrounds and our contributions in detail.

We define the attention matrix in self-attention mechanism as below:

where (1) D:=diag⁡(exp⁡(XWX⊤)⋅1n)D:=\operatorname{diag}(\exp(XWX^{\top})\cdot{\bf 1}_{n}), (2) exp⁡\exp denotes the exponential function and is applied entry-wisely, (3) diag⁡()\operatorname{diag}() operation takes a vector and outputs a diagonal matrix with the entries of that vector, and (4) 1n{\bf 1}_{n} denotes the length-nn all ones vector.

Further, we introduce the problem setup of our Attention Weights Pruning. By selectively reducing the number of non-zero elements in the attention weight matrix WW in Definition 1.1, we can preserve model performance while lowering computational cost and GPU memory usage. Below, we formally define the Attention Weights Pruning problem and the corresponding loss function:

Thus, the Attention Weights Pruning optimization problem is min⁡M L(M)\min_{M}~{}\mathcal{L}(M).

2 Our Contributions

This is the first work studying the Attention Weights Pruning problem, which is an approximation problem to a non-linear function. We provide an algorithm to obtain the near-optimal pruning mask based on Gradient Descent (GD) with convergence guarantee.

For any ϵ>0\epsilon>0, our Algorithm 1 can converge to the near-optimal pruning mask for the Attention Weights Pruning problem (Definition 1.2) in O(dpoly⁡(n)/ϵ)O(d\operatorname{poly}(n)/\epsilon) time with O(ξ+ϵ)O(\xi+\epsilon) error, where ξ\xi is a small term depending on intrinsic property of the data and weights.

In the above theorem, ξ\xi can be arbitrarily small as ξ→0\xi\rightarrow 0 when the regularization coefficient λ→0\lambda\rightarrow 0. So our analysis shows that although the objective function is highly non-linear, the GD training can converge to a near-optimal pruning mask solution, supported by our experiments in Section 6.

This is the first work that analyzes the weights pruning problem based on Softmax\mathsf{Softmax} attention, which is a non-linear function.

We provide the closed form of the gradient of Attention Weights Pruning loss function (Theorem 5.3), and Lipschitz of that gradient (Theorem 5.4),

We provide Gradient Descent based Algorithm 1 to obtain the near-optimal pruning mask and its convergence guarantee (Theorem 4.1).

We conduct preliminary experiments to verify the effectiveness of our method (Section 6).

Our paper is organized as follows. In Section 2, we review the related work. Section 3 introduces key concepts and definitions essential for the subsequent sections. In Section 4, we present our main result. Section 5 offers a technical overview of the methods employed. Experimental results are discussed in Section 6. Finally, Section 7 summarizes our findings and offers concluding remarks.

Related Work

Model compression plays a critical role in improving the efficiency and deployment of large language models (LLMs) for its effectiveness in reducing computational overhead while preserving performance. Common compression techniques include quantization , pruning , and knowledge distillation . Specifically, pruning techniques have been developed extensively, such as unstructured pruning, which removes individual weights , and structured pruning, which eliminates entire components like neurons or attention heads . proposed Wanda, a novel unstructured pruning technique that uses weight-activation products to induce up to 50% sparsity in LLMs without retraining, achieving competitive results with significantly lower computational cost. SparseGPT introduced a one-shot pruning method that achieves up to 60% sparsity in large GPT-family models with minimal impact on performance. A follow-up work improved the complexity analysis of SparseGPT, reducing the running time from O(d3)O(d^{3}) to O(d2.53)O(d^{2.53}), enabling faster pruning on LLMs. These techniques together contribute to more scalable and resource-efficient LLMs, maintaining competitive performance while having substantial reductions in computational resources.

2 Attention Acceleration

Attention mechanism has faced criticism due to its quadratic time complexity with respect to context length . Addressing this criticism, a variety of approaches are employed, including sparse attention , low-rank approximations , and kernel-based methods , to reduce computational overhead and improve scalability. enable the derivation of a low-rank representation of the attention matrix, which accelerates both the training and inference processes of single attention layer, tensor attention, and multi-layer transformer, achieving almost linear time complexity . Other approaches like Mamba , Linearizing Transformers , Hopfield Models , and PolySketchFormer focus on architectural modifications and implementation optimizations to enhance performance. System-level optimizations such as FlashAttention and block-wise parallel decoding further improve efficiency. Collectively, these innovations have significantly augmented transformer models’ ability to handle longer input sequences, unlocking broader applications across multiple sectors .

Preliminary

In this section, we introduce some basic concepts and key definitions. In Section 3.1, we introduce some basic notations we use in this paper. In Section 3.2, we provide the definition of attention weights pruning.

2 Attention Weights Pruning

The causal attention mask ensures that each token in the sequence can attend only to itself and preceding tokens. Here, we provide the formal definition for the causal attention mask:

We define the causal attention mask as Mc∈{0,1}n×nM_{c}\in\{0,1\}^{n\times n}, where (Mc)i,j=1(M_{c})_{i,j}=1 if i≥ji\geq j and (Mc)i,j=0(M_{c})_{i,j}=0 otherwise.

Now, we incorporate Attention Weights Pruning (see Definition 1.2) with causal attention mask McM_{c}.

Let Mc∈{0,1}n×nM_{c}\in\{0,1\}^{n\times n} be the causal attention mask defined in Definition 3.1. Let A:=exp⁡(XWX⊤)∘McA:=\exp(XWX^{\top})\circ M_{c} and A~:=exp⁡(X(M∘W)X⊤)∘Mc\widetilde{A}:=\exp(X(M\circ W)X^{\top})\circ M_{c}. Let D:=diag⁡(A⋅1n)D:=\operatorname{diag}(A\cdot{\bf 1}_{n}) and D~:=diag⁡(A~⋅1n)\widetilde{D}:=\operatorname{diag}(\widetilde{A}\cdot{\bf 1}_{n}). We define Attention Weights Pruning with Causal Attention Mask loss function to be

Main Results

In this section, we provide our main results. We provide an Algorithm 1 for Attention Weights Pruning problem based on Gradient Descent (GD). We also prove the convergence for our GD algorithm in Theorem 4.1.

Let μ=2min⁡i,j∈[d]{∣Wi,j∣}⋅β⋅δ\mu=2\min_{i,j\in[d]}\{|W_{i,j}|\}\cdot\beta\cdot\delta.

Let ξ=12n−1.5max⁡i,j∈[d]{∣Wi,j∣}⋅∥X∥F2⋅λd/μ\xi=12n^{-1.5}\max_{i,j\in[d]}\{|W_{i,j}|\}\cdot\|X\|_{F}^{2}\cdot\lambda d/\mu.

Then, for any ϵ>0\epsilon>0, provided η<1/L\eta<1/L where LL is the Lipschitz constant for ∇ML(M)\nabla_{M}\mathcal{L}(M) (see Theorem 5.4), GD (Algorithm 1) with fixed step size η\eta and run for T=4L(M(0))/(ημϵn2)T=4\mathcal{L}(M^{(0)})/(\eta\mu\epsilon n^{2}) iterations results in the following guarantee,

The two assumptions in Theorem 4.1 are practical. The first assumption of the positive definite matrix is widely used in theoretical deep learning analysis, e.g., . The second assumption is natural, as D~−1A~\widetilde{D}^{-1}\widetilde{A} is the pruned attention matrix, where each entry is

After solving the optimization problem, we obtain a pruning mask with real-valued entries. In practice, however, this pruning mask must be converted into a binary form, specifically M∈{0,1}d×dM\in\{0,1\}^{d\times d}. We define the pruning ratio ρ\rho as the percentage of weights to be pruned. We apply this ratio by setting the pruning mask entries to zero for weights that fall below the ρ\rho-th percentile and to one for those above. This ensures that only the specified proportion of weights are pruned.

Technique Overview

In Section 5.1, we introduce some useful tools from previous work. In Section 5.2, we derive the close form of the gradient of Attention Weights Pruning. In Section 5.3, we calculate the Lipschitz constant of that gradient. In Section 5.4, we prove the PL inequality for our loss function.

To analyze the convergence behavior of GD for our optimization problem (Definition 3.2), we first introduce the concept of gg-proxy, ξ\xi-optimal Polyak–Łojasiewicz(PL) inequality , under which GD will converge:

PL inequality is a powerful tool for studying non-convex optimization, and it has been used in recent studies on provable guarantees for neural networks trained by gradient descent . It provides a proxy convexity property, although the objective function is non-convex. In detail, for a function with good smoothness property, we can find some proxy functions and show the convergence by utilizing these proxy functions.

Leveraging this PL inequality, derives the following GD convergence guarantees.

The above theorem establishes that under the (g,ξ,α,μ)(g,\xi,\alpha,\mu)-PL inequality and Lipschitz continuity of the gradient, GD converges to a point where the proxy function g(w)g(w) is within ϵ\epsilon of ξ\xi. To apply this result to our specific problem, we need to verify these conditions for our loss function L(M){\cal L}(M).

2 Closed Form of Gradient

As a first step, we compute the close form of the gradient ∇ML(M)\nabla_{M}{\cal L}(M). The pruning mask MM is inside a non-linear function Softmax, which complicates our calculation. We defer the proof to Section B.

Based on Theorem 5.3, we calculate the gradient of the pruning mask from Line 10 to Line 15 in our Algorithm 1.

3 Lipschitz of Gradient

Having obtained the close form of gradient, we proceed to investigate its Lipschitz continuity. We aim to show that the gradient ∇ML(M)\nabla_{M}{\cal L}(M) is Lipschitz continuous with respect to MM.

We defer the proof to Section E. Establishing the Lipschitz continuity of the gradient satisfies one of the necessary conditions for applying Theorem 5.2. The above theorem implicates that the gradient for MM is upper bounded, providing a way to choose step size.

4 PL Inequality of Gradient

Next, we need to verify that our loss function satisfies the PL inequality with appropriate parameters. To complete the verification of the conditions required for convergence, we demonstrate that L(M){\cal L}(M) satisfies the PL inequality. We show that ∇ML(M)\nabla_{M}\mathcal{L}(M) satisfies the PL inequality in this lemma:

Let μ=2min⁡i,j∈[d]{∣Wi,j∣}⋅β⋅δ\mu=2\min_{i,j\in[d]}\{|W_{i,j}|\}\cdot\beta\cdot\delta.

Let ξ=12nmax⁡i,j∈[d]{∣Wi,j∣}⋅∥X∥F2⋅λd/μ\xi=12\sqrt{n}\max_{i,j\in[d]}\{|W_{i,j}|\}\cdot\|X\|_{F}^{2}\cdot\lambda d/\mu.

We defer the proof to Section F. By confirming that L(M){\cal L}(M) satisfies the PL inequality and that its gradient is Lipschitz continuous, we then apply Theorem 5.2 to conclude that GD will converge to a solution within our desired error tolerance, and further prove Theorem 4.1.

To prove the PL inequality, we also need the following two key Lemmas, which introduce our two assumptions in our Theorem 4.1, XX⊤⪰βIXX^{\top}\succeq\beta I and min⁡i,j∈[n](D~−1A~)i,j≥δ>0\min_{i,j\in[n]}(\widetilde{D}^{-1}\widetilde{A})_{i,j}\geq\delta>0.

Experiment

In this section, we discuss the experiments conducted to illustrate the effectiveness of our Algorithm 1. We first introduce our settings in Section 6.1. Then, we present our results in Section 6.2.

We implement our method following the pseudocode in Algorithm 1, using NumPy and JAX for acceleration. We evaluate our method on unstructured sparsity, meaning that zeros can occur anywhere within the attention weight matrix WW. Specifically, we use Definition 3.2 as our loss function, optimizing over the pruning mask MM using gradient descent based on the closed-form expression derived in Theorem 5.3. To accelerate convergence, we leverage momentum into the optimization process and fix the momentum parameter at 0.90.9. After obtaining the optimal pruning mask, we convert MM to a binary pruning mask to prune WW, maintaining sparsity at the desired pruning ratio ρ\rho. We use the relative error as our evaluation metric, which is defined as

where D~\widetilde{D}, A~\widetilde{A}, DD, AA are defined in Definition 3.2.

Baselines.

We compare our method with two linear pruning approaches, namely Wanda and SparseGPT . Wanda is a pruning method that removes weights with the smallest magnitudes multiplied by the corresponding input activations, achieving sparsity without requiring retraining or weight updates. SparseGPT is a second-order pruning method that utilizes the Hessian matrix to prune a portion of the weight matrix while simultaneously updating the remaining parameters. We implement Wanda and SparseGPT as described in their respective papers. Notably, since the settings of SparseGPT and Wanda are linear, we do not prune the fused weight matrix WW directly; instead, we prune WQW_{Q} and WKW_{K} separately (see Figure 1).

Data.

Setup.

2 Results

Overall, the results in Figure 2 show that our Algorithm 1 outperforms Wanda and SparseGPT with a large margin, which supports our theoretical analysis in Theorem 4.1. In the following, we will discuss each setting in detail.

The leftmost column of Figure 2 investigates the impact of the regularization coefficient λ\lambda on relative error. As λ\lambda increases from very small values, the relative error initially decreases sharply for our algorithm, reaching a minimum before gradually rising again, which forms a UU shape curve. This behavior indicates that there is an optimal λ\lambda where our algorithm achieves its best performance around 2−42^{-4}. The U-shape curve phenomena are well-known in most hyper-parameter choosing, e.g., regularization coefficient.

Relation with input sequence length n𝑛n.

The center column of Figure 2 explores how the relative error changes with respect to the input sequence length nn. As nn increases, the relative error for all three methods grows, though at different rates. Our method demonstrates a slower increase, maintaining a significant margin over both Wanda and SparseGPT, particularly for larger values of nn. Wanda, while showing better performance than SparseGPT for larger sequence lengths, becomes comparable to SparseGPT as nn is relatively small.

Relation with pruning ratio ρ𝜌\rho.

The rightmost column of Figure 2 illustrates the relationship between the relative error and the pruning ratio ρ\rho for the three methods under comparison: our algorithm, Wanda, and SparseGPT. As the pruning ratio ρ\rho increases, all methods exhibit a rise in relative error, indicating a degradation in approximation accuracy. However, our algorithm consistently outperforms both Wanda and SparseGPT across the range of ρ\rho, with a lower relative error. SparseGPT and Wanda follow a similar trend, closely tracking each other.

Conclusion

This paper introduced a novel approach to LLM weight pruning that directly optimizes for approximating the attention matrix. We provided theoretical guarantees for the convergence of our Gradient Descent-based algorithm to a near-optimal pruning mask solution. Preliminary results demonstrated the method’s effectiveness in maintaining model performance while reducing computational costs. This work establishes a new theoretical foundation for pruning algorithm design in LLMs, potentially enabling more efficient inference on resource-constrained devices.

Acknowledgement

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

References

Roadmap.

The appendix is organized as follows. In Section A, we give the preliminary of our paper. In Section B, we provide detailed gradient analysis of loss function. In Section C, we provide details about how we integrate the gradient of loss function into matrix form. In Section D, we bound some basic functions to be used later. In Section E, we provide proof for the Lipschitz property of the gradient of the loss function. In Section F, we provide proof of convergence for GD.

Appendix A Preliminary

In Section A.1, we introduce some notations we use in this paper. In Section A.2, we provide some basic facts.

A.2 Facts

x⏟n×1=u∘v⏟n×1=diag⁡(u)⏟n×nv⏟n×1=diag⁡(v)⏟n×nu⏟n×1\underbrace{x}_{n\times 1}=\underbrace{u\circ v}_{n\times 1}=\underbrace{\operatorname{diag}(u)}_{n\times n}\underbrace{v}_{n\times 1}=\underbrace{\operatorname{diag}(v)}_{n\times n}\underbrace{u}_{n\times 1}

⟨u,v⟩=⟨v,u⟩=u⊤v=v⊤u\langle u,v\rangle=\langle v,u\rangle=u^{\top}v=v^{\top}u

u∘v=v∘u=diag⁡(u)v=diag⁡(v)uu\circ v=v\circ u=\operatorname{diag}(u)v=\operatorname{diag}(v)u

⟨u,v⟩=⟨u∘v,1n⟩\langle u,v\rangle=\langle u\circ v,{\bf 1}_{n}\rangle

⟨u∘v,w⟩=⟨u∘w,v⟩=⟨w∘v,u⟩\langle u\circ v,w\rangle=\langle u\circ w,v\rangle=\langle w\circ v,u\rangle

u⊤(v∘w)=u⊤diag⁡(v)wu^{\top}(v\circ w)=u^{\top}\operatorname{diag}(v)w

(X∘Y)⊤=X⊤∘Y⊤(X\circ Y)^{\top}=X^{\top}\circ Y^{\top}

X∘eiej⊤=Xi,jeiej⊤X\circ e_{i}e_{j}^{\top}=X_{i,j}e_{i}e_{j}^{\top}

diag⁡(u)Zdiag⁡(v)=(uv⊤)∘Z\operatorname{diag}(u)Z\operatorname{diag}(v)=(uv^{\top})\circ Z

XY⊤=∑i∈[d]X∗,iY∗,i⊤XY^{\top}=\sum_{i\in[d]}X_{*,i}Y^{\top}_{*,i}

∑j∈[n]u∘A∗,j=u∘∑j∈[n]A∗,j\sum_{j\in[n]}u\circ A_{*,j}=u\circ\sum_{j\in[n]}A_{*,j}

∥X∥F2=tr⁡[XX⊤]\|X\|_{F}^{2}=\operatorname{tr}[XX^{\top}]

tr⁡[XY⊤]=tr⁡[Y⊤X]\operatorname{tr}[XY^{\top}]=\operatorname{tr}[Y^{\top}X]

∥diag⁡(u)∥F=∥u∥2\|\operatorname{diag}(u)\|_{F}=\|u\|_{2}

∥aX∥F=∣a∣∥X∥F\|aX\|_{F}=|a|\|X\|_{F} (absolute homogeneity).

∥X+Y∥F≤∥X∥F+∥Y∥F\|X+Y\|_{F}\leq\|X\|_{F}+\|Y\|_{F} (triangle inequality).

∣⟨X,Y⟩∣≤∥X∥F⋅∥Y∥F|\langle X,Y\rangle|\leq\|X\|_{F}\cdot\|Y\|_{F} (Cauchy–Schwarz inequality).

∥X∘Y∥F≤∥X∥F⋅∥Y∥F\|X\circ Y\|_{F}\leq\|X\|_{F}\cdot\|Y\|_{F}.

For any i∈[n]i\in[n], j∈[d]j\in[d], we have ∣Xi,j∣≤∥X∥F|X_{i,j}|\leq\|X\|_{F}.

∥X∥≤∥X∥F≤k∥X∥\|X\|\leq\|X\|_{F}\leq\sqrt{k}\|X\| where kk is the rank of XX.

∥Y⋅Z∥F≤∥Y∥F⋅∥Z∥F\|Y\cdot Z\|_{F}\leq\|Y\|_{F}\cdot\|Z\|_{F}.

where the first step, the second step and the third step follow from Fact A.3, the fourth step follows from Fact A.4.

Appendix B Gradient Calculation

We define i0i_{0}-th row of u~(M)\widetilde{u}(M) as follows

Let Mc∈{0,1}n×nM_{c}\in\{0,1\}^{n\times n} be the causal attention mask defined in Definition 3.1.

Let uu, u~(M)\widetilde{u}(M) be defined as Definition B.1.

We define i0i_{0}-th entry of α~(M)\widetilde{\alpha}(M) as follows

Then, we introduce the Softmax probability function.

Let Mc∈{0,1}n×nM_{c}\in\{0,1\}^{n\times n} be the causal attention mask defined in Definition 3.1.

Let uu, u~(M)\widetilde{u}(M) be defined as Definition B.1.

Let α\alpha, α~(M)\widetilde{\alpha}(M) be defined as Definition B.2.

We define i0i_{0}-th row of f~(M)\widetilde{f}(M) as follows

We define the entry in i0i_{0}-th row, j0j_{0}-th column of f~(M)\widetilde{f}(M) as follows

Then, we introduce the one unit loss function.

Let ff, f~\widetilde{f} be defined in Definition B.3.

We define i0i_{0}-th row of c(M)c(M) as follows

We define j0j_{0}-th column of c(M)c(M) as follows

We define the entry in i0i_{0}-th row, j0j_{0}-th column of c(M)c(M) as follows

Then, we introduce the reconstruction error.

Then, we introduce the regularization term.

Finally, we introduce the overall loss function.

We introduce the Lemma of gradient for each row of X(M∘W)X⊤X(M\circ W)X^{\top}.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], we have

We can simplify the derivative expression

where the first and second step follows from Fact A.1, the third step follows from Fact A.3.

where the first follows from Fact A.4, the second step follows from for any matrix XX, Xi,j=(X⊤)j,iX_{i,j}=(X^{\top})_{j,i}, the third step follows from Fact A.7, and the fourth step follows from for any matrix XX, Xi,j=(X⊤)j,iX_{i,j}=(X^{\top})_{j,i}.

where the first step follows from Eq. (B.2) and Eq. (B.2), and the second step and the third step follows from basic algebra. ∎

We introduce the Lemma of gradient for each row of u~(M)\widetilde{u}(M).

B.3 Gradient for Each Row of u~​(M)~𝑢𝑀\widetilde{u}(M)

Let u~(M)\widetilde{u}(M) be defined in Definition B.1.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], we have

where the first step follows from Definition B.1, the second step follows from Fact A.3, the third step follows from Definition B.1, and the fourth step follows from Lemma B.8. ∎

B.4 Gradient for Each Entry of α~​(M)~𝛼𝑀\widetilde{\alpha}(M)

We introduce the Lemma of gradient for each entry of α~(M)\widetilde{\alpha}(M).

Let u~(M)\widetilde{u}(M) be defined in Definition B.1.

Let α~(M)\widetilde{\alpha}(M) be defined in Definition B.2.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], we have

where the first step follows from Definition B.2, the second step follows from product rule of inner product in Fact A.3, the third step follows from product rule of Hadamard product in Fact A.3, the fourth step follows from Lemma B.9, and the last step follows from Fact A.4. ∎

B.5 Gradient for Each Entry of f~​(M)~𝑓𝑀\widetilde{f}(M)

We introduce the Lemma of gradient for each entry of f~(M)\widetilde{f}(M).

Let u~(M)\widetilde{u}(M) be defined in Definition B.1.

Let α~(M)\widetilde{\alpha}(M) be defined in Definition B.2.

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], we have

where the first step follows from Definition B.3, and the second step follows from Fact A.3.

In the following part, we compute the two terms separately.

where the first step follows from Fact A.3, the second step follows from Lemma B.10, the third step follows from basic algebra, and the fourth step follows from Definition B.3.

where the first step and the second step follow from basic algebra, the third step follows from Lemma B.9, and the fourth step follows from Definition B.3.

where the first step follows from Eq. (B.5), and the second step follows from Eq. (B.5) and Eq. (B.5). ∎

B.6 Gradient for Each Entry of c(M)

We introduce the Lemma of gradient for each entry of c(M)c(M).

Let f~(M)\widetilde{f}(M), ff be defined in Definition B.3.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], we have

where the first step follows from Definition B.4, the second step follows from Fact A.3, and the third step follows from Lemma B.11. ∎

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], we have

B1(M):=c(M)i0,j0f~(M)i0,j0Wi1,j1Xi0,i1Xj0,j1B_{1}(M):=c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{j_{0},j_{1}}

B2(M):=−c(M)i0,j0f~(M)i0,j0⟨f~(M)i0,Wi1,j1Xi0,i1X∗,j1⟩B_{2}(M):=-c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}\langle\widetilde{f}(M)_{i_{0}},W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{*,j_{1}}\rangle

where the first step follows from Definition B.5, the second step follows from the definition of Frobenius norm of matrix, the third step follows from Fact A.3, and the fourth step follows from Fact A.3.

where the second step follows from basic algebra. ∎

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], we have

where the first step follows from Definition B.6, the second step follows from the definition of Frobenius norm of matrix, and the third step follows from Fact A.3. ∎

B.9 Gradient for ℒ​(M)ℒ𝑀{\cal L}(M)

We introduce the Lemma of gradient for L(M){\cal L}(M).

Let u~(M)\widetilde{u}(M) be defined in Definition B.1.

Let α~(M)\widetilde{\alpha}(M) be defined in Definition B.2.

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

Let L(M){\cal L}(M) be defined in Definition B.7.

Let i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], we have

B1(M):=c(M)i0,j0f~(M)i0,j0Wi1,j1Xi0,i1Xj0,j1B_{1}(M):=c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{j_{0},j_{1}}

B2(M):=−c(M)i0,j0f~(M)i0,j0⟨f~(M)i0,Wi1,j1Xi0,i1X∗,j1⟩B_{2}(M):=-c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}\langle\widetilde{f}(M)_{i_{0}},W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{*,j_{1}}\rangle

where the first step follows from Definition B.7, the second step follows from Fact A.3, and the third step follows from Lemma B.13 and Lemma B.14. ∎

Appendix C Matrix Form

Given the matrix form, we define pp to simplify the calculation.

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

We define the j0j_{0}-th column of p1p_{1} as follows

We define the i0i_{0}-th row of p2p_{2} as follows

We introduce the matrix view of B1(M)B_{1}(M) and its summation.

Let B1(M,i1,j1):=c(M)i0,j0f~(M)i0,j0Wi1,j1Xi0,i1Xj0,j1B_{1}(M,i_{1},j_{1}):=c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{j_{0},j_{1}}, which is defined in Lemma B.15

where the first step follows from the definition of C1C_{1}, the second step follows from Fact A.1, and the third step follows from Fact A.2.

Following from Fact A.1, we can get j1j_{1}-th column of C1C_{1}

where the second step follows from Fact A.2.

Following from Fact A.1, we can get C1(M)C_{1}(M)

where the second step follows from Fact A.4.

Part 2. We further compute the summation of C1(M)C_{1}(M).

where the first step follows from Eq. (C.1), the second step follows from basic algebra, the third step follows from Fact A.4, and the fourth step follows from Definition C.1.

We introduce the matrix view of B2(M)B_{2}(M) and its summation.

Let B2(M,i1,j1):=−c(M)i0,j0f~(M)i0,j0⟨f~(M)i0,Wi1,j1Xi0,i1X∗,j1⟩B_{2}(M,i_{1},j_{1}):=-c(M)_{i_{0},j_{0}}\widetilde{f}(M)_{i_{0},j_{0}}\langle\widetilde{f}(M)_{i_{0}},W_{i_{1},j_{1}}X_{i_{0},i_{1}}X_{*,j_{1}}\rangle be defined in Lemma B.15.

where the first step follows from the definition of C2C_{2}, the second step, the third step and the fourth step follows from Fact A.4.

Following from Fact A.1, we can get j1j_{1}-th column of C2C_{2}

where the second step and the fourth step follows from Fact A.4, and the third step follows from Fact A.1.

Following from Fact A.1, we can get C2C_{2}.

Part 2. We further compute the summation of C2C_{2}

where the first step follows from Eq. (7), the second step and the third step follows from Fact A.4.

where the third step follows from Definition C.1. ∎

We introduce the matrix view of B3(M)B_{3}(M).

Let B3(M,i1,j1):=λMi1,j1B_{3}(M,i_{1},j_{1}):=\lambda M_{i_{1},j_{1}} be defined in Lemma B.15.

The proof is straightforward. By the definition of C3(M)C_{3}(M), for all i1,j1∈[d]i_{1},j_{1}\in[d], the (i1,j1)(i_{1},j_{1})-th entry of C3(M)C_{3}(M) is given by C3(i1,j1)=B3(M,i1,j1)=λMi1,j1C_{3}(i_{1},j_{1})=B_{3}(M,i_{1},j_{1})=\lambda M_{i_{1},j_{1}}. Thus, the entire matrix C3(M)C_{3}(M) has entries that correspond to those of λM\lambda M. Therefore, we can conclude that C3(M)=λMC_{3}(M)=\lambda M as required. ∎

We introduce the matrix form of overall loss function.

Let L(M){\cal L}(M) be defined in Definition B.7.

where the first step follows from Lemma B.15, the second step follows from Lemma C.2, Lemma C.3, and Lemma C.4, the third step follows from basic algebra, and the fourth step follows from Definition C.1. ∎

Appendix D Bounds for Basic Functions

Here we introduce our bounded parameters assumption.

Let RR be some fixed constant satisfies R>1R>1.

Here we present the lemma of bounds for MM and McM_{c}.

Let M∈d×dM\in^{d\times d} and Mc∈{0,1}n×nM_{c}\in\{0,1\}^{n\times n} be the causal attention mask defined in Definition 3.1. For MM, we have

This Lemma simply follows from the definition of Frobenius norm, given that the max value of each entry in MM and McM_{c} is 11. ∎

D.2 Bounds for Basic Functions

We first introduce the lemma of bounds for basic function.

Under Assumption D.1, for all i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], i1∈[d]i_{1}\in[d], j1∈[d]j_{1}\in[d], we have the following bounds

Proof of Part 1. Each entry in f~(M)\widetilde{f}(M) present a probability, thus for i0∈[n]i_{0}\in[n], j0∈[n]j_{0}\in[n], we have

For any i0i_{0}-th row of f~(M)\widetilde{f}(M), following from the definition of Softmax function, we know

which follows from f~(M)i0,j0≤(f~(M)i0,j0)2\widetilde{f}(M)_{i_{0},j_{0}}\leq(\widetilde{f}(M)_{i_{0},j_{0}})^{2}. Then, we can show

Proof of Part 2. Following from Part 1., we can show

where the first step follows from Definition B.4, the second step follows from triangle inequality.

Proof of Part 3. We have 0≤f~(M)i0,j0≤10\leq\widetilde{f}(M)_{i_{0},j_{0}}\leq 1, so we have

where the second step follows from Part 2..

Proof of Part 5. The proof simply follows from Assumption D.1 and Fact A.5.

Proof of Part 6. The proof simply follows from Assumption D.1 and Fact A.5.

Proof of Part 8. The proof simply follows from Part 4., Part 5., Part 6. and Part 7..

Proof of Part 9. The proof simply follows from Part 4., Part 5., Part 6. and Part 7..

where the first step follows from Fact A.4 the second step follows from Fact A.5, the third step follows from Part 3., and the last step follows from simple algebra. ∎

D.3 Bounds for Gradient of f~​(M)~𝑓𝑀\widetilde{f}(M)

We introduce the lemma of bounds for gradient of f~(M)\widetilde{f}(M).

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

Appendix E Lipschitz of Gradient

Here we introduce the fact of mean value theorem for matrix function.

For the convenience of proof, we define xx and yy as follows:

x:=vec⁡(X)x:=\operatorname{vec}(X) and y:=vec⁡(Y)y:=\operatorname{vec}(Y).

h(x):=vec⁡(g(X))h(x):=\operatorname{vec}(g(X)) and h(y):=vec⁡(g(Y))h(y):=\operatorname{vec}(g(Y)).

Assume we have 11-variable function γ(c)=f(x+c(y−x))\gamma(c)=f(x+c(y-x)), we can apply Mean Value Theorem:

where t∈t\in. Let G(c):=(h(y)−h(x))⊤h(c)G(c):=(h(y)-h(x))^{\top}h(c), we have

where the second step follows from Eq. (8), the third step follows from chain rule, the fourth step follows from Cauchy-Schwartz inequality.

Here we introduce the fact of Lipschitz for product of functions.

Let {fi(x)}i=1n\{f_{i}(x)\}_{i=1}^{n} be a sequence of function with same domain and range.

fi(x)f_{i}(x) is bounded: ∀x\forall x, ∥fi(x)∥F≤Ri\|f_{i}(x)\|_{F}\leq R_{i} with Ri≥1R_{i}\geq 1.

fi(x)f_{i}(x) is Lipschitz continuous: ∀x,y\forall x,y, ∥fi(x)−fi(y)∥F≤Li∥x−y∥F\|f_{i}(x)-f_{i}(y)\|_{F}\leq L_{i}\|x-y\|_{F}.

E.2 Lipschitz of f~​(M)~𝑓𝑀\widetilde{f}(M)

We introduce the lemma about Lipschitz of f~(M)\widetilde{f}(M).

Let f~(M)\widetilde{f}(M) be defined as Definition B.3.

where the first step follows from Fact E.1, the second step follows from Lemma D.4. ∎

E.3 Lipschitz of c​(M)𝑐𝑀c(M)

We introduce the lemma about Lipschitz of c(M)c(M).

where the first step follows from Fact E.1, the second step follows from Lemma B.12, the third step follows from Lemma D.4. ∎

E.4 Lipschitz of f~​(M)∘c​(M)~𝑓𝑀𝑐𝑀\widetilde{f}(M)\circ c(M)

We introduce the lemma about Lipschitz of f~(M)∘c(M)\widetilde{f}(M)\circ c(M).

Let f~(M)\widetilde{f}(M) be defined as Definition B.3.

where the first step follows from triangle inequality, the second step follows from Fact A.5, the third step follows from Lemma D.3, the fourth step follows from Lemma E.4 and Lemma E.3.

We introduce the lemma about Lipschitz of diag⁡((f~(M)∘c(M))⋅1n)\operatorname{diag}((\widetilde{f}(M)\circ c(M))\cdot{\bf 1}_{n}).

Let f~(M)\widetilde{f}(M) be defined as Definition B.3.

where the first step follows from Fact A.4, the second step follows from basic algebra, the third step follows from Fact A.5, and the fourth step follows from ∥1n∥2=n\|{\bf 1}_{n}\|_{2}=\sqrt{n}.

Following Eq. (E.5), Eq. (10) and Lemma E.5, we have

We introduce the lemma about Lipschitz of diag⁡((f~(M)∘c(M))⋅1n)f~(M)\operatorname{diag}((\widetilde{f}(M)\circ c(M))\cdot{\bf 1}_{n})\widetilde{f}(M).

Let f~(M)\widetilde{f}(M) be defined as Definition B.3.

where we have the upper bound in Lemma D.3, the Lipschitz of diag⁡((f~(M)∘c(M))\operatorname{diag}((\widetilde{f}(M)\circ c(M)) and f~(M)\widetilde{f}(M) in Lemma E.3 and Lemma E.6. ∎

E.7 Lipschitz of Gradient

We introduce the lemma about Lipschitz of the gradient.

We can show ∇ML(M)\nabla_{M}\mathcal{L}(M) is LL-Lipschitz.

Let f~(M)\widetilde{f}(M) be defined as Definition B.3.

where the first step follows from Theorem C.5, and the second step follows from triangle inequality. Now we proof these two terms separately.

where the first step and the second step follows from Fact A.5, the third step follows from triangle inequality, and the fourth step follows from Assumption D.1.

where the first step follows from Lemma E.5 and Lemma E.7, the second step follows from n≥1n\geq 1.

which follows from Eq. (E.7), Eq. (E.7), and Eq. (13). ∎

Appendix F Convergence of Gradient Descent

Here, we present useful fact that we use to prove our convergence result.

Proof of Part 1. Square both side of the inequality in Part 1., we have

Proof of Part 2. Square both side of the inequality in Part 2., we have

which is hold because ∣a∣>∣b∣|a|>|b| and ∣b∣≥0|b|\geq 0. ∎

F.2 Lower Bound on Frobenius Norm

We present the lemma for the lower bound on the Frobenius norm in this section.

where the first step follows from Fact A.6, the second step follows from Fact A.5, the third step follows from the upper bound of ∥B∥F\|B\|_{F} and ∥M∥F\|M\|_{F}.

Proof of Part 2. Taking the square root on both sides, we get

where the first step follows from Eq. (F.2), the second step follows from Part 2. of Fact F.1, and the third step follows from Part 1. of Fact F.1.

F.3 Sandwich Lower Bound on Frobenius Norm

Here, we introduce a sandwich trace fact.

If A⪰βIA\succeq\beta I, then tr⁡[B⊤AB]≥βtr⁡[B⊤B]\operatorname{tr}[B^{\top}AB]\geq\beta\operatorname{tr}[B^{\top}B].

As A⪰βIA\succeq\beta I, we have A−βI⪰0A-\beta I\succeq 0. Multiplying both sides by B⊤B^{\top} on the left and BB on the right (noting that these operations preserve the positive semidefiniteness), we have

Taking the trace and utilizing the property that the trace of a positive semidefinite matrix is non-negative, we have

We establish a sandwich lower bound on the Frobenius norm.

where the first step, the third step and the fifth step follows from Fact A.4, the second step and the fourth step follows from Fact F.3 and XX⊤⪰βIXX^{\top}\succeq\beta I.

Taking the square root of both side, we finish the proof. ∎

F.4 Lower Bound on Hadamard Product Between Two Matrices

We present the lemma for lower bound on Hadamard product between two matrices in this section.

The proof directly follows from the definition of the Frobenius norm. ∎

F.5 Final Bound

We introduce some useful lemmas that we use to prove the final bound.

Let f∈[δ,1]nf\in[\delta,1]^{n} and ⟨f,1n⟩=1\langle f,{\bf 1}_{n}\rangle=1.

Note that ⟨b,1n⟩=0\langle b,{\bf 1}_{n}\rangle=0 so that bb and 1n{\bf 1}_{n} are orthogonal with each other. Then, we have

where the second step is from Pythagorean theorem. ∎

We present our final bound for proving the PL inequality.

Let f~(M)∈n×n\widetilde{f}(M)\in^{n\times n} and each row summation is 1, i.e., f~(M)⋅1n=1n\widetilde{f}(M)\cdot{\bf 1}_{n}={\bf 1}_{n}.

Assume that min⁡i,j∈[n]f~(M)i,j≥δ>0\min_{i,j\in[n]}\widetilde{f}(M)_{i,j}\geq\delta>0.

F.6 PL Inequality

Here we present the bound for one unit loss function.

where the first step follows from Definition B.4 and triangle inequality, the second step follows x12+⋯+xn2≤(x1+⋯+xn)2x_{1}^{2}+\dots+x_{n}^{2}\leq(x_{1}+\dots+x_{n})^{2} when xi≥0x_{i}\geq 0 for any i∈[n]i\in[n]. ∎

We present the lemma for proving the PL inequality.

Let f~(M)\widetilde{f}(M) be defined in Definition B.3.

where the first step is by Frobenius norm definition and the second step follows from ⟨f~(M)i,1n⟩=1\langle\widetilde{f}(M)_{i},{\bf 1}_{n}\rangle=1 and c(M)i∈nc(M)_{i}\in^{n} for any i∈[n]i\in[n]. ∎

Finally, we can show the lemma for PL inequality.

Assume that min⁡i,j∈[n]f~(M)i,j≥δ>0\min_{i,j\in[n]}\widetilde{f}(M)_{i,j}\geq\delta>0.

Let L(M)\mathcal{L}(M) be defined in Definition B.7.

Let μ=2min⁡i,j∈[d]{∣Wi,j∣}⋅β⋅δ\mu=2\min_{i,j\in[d]}\{|W_{i,j}|\}\cdot\beta\cdot\delta.

Let ξ=12nmax⁡i,j∈[d]{∣Wi,j∣}⋅∥X∥F2⋅λd/μ\xi=12\sqrt{n}\max_{i,j\in[d]}\{|W_{i,j}|\}\cdot\|X\|_{F}^{2}\cdot\lambda d/\mu.

We have f~(M)⋅1n=1n\widetilde{f}(M)\cdot{\bf 1}_{n}={\bf 1}_{n} and f⋅1n=1nf\cdot{\bf 1}_{n}={\bf 1}_{n} by Definition B.3. Note that c(M)=f~(M)−fc(M)=\widetilde{f}(M)-f by Definition B.4. Thus, we have c(M)⋅1n=0nc(M)\cdot{\bf 1}_{n}={\bf 0}_{n}.

where the first and forth steps follow Lemma F.5, the second step follows from Frobenius norm property, the third step follows from triangle inequality, the fifth step follows from Lemma F.8 and Lemma F.9.

where the second step follows from Lemma F.2 and α1=6nmax⁡i,j∈[d]{∣Wi,j∣}⋅∥X∥F2⋅λd\alpha_{1}=6\sqrt{n}\max_{i,j\in[d]}\{|W_{i,j}|\}\cdot\|X\|_{F}^{2}\cdot\lambda d, the third step follows from Lemma F.5 and α2=min⁡i,j∈[d]{∣Wi,j∣}\alpha_{2}=\min_{i,j\in[d]}\{|W_{i,j}|\}, the fourth step follows from Lemma F.4 and α3=β\alpha_{3}=\beta, the fifth step follows from Lemma F.7 and α4=δ\alpha_{4}=\delta, and the last step follows from μ=2α2⋅α3⋅α4\mu=2\alpha_{2}\cdot\alpha_{3}\cdot\alpha_{4} and ξ=2α1/μ\xi=2\alpha_{1}/\mu.