Structured Pruning Method for Large Language Models and Related Devices
By calculating the fluctuation metric matrix in a large language model and generating a mask matrix for pruning, the problem that traditional methods are not suitable for special large language models is solved, and the accuracy and adaptability of model compression are improved.
Patent Information
- Application Number
- CN202510309293.X
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2025-03-17
- Publication Date
- 2025-06-10
- Estimated Expiration
- 2045-03-17
AI Technical Summary
The traditional model compression method is not suitable for large language models with special structures such as LLaMA-3, resulting in poor model compression accuracy.
A structured pruning method for large language models is proposed. By obtaining the weight matrix of the attention module and perception module of the initial large language model, the fluctuation metric matrix is calculated to determine the global pruning threshold, and the mask matrix of each module is generated based on the threshold for pruning.
The model compression accuracy of large language models is improved, and the output dimension differences of different network layers are adapted to ensure the flexibility and accuracy of pruning operations.
Smart Images

Figure CN119849578B_ABST
Abstract
Description
Technical Field
[0001] This application relates to the technical field of neural network lightweighting, and particularly relates to a structured pruning method for large language models and related devices. Background Art
[0002] A large language model (LLM) refers to a deep learning model trained using a large amount of text data. To reduce the computational resource requirements and storage space of the model without significantly degrading its performance, it is necessary to perform model compression on the large language model.
[0003] In related technologies, model compression is achieved by uniformly pruning all network layers of the large language model. However, for some large language models with special structures such as LLaMA-3, the dimensions between its various network layers are different. Therefore, traditional model compression methods are no longer applicable, resulting in poor model compression accuracy for large language models. Summary of the Invention
[0004] The main purpose of the embodiments of this application is to propose a structured pruning method for large language models and related devices, aiming to improve the model compression accuracy of large language models.
[0005] To achieve the above objective, the first aspect of the embodiments of this application proposes a structured pruning method for large language models, including:
[0006] Obtain an initial large language model to be pruned, and obtain corresponding first weight matrices from multiple attention modules of the initial large language model and corresponding second weight matrices from multiple perception modules of the initial large language model;
[0007] Determine a first fluctuation metric matrix based on the first weight matrix, determine a second fluctuation metric matrix based on the second weight matrix, and determine a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix;
[0008] Based on the global pruning threshold and the first fluctuation metric matrix, determine corresponding key-value mask matrices and query mask matrices for each attention module, and based on the global pruning threshold and the second fluctuation metric matrix, determine corresponding perception mask matrices for each perception module;
[0009] Use the key-value mask matrix and query mask matrix of each attention module to prune the corresponding first weight matrix to obtain a first target matrix, and use the perception mask matrix of each perception module to prune the corresponding second weight matrix to obtain a second target matrix;
[0010] Determine the large language model after pruning the initial large language model based on the first target matrix and the second target matrix.
[0011] In the embodiments of the present application, each attention module includes a plurality of first weight matrices, and each first weight matrix includes an initial key matrix, an initial query matrix, an initial value matrix, and a first output matrix. Moreover, the output dimensions among the initial key matrix, the initial query matrix, the initial value matrix, and the first output matrix are not exactly the same;
[0012] Determining the first fluctuation metric matrix based on the first weight matrix includes:
[0013] Obtain an initial sample and input the initial sample into the initial large language model;
[0014] For each first weight matrix, sequentially perform embedding processing on the initial sample based on the initial key matrix, the initial query matrix, the initial value matrix, and the first output matrix to obtain corresponding first output samples;
[0015] Determine a first original matrix based on the first output samples, and perform mean processing on multiple first original matrices belonging to the same attention module to obtain a first mean matrix;
[0016] Stack the first mean matrices corresponding to all attention modules to obtain the first fluctuation metric matrix.
[0017] In the embodiments of the present application, performing mean processing on multiple first original matrices belonging to the same attention module to obtain a first mean matrix includes:
[0018] Perform mean processing on the matrix elements in each first original matrix respectively to obtain the metric means corresponding to each first original matrix;
[0019] Calculate the metric standard deviation values corresponding to each first original matrix based on the matrix elements in each first original matrix;
[0020] Update the matrix elements in the corresponding first original matrix based on the metric means and the metric standard deviation values to obtain a first normalization matrix;
[0021] Perform mean processing on multiple first normalization matrices belonging to the same attention module to obtain a first mean matrix.
[0022] In the embodiments of the present application, each attention module is connected to a corresponding perception module, and each perception module includes a plurality of second weight matrices, and each second weight matrix includes a mapping matrix and a second output matrix;
[0023] Determining the second fluctuation metric matrix based on the second weight matrix includes:
[0024] For each first weight matrix, based on the mapping matrix and the second output matrix in sequence, perform embedding processing on the first output sample to obtain the corresponding second output sample;
[0025] Determine the second original matrix based on the second output sample, and perform mean processing on multiple second original matrices belonging to the same perception module to obtain the second mean matrix;
[0026] Stack the second mean matrices corresponding to all perception modules to obtain the second fluctuation metric matrix.
[0027] In the embodiments of the present application, determining the global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix includes:
[0028] Combine the first fluctuation metric matrix and the second fluctuation metric matrix to obtain the global metric matrix;
[0029] Construct an initial matrix with the same matrix size as the global metric matrix, where the initial matrix includes multiple initial elements;
[0030] Based on the preset pruning rate, preset cumulative function, and global metric matrix, update multiple initial elements to obtain the corresponding multiple updated elements;
[0031] Determine the target updated element with the smallest element value from multiple updated elements, and determine the first target position where the target updated element is located;
[0032] Based on the second target position matching the first target position of the global metric matrix, determine the global pruning threshold.
[0033] In the embodiments of the present application, based on the global pruning threshold and the first fluctuation metric matrix, determining the key-value mask matrix and query mask matrix corresponding to each attention module includes:
[0034] Based on the grouping parameters preset by the initial large language model, adjust the matrix shape of each first mean matrix to obtain the initial mask matrix;
[0035] According to the global pruning threshold, update the matrix elements of the initial mask matrix to obtain the initial key-value mask matrix;
[0036] According to the preset local pruning threshold, perform masking processing on the matrix elements of the initial key-value mask matrix to obtain the key-value mask matrices corresponding to each attention module;
[0037] Based on the grouping parameters, determine the number of attention heads pruned for any attention module, and update the global pruning threshold based on the number of attention heads to obtain the updated global pruning threshold;
[0038] Based on the updated global pruning threshold, mask the matrix elements of the initial mask matrix to obtain the corresponding query mask matrix for each attention module.
[0039] In the embodiments of the present application, the local pruning threshold includes a first local pruning threshold and a second local pruning threshold;
[0040] According to the preset local pruning threshold, mask the matrix elements of the initial key-value mask matrix to obtain the corresponding key-value mask matrix for each attention module, including:
[0041] Sum all the matrix elements of the initial key-value mask matrix to obtain the actual total value;
[0042] If the actual total value does not meet the preset total value, use the first local pruning threshold to mask the initial key-value mask matrix to obtain the corresponding key-value mask matrix for each attention module;
[0043] If the actual total value meets the preset total value, use the second local pruning threshold to mask the initial key-value mask matrix to obtain the corresponding key-value mask matrix for each attention module.
[0044] In the embodiments of the present application, updating the global pruning threshold based on the number of attention heads to obtain the updated global pruning threshold includes:
[0045] Sort the matrix elements in the first mean matrix to obtain the updated first mean matrix;
[0046] Determine that the matrix element at the third target position matching the number of attention heads in the updated first mean matrix is the updated global pruning threshold.
[0047] In the embodiments of the present application, after determining the large language model after pruning the initial large language model, it further includes:
[0048] Obtain the evaluation samples and the verification samples representing the expected output;
[0049] Input the evaluation samples into the large language model, and perform embedding processing on the evaluation samples based on the first target matrix and the second target matrix to obtain the actual samples representing the actual output;
[0050] Calculate the sample difference between the actual samples and the verification samples, and perform pruning adjustment on the large language model based on the sample difference to obtain the large language model after pruning adjustment.
[0051] To achieve the above object, the second aspect of the embodiments of the present application proposes a structured pruning device for a large language model, including:
[0052] An acquisition module, configured to acquire an initial large language model to be pruned, and acquire corresponding first weight matrices from multiple attention modules of the initial large language model and acquire corresponding second weight matrices from multiple perception modules of the initial large language model;
[0053] A fluctuation metric matrix determination module, configured to determine a first fluctuation metric matrix based on the first weight matrix, determine a second fluctuation metric matrix based on the second weight matrix, and determine a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix;
[0054] A global pruning threshold determination module, configured to determine corresponding key-value mask matrices and query mask matrices for each attention module based on the global pruning threshold and the first fluctuation metric matrix, and determine corresponding perception mask matrices for each perception module based on the global pruning threshold and the second fluctuation metric matrix;
[0055] A mask matrix determination module, configured to use the key-value mask matrix and the query mask matrix of each attention module to perform pruning processing on the corresponding first weight matrix to obtain a first target matrix, and use the perception mask matrix of each perception module to perform pruning processing on the corresponding second weight matrix to obtain a second target matrix;
[0056] A target module, configured to determine a large language model after pruning the initial large language model based on the first target matrix and the second target matrix.
[0057] To achieve the above object, a third aspect of the embodiments of the present application proposes an electronic device, which includes a memory and a processor. The memory stores a computer program, and when the processor executes the computer program, it implements the structured pruning method for the large language model in the first aspect above.
[0058] To achieve the above object, a fourth aspect of the embodiments of the present application proposes a computer-readable storage medium, which stores a computer program, and when the computer program is executed by a processor, it implements the structured pruning method for the large language model in the first aspect above.
[0059] The structured pruning method for large language models and related devices proposed in this application, the method includes: obtaining an initial large language model to be pruned, and obtaining corresponding first weight matrices from multiple attention modules of the initial large language model, and obtaining corresponding second weight matrices from multiple perception modules of the initial large language model; determining a first fluctuation metric matrix based on the first weight matrix, determining a second fluctuation metric matrix based on the second weight matrix, and determining a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix; based on the global pruning threshold and the first fluctuation metric matrix, determining corresponding key-value mask matrices and query mask matrices for each attention module, and based on the global pruning threshold and the second fluctuation metric matrix, determining corresponding perception mask matrices for each perception module; the mask matrices of different output dimensions can ensure that the pruning operation adapts to all processing layers of the initial large language model to flexibly meet the requirements of different output dimensions between layers in each attention module of the initial large language model; then, using the key-value mask matrix and query mask matrix of each attention module, pruning the corresponding first weight matrix to obtain a first target matrix, and using the perception mask matrix of each perception module, pruning the corresponding second weight matrix to obtain a second target matrix; based on the first target matrix and the second target matrix, determining the large language model after pruning the initial large language model. In this way, generating corresponding mask matrices according to the actual situations of each attention module and each perception module in each processing layer can improve the model compression accuracy of the large language model while adapting to the special structure of the initial large language model. Description of the Drawings
[0060] Figure 1 It is a schematic diagram of an optional implementation environment of the structured pruning device for large language models provided by an embodiment of this application;
[0061] Figure 2 It is an optional flowchart of the structured pruning method for large language models provided by an embodiment of this application;
[0062] Figure 3 It is another optional flowchart of the structured pruning method for large language models provided by an embodiment of this application;
[0063] Figure 4 It is yet another optional flowchart of the structured pruning method for large language models provided by an embodiment of this application;
[0064] Figure 5 It is an optional module schematic diagram of the structured pruning device for large language models provided by an embodiment of this application;
[0065] Figure 6 It is a schematic diagram of the hardware structure of the electronic device provided by an embodiment of this application. Detailed implementation manners
[0066] In order to make the objectives, technical solutions and advantages of the present application clearer and more understandable, the present application will be further described in detail below with reference to the accompanying drawings and embodiments. It should be understood that the specific embodiments described herein are only used to explain the present application and are not used to limit the present application.
[0067] It should be noted that although the functional modules are divided in the device schematic diagram and the logical sequence is shown in the flowchart, in some cases, the steps shown or described may be executed in a different module division in the device or a different sequence in the flowchart. The terms "first", "second", etc. in the description, claims and the above-mentioned drawings are used to distinguish similar objects and do not necessarily need to describe a specific order or sequence.
[0068] Unless otherwise defined, all technical and scientific terms used herein have the same meaning as commonly understood by those skilled in the technical field to which this application belongs. The terms used herein are only for the purpose of describing the embodiments of this application and are not intended to limit this application.
[0069] First, several nouns involved in the present application are analyzed:
[0070] Deep Learning is a machine learning method based on neural networks. It simulates the working mode of human brain neurons by constructing a multi-layer neural network model, enabling the computer to autonomously learn and extract high-level features from data. The core idea of deep learning is to gradually abstract data into higher-level feature representations through layer-by-layer non-linear transformations, thereby showing excellent performance in complex tasks.
[0071] Natural Language Processing (NLP): NLP uses a computer to process, understand, and apply human languages (such as Chinese, English, etc.). NLP belongs to a branch of artificial intelligence and is an interdisciplinary field of computer science and linguistics, often referred to as computational linguistics. Natural language processing includes syntactic analysis, semantic analysis, discourse understanding, etc. Natural language processing is commonly used in technical fields such as machine translation, handwritten and printed character recognition, speech recognition and text-to-speech conversion, information intention recognition, information extraction and filtering, text classification and clustering, public opinion analysis, and opinion mining. It involves data mining, machine learning, knowledge acquisition, knowledge engineering, artificial intelligence research related to language processing, and linguistic research related to language computing.
[0072] Artificial Intelligence (AI): It is a new technical science that studies and develops theories, methods, technologies, and application systems for simulating, extending, and expanding human intelligence; Artificial Intelligence is a branch of computer science. It attempts to understand the essence of intelligence and produce a new intelligent machine that can react in a way similar to human intelligence. The research in this field includes robots, speech recognition, image recognition, natural language processing, and expert systems, etc. Artificial Intelligence can simulate the information process of human consciousness and thinking. Artificial Intelligence also refers to the theory, method, technology, and application system that uses a digital computer or a machine controlled by a digital computer to simulate, extend, and expand human intelligence, perceive the environment, acquire knowledge, and use knowledge to obtain the best results.
[0073] Large language models refer to deep learning models trained with a large amount of text data, enabling the model to generate natural language text or understand the meaning of language text. Large language models can provide in-depth knowledge and language production on various topics by training on a vast dataset, learn the patterns and structures of natural language through large-scale unsupervised training, and to a certain extent simulate the human language cognition and generation process.
[0074] Among them, large language models have a huge number of parameters, often starting from billions. Therefore, their requirements for computing power and storage overhead are huge. In practical applications, in order to reduce the computing resource requirements and storage space of the model without significantly reducing the model performance, it is usually necessary to perform model compression processing on large language models.
[0075] Among them, structured pruning is a model compression technology aimed at reducing the number of model parameters and computational complexity by removing some specific structures in the neural network. For traditional large language models such as LLaMA and LLaMA2, the output dimensions of the query layer, key-value layer, and output layer of the attention module are usually the same. Therefore, in related technologies, unified pruning processing is performed on all network layers of the large language model by using the same mask matrix to achieve model compression, where the network layer includes the query layer, key-value layer, and output layer. However, for some special large language models different from the traditional model structure, such as LLaMA-3 (LLaMA3), the output dimension of the key layer / value layer is different from that of the query layer. Therefore, the traditional model compression method is no longer applicable, resulting in poor model compression accuracy of large language models.
[0076] Based on this, the embodiments of this application provide a structured pruning method and related devices for large language models, aiming to improve the model compression accuracy of large language models.
[0077] Exemplarily, as Figure 1 shown, Figure 1FIG. 0 is a schematic diagram of an optional implementation environment of the structured pruning device for large language models provided by an embodiment of the present application. The implementation environment includes a terminal 11 and a server side 12. At least one terminal 11 and the server side 12 are connected through a communication network.
[0078] Further, the server side 12 deploys the structured pruning device for large language models proposed by the embodiment of the present application (for the convenience of description, hereinafter it may also be simply referred to as the "structured pruning device"). After the server side 12 obtains the initial large language model to be pruned, and obtains the corresponding first weight matrix from multiple attention modules of the initial large language model and the corresponding second weight matrix from multiple perception modules of the initial large language model; it determines the first fluctuation metric matrix based on the first weight matrix, determines the second fluctuation metric matrix based on the second weight matrix, and determines the global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix; then, based on the global pruning threshold and the first fluctuation metric matrix, it determines the corresponding key-value mask matrix and query mask matrix for each attention module, and based on the global pruning threshold and the second fluctuation metric matrix, it determines the corresponding perception mask matrix for each perception module; then, using the key-value mask matrix and query mask matrix of each attention module, it performs pruning processing on the corresponding first weight matrix to obtain the first target matrix, and using the perception mask matrix of each perception module, it performs pruning processing on the corresponding second weight matrix to obtain the second target matrix; finally, based on the first target matrix and the second target matrix, it determines the large language model obtained after pruning the initial large language model.
[0079] It can be understood that the pruned large language model requires less computational effort and storage space. Therefore, after the server side 12 determines the pruned large language model, it can send the large language model to the terminal 11 so that the pruned large language model can run on the terminal 11. Among them, the initial large language model obtained by the server side 12 can be sent by the terminal 11 that finally receives the large language model, or can be sent by other terminals or the server side.
[0080] In addition, the above structured pruning method can also be performed only on the terminal 11 side.
[0081] Among them, the server 12 can be an independent physical server, a server cluster or a distributed system composed of multiple physical servers, or a cloud server that provides basic cloud computing services such as cloud services, cloud databases, cloud computing, cloud functions, cloud storage, network services, cloud communications, middleware services, domain name services, security services, content delivery network (CDN), and big data and artificial intelligence platforms. Additionally, the server 12 can also be a node server in a blockchain network. The terminal 11 can be a mobile phone, a computer, a smart voice interaction device, a smart wearable device, a smart home appliance, a vehicle-mounted terminal, etc., but is not limited thereto. The terminal 11 and the server 12 can be directly or indirectly connected through wired or wireless communication means, and the embodiments of the present application do not limit this here.
[0082] It should be noted that in the embodiments of the present application, when it comes to information related to user characteristics such as user basic information or user identity, the user's permission or consent will be obtained first. Moreover, the collection, use, and processing of these data will comply with relevant laws, regulations, and standards. In addition, when the embodiments of the present application need to obtain the user's sensitive personal information, the user's separate permission or separate consent will be obtained first. After clearly obtaining the user's separate permission or separate consent, the necessary data for the normal operation of the embodiments of the present application will be obtained. For example, when the embodiments of the present application obtain the initial large language model to be pruned, the authorization or consent of relevant personnel will be obtained first, otherwise the initial large language model and other relevant data that cannot be applied to the embodiments of the present application will be obtained. Additionally, other relevant data obtained by the structured pruning device of the present application are all authorized data, which will not be elaborated here one by one.
[0083] In the embodiments of the present application, the description will be made from the dimension of the structured pruning device, and this structured pruning device can be integrated in a computer device, such as a server. As Figure 2 shown, Figure 2 is an optional flowchart of the structured pruning method for the large language model provided by the embodiments of the present application. Figure 2 The flowchart shown can include but is not limited to the following steps 101 to step 105. When the structured pruning device executes the structured pruning method for the large language model (for the sake of convenience of description, it can also be simply referred to as the "structured pruning method" hereinafter), the specific process is as follows. It should be noted first that the present embodiment does not make a specific limitation on the Figure 2 sequence of steps 101 to 105, and the order of steps can be adjusted according to actual needs, or some steps can be reduced or added.
[0084] Step 101: Obtain the initial large language model to be pruned, and obtain the corresponding first weight matrix from multiple attention modules of the initial large language model and the corresponding second weight matrix from multiple perception modules of the initial large language model.
[0085] The following provides a detailed description of Step 101.
[0086] Among them, the large language model includes multiple processing layers. To overcome the limitations of traditional models in dealing with long-distance dependencies and improve the performance of the large language model in various natural language processing tasks, many large language models have introduced the multi-head self-attention mechanism. In such a case, each processing layer of the large language model includes at least one set of interconnected attention modules and perception modules; each attention module includes multiple sets of interconnected key layers (Key, K), query layers (Query, Q), value layers (Value, V), and a first output layer. The multi-head self-attention mechanism dynamically calculates the correlations between different positions in the input sequence through the cooperation between the key layer, query layer, value layer, and the first output layer, enabling the large language model to capture complex dependencies and improve its performance in various natural language processing tasks.
[0087] Among them, to achieve different levels of processing of the input data, the key layer, query layer, value layer, and the first output layer are each provided with corresponding weight parameter matrices. In the embodiments of the present application, the initial large language model refers to the original large-scale deep learning model without any compression or optimization processing. Moreover, the key layer of the initial large language model is provided with an initial key matrix, the query layer is provided with an initial query matrix, the value layer is provided with an initial value matrix, and the first output layer is provided with a first output matrix. Each set of interconnected initial value matrix, initial query matrix, initial value matrix, and first output matrix is called the first weight matrix; the output dimensions between the initial key matrix, initial query matrix, initial value matrix, and the first output matrix are not exactly the same. In addition, each perception module includes multiple sets of interconnected mapping layers and a second output layer. The mapping layer is provided with a mapping matrix, and the second output layer is provided with a second output matrix.
[0088] Among them, the attention module is used to capture the dependencies between different positions in the input data and dynamically focus on different parts of the input data; the perception module is used to further analyze and process the output vector of the attention module to capture more complex data relationships, thereby enhancing the expressive ability of the large language model.
[0089] Exemplarily, the initial large language model obtained in the embodiments of the present application is LLaMA3, which is a large language model released by Meta on April 2024 local time. Among them, the output dimensions of the key layer (k_proj) / value layer (v_proj) of LLaMA3 are different from those of the query layer (q_proj). It should be noted that the special structure of LLaMA3 makes traditional structured pruning methods inapplicable to it.
[0090] The embodiments of the present application will take LLaMA3 as an example of the initial large language model to explain the structured pruning method proposed in the embodiments of the present application, and the beneficial effects of the embodiments of the present application will also be gradually revealed in the subsequent description.
[0091] Step 102: Determine a first fluctuation metric matrix based on the first weight matrix, determine a second fluctuation metric matrix based on the second weight matrix, and determine a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix.
[0092] The following provides a detailed description of step 102.
[0093] Among them, the first fluctuation metric matrix is a metric matrix calculated based on multiple first weight matrices of each attention module, and the first fluctuation metric matrix is used to evaluate the influence degree of each first weight matrix on the performance of the initial large language model; the second fluctuation metric matrix is a metric matrix calculated based on multiple second weight matrices of each perception module, and the second fluctuation metric matrix is used to evaluate the influence degree of each second weight matrix on the performance of the initial large language model.
[0094] Among them, the global pruning threshold is used to generate the corresponding key-value mask matrix and query mask matrix for each attention module, and the corresponding perception mask matrix for each perception module in the subsequent process.
[0095] In the embodiments of the present application, determining the first fluctuation metric matrix based on the first weight matrix includes the following steps:
[0096] (102.a.1) Obtain an initial sample and input the initial sample into the initial large language model.
[0097] (102.a.2) For each of the first weight matrices, perform embedding processing on the initial sample based on the initial key matrix, initial query matrix, initial value matrix, and first output matrix in sequence to obtain the corresponding first output sample.
[0098] (102.a.3) Determine a first original matrix based on the first output sample, and perform mean processing on multiple first original matrices belonging to the same attention module to obtain a first mean matrix.
[0099] (102.a.4) Stack all the corresponding first mean matrices of the attention modules to obtain a first fluctuation metric matrix.
[0100] The following provides a detailed description of steps (102.a.1) to (102.a.4).
[0101] Among them, the initial samples refer to the original data input into the large language model, and these data can be text sequences, images, or other types of data, depending on the actual application scenarios of LLaMA3. For example, in natural language processing tasks, the initial samples are usually one or more pieces of text, such as sentences or paragraphs. Additionally, the initial samples can be obtained from open-source databases or input in real-time by relevant personnel. The embodiments of this application do not limit the acquisition sources and quantities of the initial samples.
[0102] Furthermore, the attention modules on any processing layer of LLaMA3 usually include multiple attention heads. Each attention head is used to independently calculate attention scores, and each attention head corresponds to a first weight matrix. Then, for each attention head, based on the corresponding initial key matrix, initial query matrix, initial value matrix, and first output matrix in sequence, the initial samples are embedded to obtain corresponding first output samples. Among them, the attention head is the basic unit of the multi-head attention mechanism in the initial large language model architecture.
[0103] Furthermore, through the following formula <1>, the corresponding fluctuation metric (first raw matrix) of any attention head in each attention module is calculated:
[0104] <1>
[0105] Among them, N is the number of input initial samples; is the layer number, representing the th processing layer; represents the vector value (first output sample) of the nth sample after being processed by the jth attention head in the th layer; represents the current average value of the jth column in the th layer; represents the historical average value of the jth column of the input in the th layer; represents the square of the 2-norm of the jth column of the first weight matrix of the input in the th layer.
[0106] Among them, the 2-norm is also called the Euclidean norm or L2 norm, which is a measurement method for vectors or matrices and is used to measure the "length" of vectors or the "size" of matrices.
[0107] In the embodiment of the present application, mean processing is performed on multiple first original matrices belonging to the same attention module to obtain a first mean matrix, including the following steps:
[0108] (A.1) Perform mean processing on the matrix elements in each first original matrix to obtain the corresponding metric mean of each first original matrix.
[0109] (A.2) Based on the matrix elements in each first original matrix, calculate the corresponding metric standard deviation of each first original matrix.
[0110] (A.3) Based on the metric mean and the metric standard deviation, update the matrix elements in the corresponding first original matrix to obtain a first normalized matrix.
[0111] (A.4) Perform mean processing on multiple first normalized matrices belonging to the same attention module to obtain a first mean matrix.
[0112] The following provides a detailed description of steps (A.1) to (A.4).
[0113] In the embodiment of the present application, to make different features comparable, improve the model performance, and accelerate the training process, the present application embodiment also performs normalization processing on the first original matrix. Specifically, the first normalized matrix after the normalization processing is calculated through the following formula <2> :
[0114]
[0115] <2>
[0116] where, represents the i-th first original matrix; represents the i-th metric mean, and the metric mean is obtained by performing mean processing on the matrix elements in the i-th first original matrix; std is the standard deviation function, is the metric standard deviation.
[0117] Further, mean processing is performed on multiple first normalized matrices belonging to the same attention module to obtain a first mean matrix; then, stack the first mean matrices corresponding to all attention modules to obtain a first fluctuation metric matrix.
[0118] In the embodiment of the present application, determining a second fluctuation metric matrix based on a second weight matrix includes the following steps:
[0119] (102.b.1) For each first weight matrix, sequentially perform embedding processing on the first output sample based on the mapping matrix and the second output matrix to obtain the corresponding second output sample.
[0120] (102.b.2) Determine the second original matrix based on the second output samples, and perform a mean processing on multiple second original matrices belonging to the same perception module to obtain a second mean matrix.
[0121] (102.b.3) Stack the corresponding second mean matrices of all perception modules to obtain a second fluctuation metric matrix.
[0122] The following gives a detailed description of steps (102.b.1) to (102.b.3).
[0123] In the embodiments of the present application, each attention module is connected to a corresponding perception module. Each perception module includes multiple second weight matrices, and each second weight matrix includes a mapping matrix and a second output matrix. Among them, the mapping matrix of LLaMA3 includes a control gate matrix (gate_proj) and an ascending matrix (up_proj), and the second output matrix is a descending matrix (down_proj).
[0124] Among them, the corresponding fluctuation metric (second original matrix) of each perception module is calculated through the following formula <3>:
[0125] <3>
[0126] where N is the first output sample from the attention module; is the layer number, representing the th processing layer; represents the nth first output sample, which is the vector value (second output sample) output after being processed by the jth column mapping matrix or the second output matrix of the th layer; represents the current average value of the jth column of the th layer; represents the historical average value of the jth column of the input of the th layer; represents the square of the 2-norm of the first weight matrix of the jth column of the input of the th layer.
[0127] Further, formula <2> calculates the square of the product of the input data sample variance and the 2-norm of the weight matrix, and formula <3> calculates the product of the input data sample variance and the 2-norm of the weight matrix. The difference between formula <3> and formula <2> is that no square processing is performed on the calculation result, so that the attention module can control the value range through square calculation and make the outliers more prominent. Outliers refer to the data points in the data set that are significantly different from other data points.
[0128] Further, perform a mean processing on multiple second original matrices belonging to the same perception module to obtain a second mean matrix; then, stack the corresponding second mean matrices of all perception modules to obtain a second fluctuation metric matrix.
[0129] Alternatively, after obtaining the second original matrix, perform a normalization process on the second original matrix by a method similar to steps (A.1) to (A.4) to obtain a second normalized matrix; then, perform a mean processing on multiple second normalized matrices belonging to the same perception module to obtain a second mean matrix; then, stack the corresponding second mean matrices of all perception modules to obtain a second fluctuation metric matrix.
[0130] In this way, accurately evaluate the importance of each attention module and perception module in the initial large language model through the first fluctuation metric matrix and the second fluctuation metric matrix, identify the parts that contribute less to the performance of the initial large language model from a global perspective, facilitate setting a global pruning threshold for the initial large language model based on the first fluctuation metric matrix and the second fluctuation metric matrix, and then implement subsequent model pruning processing.
[0131] In the embodiments of the present application, determining the global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix includes the following steps:
[0132] (102.c.1) Combine the first fluctuation metric matrix and the second fluctuation metric matrix to obtain a global metric matrix.
[0133] (102.c.2) Construct an initial matrix with the same matrix size as the global metric matrix, and the initial matrix includes multiple initial elements.
[0134] (102.c.3) Update multiple initial elements based on a preset pruning rate, a preset cumulative function, and the global metric matrix to obtain corresponding multiple updated elements.
[0135] (102.c.4) Determine a target updated element with the smallest element value from the multiple updated elements, and determine the first target position where the target updated element is located.
[0136] (102.c.5) Determine the global pruning threshold based on the second target position matching the first target position in the global metric matrix.
[0137] The following will describe steps (102.c.1) to (102.c.5) in detail.
[0138] Among them, the global metric matrix can be obtained by first flattening and concatenating the matrix elements in the first fluctuation metric matrix and the second fluctuation metric matrix. Matrix elements refer to the specific numerical values at each position in the matrix. Of course, the way to merge the first fluctuation metric matrix and the second fluctuation metric matrix can be specifically set according to the actual situation, and the embodiments of this application do not limit this.
[0139] Further, to determine which weight parameters of the initial large language model should be retained or pruned, it is necessary to construct an intermediate variable, that is, an initial matrix, to help determine the global pruning threshold subsequently. Among them, the initial matrix is usually initialized to zero or other constant values (initial elements), indicating that all weight parameters have not been evaluated or determined whether to be pruned. It should be noted that the essence of the initial element is also a matrix element, which is called the initial element here for the convenience of distinction.
[0140] Further, sort the matrix elements in the global metric matrix from smallest to largest to obtain the second target element in the sorted global metric matrix and the corresponding second target position of the second target element. Then, set the positions in the initial matrix that are less than the number of elements in the first fluctuation metric matrix to true, and the others to false, and set the initial element value at the true position to 512.0 / 3 to obtain the first fluctuation metric matrix with updated matrix elements.
[0141] Further, for the updated first fluctuation metric matrix, update multiple initial elements through the following formulas <4> to <5> and obtain the corresponding multiple updated elements cw:
[0142] <4>
[0143] <5>
[0144] Among them, pruning_ratio is the preset pruning rate; cumsum is the preset cumulative function; sum is the summation function; compression_weight is the initial element.
[0145] Further, through the following formula <6>, determine the target updated element with the smallest element value from multiple updated elements and determine the first target position pos where the target updated element is located:
[0146] <6>
[0147] Among them, argmin is the minimum value function, which is used to return the index of the smallest element in the array.
[0148] Further, through the following formula <7>, determine the global pruning threshold threshold:
[0149] <7>
[0150] Among them, sorted_prune is the matrix element in the updated first fluctuation metric matrix. That is, the first target position pos where the minimum target update element is determined from multiple updated elements of the initial matrix is determined, and the matrix element at the corresponding second target position of the global metric matrix is used as the global pruning threshold.
[0151] Among them, the global pruning threshold is used to unify the pruning standard, which can lay a foundation for separately determining the mask matrices of the attention module and the perception module subsequently, and avoid local optimal solutions or inconsistent pruning decisions when performing pruning processing based on the mask matrices subsequently.
[0152] Step 103: Based on the global pruning threshold and the first fluctuation metric matrix, determine the corresponding key-value mask matrix and query mask matrix of each attention module, and based on the global pruning threshold and the second fluctuation metric matrix, determine the corresponding perception mask matrix of each perception module.
[0153] The following provides a detailed description of step 103.
[0154] Among them, since the output dimensions of the matrices corresponding to each layer in the attention module of the initial large language model are not completely the same, it is impossible to directly perform pruning processing on the matrices of the attention module and the perception module according to the global pruning threshold. Based on this, in the embodiment of the present application, after obtaining the global pruning threshold, the mask matrices of each attention module and each perception module are determined respectively based on the global pruning threshold, and each attention module includes two mask matrices with different output dimensions.
[0155] Furthermore, the mask matrices with different output dimensions can ensure that the pruning operation adapts to all processing layers of the initial large language model to flexibly meet the requirements of different output dimensions among the layers in each attention module of the initial large language model; generating corresponding mask matrices according to the actual situations of each attention module and each perception module in each processing layer can improve the pruning accuracy of the large language model under the condition of adapting to the special structure of the initial large language model, and further improve the model compression accuracy of the large language model.
[0156] In the embodiment of the present application, determining the corresponding key-value mask matrix and query mask matrix of each attention module based on the global pruning threshold and the first fluctuation metric matrix includes the following steps:
[0157] (103.a.1) Adjust the matrix shape of each first mean matrix based on the preset grouping parameters of the initial large language model to obtain the initial mask matrix.
[0158] (103.a.2)Update the matrix elements of the initial mask matrix according to the global pruning threshold to obtain the initial key-value mask matrix.
[0159] (103.a.3)Perform masking on the matrix elements of the initial key-value mask matrix according to the preset local pruning threshold to obtain the corresponding key-value mask matrices for each attention module.
[0160] (103.a.4)Determine the number of attention heads pruned for any attention module based on the grouping parameter, and update the global pruning threshold based on the number of attention heads to obtain the updated global pruning threshold.
[0161] (103.a.5)Perform masking on the matrix elements of the initial mask matrix based on the updated global pruning threshold to obtain the corresponding query mask matrices for each attention module.
[0162] The following gives a detailed description of steps (103.a.1) to (103.a.5).
[0163] In the embodiment of the present application, the attention mechanism of the initial large language model groups each query layer, so that each group of query layers shares the same key weight matrix and value weight matrix, thereby reducing the computational complexity while maintaining more diversity. The parameter used to control the number of groups is the grouping parameter, and the grouping parameter determines the actual output dimension sizes of each query weight matrix, key weight matrix, value weight matrix, and output weight matrix after grouping.
[0164] Further, for the corresponding first mean matrix of each attention module, use the grouping parameter to adjust its shape. Specifically, adjust the shape of the first mean matrix (attn_metric_i) to (num_head / num_key_value_group, num_key_value_group) to obtain the initial mask matrix (attn_metric_kv), where num_head is the number of attention heads and num_key_value_group is the grouping parameter.
[0165] Further, set the matrix elements in the initial mask matrix that are equal to or greater than the global pruning threshold to True, and set the matrix elements in the initial mask matrix that are less than the global pruning threshold to False, to obtain the initial key-value mask matrix. In order to further refine the pruning operation of each attention module on the basis of the global pruning strategy, it is also necessary to further update the initial key-value mask matrix according to the preset local pruning threshold. Among them, the local pruning threshold is obtained by the user's customization, or is autonomously determined by the initial large language model or other related models based on historical data or current data. The specific value of the local pruning threshold can be adaptively adjusted according to the actual situation, and the embodiments of the present application do not limit this.
[0166] Specifically, sum the elements of each row of the initial mask matrix to obtain a one-dimensional array of size num_key_value_group, which represents its importance measure; then, mask the matrix elements in this one-dimensional array based on the local pruning threshold to obtain the key-value mask matrix.
[0167] Further, since the output dimensions of the query layer and the key-value layer are different, an updated global pruning threshold is obtained by dynamically adjusting the global pruning threshold, and the query mask matrix is determined by the updated global pruning threshold. In this way, although the dimensions of each layer in the attention module of the initial large language model are different, the structured pruning method proposed in the embodiments of the present application ensures consistent pruning operations for each module in its respective dimension by generating independent mask matrices respectively, making the pruning strategy better adapt to the model structure of LLaMA3, thereby improving the pruning refinement degree of the initial large language model, and further improving the pruning accuracy of the initial large language model.
[0168] Among them, the key-value mask matrix represents a binary matrix (0 means pruning, 1 means retaining) indicating whether the initial key matrix and the initial value matrix in each attention module need to be pruned. The query mask matrix represents a binary matrix indicating whether the initial query matrix in each attention module needs to be pruned. The number of pruned attention heads represents the number of attention heads removed after pruning, which is used to dynamically adjust the global pruning threshold to obtain the adjusted updated global pruning threshold.
[0169] In the embodiments of the present application, masking the matrix elements of the initial key-value mask matrix according to the preset local pruning threshold to obtain the corresponding key-value mask matrix for each attention module includes the following steps:
[0170] (B.1) Sum all the matrix elements of the initial key-value mask matrix to obtain the actual total value.
[0171] (B.2) If the actual total value does not meet the preset total value, mask the initial key-value mask matrix using the first local pruning threshold to obtain the corresponding key-value mask matrices for each attention module.
[0172] (B.3) If the actual total value meets the preset total value, mask the initial key-value mask matrix using the second local pruning threshold to obtain the corresponding key-value mask matrices for each attention module.
[0173] The following provides a detailed description of steps (B.1) to (B.3).
[0174] In the embodiments of the present application, the local pruning threshold includes a first local pruning threshold and a second local pruning threshold. Among them, the preset total value can be set to 0. In this way, for the current attention module, if the sum of all matrix elements in the initial key-value mask matrix is not 0, set the matrix elements equal to or greater than the first local pruning threshold to 1, and set the matrix elements less than the first local pruning threshold to 0 to obtain the key-value mask matrix.
[0175] Or, if the sum of all matrix elements in the initial key-value mask matrix is 0, set the matrix elements equal to or greater than the second local pruning threshold to 1, and set the matrix elements less than the second local pruning threshold to 0 to obtain the key-value mask matrix.
[0176] It should be noted that the preset total value, the first local pruning threshold, and the second local pruning threshold can all be set according to the actual situation. This is only an example here and does not represent a limitation in the embodiments of the present application.
[0177] In the embodiments of the present application, updating the global pruning threshold based on the number of attention heads to obtain the updated global pruning threshold includes the following steps:
[0178] (C.1) Sort each matrix element in the first mean matrix to obtain the updated first mean matrix.
[0179] (C.2) Determine that the matrix element at the third target position matching the number of attention heads in the updated first mean matrix is the updated global pruning threshold.
[0180] The following provides a detailed description of steps (C.1) to (C.2).
[0181] In the embodiments of the present application, first determine the number of pruned attention groups num_pruned_groups of any attention module through the following formula <8>:
[0182] num_pruned_groups = num_key_value_group - sum(attn_metric_kv) <8>
[0183] Among them, attn_metric_kv is the initial key-value mask matrix; num_key_value_group is the grouping parameter.
[0184] Further, the number of pruned attention heads pruned_num_heads of any attention module is determined by the following formulas <9> to <10>:
[0185] group_size = head_num / / num_key_value_group <9>
[0186] <10>
[0187] Among them, head_num represents the total number of attention heads of the current attention module; / / represents integer division.
[0188] Further, after determining the number of attention heads, the matrix elements in the first mean matrix can be sorted from small to large to obtain the updated first mean matrix. Then, the updated global pruning threshold threshold' is determined by the following formula <11>:
[0189] threshold' = sort[pruned_num_heads] <11>
[0190] Among them, sort is the sorting function; that is, it is determined that the matrix element at the third target position matching the number of pruned attention heads pruned_num_heads in the sorted and updated first mean matrix is the updated global pruning threshold.
[0191] Step 104, using the key-value mask matrix and query mask matrix of each attention module, pruning the corresponding first weight matrix to obtain a first target matrix, and using the perception mask matrix of each perception module, pruning the corresponding second weight matrix to obtain a second target matrix.
[0192] The following is a detailed description of step 104.
[0193] In the embodiment of the present application, for each attention module of the initial large language model, the output channel numbers of the initial key matrix and initial value matrix are pruned using the key-value mask matrix; and the output channel number of the initial query matrix is pruned using the query mask matrix, and the input channel number of the first output matrix is pruned using the query mask matrix; and then the first target matrix after pruning the first weight matrix is obtained.
[0194] Further, for each perception module of the initial large language model, the input channel number of the mapping matrix is pruned using a perception mask matrix, and the output channel number of the second output matrix is pruned using the perception mask matrix; thereby obtaining a second target matrix after pruning the second weight matrix.
[0195] Step 105: Based on the first target matrix and the second target matrix, determine the large language model after pruning the initial large language model.
[0196] The following provides a detailed description of step 105.
[0197] Among them, the large language model after pruning the initial large language model can be obtained by replacing the original first weight matrix with the first target matrix and replacing the original second weight matrix with the second target matrix, thereby achieving high-accuracy model compression of the initial large language model.
[0198] In the embodiment of the present application, after determining the large language model after pruning the initial large language model, the following steps are further included:
[0199] (105.a.1) Obtain an evaluation sample and a verification sample representing the expected output.
[0200] (105.a.2) Input the evaluation sample into the large language model, and perform embedding processing on the evaluation sample based on the first target matrix and the second target matrix to obtain an actual sample representing the actual output.
[0201] (105.a.3) Calculate the sample difference between the actual sample and the verification sample, and perform pruning adjustment on the large language model based on the sample difference to obtain the large language model after pruning adjustment.
[0202] The following provides a detailed description of steps (105.a.1) to (105.a.3).
[0203] In the embodiment of the present application, to verify the model performance of the large language model obtained after pruning, the result output after processing the input evaluation sample by the large language model can be compared with the preset expected value; if the difference between the two is within the expected range, it indicates that the large language model after pruning meets the model compression requirements, otherwise, the large language model needs to be further pruned and adjusted until a satisfactory large language model is obtained.
[0204] Among them, the evaluation sample is a sample used to test and evaluate the performance of the large language model, and the evaluation sample can be used as input data to test the performance of the large language model after pruning in actual applications. The specific forms of the evaluation sample include, but are not limited to, text, images, etc., depending on the task type of the large language model. The verification sample represents the result expected to be output after the evaluation sample is processed by the large language model.
[0205] Among them, the sample difference can be determined by calculating the mean squared error, cross-entropy loss, cosine similarity, etc. between the actual sample and the verification sample. The specific method for determining the sample difference can be set according to the actual situation. In addition, the threshold used to compare with the sample difference to determine whether the large language model needs further pruning adjustment can also be set according to the actual situation, and the embodiments of the present application do not limit this.
[0206] Among them, the pruning adjustment can be an adjustment of the pruning rate, or the specific method of pruning adjustment can be set according to the actual situation, and the embodiments of the present application do not limit this.
[0207] Such as Figure 3 、 Figure 4 shown, Figure 3 is another optional flowchart of the structured pruning method for the large language model provided by the embodiments of the present application, Figure 4 is yet another optional flowchart of the structured pruning method for the large language model provided by the embodiments of the present application. To enable readers to more clearly understand the structured pruning method proposed by the embodiments of the present application, the following will take LLaMA3 as an example and combine Figure 3 、 Figure 4 for a complete example description.
[0208] (1) Load the LLaMA3 large model, where LLaMA3 includes multiple processing layers, and each processing layer includes a group of interconnected Attention modules and MLP modules; the Attention module includes multiple attention heads, and each attention head includes a key layer (k_proj), a value layer (v_proj), a query layer (q_proj), a first output layer (o_proj), and their corresponding weight matrices.
[0209] (2) Obtain the initial sample, and perform the following operations on each processing layer of the LLaMA3 model: process the initial sample based on the Attention module and MLP module of the LLaMA3 model, and then calculate the fluctuation metric attn_metric of the current processing layer's attention module (Attention) and the fluctuation metric mlp_metric of the current processing layer's perception module (MLP).
[0210] (3) Store the fluctuation metric attn_metric of the Attention module in the attn_metric_list linked list, and store the fluctuation metric mlp_metric of the MLP module in the mlp_metric_list linked list; also store the input sample mean attn_mean obtained after processing the initial sample by the Attention module in the attn_mean_list linked list, and store the input sample mean mlp_mean obtained after processing the initial sample by the MLP module in the mlp_mean_list linked list, so as to update the elements of the correlation matrix based on the attn_mean_list linked list and the mlp_mean_list linked list later.
[0211] (4) After calculating the fluctuation metrics of all processing layers, stack all the metric vectors in attn_metric_list into the first fluctuation metric matrix attn_metric; stack all the metric vectors of all layers in mlp_metric_list into the second fluctuation metric matrix mlp_metric.
[0212] (5) Perform normal standardization on the first fluctuation metric matrix attn_metric and the second fluctuation metric matrix mlp_metric respectively; and adjust the shape of the normalized attn_metric matrix to (nlayers, head_num, 128), where nlayers represents the number of processing layers, head_num is the number of attention heads. For example, for the LLaMA3-70B model, head_num = 64; 128 represents the dimension of each attention head; it should be noted that the specific shape after adjusting the attn_metric matrix can be set according to the actual situation.
[0213] (6) Take the average of the second dimension (counting from 0) of the attn_metric matrix to obtain the importance scores of each head (still named attn_metric), whose shape is (nlayers, head_num). This new attn_metric will be used as the input of the get_layer_mask function.
[0214] (7) Call the get_layer_mask function to calculate three mask matrices: the key-value mask matrix (attn_mask_kv), the query mask matrix (attn_mask_q), and the perception mask matrix (mlp_mask). Compared with the traditional structured pruning method that uses the same mask matrix, in the embodiments of the present application, by separately determining the independent mask matrices for each processing layer, it can better adapt to the special structure of LLaMA3, thereby refining the pruning fineness and improving the pruning accuracy of LLaMA3.
[0215] (8) For each processing layer, based on attn_mask_q and attn_mask_kv, call the compress function to prune the weight matrices of the q_proj layer, k_proj layer, v_proj layer, and o_proj layer:
[0216] For the Attention module of the current processing layer, use attn_mask_q to prune the weight matrices of the q_proj layer and o_proj layer, and use attn_mask_kv to prune the weight matrices of the k_proj layer and v_proj layer. Specifically, use attn_mask_q to prune the number of output channels of the weight matrix of the q_proj layer, and use attn_mask_kv to prune the number of output channels of the k_proj layer and v_proj layer: update the number of output channels out_features of the q_proj layer to the number of non-zero elements of attn_mask_q; and update the number of output channels out_features of the k_proj layer and v_proj layer to the number of non-zero elements of attn_mask_kv. Use attn_mask_q to prune the number of input channels of the o_proj layer, update the number of input channels of the o_proj layer to the number of non-zero elements of attn_mask_q, and update the bias of the o_proj layer to the compensation value vector.
[0217] For the MLP module of the current processing layer, use mlp_mask to prune the number of input channels of the weight matrices of the up_proj layer and gate_proj layer of the MLP module, and prune the number of output channels of the weight matrix of the down_proj layer; update the bias of the down_proj layer to the compensation value vector, synchronously update the number of output channels out_features of the up_proj layer and gate_proj layer, and update the parameters of the intermediate hidden layer of the MLP module to the number of non-zero elements of mlp_mask.
[0218] (9) According to the obtained evaluation samples, perform verification processing on the large language model obtained after pruning. When the verification passes, save the corresponding large language model.
[0219] Furthermore, if Figure 4 As shown, step (7) specifically also includes the following steps:
[0220] (7.1) Flatten the attn_metric and mlp_metric matrices and concatenate them to obtain the global metric matrix prune_metric.
[0221] (7.2) Sort the matrix vectors in prune_metric from large to small to obtain the sorted vector sorted_prune and the corresponding index vector indices of each vector.
[0222] (7.3) Calculate the global pruning threshold threshold: Initialize a vector compression_weight, which is the same size as indices and all values are 1. Set the positions in indices that are less than the number of attn_metric elements to True, and the others to False, to obtain a bool vector indices_b; set the value of compression_weight in the True position in indices_b to 512.0 / 3. Then, calculate the threshold based on compression_weight and the preset pruning rate pruning_ratio.
[0223] (7.4) Calculate the mask matrix mlp_mask of the MLP module: set the positions where the value of mlp_metric is greater than threshold to True, and set the others to False. At this point, the mask matrix mlp_mask corresponding to the MLP module is determined.
[0224] (7.5) Calculate the mask matrices attn_mask_q and attn_mask_kv of the Attention module based on attn_metric and threshold. The shape of attn_metric is [nlayers, head_num], which stores the head importance scores of all processing layers. For each layer of attn_metric, perform the following processing on attn_metric_i to obtain attn_mask_q and attn_mask_kv:
[0225] Furthermore, step (7.5) specifically includes the following steps:
[0226] (7.5.1) Reshape attn_metric_i to (num_head / num_key_value_group, num_key_value_group) to obtain a new matrix attn_metric_kv; here, multiple heads are divided into num_key_value_group groups, and the value of num_key_value_group is a parameter of the Attention module.
[0227] (7.5.2) Set the places in attn_metric_kv greater than threshold to True, and set the places less than or equal to threshold to False.
[0228] (7.5.3) Sum each row of attn_metric_kv to obtain a one-dimensional array of size num_key_value_group representing the importance metric; if the sum of all elements of attn_mask_kv is not 0, set the values in this one-dimensional array greater than or equal to the first local pruning threshold (group_prune_thr1) to 1, and set the values less than group_prune_thr1 to 0 to obtain the mask matrix attn_mask_kv.
[0229] If the sum of all elements of attn_mask_kv is 0, recalculate attn_mask_kv using the second local pruning threshold (group_prune_thr2) to avoid all attention groups being pruned. In an optional embodiment, group_prune_thr1 and group_prune_thr2 can be set to 0.25 and 0 respectively.
[0230] (7.5.4) Determine the updated global pruning threshold threshold' based on threshold.
[0231] (7.5.5) Calculate the mask vector attn_mask_q: Set the values in attn_metric_i (referring to the initial attn_metric_i, unsorted) greater than threshold to True, and set the others to False to obtain the mask vector attn_mask_q.
[0232] (7.5.6) Return mlp_mask, attn_mask_kv, and attn_mask_q.
[0233] It should be noted that the parts involving formula calculations in this example will not be elaborated in detail. For details, please refer to steps 101 to 105.
[0234] As Figure 5 shown Figure 5 is an optional module schematic diagram of the structured pruning device for large language models provided by an embodiment of the present application, including:
[0235] An acquisition module 201, configured to acquire an initial large language model to be pruned, and acquire a corresponding first weight matrix from multiple attention modules of the initial large language model and a corresponding second weight matrix from multiple perception modules of the initial large language model.
[0236] A fluctuation metric matrix determination module 202, configured to determine a first fluctuation metric matrix based on the first weight matrix, determine a second fluctuation metric matrix based on the second weight matrix, and determine a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix.
[0237] A global pruning threshold determination module 203, configured to determine a corresponding key-value mask matrix and query mask matrix for each attention module based on the global pruning threshold and the first fluctuation metric matrix, and determine a corresponding perception mask matrix for each perception module based on the global pruning threshold and the second fluctuation metric matrix.
[0238] A mask matrix determination module 204, configured to perform pruning processing on the corresponding first weight matrix by using the key-value mask matrix and query mask matrix of each attention module to obtain a first target matrix, and perform pruning processing on the corresponding second weight matrix by using the perception mask matrix of each perception module to obtain a second target matrix.
[0239] A target module 205, configured to determine a large language model after pruning the initial large language model based on the first target matrix and the second target matrix.
[0240] The structured pruning method and related devices for large language models proposed in this application. The method includes: obtaining an initial large language model to be pruned, and obtaining corresponding first weight matrices from multiple attention modules of the initial large language model and obtaining corresponding second weight matrices from multiple perception modules of the initial large language model; determining a first fluctuation metric matrix based on the first weight matrix, determining a second fluctuation metric matrix based on the second weight matrix, and determining a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix; determining corresponding key-value mask matrices and query mask matrices for each attention module based on the global pruning threshold and the first fluctuation metric matrix, and determining corresponding perception mask matrices for each perception module based on the global pruning threshold and the second fluctuation metric matrix; mask matrices with different output dimensions can ensure that the pruning operation adapts to all processing layers of the initial large language model to flexibly meet the requirements of different output dimensions between layers in each attention module of the initial large language model; then, using the key-value mask matrix and query mask matrix of each attention module to perform pruning processing on the corresponding first weight matrix to obtain a first target matrix, and using the perception mask matrix of each perception module to perform pruning processing on the corresponding second weight matrix to obtain a second target matrix; determining the large language model after pruning the initial large language model based on the first target matrix and the second target matrix. In this way, generating corresponding mask matrices according to the actual situations of each attention module and each perception module in each processing layer can improve the model compression accuracy of the large language model while adapting to the special structure of the initial large language model.
[0241] The specific implementation manner of this structured pruning device is basically the same as the specific embodiment of the above structured pruning method, and will not be elaborated here.
[0242] This application embodiment also provides an electronic device. The electronic device includes a memory and a processor. The memory stores a computer program, and when the processor executes the computer program, it implements the above structured pruning method. The electronic device can be any intelligent terminal including a tablet computer, an in-vehicle computer, etc.
[0243] As Figure 6 shown, Figure 6 is the hardware structure schematic diagram of the electronic device provided by this application embodiment. The electronic device includes:
[0244] A processor 301, which can be implemented by using a general-purpose CPU (Central Processing Unit), a microprocessor, an application-specific integrated circuit (ASIC), or one or more integrated circuits, etc., and is used to execute relevant programs to implement the technical solutions provided by this application embodiment;
[0245] The memory 302 can be implemented in the form of a read-only memory (ROM), a static storage device, a dynamic storage device, or a random access memory (RAM), etc. The memory 302 can store an operating system and other application programs. When implementing the technical solutions provided in the embodiments of this specification through software or firmware, the relevant program codes are stored in the memory 302, and the processor 301 is used to call and execute the structured pruning method of the embodiments of this application;
[0246] The input / output interface 303 is used to implement information input and output;
[0247] The communication interface 304 is used to implement communication and interaction between this device and other devices. Communication can be achieved through wired means (such as USB, network cable, etc.) or through wireless means (such as mobile network, WIFI, Bluetooth, etc.);
[0248] The bus 305 transmits information between various components of the device (such as the processor 301, the memory 302, the input / output interface 303, and the communication interface 304);
[0249] Among them, the processor 301, the memory 302, the input / output interface 303, and the communication interface 304 achieve communication connections with each other inside the device through the bus 305.
[0250] The embodiments of this application also provide a computer-readable storage medium. The computer-readable storage medium stores a computer program, and when the computer program is executed by a processor, the above-mentioned structured pruning method is implemented.
[0251] As a non-transitory computer-readable storage medium, the memory can be used to store non-transitory software programs and non-transitory computer-executable programs. In addition, the memory can include high-speed random access memory, and can also include non-transitory memory, such as at least one magnetic disk storage device, a flash memory device, or other non-transitory solid-state storage devices. In some embodiments, the memory optionally includes a memory remotely set relative to the processor, and these remote memories can be connected to the processor through a network. Examples of the above-mentioned network include but are not limited to the Internet, an enterprise intranet, a local area network, a mobile communication network, and combinations thereof.
[0252] The embodiments described in the embodiments of this application are for more clearly illustrating the technical solutions of the embodiments of this application, and do not constitute a limitation on the technical solutions provided by the embodiments of this application. Those skilled in the art know that with the evolution of technology and the emergence of new application scenarios, the technical solutions provided by the embodiments of this application are equally applicable to similar technical problems.
[0253] Those skilled in the art can understand that the technical solutions shown in the figures do not constitute a limitation on the embodiments of the present application, and may include more or fewer steps than those shown, or combine certain steps, or different steps.
[0254] The device embodiments described above are merely illustrative. The units described as separate components may or may not be physically separated, that is, they may be located in one place, or may be distributed to multiple network units. Some or all of the modules may be selected according to actual needs to achieve the purpose of the solution of this embodiment.
[0255] Those of ordinary skill in the art can understand that all or some of the steps in the methods disclosed above, and the functional modules / units in the systems and devices, can be implemented as software, firmware, hardware, and appropriate combinations thereof.
[0256] The terms "first", "second", "third", "fourth", etc. (if any) in the specification of this application and the above-mentioned drawings are used to distinguish similar objects, and do not necessarily need to describe a specific order or sequence. It should be understood that the data used in this way can be interchanged under appropriate circumstances, so that the embodiments of the present application described here can be implemented in an order other than those illustrated or described here. In addition, the terms "including" and "having" and any variations thereof are intended to cover non-exclusive inclusion. For example, a process, method, system, product, or device that includes a series of steps or units does not necessarily have to be limited to those steps or units clearly listed, but may include other steps or units not clearly listed or inherent to these processes, methods, products, or devices.
[0257] It should be understood that in this application, "at least one (item)" means one or more, and "a plurality" means two or more. "And / or" is used to describe the association relationship of associated objects, indicating that three relationships may exist. For example, "A and / or B" may mean: only A exists, only B exists, and both A and B exist at the same time. Among them, A and B can be singular or plural. The character " / " generally means that the associated objects before and after are in an "or" relationship. "At least one (one) of the following" or a similar expression means any combination of these items, including any combination of single items (ones) or plural items (ones). For example, at least one (one) of a, b, or c may mean: a, b, c, "a and b", "a and c", "b and c", or "a and b and c", where a, b, and c can be single or multiple.
[0258] In several embodiments provided by the present application, it should be understood that the disclosed devices and methods can be implemented in other ways. For example, the device embodiments described above are merely illustrative. For example, the division of the above-mentioned units is only a logical function division. In actual implementation, there may be other division methods. For example, multiple units or components can be combined or integrated into another system, or some features can be ignored or not executed. Another point is that the displayed or discussed couplings or direct couplings or communication connections between each other can be through some interfaces. The indirect couplings or communication connections of devices or units can be in electrical, mechanical or other forms.
[0259] The units described above as separate components may or may not be physically separated. The components displayed as units may or may not be physical units, that is, they can be located in one place or distributed to multiple network units. Some or all of the units can be selected according to actual needs to achieve the purpose of the solution of this embodiment.
[0260] In addition, in each embodiment of the present application, each functional unit can be integrated in a processing unit, or each unit can exist physically alone, or two or more units can be integrated in one unit. The above-mentioned integrated units can be implemented in the form of hardware or in the form of software functional units.
[0261] If the integrated unit is implemented in the form of a software functional unit and sold or used as an independent product, it can be stored in a computer-readable storage medium. Based on such an understanding, the technical solution of the present application, in essence, or the part that contributes to the prior art, or all or part of this technical solution, can be embodied in the form of a software product. This computer software product is stored in a storage medium and includes multiple instructions for causing a computer device (which can be a personal computer, a server, or a network device, etc.) to execute all or part of the steps of the methods in each embodiment of the present application. The foregoing storage medium includes: various media such as USB flash drives, mobile hard disks, read-only memories (ROM), random access memories (RAM), magnetic disks, or optical discs that can store programs.
[0262] The preferred embodiments of the embodiments of the present application have been described above with reference to the drawings, but this does not limit the scope of rights of the embodiments of the present application. Any modifications, equivalent replacements, and improvements made by those skilled in the art without departing from the scope and essence of the embodiments of the present application shall be within the scope of rights of the embodiments of the present application.
Claims
1. A structured pruning method for a large language model, characterized in that: include: Obtain an initial large language model to be pruned, and obtain a corresponding first weight matrix from multiple attention modules of the initial large language model, and obtain a corresponding second weight matrix from multiple perception modules of the initial large language model; Determine a first volatility measurement matrix based on the first weight matrix, determine a second volatility measurement matrix based on the second weight matrix, and determine a global pruning threshold according to the first volatility measurement matrix and the second volatility measurement matrix; Based on the global pruning threshold and the first fluctuation metric matrix, determining a key value mask matrix and a query mask matrix corresponding to each of the attention modules, and based on the global pruning threshold and the second fluctuation metric matrix, determining a perception mask matrix corresponding to each of the perception modules; Using the key-value mask matrix and the query mask matrix of each of the attention modules, pruning the corresponding first weight matrix to obtain a first target matrix, and using the perception mask matrix of each of the perception modules, pruning the corresponding second weight matrix to obtain a second target matrix; Determining a large language model after pruning the initial large language model based on the first target matrix and the second target matrix; Each of the attention modules includes a plurality of first weight matrices, each of the first weight matrices includes an initial key matrix, an initial query matrix, an initial value matrix and a first output matrix, and output dimensions of the initial key matrix, the initial query matrix, the initial value matrix and the first output matrix are not completely the same; The determining a first volatility measurement matrix based on the first weight matrix comprises: Acquire an initial sample, and input the initial sample into the initial large language model; For each of the first weight matrices, embedding processing is performed on the initial samples based on the initial key matrix, the initial query matrix, the initial value matrix and the first output matrix in sequence to obtain corresponding first output samples; Determine a first original matrix based on the first output sample, and perform mean processing on multiple first original matrices belonging to the same attention module to obtain a first mean matrix; The first mean matrices corresponding to all the attention modules are stacked to obtain a first fluctuation measure matrix.
2. The structured pruning method for a large language model according to claim 1, characterized in that: The step of performing mean processing on a plurality of the first original matrices belonging to the same attention module to obtain a first mean matrix includes: Performing mean processing on matrix elements in each of the first original matrices respectively to obtain a corresponding metric mean of each of the first original matrices; Based on the matrix elements in each of the first original matrices, calculate and obtain the metric standard difference value corresponding to each of the first original matrices; Based on the metric mean and the metric standard deviation, update the matrix elements in the corresponding first original matrix to obtain a first normalized matrix; Perform mean processing on multiple first normalized matrices belonging to the same attention module to obtain a first mean matrix.
3. The structured pruning method for a large language model according to claim 1, characterized in that: Each of the attention modules is connected to the corresponding perception module, each of the perception modules includes a plurality of the second weight matrices, and each of the second weight matrices includes a mapping matrix and a second output matrix; The determining a second volatility measurement matrix based on the second weight matrix comprises: For each of the first weight matrices, embedding processing is performed on the first output samples based on the mapping matrix and the second output matrix in sequence to obtain corresponding second output samples; Determine a second original matrix based on the second output sample, and perform mean processing on a plurality of the second original matrices belonging to the same perception module to obtain a second mean matrix; The second mean matrices corresponding to all the perception modules are stacked to obtain a second fluctuation measurement matrix.
4. The structured pruning method for a large language model according to claim 1, characterized in that: The determining of a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix comprises: Combining the first volatility metric matrix and the second volatility metric matrix to obtain a global metric matrix; Constructing an initial matrix having the same matrix size as the global metric matrix, wherein the initial matrix includes a plurality of initial elements; Based on a preset pruning rate, a preset accumulation function and the global metric matrix, updating the plurality of initial elements to obtain a corresponding plurality of updated elements; Determine a target update element with the smallest element value from the multiple update elements, and determine a first target position where the target update element is located; The global pruning threshold is determined based on a second target position that matches the global metric matrix with the first target position.
5. The structured pruning method for a large language model according to claim 1, characterized in that: The determining, based on the global pruning threshold and the first fluctuation metric matrix, a key value mask matrix and a query mask matrix corresponding to each of the attention modules comprises: Based on the grouping parameters preset in the initial large language model, adjusting the matrix shape of each of the first mean matrices to obtain an initial mask matrix; According to the global pruning threshold, updating the matrix elements of the initial mask matrix to obtain an initial key-value mask matrix; According to a preset local pruning threshold, masking is performed on the matrix elements of the initial key-value mask matrix to obtain the key-value mask matrix corresponding to each of the attention modules; Determine the number of pruned attention heads of any of the attention modules based on the grouping parameters, and update the global pruning threshold based on the number of attention heads to obtain an updated global pruning threshold; Based on the updated global pruning threshold, the matrix elements of the initial mask matrix are masked to obtain the query mask matrix corresponding to each attention module.
6. The structured pruning method for a large language model according to claim 5, characterized in that: The local pruning threshold includes a first local pruning threshold and a second local pruning threshold; The masking process is performed on the matrix elements of the initial key-value mask matrix according to the preset local pruning threshold to obtain the key-value mask matrix corresponding to each attention module, including: Summing all matrix elements of the initial key value mask matrix to obtain an actual total value; If the actual total value does not meet the preset total value, masking the initial key-value mask matrix using the first local pruning threshold to obtain the key-value mask matrix corresponding to each attention module; If the actual total value meets the preset total value, the initial key-value mask matrix is masked using the second local pruning threshold to obtain the key-value mask matrix corresponding to each attention module.
7. The structured pruning method for a large language model according to claim 5, characterized in that: The updating of the global pruning threshold based on the number of attention heads to obtain an updated global pruning threshold includes: Sorting each matrix element in the first mean matrix to obtain an updated first mean matrix; Determine in the updated first mean matrix that the matrix element at the third target position that matches the number of attention heads is the updated global pruning threshold.
8. The structured pruning method for a large language model according to claim 6, characterized in that: After determining the large language model after pruning the initial large language model, the method further includes: Obtain evaluation samples and validation samples that represent the expected output; Inputting the evaluation sample into the large language model, and performing embedding processing on the evaluation sample based on the first target matrix and the second target matrix to obtain an actual sample representing an actual output; A sample difference between the actual sample and the verification sample is calculated, and the large language model is pruned and adjusted based on the sample difference to obtain the pruned and adjusted large language model.
9. A structured pruning device for a large language model, characterized in that: include: An acquisition module, used to acquire an initial large language model to be pruned, and acquire a corresponding first weight matrix from multiple attention modules of the initial large language model, and acquire a corresponding second weight matrix from multiple perception modules of the initial large language model; a fluctuation metric matrix determination module, configured to determine a first fluctuation metric matrix based on the first weight matrix, determine a second fluctuation metric matrix based on the second weight matrix, and determine a global pruning threshold according to the first fluctuation metric matrix and the second fluctuation metric matrix; A global pruning threshold determination module, used to determine a key value mask matrix and a query mask matrix corresponding to each of the attention modules based on the global pruning threshold and the first fluctuation metric matrix, and to determine a perception mask matrix corresponding to each of the perception modules based on the global pruning threshold and the second fluctuation metric matrix; A mask matrix determination module, configured to use the key-value mask matrix and the query mask matrix of each of the attention modules to perform pruning processing on the corresponding first weight matrix to obtain a first target matrix, and use the perception mask matrix of each of the perception modules to perform pruning processing on the corresponding second weight matrix to obtain a second target matrix; A target module, which determines a large language model after pruning the initial large language model based on the first target matrix and the second target matrix; Each of the attention modules includes a plurality of first weight matrices, each of the first weight matrices includes an initial key matrix, an initial query matrix, an initial value matrix and a first output matrix, and output dimensions of the initial key matrix, the initial query matrix, the initial value matrix and the first output matrix are not completely the same; The determining a first volatility measurement matrix based on the first weight matrix comprises: Acquire an initial sample, and input the initial sample into the initial large language model; For each of the first weight matrices, embedding processing is performed on the initial samples based on the initial key matrix, the initial query matrix, the initial value matrix and the first output matrix in sequence to obtain corresponding first output samples; Determine a first original matrix based on the first output sample, and perform mean processing on multiple first original matrices belonging to the same attention module to obtain a first mean matrix; The first mean matrices corresponding to all the attention modules are stacked to obtain a first fluctuation measure matrix.
10. An electronic device, characterized in that: The electronic device includes a memory and a processor, the memory stores a computer program, and the processor implements the structured pruning method for a large language model according to any one of claims 1 to 8 when executing the computer program.
11. A computer-readable storage medium storing a computer program, characterized in that: When the computer program is executed by a processor, the structured pruning method for a large language model according to any one of claims 1 to 8 is implemented.
Citation Information
Patent Citations
Structured pruning method and system
CN115222042A
Large language model training method and device
CN118445379A