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) , (2) denotes the exponential function and is applied entry-wisely, (3) operation takes a vector and outputs a diagonal matrix with the entries of that vector, and (4) denotes the length- 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 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 .
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 , our Algorithm 1 can converge to the near-optimal pruning mask for the Attention Weights Pruning problem (Definition 1.2) in time with error, where is a small term depending on intrinsic property of the data and weights.
In the above theorem, can be arbitrarily small as when the regularization coefficient . 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 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 to , 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 , where if and otherwise.
Now, we incorporate Attention Weights Pruning (see Definition 1.2) with causal attention mask .
Let be the causal attention mask defined in Definition 3.1. Let and . Let and . 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 .
Let .
Then, for any , provided where is the Lipschitz constant for (see Theorem 5.4), GD (Algorithm 1) with fixed step size and run for 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 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 . We define the pruning ratio 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 -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 -proxy, -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 -PL inequality and Lipschitz continuity of the gradient, GD converges to a point where the proxy function is within of . To apply this result to our specific problem, we need to verify these conditions for our loss function .
2 Closed Form of Gradient
As a first step, we compute the close form of the gradient . The pruning mask 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 is Lipschitz continuous with respect to .
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 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 satisfies the PL inequality. We show that satisfies the PL inequality in this lemma:
Let .
Let .
We defer the proof to Section F. By confirming that 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, and .
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 . Specifically, we use Definition 3.2 as our loss function, optimizing over the pruning mask 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 . After obtaining the optimal pruning mask, we convert to a binary pruning mask to prune , maintaining sparsity at the desired pruning ratio . We use the relative error as our evaluation metric, which is defined as
where , , , 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 directly; instead, we prune and 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 on relative error. As increases from very small values, the relative error initially decreases sharply for our algorithm, reaching a minimum before gradually rising again, which forms a shape curve. This behavior indicates that there is an optimal where our algorithm achieves its best performance around . 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 . As 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 . Wanda, while showing better performance than SparseGPT for larger sequence lengths, becomes comparable to SparseGPT as 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 for the three methods under comparison: our algorithm, Wanda, and SparseGPT. As the pruning ratio 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 , 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
(absolute homogeneity).
(triangle inequality).
(Cauchy–Schwarz inequality).
.
For any , , we have .
where is the rank of .
.
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 -th row of as follows
Let be the causal attention mask defined in Definition 3.1.
Let , be defined as Definition B.1.
We define -th entry of as follows
Then, we introduce the Softmax probability function.
Let be the causal attention mask defined in Definition 3.1.
Let , be defined as Definition B.1.
Let , be defined as Definition B.2.
We define -th row of as follows
We define the entry in -th row, -th column of as follows
Then, we introduce the one unit loss function.
Let , be defined in Definition B.3.
We define -th row of as follows
We define -th column of as follows
We define the entry in -th row, -th column of 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 .
Let , , , 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 , , the third step follows from Fact A.7, and the fourth step follows from for any matrix , .
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 .
B.3 Gradient for Each Row of u~(M)~𝑢𝑀\widetilde{u}(M)
Let be defined in Definition B.1.
Let , , , 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 .
Let be defined in Definition B.1.
Let be defined in Definition B.2.
Let , , , 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 .
Let be defined in Definition B.1.
Let be defined in Definition B.2.
Let be defined in Definition B.3.
Let , , , , 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 .
Let , be defined in Definition B.3.
Let , , , , 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 be defined in Definition B.3.
Let , , , , we have
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 , , 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 .
Let be defined in Definition B.1.
Let be defined in Definition B.2.
Let be defined in Definition B.3.
Let be defined in Definition B.7.
Let , , , , we have
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 to simplify the calculation.
Let be defined in Definition B.3.
We define the -th column of as follows
We define the -th row of as follows
We introduce the matrix view of and its summation.
Let , which is defined in Lemma B.15
where the first step follows from the definition of , 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 -th column of
where the second step follows from Fact A.2.
Following from Fact A.1, we can get
where the second step follows from Fact A.4.
Part 2. We further compute the summation of .
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 and its summation.
Let be defined in Lemma B.15.
where the first step follows from the definition of , the second step, the third step and the fourth step follows from Fact A.4.
Following from Fact A.1, we can get -th column of
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 .
Part 2. We further compute the summation of
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 .
Let be defined in Lemma B.15.
The proof is straightforward. By the definition of , for all , the -th entry of is given by . Thus, the entire matrix has entries that correspond to those of . Therefore, we can conclude that as required. ∎
We introduce the matrix form of overall loss function.
Let 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 be some fixed constant satisfies .
Here we present the lemma of bounds for and .
Let and be the causal attention mask defined in Definition 3.1. For , we have
This Lemma simply follows from the definition of Frobenius norm, given that the max value of each entry in and is . ∎
D.2 Bounds for Basic Functions
We first introduce the lemma of bounds for basic function.
Under Assumption D.1, for all , , , , we have the following bounds
Proof of Part 1. Each entry in present a probability, thus for , , we have
For any -th row of , following from the definition of Softmax function, we know
which follows from . 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 , 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 .
Let 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 and as follows:
and .
and .
Assume we have -variable function , we can apply Mean Value Theorem:
where . Let , 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 be a sequence of function with same domain and range.
is bounded: , with .
is Lipschitz continuous: , .
E.2 Lipschitz of f~(M)~𝑓𝑀\widetilde{f}(M)
We introduce the lemma about Lipschitz of .
Let 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 .
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 .
Let 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 .
Let 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 .
Following Eq. (E.5), Eq. (10) and Lemma E.5, we have
We introduce the lemma about Lipschitz of .
Let be defined as Definition B.3.
where we have the upper bound in Lemma D.3, the Lipschitz of and 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 is -Lipschitz.
Let 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 .
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 and . ∎
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 and .
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 , then .
As , we have . Multiplying both sides by on the left and 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 .
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 and .
Note that so that and 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 and each row summation is 1, i.e., .
Assume that .
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 when for any . ∎
We present the lemma for proving the PL inequality.
Let be defined in Definition B.3.
where the first step is by Frobenius norm definition and the second step follows from and for any . ∎
Finally, we can show the lemma for PL inequality.
Assume that .
Let be defined in Definition B.7.
Let .
Let .
We have and by Definition B.3. Note that by Definition B.4. Thus, we have .
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 , the third step follows from Lemma F.5 and , the fourth step follows from Lemma F.4 and , the fifth step follows from Lemma F.7 and , and the last step follows from and .