Less-forgetting Learning in Deep Neural Networks
Heechul Jung, Jeongwoo Ju, Minju Jung, Junmo Kim
I Introduction
Deep neural networks (DNNs) have grown to nearly human levels of recognition in identifying objects, faces, and speeches . Despite this advancement of deep learning, remaining issues still exist; a catastrophic forgetting problem is the one of these remaining issues . The problem is an important issue in DNNs since this enables an improvement in the performance of DNNs in several important applications such as domain adaptation and incremental learning.
A catastrophic forgetting phenomenon is often observed when performing domain adaptation because the distribution of source data is significantly different from the distribution of target data. For example, consider the traditional transfer learning protocol (weight copy fine-tuning) for adapting to a new domain. Usually, the network pre-trained using the original data (source domain) is used as initial weights for adapting to the new data (target domain) . During learning new data from the target domain, it is natural that the network forgets the previously learned information from the source domain. Even if source and target domains are nearly homogeneous, the network forgets the information about source data.
Several researches have been performed for alleviating such problem. Srivastava et al. proposed a local winner-take-all (LWTA) activation function that helps to prevent the catastrophic forgetting . This activation function has the effectiveness of implicit long-term memory. Recently, several experiments of a catastrophic forgetting problem in DNNs were empirically performed in . The paper shows a dropout method with a Maxout activation function is helpful for forgetting less of the learned information. However, these kinds of methods are not explicitly used for the unforgetting of previously learned information, which means that it does not guarantee the unforgetting ability.
An unsupervised approach was proposed in . Goodrich et. al extended this method to a recurrent neural network . These methods compute cluster centroids while learning the training data in source domain. They also use the computed centroids for the target domain. Consequently, these are not generally applicable for the pre-trained model, and these are not adequate for our problem settings. and have similar issues mentioned above.
In this paper, we try to solve a catastrophic forgetting problem in DNNs by using a new learning method. The proposed method has the ability to maintain the original feature space of the source domain, even if the network does not see any previous training data. Therefore, the method forgets less of the information obtained from the source data compared to the traditional transfer learning method. Figure 1 shows the feature spaces of the traditional transfer learning method and the proposed method. In the proposed method, features of the same source class and target data are well clustered, even if re-training only using the target data is finished.
Our proposed method is also applicable to the learning method from scratch, which is a general training protocol using stochastic gradient descent methods. We observed that the forgetting problem also arises between mini-batches because mini-batches are small datasets subsampled from a large amount of whole data. We also deal with such problems. We summarize our main contributions as follows:
We propose a less-forgetting learning method to alleviate a catastrophic forgetting problem in DNNs.
We observe that a catastrophic forgetting problem also occurs between mini-batches when using a stochastic gradient descent learning.
Furthermore, we show our less-forgetting learning is effective to solve the problem, and it gives better generalization performance.
To theoretically show the less-forgetting problem, we assume that we have the weight parameters of the pre-trained model for the source domain, and training data are given for the target domain (target data). Consequently, data for the source domain (source data) are not accessible. For more clarity, we explain the problem using mathematical expressions as follows.
II Less-forgetting Learning
In DNNs, the lower layer is considered as a feature extractor, and the top layer is regarded as a linear classifier. This means that the weights of the softmax function represent a decision boundary for classifying the features. Due to the linear classifier on the top layer, the features extracted from the top hidden layer are usually linearly separable. Using this knowledge, we propose a new learning scheme that satisfies the following two properties to reduce the forgetting problem of the information learned from the source domain:
Property 1. The decision boundaries should be unchanged. Property 2. The features extracted from source data by the target network should be present in a position close to the features extracted from source data by the source network.
We build a less-forgetting learning algorithm based on two properties. The first property is easily implemented by setting the learning rates of the boundary to zero, but satisfying the second property is not trivial since we cannot access to the source data. Instead of using the source data, we use the target data and show it is helpful to satisfy Property 2. Figure 2 briefly explains our algorithm, and the details are as follows.
Initially, like a traditional transfer learning method, we reuse the weights of the source network as the initial weights of the target network. Next, we freeze the weights of the softmax layer to maintain the boundaries of the classifier. Then, we train the network where to simultaneously minimize the total loss function as follows:
where is the -th value of the ground truth label, and is the -th output value of the softmax of the target network. is the total number of class. In other words, this loss function helps the network to classify the input data correctly. In order to satisfy the second property, is defined as follows:
where is the total number of hidden layers, and is a feature vector of layer . Using the loss function, the target network learns to extract features which are similar to the features extracted by source network.
Finally, we build a less-forgetting learning algorithm, as shown in Algorithm 1. and in the algorithm denote the number of iterations and the size of mini-batch.
III Less-forgetting for General Learning Cases
It is well known the forgetting problem occurs when learning a new task . In addition, we show that the forgetting problem also occurs when performing general learning process using gradient descent methods. To observe the forgetting phenomenon of the network, we probed the training loss values of a particular mini-batch set, as shown in Figure 3. The observing procedure is as follows:
Figure 3 shows our less-forgetting method has more smooth graph than traditional learning method. It seems to alleviate the forgetting problem, but the value of training loss is higher than traditional learning method. This is due to the first property in Section II, and freezing the boundary is an obstacle to learn new data. In the modified less-forgetting algorithm, we often unfreeze the boundary of the network. In addition, we update parameters of the source network using the parameters of the target network. Finally, as shown in Figure 3, the green line is more smooth than the black line, and the green line has lower loss values than the red line.
Using this observation, we present a less-forgetting algorithm for general learning cases, as shown in Algorithm 2. Our algorithm has two main parts (line 35 and line 612). The first part is to switch the parameters of source network which means the network that we want to forget less, and the second part is to unfreeze repeatedly the boundary of the network. If the value of is large, the network does not adapt a new data well. Further, the parameter of plays a role similar to . As a result, our algorithm has an ability to forget less the information learned previously. We set the value of is smaller than , and we set and for all the experiments.
IV Experiments
We establish two different recognition experiments: unforgetting and generalization tests for Algorithm 1 and 2, respectively. In the unforgetting experiment, we manually set up a source domain and a target domain using an object recognition dataset (CIFAR-10) and considered two different digit number recognition datasets (MNIST, SVHN) to belong to the source domain and the target domain, respectively. In the generalizaton experiment, we used CIFAR-10 dataset. In addition, we used a dropout method for Maxout and LWTA experiments because we want to raise their maximum performance.
Table I includes recognition rates of source networks, which are trained only on source domain and tested—with different activation functions—on each domain separately. The large accuracy gap between them certainly makes sense, since the DNNs has never seen examples from the target domain during training. As opposed to source network, we resume training with different methods on the target domain after learning is finished on the source domain nunder the constraint that during, target domain learning, DNNs cannot access examples from source domain. The results are also listed in Table I. Two values are selected based on observing Figure 4. Our method predicts true labels more accurately than previous works, such as the traditional transfer learning method, LWTA, and Maxout. In addition, this results provide compelling evidence that an alternating activation function is not the proper way to prevent a network from forget as little previously learned information as possible.
Next, we observe the relation between target accuracy and domain accuracy in accordance with a value of , as displayed in Figure 4. It should be noted that the closer the curve comes to the top right corner, the better performance. Note that the top target recognition rate of less forgetting is even higher than that of transfer learning, LWTA, and Maxout in CIFAR-10. Therfore, it could be inferred that dissimilarity between source(color) and target domain(grayscale) is not that much, and remembering previous domain help more generlization.
IV-B Generalization Test
As explained in Section III, we adopt Algorithm 2 on the object recognition and make a discovery that less forgetting learning shows better performance, i.e. more generalization. Table II shows the accuracy of original, batch normalization, batch normalization plus less forgetting, and less forgetting learning as the number of iteratioin increase. It is obvious that all cases equipped with less forgetting show an improvement over ones without it.
V Conclusion
We proposed a less-forgetting learning method to alleviate a catastrophic forgetting problem in DNNs. Also, we observed that the forgetting phenomenon occurs when performing general learning method like a stochastic gradient descent method. Our method is also effective to mitigate this problem. Finally, we showed that our method is useful for improving generalization ability of DNNs.