Model Selection Learning for Knowledge Distillation

Through knowledge distillation methods and reinforcement learning technology, reference models and update strategy parameters are dynamically selected, which solves the problem of training and inference time-consuming caused by the large amount of parameters of deep pre-trained models, and realizes efficient model training and deployment.

CN113822434BActive Publication Date: 2025-05-13MICROSOFT TECHNOLOGY LICENSING LLC
View PDF 0 Cites 1 Cited by

Patent Information

Application Number
CN202010561319.7
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2020-06-18
Publication Date
2025-05-13
Estimated Expiration
2040-06-18

AI Technical Summary

Technical Problem

The existing deep pre-trained models are difficult to apply in actual business scenarios due to the large number of parameters and the long training and inference.

Method used

Through the knowledge distillation method, a set of candidate reference models is used to train the target model, the appropriate reference model is dynamically selected for each training sample, and the policy parameters are updated through reinforcement learning and policy functions to improve the performance of the target model.

Benefits of technology

It effectively reduces the training and inference time of the target model, improves the performance and deployment efficiency of the model, and is suitable for actual business scenarios.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN113822434B_ABST
    Figure CN113822434B_ABST
Patent Text Reader

Abstract

The present disclosure provides a method and apparatus for obtaining a target model based on knowledge distillation. A data set and a set of candidate reference models may be obtained. A set of selected reference models selected from the set of candidate reference models may be determined for each training sample in the data set. A set of target probability distributions output by the set of selected reference models for the training sample may be obtained. The target model may be trained using the set of target probability distributions.
Need to check novelty before this filing date? Find Prior Art

Description

Background Art

[0001] With the development of deep learning technology, various deep pre-trained models have been continuously developed and have performed well in fields such as natural language processing and computer vision. For example, in the field of natural language processing, deep pre-trained models such as the Bidirectional Encoder Resentations from Transformers (BERT) model and the Generative Pre-trained Transformer (GPT) model have been shown to have good results. Such deep pre-trained models are often complex models that rely on deep networks with huge parameters. For example, the BERT model may contain 24 transformer layers with a total of 340 million parameters, and the GPT model may contain 48 transformer layers with a total of 1.5 billion parameters. Training such complex models and using such complex models for inference are very time-consuming, making it difficult to apply them to actual business scenarios. Model compression methods are usually used to obtain simple models that can be deployed with fewer parameters than complex models. Summary of the invention

[0002] This Summary is provided to introduce a group of concepts that will be further described in the following Detailed Description. This Summary is not intended to identify key features or essential features of the claimed subject matter, nor is it intended to be used to limit the scope of the claimed subject matter.

[0003] Embodiments of the present disclosure provide a method and apparatus for obtaining a target model based on knowledge distillation. A data set and a set of candidate reference models may be obtained. A set of selected reference models selected from the set of candidate reference models may be determined for each training sample in the data set. A set of target probability distributions output by the set of selected reference models for the training sample may be obtained. The target model may be trained using the set of target probability distributions.

[0004] It should be noted that one or more of the above aspects include the features specifically pointed out in the following detailed description and claims. The following description and drawings set forth in detail certain illustrative features of the one or more aspects. These features are merely indicative of the various ways in which the principles of various aspects may be implemented, and the present disclosure is intended to include all such aspects and their equivalents. BRIEF DESCRIPTION OF THE DRAWINGS

[0005] The disclosed aspects will be described below in conjunction with the accompanying drawings, which are provided to illustrate rather than limit the disclosed aspects.

[0006] Figure 1An exemplary process for obtaining a target model according to an embodiment of the present disclosure is shown.

[0007] Figure 2 An exemplary process for selecting a reference model according to an embodiment of the present disclosure is shown.

[0008] Figure 3 An exemplary process for training a target model according to an embodiment of the present disclosure is shown.

[0009] Figure 4 A specific example for training a target model according to an embodiment of the present disclosure is shown.

[0010] Figure 5 An exemplary process for updating policy parameters according to an embodiment of the present disclosure is shown.

[0011] Figure 6 An exemplary process for initializing policy parameters according to an embodiment of the present disclosure is shown.

[0012] Figure 7 is a flowchart of an exemplary method for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure.

[0013] Figure 8 An exemplary apparatus for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure is shown.

[0014] Fig. 9 An exemplary apparatus for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure is shown. DETAILED DESCRIPTION

[0015] The present disclosure will now be discussed with reference to several exemplary embodiments. It should be understood that the discussion of these embodiments is only used to enable those skilled in the art to better understand and thereby implement the embodiments of the present disclosure, and does not teach any limitation on the scope of the present disclosure.

[0016] A commonly used model compression method can be based on knowledge distillation. This method usually transfers knowledge from a complex model to a simple model by learning the output distribution of a complex model with a simple model. Knowledge can be considered as the parameters of a complex model and the mapping of inputs to outputs implemented by a complex model. This method is based on a teacher-student architecture, in which the model that provides knowledge can be considered as a teacher model, and the model that learns knowledge can be considered as a student model. Specifically, when training a student model, the student model is provided with training data that not only has real annotations, such as annotations provided by humans, but also has the probability distribution of the output of the teacher model. Therefore, the student model can optimize its model parameters by learning the probability distribution of the output of the teacher model in an attempt to achieve the effect of the teacher model. Based on the number of teacher models and student models, the current knowledge distillation methods can be divided into one-to-one methods, many-to-many methods, and many-to-one methods. In this article, the one-to-one method refers to a teacher model providing knowledge to a student model, the many-to-many method refers to multiple teacher models providing knowledge to multiple student models and combining these multiple student models into a student model set when applied, and the many-to-one rule refers to multiple teacher models providing knowledge to a student model. Recent studies and experiments have shown that using a many-to-one approach to train the student model can more effectively improve the performance of the student model.

[0017] Current many-to-one approaches usually assign the same weight to each teacher model, or assign different weights to each teacher model, but these weights are fixed throughout the knowledge distillation. However, even if a set of training data for the same task is used to train the student model, the performance of each teacher model for different training samples in the set of training data is different. Taking two training samples for predicting semantic equivalence between two sentences as an example, for the first training sample, the performance of teacher model A may be better than that of teacher model B; while for the second training sample, the performance of teacher model B may be better than that of teacher model A. In addition, at each stage of knowledge distillation, the performance of the student model is gradually improved, and the teacher model used to train the student model should also change accordingly. For example, in the early stage of knowledge distillation, the performance of the student model is weak, and its learning effect from the complex teacher model may not be good. This is because the complex teacher model captures the finer-grained patterns in the training data, which may cause the student model to overfit some parts of the training data. As the training process progresses, the student model has a strong performance, and it may be difficult to achieve significant results if it is trained by a teacher model with a small performance gap. Therefore, assigning the same or fixed weights to each teacher model during knowledge distillation may limit the performance improvement of the student model.

[0018] The embodiments of the present disclosure propose to improve the performance of the target model through an improved training process. For example, the target model can be trained using a set of reference models through knowledge distillation. In this article, the target model refers to a model that is expected to be trained with a simple structure and can be deployed, which can also be called a student model, and the reference model refers to a model that can be used to assist in training the target model and has a higher complexity than the target model, which can also be called a teacher model.

[0019] In one aspect, embodiments of the present disclosure propose to select a reference model for training a target model from a set of candidate reference models through reinforcement learning. For example, different weights can be dynamically assigned to each reference model for each training sample in a data set used to train the target model. In this article, the weight assigned to a specific reference model can be implemented as a sampling probability corresponding to the reference model, which can be used to determine whether to select the reference model to train the target model for the training sample.

[0020] In another aspect, the embodiments of the present disclosure propose to determine the sampling probability of each reference model for the current training sample through a strategy function. The strategy function for each reference model includes, for example, strategy parameters, and information related to the current training sample and the performance of the reference model for the current training sample. Using such a strategy function can help select a reference model that performs well for the current training sample.

[0021] In another aspect, embodiments of the present disclosure propose to update policy parameters in a policy function based on the performance of a target model. For example, a data set used to train a target model may be divided into multiple data subsets. After training a target model using a data subset, the policy parameters in the policy function may be updated based on the performance of the trained target model, thereby affecting the sampling probability of each reference model for the next data subset. Updating the policy parameters in the policy function based on the performance of the target model may help select a reference model that matches the performance of the current target model.

[0022] In another aspect, the embodiments of the present disclosure propose to initialize the policy parameters in the policy function before determining the sampling probability of each reference model through the policy function. For example, a group of reference models can be selected from a group of candidate reference models through the policy function, and the policy parameters in the policy function can be initialized according to the average performance of the selected group of reference models.

[0023] In another aspect, the embodiments of the present disclosure propose to pre-train the target model before using a set of reference models to train the target model. For example, all reference models in the set of reference models can be used to score a data set, and the scored data set can be used to pre-train the target model.

[0024] Figure 1 An exemplary process 100 for obtaining a target model according to an embodiment of the present disclosure is shown. The target model is, for example, Figure 1 The target model 160 in , which can be a BERT model with 3 or 6 layers of transformers.

[0025] First, a data set for training the target model 160 can be obtained. Data Collection Can be divided into multiple data subsets Where M represents the number of data subsets. Each data subset may include multiple training samples. In this article, the samples used to train the target model 160 are referred to as training samples. For example, it can include m training samples For example, training sample i 102(x i ,y i ), where x i is the i-th input, and y i Is for x i , such as human-provided annotations.

[0026] A set of candidate reference models 110 may be obtained for training a target model 160. The candidate reference model may be a model with higher complexity than the target model 160, such as a BERT model with 12 layers of transformers. The candidate reference model may be obtained by optimizing, such as fine-tuning, a pre-trained model using training data for a specific task.

[0027] A representation model 120 may also be obtained, which may be a representation model that can effectively represent x i Any pre-trained model of the content.

[0028] The training sample i 102 may be provided as an input to each reference model in a set of candidate reference models 110 and the representation model 120 to obtain a set of state information. The set of state information may at least include information related to the training sample i 102 and the target probability distribution output by each reference model for the training sample i 102. Figure 2 To explain the specific process of obtaining status information.

[0029] Subsequently, at 130, for each candidate reference model in a set of candidate reference models 110, reinforcement learning may be used to determine whether to select the candidate reference model for training the target model 160. For example, the policy function π θ 132 to determine whether to select the candidate reference model. The selected reference models can be combined into a set of selected reference models 140. Figure 2 To explain the specific process of selecting a reference model.

[0030] Next, the target probability distribution output by each reference model in the set of selected reference models 140 for the training sample i 102 may be obtained to obtain a set of target probability distributions 150. A set of state information used to determine a set of selected reference models 140 may include information related to the target probability distribution output by each reference model for the training sample i 102. A set of target probability distributions 150 corresponding to each reference model in the set of selected reference models 140 may be extracted from the set of state information.

[0031] The target model 160 may be trained using the training sample i 102 and the set of target probability distributions 150. In one embodiment, after performing the above process using a single training sample, the parameters of the target model 160 may be optimized to obtain a trained target model 170. In another embodiment, the target model 170 may be trained using a data subset, such as a data subset. After the above process is performed for all training samples in the data subset, the parameters of the target model 160 are optimized to obtain the trained target model 170. In this case, for all training samples in the same data subset, the parameters of the target model 160 remain unchanged. Figure 3 and Figure 4 To explain the specific process of training target model 160.

[0032] Subsequently, the performance of the trained target model 170 may be evaluated. The performance of the trained target model 170 may be evaluated using a validation sample 180. In this document, the sample used to evaluate the performance of the target model is referred to as a validation sample, which may be the same as or different from the training sample. The evaluated performance may be converted into a reward 190. The reward 190 may then be used to update the policy function π θ The policy parameter θ in 132 will be combined later Figure 5 To explain the specific process of updating strategy parameters.

[0033] In the case of using a single training sample to obtain a trained target model 170, updating the strategy parameters based on the performance of the trained target model 170 can affect the sampling probability of each reference model for the next training sample. In the case of obtaining a trained target model 170 based on all training samples in the same data subset, updating the policy parameters based on the performance of the trained target model 170 can affect the sampling probability of each reference model for the next data subset. In this case, for all training samples in the same data subset, the policy function π θ The policy parameters θ remain unchanged.

[0034] The process 100 mainly includes a process of training a target model and a process of updating policy parameters, and these two processes can be iteratively performed until the performance of the target model converges. During the process of training the target model, the parameters of the target model can be optimized while the policy parameters are fixed; and during the process of updating the policy parameters, the policy parameters can be optimized while the parameters of the target model are fixed. The current parameters of the target model can be set as And set the current policy parameter to θ b . We can first fix the policy parameters to θ b And using the data subset After training the target model, update the parameters of the target model to Next, the parameters of the target model can be fixed as Based on the parameters The performance of the target model is used to update the policy parameters to update the policy parameters to θ b+1 ;etc.

[0035] Figure 2 An exemplary process 200 for selecting a reference model according to an embodiment of the present disclosure is shown. The process 200 may correspond to Figure 1 Step 130 in the above example. First, training sample i 202 (x i ,y i ), which can correspond to Figure 1 A set of candidate reference models 210 and a representation model 220 may also be obtained. The set of candidate reference models 210 may include K reference models, such as reference model 210-1, reference model 210-2, ..., reference model 210-K. The set of candidate reference models 210 may correspond to Figure 1 A set of candidate reference models 110 in , and the representation model 220 may correspond to Figure 1 The representation model 120 in .

[0036] The process 200 may encode the state corresponding to each reference model for the training sample i 202 as state information, and determine whether to select the reference model based on the state information. The state information may include information related to the training sample i 202 and the performance of the reference model for the training sample i 202. Taking the reference model 210-k (1≤k≤K) as an example, the state for the reference model 210-k may be represented as s jk , and will target state s jk The state information is expressed as F(s jk ). Status information F(s) jk ) can be implemented as a real-valued vector, which, for example, comprises the concatenation of three features.

[0037] The first feature can be the training sample i 202(x i ,y i ) i x can be obtained, for example, by representing the model 220 i Vector representation of Where d is the hidden size.

[0038] The second feature may be the probability distribution output by the reference model 210-k for the training sample i 202. Taking the training sample i 202 as an example of a sample for a classification task, the probability distribution output by the reference model 210-k may be expressed as in is the x output of reference model 210-k i The probability of belonging to category c, where c is an integer between 1 and C, where C is the number of categories, and are parameters of the reference model 210 - k.

[0039] The third feature may be a prediction loss corresponding to the probability distribution. In one embodiment, the prediction loss may be calculated by a cross entropy function. For example, the probability distribution for training sample i 202 output by the reference model 210-k may be calculated by the following formula: Predicting Losses

[0040]

[0041] in, is from the true annotation y i The one-hot vector of .

[0042] The probability distribution output by the reference model 210 - k for the training sample i 202 and the prediction loss corresponding to the probability distribution can be regarded as the performance of the reference model 210 - k for the training sample i 202 .

[0043] The vector representation output by the representation model 220 can be represented as Probability distribution of reference model 210-k output and probability distribution The corresponding prediction loss Cascading is performed to obtain the state information 230-k F(s) for the reference model 210-k jk ).

[0044] Through the above process, a set of state information 230 for each reference model in the set of candidate reference models 210 may be obtained, which includes, for example, state information 230 - 1 , state information 230 - 2 , . . . , state information 230 -K.

[0045] Policy function π θ 240 may determine the sampling probability 250-k of the reference model 210-k based on the state information 230-k of the reference model 210-k. For example, a logic function may be used as the strategy function, as shown in the following formula:

[0046]

[0047] in, is the state information, σ(·) is a trainable parameter The sigmoid function, and P θ (a jk |s jk ) is the sampling probability, which means that in state s jk Next select action value a jk The probability of jk ∈{0,1}. We can use P θ (a jk |s jk ) to the action value a jk Sampling is performed. jk When a is sampled as a value of "0", it indicates that the reference model 210-k is not selected; jk When sampled as a value of “1”, it indicates that the reference model 210 - k is selected.

[0048] The above process can be used to obtain a set of sampling probabilities 250 for each reference model in the group of candidate reference models 210, which include, for example, sampling probability 250-1, sampling probability 250-2, ..., sampling probability 250-K, and a set of action values ​​260, which include, for example, action value 260-1, action value 260-2, ..., action value 260-K.

[0049] After determining a set of action values ​​260, at 270, a set of selected reference models 280 may be determined based on the set of reference models 210 and the set of action values ​​260, which may include, for example, a reference model 280-1, a reference model 280-2, ..., a reference model 280-K' (0≤K'≤K). Each reference model in the set of selected reference models 280 is, for example, a reference model whose action value is sampled as "1".

[0050] It should be understood that, although the foregoing discussion and the following discussion may involve selecting at least one reference model to train the target model, it is also possible that no reference model is selected. For example, for some training samples, all reference models perform poorly on them, so the sampling probabilities of all reference models are low, and further, the action values ​​sampled according to these sampling probabilities may all be "0", which will result in no reference model being selected.

[0051] After determining a set of selected reference models 280, a set of target probability distributions output by the set of selected reference models 280 for training sample i 202 may be obtained. The set of state information 230 determined above includes target probability distributions output by each reference model in a set of candidate reference models for training sample i 202. A set of target probability distributions corresponding to each reference model in the set of selected reference models 280 may be extracted from these target probability distributions. This set of target probability distributions may be used to train a target model.

[0052] Figure 3 An exemplary process 300 for training a target model according to an embodiment of the present disclosure is shown. Process 300 can train the target model using training sample i and a set of target probability distributions output by a set of selected reference models for training sample i. Training sample i can include true annotations.

[0053] At 310 , the target model may score the training sample i to obtain a predicted probability distribution of the training sample i.

[0054] At 320, a sub-prediction loss corresponding to the target probability distribution may be calculated based on the prediction probability distribution and each target probability distribution in a set of target probability distributions to obtain a set of sub-prediction losses. In one embodiment, the sub-prediction loss may be calculated by a cross entropy function.

[0055] At 330, a first prediction loss corresponding to the training sample i may be calculated based on the number of the set of selected reference models and the set of sub-prediction losses. In one embodiment, the first prediction loss may be calculated by first summing the set of sub-prediction losses to obtain an intermediate prediction loss, and then dividing the intermediate prediction loss by the number of the set of selected reference models.

[0056] At 340, a second prediction loss corresponding to the training sample i may be calculated based on the predicted probability distribution and the true label in the training sample i. In one embodiment, the second prediction loss may be calculated by a cross entropy function.

[0057] At 350, a comprehensive prediction loss corresponding to the training sample i may be calculated based on the first prediction loss and the second prediction loss. In one embodiment, the comprehensive prediction loss may be calculated by weighted summing the first prediction loss and the second prediction loss.

[0058] At 360 , the target model may be optimized by minimizing the composite prediction loss.

[0059] Figure 4 FIG. 4 shows a specific example 400 for training a target model according to an embodiment of the present disclosure. In example 400, the training sample for training the target model may be, for example, a training sample 410 (x i ,y i ), where x i is the input, y i is the real label. A set of target probability distributions output by a set of selected reference models 420 for the training sample 410 can be obtained. To train the target model 430, where The set of selected reference models 420 includes, for example, K′ reference models numbered as reference model 420 - 1 , reference model 420 - 2 , . . . , reference model 420 -K′.

[0060] The target model 430 can score the training sample 410 to obtain the predicted probability distribution of the training sample 410. Among them, P s (y i =c|x i ; Θ s ) represents the output x of the target model 430 i The probability of belonging to category c, c is an integer between 1 and C, C is the number of categories, and Θ s are parameters of the target model 430 .

[0061] Then we can predict the probability distribution based on and a set of target probability distributions Each target probability distribution in To calculate the target probability distribution The corresponding sub-prediction loss To obtain a set of sub-prediction losses In one embodiment, the sub-prediction loss can be calculated by the cross entropy function As shown in the following formula:

[0062]

[0063] Next, a first prediction loss for training sample i may be calculated based on the number K′ of reference models in the set of selected reference models 420 and the set of sub-prediction losses: In one embodiment, the first prediction loss may be calculated by first summing the group of sub-prediction losses to obtain an intermediate prediction loss, and then dividing the intermediate prediction loss by the number K' of the group of selected reference models, as shown in the following formula:

[0064]

[0065] Then, the predicted probability distribution output by the target model can be and the true annotation y in training sample i i To calculate the second prediction loss corresponding to training sample i In one implementation, the second prediction loss may be calculated using a cross entropy function, as shown in the following formula:

[0066]

[0067] After obtaining the first prediction loss corresponding to training sample i and the second prediction loss Afterwards, the comprehensive prediction loss corresponding to training sample i can be calculated In one implementation, the comprehensive prediction loss may be calculated by weighted summing the first prediction loss and the second prediction loss, as shown in the following formula:

[0068]

[0069] Here, α is a hyperparameter used to balance the first prediction loss and the second prediction loss.

[0070] The comprehensive prediction loss can be Minimize to optimize the target model.

[0071] Combination of the above Figure 3 and Figure 4The illustrated process trains the target model by minimizing the comprehensive prediction loss corresponding to a single training sample. Alternatively, to improve training efficiency, a subset of the data, such as a subset of the data, can be trained. Perform the above process on all training samples in and obtain the data subset The corresponding comprehensive prediction loss This combined prediction loss can be Minimize to optimize the target model. In this case, the parameters of the target model remain unchanged for all training samples in the same data subset. The corresponding comprehensive prediction loss It can be obtained by the following formula:

[0072]

[0073] According to an embodiment of the present disclosure, after the target model is trained using a training sample or a data subset, the strategy function π may be updated based on the performance of the trained target model. θ The strategy parameter θ in , thereby affecting the sampling probability of each reference model for the next training sample or the next data subset. Figure 5 FIG. 5 shows an exemplary process 500 for updating policy parameters according to an embodiment of the present disclosure. The process 500 can use the validation sample (x′, y′) to evaluate the performance of the trained target model. The validation sample can be compared with the training sample (x′, y′) used to train the target model. i ,y i ) are the same or different. The evaluated performance can be converted into a reward. The reward can then be used to update the policy function π θ The policy parameter θ in .

[0074] At 510, the validation sample (x', y') may be scored by the target model to obtain a predicted probability distribution of the validation sample (x', y') where Θ s are the current parameters of the target model.

[0075] At 520, the true annotation y′ in the validation sample (x′, y′) and the predicted probability distribution can be used to obtain the predicted value. To calculate the prediction loss corresponding to the validation sample In one implementation, the second prediction loss may be calculated using a cross entropy function, as shown in the following formula:

[0076]

[0077] At 530, based on the predicted loss To calculate the reward v corresponding to the validation samplej In one embodiment, the reward v j Calculated as prediction loss The opposite of , as shown in the following formula:

[0078]

[0079] At 540, the reward v j To update the policy parameter θ. In one embodiment, the policy parameter θ can be updated by a standard policy gradient method, such as a Monte-Carlo-based policy gradient method, as shown in the following formula:

[0080]

[0081] where β is the learning rate and π θ (s jk , a jk ) is the policy function for the kth reference model.

[0082] According to an embodiment of the present disclosure, before determining the sampling probability of each reference model through a policy function, the policy parameters in the policy function may be initialized. For example, at least one reference model may be selected from a set of candidate reference models through a policy function, and the policy parameters in the policy function may be initialized according to the average performance of the selected reference model.

[0083] Figure 6 An exemplary process 600 for initializing policy parameters according to an embodiment of the present disclosure is shown. Process 600 can use an initialization sample to initialize the policy parameters. In the text, the sample used to initialize the policy parameters is referred to as an initialization sample, which can be the same as or different from the training sample used to train the target model. The initialization sample can include an input and a true annotation corresponding to the input.

[0084] At 610, a set of selected reference models selected from a set of candidate reference models may be determined by a strategy function for the initialization sample. The strategy function may have an original strategy parameter.

[0085] At 620 , the initialization samples may be scored respectively by the set of selected reference models to obtain a set of probability distributions of the initialization samples.

[0086] At 630, a set of prediction losses corresponding to the initialization sample may be calculated based on the true annotation in the initialization sample and the set of probability distributions. For example, a sub-prediction loss for each probability distribution may be calculated based on the true annotation and each probability distribution in the set of probability distributions to obtain a set of sub-prediction losses. In one embodiment, the prediction loss for each probability distribution may be calculated by a cross entropy function.

[0087] At 640, a prediction loss corresponding to the initialization sample may be calculated based on the number of a set of candidate reference models and the set of sub-prediction losses. In one embodiment, the prediction loss may be calculated by first summing the set of sub-prediction losses to obtain an intermediate prediction loss, and then dividing the intermediate prediction loss by the number of the set of selected reference models.

[0088] At 650, a reward corresponding to the initialization sample may be calculated based on the prediction loss. In one embodiment, the reward may be calculated as the inverse of the prediction loss.

[0089] At 660, the policy parameters may be initialized based on the reward. In one embodiment, the policy parameters may be initialized by updating the original policy parameters using a standard policy gradient method, as shown in equation (10) above.

[0090] According to an embodiment of the present disclosure, before using a set of candidate reference models to train a target model, the target model may be pre-trained. In one embodiment, all reference models in the set of candidate reference models may be used to score a pre-training data set, and the scored pre-training data set may be used to pre-train the target model. In this document, the data set used to pre-train the target model is referred to as a pre-training data set. The process for pre-training the target model may be combined with Figure 3 and Figure 4 The process explained for training the target model is similar, except that the set of target probability distributions involved is the probability distribution output by all reference models in the set of candidate reference models rather than a selected reference model in the set of candidate reference models for the pre-training samples in the pre-training data set.

[0091] Figure 7 is a flowchart of an exemplary method 700 for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure.

[0092] At step 710 , a data set and a set of candidate reference models may be obtained.

[0093] At step 720, a set of selected reference models selected from the set of candidate reference models may be determined for each training sample in the data set.

[0094] At step 730, a set of target probability distributions output by the set of selected reference models for the training samples may be obtained.

[0095] At step 740, the target model may be trained using the set of target probability distributions.

[0096] In one embodiment, determining a set of selected reference models may include: for each candidate reference model in the set of candidate reference models, determining whether to select the candidate reference model by reinforcement learning.

[0097] The determining whether to select the candidate reference model may include: determining a sampling probability of the candidate reference model based on a policy function; sampling an action value of the candidate reference model based on the sampling probability; and selecting the candidate reference model based on the sampled action value.

[0098] The strategy function may have a strategy parameter. The determining whether to select the candidate reference model may also include: updating the strategy parameter based on the performance of the target model.

[0099] Updating the strategy parameters may include: scoring the verification sample through the target model to obtain the predicted probability distribution of the verification sample; calculating the prediction loss corresponding to the verification sample based on the true annotation in the verification sample and the predicted probability distribution; calculating the reward corresponding to the verification sample based on the prediction loss; and updating the strategy parameters based on the reward.

[0100] The data set may include multiple data subsets. For all training samples in the same data subset, the policy parameters of the policy function may remain unchanged.

[0101] The determining of the sampling probability may be performed on state information. The state information may include at least: a representation of the training sample, a probability distribution output by the candidate reference model for the training sample, and a prediction loss corresponding to the probability distribution.

[0102] The strategy function may have strategy parameters. The strategy parameters may be initialized by the following operations: for the initialization sample, determining a set of selected reference models selected from the set of candidate reference models by the strategy function; scoring the initialization samples respectively by the set of selected reference models to obtain a set of probability distributions of the initialization samples; calculating the prediction loss corresponding to the initialization sample based on the true annotations in the initialization sample and the set of probability distributions; calculating the reward corresponding to the initialization sample based on the prediction loss; and initializing the strategy parameters based on the reward.

[0103] In one embodiment, the training sample may include a true annotation. The training of the target model may include: scoring the training sample by the target model to obtain a predicted probability distribution of the training sample; calculating a first prediction loss corresponding to the training sample based on the predicted probability distribution and the set of target probability distributions; calculating a second prediction loss corresponding to the training sample based on the predicted probability distribution and the true annotation; calculating a comprehensive prediction loss corresponding to the training sample based on the first prediction loss and the second prediction loss; and optimizing the target model by minimizing the comprehensive prediction loss.

[0104] The calculating of the first prediction loss may include: respectively calculating the sub-prediction losses corresponding to the target probability distribution based on the prediction probability distribution and each target probability distribution in the set of target probability distributions to obtain a set of sub-prediction losses; and calculating the first prediction loss based on the number of reference models in the set of selected reference models and the set of sub-prediction losses.

[0105] The data set may include multiple data subsets. For all training samples in the same data subset, the parameters of the target model may remain unchanged.

[0106] In one embodiment, method 700 may further include: scoring the pre-training data set using the set of candidate reference models; and pre-training the target model using the scored pre-training data set.

[0107] In one embodiment, the set of candidate reference models may be models having higher complexity than the target model.

[0108] It should be understood that method 700 may also include any steps / processes for obtaining a target model based on knowledge distillation according to the above-mentioned embodiments of the present disclosure.

[0109] Figure 8An exemplary device 800 for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure is shown. The device 800 may include: an obtaining module 810 for obtaining a data set and a set of candidate reference models; a reference model determination module 820 for determining, for each training sample in the data set, a set of selected reference models selected from the set of candidate reference models; a probability distribution acquisition module 830 for obtaining a set of target probability distributions output by the set of selected reference models for the training samples; and a target model training module 840 for training the target model using the set of target probability distributions.

[0110] In one implementation, the reference model determination module 820 may also be configured to: for each candidate reference model in the set of candidate reference models, determine whether to select the candidate reference model through reinforcement learning.

[0111] The determining whether to select the candidate reference model may include: determining a sampling probability of the candidate reference model based on a policy function; sampling an action value of the candidate reference model based on the sampling probability; and selecting the candidate reference model based on the sampled action value.

[0112] The strategy function may have a strategy parameter. The determining whether to select the candidate reference model may also include: updating the strategy parameter based on the performance of the target model.

[0113] The data set may include multiple data subsets. For all training samples in the same data subset, the policy parameters of the policy function may remain unchanged.

[0114] The determining of the sampling probability may be performed on state information. The state information may include at least: a representation of the training sample, a probability distribution output by the candidate reference model for the training sample, and a prediction loss corresponding to the probability distribution.

[0115] It should be understood that the apparatus 800 may also include any other modules configured to obtain a target model based on knowledge distillation according to the above-mentioned embodiments of the present disclosure.

[0116] Fig. 9 An exemplary apparatus 900 for obtaining a target model based on knowledge distillation according to an embodiment of the present disclosure is shown.

[0117] The apparatus 900 may include at least one processor 910. The apparatus 900 may also include a memory 920 connected to the processor 910. The memory 920 may store computer executable instructions, which, when executed, cause the processor 1910 to perform any operation of the method for obtaining a target model based on knowledge distillation according to the above-mentioned embodiments of the present disclosure.

[0118] Embodiments of the present disclosure may be embodied in a non-transitory computer-readable medium. The non-transitory computer-readable medium may include instructions that, when executed, cause one or more processors to perform any operation of the method for obtaining a target model based on knowledge distillation according to the embodiments of the present disclosure as described above.

[0119] It should be appreciated that all operations in the method described above are merely exemplary, and the present disclosure is not limited to any operations in the method or the order of these operations, but should cover all other equivalent changes under the same or similar concept.

[0120] It should also be appreciated that all modules in the above described device can be implemented in various ways. These modules can be implemented as hardware, software, or a combination thereof. In addition, any module in these modules can be further divided into submodules or combined together in function.

[0121] Processors have been described in conjunction with various devices and methods. These processors can be implemented using electronic hardware, computer software or any combination thereof. Whether these processors are implemented as hardware or software will depend on specific application and the overall design constraints imposed on the system. As an example, the processor provided in the present disclosure, any part of the processor or any combination of processors can be realized using a microprocessor, a microcontroller, a digital signal processor (DSP), a field programmable gate array (FPGA), a programmable logic device (PLD), a state machine, a gated logic unit, a discrete hardware circuit, and other suitable processing components configured to perform the various functions described in the present disclosure. The function of the processor provided in the present disclosure, any part of the processor or any combination of processors can be realized using software executed by a microprocessor, a microcontroller, a DSP or other suitable platforms.

[0122] Software should be broadly considered to mean instructions, instruction sets, codes, code segments, program codes, programs, subroutines, software modules, applications, software applications, software packages, routines, subroutines, objects, running threads, processes, functions, etc. Software can reside in a computer-readable medium. A computer-readable medium can include, for example, a memory, which can be, for example, a magnetic storage device (e.g., a hard disk, a floppy disk, a magnetic stripe), an optical disk, a smart card, a flash memory device, a random access memory (RAM), a read-only memory (ROM), a programmable ROM (PROM), an erasable PROM (EPROM), an electrically erasable PROM (EEPROM), a register, or a removable disk. Although the memory is shown as being separated from the processor in the multiple aspects provided in the present disclosure, the memory can also be located inside the processor, such as a cache or a register.

[0123] The above description is provided to enable any person skilled in the art to practice the various aspects described herein. Various modifications to these aspects will be apparent to those skilled in the art, and the general principles defined herein may be applied to other aspects. Therefore, the claims are not intended to be limited to the aspects shown herein. All structural and functional equivalents to the elements of the various aspects described in this disclosure that are known or to be known to those of ordinary skill in the art are expressly incorporated herein and covered by the claims.

Claims

1. A method for obtaining a target model based on knowledge distillation, wherein the target model has a simple structure, the method comprising: Obtaining a data set and a set of candidate reference models, wherein the set of candidate reference models are models with higher complexity than the target model; For each training sample in the data set, determine a set of selected reference models selected from the set of candidate reference models, wherein determining a set of selected reference models comprises: for each candidate reference model in the set of candidate reference models, determining whether to select the candidate reference model by the following operations: Determining a sampling probability of the candidate reference model based on a policy function, wherein the determining of the sampling probability is performed on state information, the state information comprising at least: a representation of the training sample, a probability distribution output by the candidate reference model for the training sample, and a prediction loss corresponding to the probability distribution; Based on the sampling probability, sampling the action value of the candidate reference model; and Selecting the candidate reference model based on the sampled action value; Obtaining a set of target probability distributions output by the set of selected reference models for the training samples; and The target model is trained using the set of target probability distributions.

2. The method according to claim 1, wherein: The strategy function has a strategy parameter, and the determining whether to select the candidate reference model further comprises: The policy parameters are updated based on the performance of the target model.

3. The method according to claim 2, wherein: The updating of the policy parameters comprises: Scoring the validation samples by using the target model to obtain a predicted probability distribution of the validation samples; Calculating a prediction loss corresponding to the validation sample based on the true annotation in the validation sample and the predicted probability distribution; Calculating a reward corresponding to the validation sample based on the prediction loss; and The policy parameters are updated based on the reward.

4. The method according to claim 1, wherein: The data set includes multiple data subsets, and for all training samples in the same data subset, the policy parameters of the policy function remain unchanged.

5. The method according to claim 1, wherein: The policy function has policy parameters, and the policy parameters are initialized by the following operations: For the initialization sample, determining a set of selected reference models selected from the set of candidate reference models by using the strategy function; Scoring the initialization samples respectively by using the set of selected reference models to obtain a set of probability distributions of the initialization samples; Calculating a prediction loss corresponding to the initialization sample based on the true annotation in the initialization sample and the set of probability distributions; Calculating a reward corresponding to the initialization sample based on the prediction loss; as well as The policy parameters are initialized based on the reward.

6. The method according to claim 1, wherein: The training samples include real annotations, and the training of the target model includes: Scoring the training samples by using the target model to obtain a predicted probability distribution of the training samples; Calculating a first prediction loss corresponding to the training sample based on the prediction probability distribution and the set of target probability distributions; Calculating a second prediction loss corresponding to the training sample based on the predicted probability distribution and the true annotation; Calculating a comprehensive prediction loss corresponding to the training sample based on the first prediction loss and the second prediction loss; and The target model is optimized by minimizing the comprehensive prediction loss.

7. The method according to claim 6, wherein: The calculating the first prediction loss comprises: respectively calculating sub-prediction losses corresponding to the target probability distribution based on the prediction probability distribution and each target probability distribution in the set of target probability distributions to obtain a set of sub-prediction losses; and The first prediction loss is calculated based on the number of reference models in the set of selected reference models and the set of sub-prediction losses.

8. The method according to claim 6, wherein: The data set includes multiple data subsets, and for all training samples in the same data subset, the parameters of the target model remain unchanged.

9. The method according to claim 1, further comprising: Scoring the pre-training data set by the set of candidate reference models; as well as The target model is pre-trained using the scored pre-training dataset.

10. A device for obtaining a target model based on knowledge distillation, wherein the target model has a simple structure, the device comprising: An obtaining module, configured to obtain a data set and a set of candidate reference models, wherein the set of candidate reference models are models having a higher complexity than the target model; A reference model determination module is used to determine, for each training sample in the data set, a set of selected reference models selected from the set of candidate reference models, wherein determining a set of selected reference models comprises: for each candidate reference model in the set of candidate reference models, determining whether to select the candidate reference model by the following operations: Determining a sampling probability of the candidate reference model based on a policy function, wherein the determining of the sampling probability is performed with respect to state information, the state information comprising at least: a representation of the training sample, a probability distribution output by the candidate reference model with respect to the training sample, and a prediction loss corresponding to the probability distribution; Based on the sampling probability, sampling the action value of the candidate reference model; and Selecting the candidate reference model based on the sampled action value; A probability distribution acquisition module, used to acquire a set of target probability distributions output by the set of selected reference models for the training samples; and A target model training module is used to train the target model using the set of target probability distributions.

11. The device according to claim 10, wherein: The strategy function has a strategy parameter, and the determining whether to select the candidate reference model further comprises: The policy parameters are updated based on the performance of the target model.

12. The device according to claim 10, wherein: The data set includes multiple data subsets, and for all training samples in the same data subset, the policy parameters of the policy function remain unchanged.

13. A device for obtaining a target model based on knowledge distillation, wherein the target model has a simple structure, the device comprising: at least one processor; as well as a memory storing computer executable instructions that, when executed, cause the at least one processor to: obtaining a data set and a set of candidate reference models, wherein the set of candidate reference models are models with higher complexity than the target model, For each training sample in the data set, determine a set of selected reference models selected from the set of candidate reference models, wherein determining a set of selected reference models comprises: for each candidate reference model in the set of candidate reference models, determining whether to select the candidate reference model by the following operations: The sampling probability of the candidate reference model is determined based on the policy function, wherein the determining of the sampling probability is performed with respect to state information, and the state information at least includes: a representation of the training sample, a probability distribution output by the candidate reference model for the training sample, and a prediction loss corresponding to the probability distribution; Based on the sampling probability, sampling the action value of the candidate reference model; and Based on the sampled action values, selecting the candidate reference model, Obtaining a set of target probability distributions output by the set of selected reference models for the training samples, and The target model is trained using the set of target probability distributions.

Citation Information

Cited By

  • Difficulty sample screening and distilling method based on speculative sampling inference difference

    CN122047388A