Grammar error correction model training method and system based on adversarial neural network
By using an adversarial neural network with CycleGan structure in syntax error correction technology, the game between the generator and the discriminator is trained, and the problems of low data enhancement quality and model complexity in the existing technology are solved, achieving efficient and low-cost syntax error correction effect.
Patent Information
- Application Number
- CN202311533473.3
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2023-11-16
- Publication Date
- 2025-05-16
AI Technical Summary
In the existing syntax error correction technology, the data enhancement method has low quality and high cost, the model architecture is complex and the inference speed is slow, making it difficult to effectively solve the problems of insufficient data and inefficiency in syntax error correction tasks.
Using the CycleGan structure based on an adversarial neural network, two generators and two discriminators are trained. Through the game between the generator and the discriminator, the generator generation quality is improved, the end-to-end training of the syntax error correction model is realized, and the cost of manual labeling is reduced.
It improves the generation quality and efficiency of the syntax error correction model, reduces labor costs, achieves higher-quality syntax error correction effects, and reduces the complexity of model training.
Smart Images

Figure CN120012768A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of adversarial neural networks and model training, and in particular to a grammatical error correction model training method and system based on adversarial neural networks. Background Art
[0002] With the widespread adoption of the internet, more and more information is being disseminated online. Statistics show that the amount of data generated by users on Facebook, Twitter, WhatsApp, YouTube, and email is staggering. Every day, 720,000 hours of new YouTube videos, 65 billion WhatsApp messages, 500 million Twitter tweets, and 40 million new Facebook posts are generated. In 2018, the world's total data volume reached 33 zettabytes (ZB), equivalent to 33 trillion GB. In 2020, this figure reached nearly 59 ZB, and is projected to increase to approximately 175 ZB by 2025—an incredible number. A significant portion of this vast amount of information consists of text, which is often riddled with errors. This erroneous content not only significantly impacts the user experience but also poses challenges to various natural language processing tasks. Consequently, grammatical correction has emerged as a hot research area in natural language processing.
[0003] Grammatical Error Collection (GEC), as the name suggests, is to modify a grammatically incorrect sentence into a grammatically correct one. Grammatical error collection technology can not only help language learners and writers automatically diagnose and correct grammatical errors, thereby improving work and learning efficiency, but is also an important prerequisite for many natural language processing tasks. For example, in upper-level applications such as information retrieval, voice interaction, and machine translation, pre-processing operations such as grammatical correction of input text can effectively improve application performance. For example, in information retrieval, the text entered by the user will be checked in the search engine to improve the accuracy of the final result; in voice interaction, the user's voice will be converted into text, and corresponding grammatical correction and error correction will be performed to improve the accuracy of intent recognition and interaction; in machine translation, grammatical correction will be performed after the source language is translated to ensure translation quality.
[0004] The Transformer model is a deep learning model that uses an attention mechanism to speed up model training. The Transformer consists of two parts: an encoder and a decoder. The encoder consists of six encoding blocks, each of which has two sublayers: a self-attention layer and a feedforward neural network layer. These layers share the same structure but different parameters. The decoder also consists of six decoding blocks, each of which has three sublayers: a self-attention layer, a feedforward neural network layer, and an encoder-decoder attention layer. This layer helps the decoder focus on relevant parts of the input sentence. These layers also share the same structure but different parameters. The self-attention mechanism is a core component of the Transformer model. It captures dependencies between different positions in the input sequence by calculating the correlation between each element and every other element.
[0005] Transfer learning is a machine learning method that applies existing mature models to new fields. The two fields may solve different problems, but they have certain similarities. Transfer learning can be divided into several different types, including model-based, instance-based, and feature-based. These focus on solving different types of problems, and the specific method used depends on the specific application scenario. Instance-based transfer learning methods focus on selecting instances from the source field that are useful to the target field; feature-based transfer learning methods mainly hope to represent the features between the source and target fields in the same way; model-based transfer learning methods focus on sharing models between the two fields, specifically sharing the prior distribution or parameters of the model.
[0006] CycleGan is a generative model that can generate new data samples based on given data. The main purpose of CycleGan is to achieve the conversion between different domain data. For example, there are two datasets X and Y that store pictures of different styles. CycleGan hopes to train a generator G, which takes an input x and outputs a y′, that is, G(x) = y ′ ,x∈X; At the same time, we also need to train a generator F, whose input y gets x′, that is, F(y)=x ′ ,y∈Y. In order to achieve this goal, it is necessary to train two discriminators D X ,D Y , used to judge the quality of the generated image. If the generator generates x ′If ,y′ differs significantly from x,y in dataset X,Y, it should be given a low score; otherwise, it should be given a high score. Furthermore, the discriminator should always give high scores to x,y in dataset X,Y. Grammatical correct and grammatically incorrect sentences can also be viewed as two different styles. This method uses CycleGan to train a generator, which can generate grammatically incorrect sentences, and the generated sentences will continuously converge to real grammatically incorrect sentences. Therefore, it can be used for data augmentation, and training the GEC model with an expanded corpus can achieve better results.
[0007] The current mainstream research methods for grammatical error correction are divided into two categories: (1) data enhancement and (2) model structure improvement.
[0008] Currently, grammatical error correction is primarily treated as a specialized machine translation task, treated as a translation task between the same language and trained using a Seq2Seq (sequence-to-sequence) network architecture. Clearly, a generative model can be trained using a large number of incorrect-correct sentence pairs to automatically correct grammatical errors. However, generative models typically require a relatively large parallel corpus. Compared to the large-scale grammatical corpora used for machine translation, the corpus size for grammatical error correction is typically smaller, only a few hundred thousand sentences. Therefore, scientific data augmentation has become a research area. Some existing methods randomly add errors to sentences to construct data or utilize large-scale manual annotation. However, the data generated by randomly adding errors differs significantly from real data and is of lower quality. Large-scale manual annotation is also prohibitively expensive. Automatically generating high-quality data presents a major challenge.
[0009] The source and target sentences often have only minor differences in grammatical correction, so using a sequence-to-sequence model to generate the entire target sentence from scratch is not the best option. To address this issue, some methods have built model structures specifically for this task. These methods modify the Seq2Seq (sequence-to-sequence) model architecture to a Seq2Edit (sequence-to-edit) architecture. The model still accepts text input, but instead of directly outputting the corrected text, it outputs the location and type of error to be corrected, transforming the generation problem into a classification problem. This approach can run in parallel during inference, resulting in improved speed compared to sequence-to-sequence frameworks. However, this approach typically requires multiple rounds of correction to modify sentences based on error type, as sentences may contain more than one error. Multiple rounds of correction can easily introduce new grammatical errors, resulting in new errors due to inappropriate modifications to previously correct locations. This is a challenge faced by current sequence-to-edit architectures.
[0010] Patent document CN114239557A discloses a grammatical error correction method and training method, apparatus, electronic device, and storage medium. The training method includes: respectively obtaining a first training corpus and a second monolingual corpus containing annotation information, wherein the annotation information is used to characterize grammatical error pairs of each training corpus in the first training corpus, wherein the grammatical error pairs include a source segment in an incorrect form and a target segment in a correct form corresponding to the source segment in the incorrect form; extracting grammatical error pairs from each training corpus in the first training corpus to construct a grammatical error pair reference set; based on the grammatical error pair reference set, corrupting the second monolingual corpus to obtain a pseudo-error corpus corresponding to the second monolingual corpus; inputting the pseudo-error corpus and the first training corpus into a preset neural network model, training the preset neural network model, and obtaining a grammatical error correction model.
[0011] Current data augmentation methods randomly add erroneous constructed data, which differs significantly from real data and results in poorly trained models. Manual labeling is also costly and impractical. Sequence-to-sequence model architectures require high data volumes and suffer from slow inference speeds. Sequence-to-modification frameworks address both of these issues to some extent, but also require more complex model design. Summary of the Invention
[0012] In view of the defects in the prior art, the purpose of the present invention is to provide a grammatical error correction model training method and system based on adversarial neural network.
[0013] The grammatical error correction model training method based on the adversarial neural network provided by the present invention includes:
[0014] Step 1: Use the pre-trained model as a generator to predict the word with the highest probability at each position to form a sentence, which is used for subsequent discriminator judgment and input into another generator for further generation;
[0015] Step 2: Build a discriminator to determine whether the input sentence is generated by the generator or real data;
[0016] Step 3: Based on two generators and two discriminators, build a CycleGan network to realize the game between the generator and the discriminator;
[0017] Step 4: Calculate the cross entropy loss and use it for back propagation to update the neural network parameters.
[0018] Preferably, the step 1 comprises:
[0019] The T5 model is selected as the generator. The sentences input into the T5 model are first processed by the word segmenter, and then the T5 model uniformly fills and truncates them, and generates the corresponding identity identification code of the input sentence;
[0020] Generate an attention mask matrix so that the T5 model only focuses on the sentence itself when performing attention calculations;
[0021] Perform tokenizer processing, padding, and truncation on the target sentence, changing all values in the target sentence to the specified padding token pad_token_id to -100;
[0022] The identity code, attention mask, and sentence label of the input sentence are input into the T5 model for forward propagation to generate the cross entropy loss between the input sentence and the target sentence, which is used for backpropagation to update the neural network parameters and the score of each word in the vocabulary at each position. After Softmax, it is the probability of each word appearing at this position.
[0023] Preferably, the step 2 includes:
[0024] The sentence input to the discriminator is first encoded to extract features and then judged;
[0025] Use the T5 encoder to encode the sentence. First, the sentence is processed by the word segmenter to obtain the corresponding input sentence identity code and attention mask, and then input it into the T5 encoder. The T5 encoder outputs a hidden layer encoding vector for each word in the sentence.
[0026] Construct a linear layer to predict the score based on the input sentence encoding vector and convert the predicted score into the corresponding probability. The probability will be calculated with the sentence label to get the loss. A Sigmoid layer is connected after the linear layer to convert the score into a probability between 0 and 1.
[0027] Preferably, the step 3 includes:
[0028] Let X be a grammatically incorrect sentence and Y be a grammatically correct sentence, then the generator expression is:
[0029] G t (x) = y ′ ,x∈X#(1)
[0030] G f (y) = x ′ ,y∈Y#(2)
[0031] Among them, the generator G t Input grammatically incorrect sentence x, generate grammatically correct sentence y ′ ; Generator G fInput grammatically correct sentence y, output grammatically incorrect sentence x ′ ;
[0032] y ′ After correcting the error again, we can restore the original sentence x, which is expressed as:
[0033] G f (G t (x))=x,x∈X#(3)
[0034] G t (G f (y))=y,y∈Y#(4)
[0035] The two discriminators are D t and D f , D t Used to determine whether a sentence is composed of G t Generated or real, D f Used to determine whether a sentence is composed of G f Real or generated; Discriminator D t G t The corrected sentence is judged as 0, and the true grammatically correct sentence y, y∈Y is judged as 1.
[0036] Preferably, step 4 includes:
[0037] When training the discriminator, G t ,G f The parameters are fixed and only D is updated t ,D f Parameters; when training the generator, D t ,D f The parameters of G are fixed and only G is updated t ,G f Parameters;
[0038] For the cross entropy loss of the generator, first G t Correct the grammatically incorrect sentence x, and calculate the loss between the corrected sentence and the target sentence. This loss is called Loss forward ; At the same time, the corrected sentence needs to make the discriminator D t It is judged as real data, and the loss is called Loss gan ; G f The grammatically incorrect sentences generated by G t The correction is consistent with the original sentence, and the loss is called Loss cycle ; Finally, for the original grammatical sentence, G t Without modification, the loss is called Loss identity ; Overall generator Gt Loss gen The calculation expression is:
[0039] Loss gen =Loss forward +Loss gan +Loss cycle +Loss identity #(5)
[0040]
[0041]
[0042]
[0043]
[0044] When training the discriminator, the discriminator classifies the sentences generated by the generator as 0 and the real sentences as 1. Its Loss dis The calculation expression is:
[0045]
[0046] After the loss is calculated, it is back-propagated to update the parameters of the neural network.
[0047] The grammatical error correction model training system based on the adversarial neural network provided by the present invention includes:
[0048] Module M1: Use the pre-trained model as a generator to predict the word with the highest probability at each position to form a sentence, which is used for subsequent discriminator judgment and input into another generator for further generation;
[0049] Module M2: Build a discriminator to determine whether the input sentence is generated by the generator or real data;
[0050] Module M3: Based on two generators and two discriminators, build a CycleGan network to realize the game between the generator and the discriminator;
[0051] Module M4: Calculates cross entropy loss for back propagation to update neural network parameters.
[0052] Preferably, the module M1 includes:
[0053] The T5 model is selected as the generator. The sentences input into the T5 model are first processed by the word segmenter, and then the T5 model uniformly fills and truncates them, and generates the corresponding identity identification code of the input sentence;
[0054] Generate an attention mask matrix so that the T5 model only focuses on the sentence itself when performing attention calculations;
[0055] Perform tokenizer processing, padding, and truncation on the target sentence, changing all values in the target sentence to the specified padding token pad_token_id to -100;
[0056] The identity code, attention mask, and sentence label of the input sentence are input into the T5 model for forward propagation to generate the cross entropy loss between the input sentence and the target sentence, which is used for backpropagation to update the neural network parameters and the score of each word in the vocabulary at each position. After Softmax, it is the probability of each word appearing at this position.
[0057] Preferably, the module M2 includes:
[0058] The sentence input to the discriminator is first encoded to extract features and then judged;
[0059] Use the T5 encoder to encode the sentence. First, the sentence is processed by the word segmenter to obtain the corresponding input sentence identity code and attention mask, and then input it into the T5 encoder. The T5 encoder outputs a hidden layer encoding vector for each word in the sentence.
[0060] Construct a linear layer to predict the score based on the input sentence encoding vector and convert the predicted score into the corresponding probability. The probability will be calculated with the sentence label to get the loss. A Sigmoid layer is connected after the linear layer to convert the score into a probability between 0 and 1.
[0061] Preferably, the module M3 includes:
[0062] Let X be a grammatically incorrect sentence and Y be a grammatically correct sentence, then the generator expression is:
[0063] G t (x) = y ′ ,x∈X#(1)
[0064] G f (y) = x ′ ,y∈Y#(2)
[0065] Among them, the generator G t Input grammatically incorrect sentence x, generate grammatically correct sentence y ′ ; Generator G f Input grammatically correct sentence y, output grammatically incorrect sentence x ′ ;
[0066] y ′ After correcting the error again, we can restore the original sentence x, which is expressed as:
[0067] G f (G t (x))=x,x∈X#(3)
[0068] G t (G f (y))=y,y∈Y#(4)
[0069] The two discriminators are D t and D f , D t Used to determine whether a sentence is composed of G t Generated or real, D f Used to determine whether a sentence is composed of G f Real or generated; Discriminator D t G t The corrected sentence is judged as 0, and the true grammatically correct sentence y, y∈Y is judged as 1.
[0070] Preferably, the module M4 includes:
[0071] When training the discriminator, G t ,G f The parameters are fixed and only D is updated t ,D f Parameters; when training the generator, D t ,D f The parameters of G are fixed and only G is updated t ,G f Parameters;
[0072] For the cross entropy loss of the generator, first G t Correct the grammatically incorrect sentence x, and calculate the loss between the corrected sentence and the target sentence. This loss is called Loss forward ; At the same time, the corrected sentence needs to make the discriminator D t It is judged as real data, and the loss is called Loss gan ; G f The grammatically incorrect sentences generated by G t The correction is consistent with the original sentence, and the loss is called Loss cycle ; Finally, for the original grammatical sentence, G t Without modification, the loss is called Loss identity ; Overall generator G t Loss gen The calculation expression is:
[0073] Loss gen =Loss forward +Lossgan +Loss cycle +Loss identity #(5)
[0074]
[0075]
[0076]
[0077]
[0078] When training the discriminator, the discriminator classifies the sentences generated by the generator as 0 and the real sentences as 1. Its Loss dis The calculation expression is:
[0079]
[0080] After the loss is calculated, it is back-propagated to update the parameters of the neural network.
[0081] Compared with the prior art, the present invention has the following beneficial effects:
[0082] (1) The present invention uses the CycleGan structure to train the corresponding generator and discriminator. The game between the two is the key to continuously improving the quality of the generator. One generator continuously generates higher-quality error sentences to expand the database, and the other generator continuously improves its error correction performance, directly improving the effect of grammatical error correction. Moreover, the generator and discriminator trained by this method are both end-to-end models. The model extracts features and generates data by itself, and does not require manual feature extraction. It only needs to provide the model with standard grammatical error-correct sentence pairs. The model can learn the difference between the two sentences and automatically convert them. It does not require manual labeling, which greatly reduces labor costs and can be more widely used in actual scenarios.
[0083] (2) The present invention uses the CycleGan structure and trains two generators. One generator changes grammatically correct sentences into grammatically incorrect sentences for data augmentation, while the other generator changes grammatically incorrect sentences into correct sentences. This is the grammatical error correction model we need. Moreover, under the action of the discriminator, the grammatical error correction model not only corrects grammatical errors, but also makes the generated sentences closer to real human expressions.
[0084] (3) The grammatical error correction model trained by the present invention directly generates grammatically correct sentences based on grammatically incorrect sentences rather than predicting error correction operations, so only one error correction process is required. Moreover, due to the specificity of the Seq2Seq architecture, the model is more inclined to generate sentences that are identical to the original ones, which means that there is a smaller chance of introducing new grammatical errors. Under the existing training, the error correction effect of the model has been verified, and its own structure can ensure that no new errors are introduced. The grammatical error correction model trained by the present invention has more outstanding performance. BRIEF DESCRIPTION OF THE DRAWINGS
[0085] Other features, objects and advantages of the present invention will become more apparent upon reading the detailed description of non-limiting embodiments with reference to the following drawings:
[0086] Figure 1 This is a system structure diagram of the present invention;
[0087] Figure 2 Schematic diagram of the discriminator structure. DETAILED DESCRIPTION
[0088] The present invention will be described in detail below with reference to specific embodiments. The following examples will help those skilled in the art to further understand the present invention, but are not intended to limit the present invention in any form. It should be noted that, for those skilled in the art, several changes and improvements can be made without departing from the scope of the present invention. These all fall within the scope of protection of the present invention.
[0089] Example 1
[0090] The present invention provides a grammatical error correction model training method based on adversarial neural network, such as Figure 1 , this method designs four components, two generators and two discriminators, which are named G t ,G f ,D t ,D f .G t Input grammatically incorrect sentences and generate grammatically correct sentences; G f Input a grammatically correct sentence and output a grammatically incorrect sentence; D t Used to determine whether a sentence is composed of G t Generated or real; D f Used to determine whether a sentence is composed of G f Real or generated. The generator hopes that the sentences it generates can be scored high by the discriminator, while the discriminator scores low scores for the sentences generated by the generator and high scores for the real sentences. The constant game between the discriminator and the generator will lead to a generator with better generation quality. This method also hopes that G t ,G fIt is only used to correct errors and randomly generate errors, and does not change the meaning of the original sentence. Therefore, the generator G t The generated sentence should also be input into G again f In the generator G, the regenerated sentence should be consistent with the original sentence. f The same is true for generated sentences. In grammatical error correction tasks, the input may be a completely error-free sentence. This means that in certain situations, the generator should output the input directly without making any modifications. This is also incorporated into the training process to improve results.
[0091] The generator's function is to transform one sentence into another, which conforms to the input and output of the Seq2Seq structure. Therefore, this method uses a Seq2Seq model to build the generator. However, the grammatical error correction corpus is relatively small, making it impractical to train a Seq2Seq generator from scratch. However, large pre-trained models are widely used in various natural language processing tasks and have achieved promising results. Therefore, it is natural to apply pre-trained models to this method. Therefore, this method uses the T5 model as the generator.
[0092] 1. Generator module design
[0093] This method uses a pre-trained model as the generator. Current mainstream pre-trained models include Bert, Gpt-2, and T5, which perform well in various natural language processing tasks. However, Bert and Gpt-2 only contain an encoder and decoder, respectively. If you want to generate text using a Seq2Seq network architecture, you need to train the corresponding decoder and encoder yourself. The small corpus with grammatical error correction poses a challenge to training the decoder and encoder from scratch. The T5 model is a complete Transformer structure, including an encoder and decoder, which can effectively extract semantic information and generate text. Therefore, this method uses the T5 model as the generator.
[0094] T5-base has approximately 220 million parameters, including a 12-layer encoder and a 12-layer decoder. Pre-trained on the C4 dataset, it already demonstrates excellent semantic extraction and text generation capabilities. Huggingface has released the Transformers library, which integrates thousands of pre-trained models for various tasks. Developers can select a model for training and fine-tuning based on their needs, including the T5-base model required for this method.
[0095] Sentences input to the T5 model must first be processed by the tokenizer. The T5Tokenizer uniformly pads and truncates a batch of data and generates corresponding input_ids. Padding a sentence simply facilitates computation; the padding does not contribute to understanding the meaning of the sentence. To ensure the model focuses solely on the sentence itself and not the padded parts, this method generates an attention mask matrix. This matrix allows the T5 model to focus solely on the sentence during attention calculations. The target sentence also undergoes tokenization and padding and truncation. However, the T5Model input does not include the attention mask for the label sentence. Therefore, this method changes all values in the target sentence to -100 for the designated padding token. Smaller values reduce the model's focus on that position. Therefore, all inputs are constructed. The input_ids, attention mask, and label are fed into the T5Model for forward propagation, generating the loss and logits. The loss is the cross-entropy loss between the target sentence and the target sentence. This loss can be used for backpropagation to update neural network parameters and is also part of the final loss. Logits are the scores of each word in the vocabulary at each position. After softmax, they are the probability of each word appearing at that position. This method selects the word with the highest probability at each position to form a sentence. This sentence is the sentence predicted by the T5Model and can be used for subsequent discriminator evaluation and input into another generator for further generation.
[0096] 2. Discriminator module design
[0097] The discriminator is used to determine whether the input sentence is generated by the generator or real data. '0' represents data generated by the generator, and '1' represents real data. The input to the discriminator is also a sentence, which needs to be encoded to extract features before the judgment is made. Pre-trained models for encoding sentences include Bert and T5Encoder (T5Encoder is the encoder of T5Model). Since we use T5Model in the generator stage and have already introduced the tokenizer T5Tokenizer, this method uses T5Encoder to encode sentences to avoid introducing a new tokenizer that takes up space.
[0098] Similar to the preprocessing for the generator, sentences are also processed by the tokenizer before being input into the T5Encoder, obtaining the corresponding input_ids and attention_mask. These variables are then fed into the encoder, and the encoder's last_hidden_state (the state of the last hidden layer) is what we need. The T5Encoder's output vector has a dimension of 768, and it outputs a hidden vector for each word in the sentence. Assuming a batch size of bs and a sentence length of T, the final T5Encoder output has a dimension of (bs, T, 768). To obtain the sentence's encoding vector, we need to take the mean of all the vectors in the sentence. This means taking the mean of the second dimension of the (bs, T, 768) matrix to obtain a matrix of size (bs, 768). The resulting vector has a dimension of 768, and the output label has a dimension of 1. This means that the network layer has an input dimension of 768 and an output dimension of 1. Therefore, we need to construct a linear layer, whose primary function is to predict the score based on the input sentence encoding vector. The output of the linear layer is the predicted score, which needs to be converted into the corresponding probability. The probability will be calculated with the sentence label to get the loss. A Sigmoid layer is connected after the linear layer to convert the score into a probability between 0 and 1. Because this task is a binary classification task, this method uses BCELoss (binary cross entropy loss) to calculate the loss. The overall discriminant structure is as follows Figure 2 shown.
[0099] The discriminator needs to classify real data as '1' and generated data as '0', so label data must be constructed. For the data generated by the generator, this method constructs a matrix of size (bs, 1) filled with all zeros, and for the real data, a matrix of size (bs, 1) filled with all ones. These two matrices form the sentence label matrices, which are input into the linear layer to obtain the loss, which is also part of the final loss.
[0100] 3. Overall network architecture
[0101] The generator and discriminator are the main components of the CycleGan network. After designing them, we can build the CycleGan network. The specific network architecture is as follows Figure 1 As shown. Assume that X is a grammatically incorrect sentence and Y is a grammatically correct sentence. Then the generator G t The main work of G t (x) = y ′ ,x∈X, we hope that y ′ Be as similar as possible to the sentence in Y, as shown in Equations 1 and 2:
[0102] G t (x) = y′ ,x∈X#(1)
[0103] G f (y) = x ′ ,y∈Y#(2)
[0104] Because correct and incorrect sentences appear in pairs in the grammatical correction task, we also hope that y ′ It is just a corrected version of the original sentence x, and we do not want to change the content of the original sentence. ′ After correcting the error again, the original sentence x can be restored, which is shown in Equations 3 and 4:
[0105] G f (G t (x))=x,x∈X#(3)
[0106] G t (G f (y))=y,y∈Y#(4)
[0107] Generator G t The grammatical error sentence is modified into a grammatically correct sentence, which is the grammatical error correction model that is ultimately needed. t In addition to correcting the error of x, it also needs to correct the error of G t (y) is restored to y, which is actually the generator G f The generated grammatical error sentences are corrected again. t Not only is the original dataset trained, but the generator G is also used f The generated corpus is additionally trained. A larger corpus will give G t Bring better training results.
[0108] CycleGan also includes the game between the generator and the discriminator. t Will G t The corrected sentence is judged as '0', and the real grammatically correct sentence y,y∈Y is judged as '1'. t Will try to generate better and higher quality sentences to try to make the discriminator D t Misjudgment. Discriminator D t With the generator G t The adversarial training between them will make the error correction quality of the generator higher and higher. f With the generator G f The adversarial training between G f The sentence that is closer to the real grammatical error is generated, and as shown in Formula 4, the sentence will be given to the generator G as a new corpus. t Training. Higher quality additional corpus will also improve the generator Gt error correction effect.
[0109] 4. Network Loss Design
[0110] We've previously introduced the CycleGan network architecture and its application in grammatical error correction tasks. Now, we'll detail the loss design used to update neural network parameters.
[0111] This method trains the generator and the discriminator separately. When training the discriminator, G t ,G f The parameters are fixed and only D is updated t ,D f Parameters; when training the generator, D t ,D f The parameters of G are fixed and only G is updated t ,G f Parameters. For the loss of the generator, there are four main parts. t As an example, first G t The grammatically incorrect sentence x needs to be corrected, and the loss can be calculated between the corrected sentence and the target sentence. This loss is called Loss forward At the same time, the corrected sentence should be able to make the discriminator D t It is judged as real data, and the loss is called Loss gan .G f The generated grammatically incorrect sentences must also pass G t The corrected sentence should be as consistent as possible with the original sentence. This loss is called Loss cycle Finally, for the original grammatical sentence, G t It should not be modified, and the loss is called Loss identity The overall generator G t Loss gen The calculation is shown in formula 5, and the specific loss calculation is shown in formulas 6, 7, 8, and 9. For the generator G f The calculation is the same idea.
[0112] Loss gen =Loss forward +Loss gan +Loss cycle +Loss identity #(5)
[0113]
[0114]
[0115]
[0116]
[0117] There are two main considerations when training the discriminator: the discriminator needs to classify the sentences generated by the generator as '0' and the real sentences as '1'. t For example, its Loss dis The calculation is shown in formula 10. For the discriminator D f The same idea:
[0118]
[0119] After the loss is calculated, this method backpropagates it to update the parameters of the neural network.
[0120] Example 2
[0121] The present invention also provides a grammatical error correction model training system based on an adversarial neural network. The grammatical error correction model training system based on an adversarial neural network can be implemented by executing the process steps of the grammatical error correction model training method based on an adversarial neural network. That is, those skilled in the art can understand the grammatical error correction model training method based on an adversarial neural network as a preferred implementation of the grammatical error correction model training system based on an adversarial neural network.
[0122] The grammatical error correction model training system based on the adversarial neural network provided by the present invention includes: module M1: using a pre-trained model as a generator to predict the word with the highest probability at each position to form a sentence, which is used for subsequent discrimination by a discriminator and input into another generator for continued generation; module M2: building a discriminator to discriminate whether the input sentence is generated by the generator or real data; module M3: building a CycleGan network based on two generators and two discriminators to realize the game between the generator and the discriminator; module M4: calculating the cross entropy loss for back propagation to update the neural network parameters.
[0123] The module M1 includes:
[0124] The T5 model is selected as the generator. The sentences input into the T5 model are first processed by the word segmenter, and then the T5 model uniformly fills and truncates them, and generates the corresponding identity identification code of the input sentence;
[0125] Generate an attention mask matrix so that the T5 model only focuses on the sentence itself when performing attention calculations;
[0126] Perform tokenizer processing, padding, and truncation on the target sentence, changing all values in the target sentence to the specified padding token pad_token_id to -100;
[0127] The identity code, attention mask, and sentence label of the input sentence are input into the T5 model for forward propagation to generate the cross entropy loss between the input sentence and the target sentence, which is used for backpropagation to update the neural network parameters and the score of each word in the vocabulary at each position. After Softmax, it is the probability of each word appearing at this position.
[0128] The module M2 includes:
[0129] The sentence input to the discriminator is first encoded to extract features and then judged;
[0130] Use the T5 encoder to encode the sentence. First, the sentence is processed by the word segmenter to obtain the corresponding input sentence identity code and attention mask, and then input it into the T5 encoder. The T5 encoder outputs a hidden layer encoding vector for each word in the sentence.
[0131] Construct a linear layer to predict the score based on the input sentence encoding vector and convert the predicted score into the corresponding probability. The probability will be calculated with the sentence label to get the loss. A Sigmoid layer is connected after the linear layer to convert the score into a probability between 0 and 1.
[0132] The module M3 includes:
[0133] Let X be a grammatically incorrect sentence and Y be a grammatically correct sentence, then the generator expression is:
[0134] G t (x) = y ′ ,x∈X#(1)
[0135] G f (y) = x ′ ,y∈Y#(2)
[0136] Among them, the generator G t Input grammatically incorrect sentence x, generate grammatically correct sentence y ′ ; Generator G f Input grammatically correct sentence y, output grammatically incorrect sentence x ′ ;
[0137] y ′ After correcting the error again, we can restore the original sentence x, which is expressed as:
[0138] G f (G t (x))=x,x∈X#(3)
[0139] G t (G f (y))=y,y∈Y#(4)
[0140] The two discriminators are D t and D f , D t Used to determine whether a sentence is composed of G t Generated or real, D f Used to determine whether a sentence is composed of G f Real or generated; Discriminator D t G t The corrected sentence is judged as 0, and the true grammatically correct sentence y, y∈Y is judged as 1.
[0141] The module M4 includes:
[0142] When training the discriminator, G t ,G f The parameters are fixed and only D is updated t ,D f Parameters; when training the generator, D t ,D f The parameters of G are fixed and only G is updated t ,G f Parameters;
[0143] For the cross entropy loss of the generator, first G t Correct the grammatically incorrect sentence x, and calculate the loss between the corrected sentence and the target sentence. This loss is called Loss forward ; At the same time, the corrected sentence needs to make the discriminator D t It is judged as real data, and the loss is called Loss gan ; G f The grammatically incorrect sentences generated by G t The correction is consistent with the original sentence, and the loss is called Loss cycle ; Finally, for the original grammatical sentence, G t Without modification, the loss is called Loss identity ; Overall generator G t Loss gen The calculation expression is:
[0144] Loss gen =Loss forward +Loss gan +Loss cycle +Loss identity #(5)
[0145]
[0146]
[0147]
[0148]
[0149] When training the discriminator, the discriminator classifies the sentences generated by the generator as 0 and the real sentences as 1. Its Loss dis The calculation expression is:
[0150]
[0151] After the loss is calculated, it is back-propagated to update the parameters of the neural network.
[0152] Those skilled in the art will appreciate that, in addition to implementing the system, device, and various modules provided by the present invention in purely computer-readable program code, it is entirely possible to implement the same program in the form of logic gates, switches, application-specific integrated circuits, programmable logic controllers, embedded microcontrollers, and the like by logically programming the method steps. Therefore, the system, device, and various modules provided by the present invention can be considered a hardware component, and the modules included therein for implementing various programs can also be considered structures within the hardware component; the modules for implementing various functions can also be considered both software programs for implementing the method and structures within the hardware component.
[0153] The above describes specific embodiments of the present invention. It should be understood that the present invention is not limited to the specific embodiments described above, and those skilled in the art may make various changes or modifications within the scope of the claims, which do not affect the essence of the present invention. The embodiments of this application and the features in the embodiments may be combined with each other in any manner unless there is a conflict.
Claims
1. A grammar error correction model training method based on adversarial neural network, characterized in that: include: Step 1: Use the pre-trained model as a generator to predict the word with the highest probability at each position to form a sentence, which is used for subsequent discriminator judgment and input into another generator for further generation; Step 2: Build a discriminator to determine whether the input sentence is generated by the generator or real data; Step 3: Based on two generators and two discriminators, build a CycleGan network to realize the game between the generator and the discriminator; Step 4: Calculate the cross entropy loss for back propagation to update the neural network parameters.
2. The grammatical error correction model training method based on adversarial neural network according to claim 1 is characterized in that: The step 1 comprises: The T5 model is selected as the generator. The sentences input into the T5 model are first processed by the word segmenter, and then the T5 model uniformly fills and truncates, and generates the identity code of the corresponding input sentence; Generate an attention mask matrix so that the T5 model only focuses on the sentence itself when performing attention calculations; Perform tokenizer processing, padding, and truncation on the target sentence, and change all values in the target sentence to the specified padding token pad_token_id -100; The identity code, attention mask and sentence label of the input sentence are input into the T5 model for forward propagation to generate the cross entropy loss with the target sentence, which is used for backpropagation to update the neural network parameters and the score of each word in the vocabulary at each position. After Softmax, it is the probability of each word appearing at this position.
3. The grammatical error correction model training method based on adversarial neural network according to claim 1 is characterized in that: The step 2 comprises: The sentence input to the discriminator is first encoded to extract features and then judged; Use the T5 encoder to encode the sentence. First, the sentence is processed by the word segmenter to obtain the corresponding input sentence identity code and attention mask, and then input into the T5 encoder. The T5 encoder outputs a hidden layer encoding vector for each word in the sentence. Construct a linear layer to predict the score based on the input sentence encoding vector and convert the predicted score into the corresponding probability, which will be calculated with the sentence label to get the loss. Connect a Sigmod layer after the linear layer to convert the score into a probability between 0 and 1.
4. The grammatical error correction model training method based on adversarial neural network according to claim 1 is characterized in that: The step 3 comprises: Let X be a grammatically incorrect sentence and Y be a grammatically correct sentence, then the expression of the generator is: G t (x)=y′,x∈X#(1) G f (y)=x′,y∈Y#(2) Among them, the generator G t Input grammatically incorrect sentence x, generate grammatically correct sentence y′; generator G f Input grammatically correct sentence y, output grammatically incorrect sentence x′; After correcting y′ again, restore it to the original sentence x, the expression is: G f (G t (x))=x,x∈X#(3) G t (G f (y))=y,y∈Y#(4) The two discriminators are D t and D f , D t Used to determine whether a sentence is composed of G t Generated or real, D f Used to determine whether a sentence is composed of G f Real or generated; Discriminator D t G t The corrected sentence is judged as 0, and the true grammatically correct sentence y, y∈Y is judged as 1.
5. The grammatical error correction model training method based on adversarial neural network according to claim 4 is characterized in that: The step 4 comprises: When training the discriminator, G t , G f The parameters are fixed and only D is updated t , D f Parameters; when training the generator, D t , D f The parameters of G are fixed and only G is updated t , G f Parameters; For the cross entropy loss of the generator, first G t Correct the grammatically incorrect sentence x, and calculate the loss between the corrected sentence and the target sentence. This loss is called Loss forward ; At the same time, the corrected sentence needs to make the discriminator D t It is judged as real data, and this loss is called Loss gan ; G f The grammatically incorrect sentences generated by G t The corrected sentence is consistent with the original sentence. This loss is called Loss cycle ; Finally, for the original grammatical sentence, G t Without modification, this loss is called Loss identity ; The overall generator G t Loss gen The calculation expression is: Loss gen =Loss forward +Loss gan +Loss cycle +Loss identity #(5) When training the discriminator, the discriminator classifies the sentences generated by the generator as 0 and the real sentences as 1. Its Loss dis The calculation expression is: After the loss is calculated, it is back-propagated to update the parameters of the neural network.
6. A grammar error correction model training system based on adversarial neural network, characterized in that: include: Module M1: Use the pre-trained model as a generator to predict the word with the highest probability at each position to form a sentence, which is used for subsequent discriminator judgment and input into another generator for further generation; Module M2: Build a discriminator to determine whether the input sentence is generated by the generator or real data; Module M3: Based on two generators and two discriminators, build a CycleGan network to realize the game between the generator and the discriminator; Module M4: Calculate the cross entropy loss for back propagation to update the neural network parameters.
7. The grammar error correction model training system based on adversarial neural network according to claim 6 is characterized in that: The module M1 comprises: The T5 model is selected as the generator. The sentences input into the T5 model are first processed by the word segmenter, and then the T5 model uniformly fills and truncates, and generates the identity code of the corresponding input sentence; Generate an attention mask matrix so that the T5 model only focuses on the sentence itself when performing attention calculations; Perform tokenizer processing, padding, and truncation on the target sentence, and change all values in the target sentence to the specified padding token pad_token_id -100; The identity code, attention mask and sentence label of the input sentence are input into the T5 model for forward propagation to generate the cross entropy loss with the target sentence, which is used for backpropagation to update the neural network parameters and the score of each word in the vocabulary at each position. After Softmax, it is the probability of each word appearing at this position.
8. The grammar error correction model training system based on adversarial neural network according to claim 6 is characterized in that: The module M2 comprises: The sentence input to the discriminator is first encoded to extract features and then judged; Use the T5 encoder to encode the sentence. First, the sentence is processed by the word segmenter to obtain the corresponding input sentence identity code and attention mask, and then input into the T5 encoder. The T5 encoder outputs a hidden layer encoding vector for each word in the sentence. Construct a linear layer to predict the score based on the input sentence encoding vector and convert the predicted score into the corresponding probability, which will be calculated with the sentence label to get the loss. Connect a Sigmod layer after the linear layer to convert the score into a probability between 0 and 1.
9. The grammar error correction model training system based on adversarial neural network according to claim 6 is characterized in that: The module M3 comprises: Let X be a grammatically incorrect sentence and Y be a grammatically correct sentence, then the expression of the generator is: G t (x)=y′,x∈X#(1) G f (y)=x′,y∈Y#(2) Among them, the generator G t Input grammatically incorrect sentence x, generate grammatically correct sentence y′; generator G f Input grammatically correct sentence y, output grammatically incorrect sentence x′; After correcting y′ again, restore it to the original sentence x, the expression is: G f (G t (x))=x,x∈X#(3) G t (G f (y))=y,y∈Y#(4) The two discriminators are D t and D f , D t Used to determine whether a sentence is composed of G t Generated or real, D f Used to determine whether a sentence is composed of G f Real or generated; Discriminator D t G t The corrected sentence is judged as 0, and the true grammatically correct sentence y, y∈Y is judged as 1.
10. The grammar error correction model training system based on adversarial neural network according to claim 9 is characterized in that: The module M4 comprises: When training the discriminator, G t , G f The parameters are fixed and only D is updated t , D f Parameters; when training the generator, D t , D f The parameters of G are fixed and only G is updated t , G f Parameters; For the cross entropy loss of the generator, first G t Correct the grammatically incorrect sentence x, and calculate the loss between the corrected sentence and the target sentence. This loss is called Loss forward ; At the same time, the corrected sentence needs to make the discriminator D t It is judged as real data, and this loss is called Loss gan ; G f The grammatically incorrect sentences generated by G t The corrected sentence is consistent with the original sentence. This loss is called Loss cycle ; Finally, for the original grammatical sentence, G t Without modification, this loss is called Loss identity ; The overall generator G t Loss gen The calculation expression is: Loss gen =Loss forward +Loss gan +Loss cycle +Loss identity #(5) When training the discriminator, the discriminator classifies the sentences generated by the generator as 0 and the real sentences as 1. Its Loss dis The calculation expression is: After the loss is calculated, it is back-propagated to update the parameters of the neural network.
Citation Information
Patent Citations
Grammar error correction method and device, training method and device, electronic equipment and storage medium
CN114239557A