RATT: Recurrent Attention to Transient Tasks for Continual Image Captioning

Riccardo Del Chiaro, Bartłomiej Twardowski, Andrew D. Bagdanov, Joost van de Weijer

Introduction

Classical supervised learning systems acquire knowledge by providing them with a set of annotated training samples from a task, which for classifiers is a single set of classes to learn. This view of supervised learning stands in stark contrast with how humans acquire knowledge, which is instead continual in the sense that mastering new tasks builds upon previous knowledge acquired when learning previous ones. This type of learning is referred to as continual learning (sometimes incremental or lifelong learning), and continual learning systems instead consume a sequence of tasks, each containing its own set of classes to be learned. Through a sequence of learning sessions, in which the learner has access only to labeled examples from the current task, the learning system should integrate knowledge from past and current tasks in order to accurately master them all in the end. A principal shortcoming of state-of-the-art learning systems in the continual learning regime is the phenomenon of catastrophic forgetting : in the absence of training samples from previous tasks, the learner is likely to forget them in the process of acquiring new ones.

Continual learning research has until now concentrated primarily on classification problems modeled with deep, feed-forward neural networks . Given the importance of recurrent networks for many learning problems, it is surprising that continual learning of recurrent networks has received so little attention . A recent study on catastrophic forgetting in deep LSTM networks observes that forgetting is more pronounced than in feed-forward networks. This is caused by the recurrent connections which amplify each small change in the weights. In this paper, we consider continual learning for captioning, where a recurrent network (LSTM) is used to produce the output sentence describing an image. Rather than having access to all captions jointly during training, we consider different captioning tasks which are learned in a sequential manner (examples of tasks could be captioning of sports, weddings, news, etc).

Most continual learning settings consider tasks that each contain a set of classes, and these sets are disjoint . A key aspect of continual learning for image captioning is the fact that tasks are naturally split into overlapping vocabularies. Task vocabularies might contain nouns and some verbs which are specific to a task, however many of the words (adjectives, adverbs, and articles) are shared among tasks. Moreover, the presence of homonyms in different tasks might directly lead to forgetting of previously acquired concepts. This transient nature of words in task vocabularies makes continual learning in image captioning networks different from traditional continual learning.

In this paper we take a systematic look at continual learning for image captioning problems using recurrent, LSTM networks. We consider three of the principal classes of approaches to exemplar-free continual learning: weight-regularization approaches, exemplified by Elastic Weight Consolidation (EWC) ; knowledge distillation approaches, exemplified by Learning without Forgetting (LwF) ; and attention-based approached like Hard Attention to the Task (HAT) . For each we propose modifications specific to their application to recurrent LSTM networks, in general, and more specifically to image captioning in the presence of transient task vocabularies.

The contributions of this work are threefold: (1) we propose a new framework and splitting methodologies for modeling continual learning of sequential generation problems like image captioning; (2) we propose an approach to continual learning in recurrent networks based on transient attention masks that reflect the transient nature of the vocabularies underlying continual image captioning; and (3) we support our conclusions with extensive experimental evaluation on our new continual image captioning benchmarks and compare our proposed approach to continual learning baselines based on weight regularization and knowledge distillation. To the best of our knowledge we are the first to consider continual learning of sequential models in the presence of transient tasks vocabularies whose classes may appear in some learning sessions, then disappear, only to reappear in later ones.

Related work

Catastrophic forgetting. Early works demonstrating the inability of networks to retain knowledge from previously task when learning new ones are and . Approaches include methods that mitigate catastrophic forgetting via replay of exemplars (iCarl , EEIL , and GEM ) or by performing pseudo-replay with GAN-generated data . Weight regularization has also been investigated . Output regularization via knowledge distillation was investigated in LwF , as well as architectures based on network growing and attention masking . For more details we refer to recent surveys on continual learning .

Image captioning. Modern captioning techniques are inspired by machine translation and usually employ a CNN image encoder and RNN text decoder to “translate” images into sentences. NIC uses a pre-trained CNN to encode the image and initialize an LSTM decoder. Differently, in image features are used at each time step, while in a two-layer LSTM is employed. Recurrent latent variable is introduced in , encoding the visual interpretation of previously-generated words and acting as a long-term visual memory during next words generation. In a spatial attention mechanism is introduced: the model is able to focus on specific regions of the image according to the previously generated words. ReviewerNet also selects in advance which part of the image will be attended, so that the decoder is aware of it from the beginning. Areas of Attention models the dependencies between image regions and generated words given the RNN state. A visual sentinel is introduced in to determine, at each decoding step, if it is important to attend the visual features. The authors of mixed bottom-up attention (implemented with an object detection network in the encoder) and a top-down attention mechanism in the LSTM decoder that attend to the visual features of the salient image regions selected by the encoder. Recently, transformer-based methods have been applied to image captioning , which eliminate the LSTM in the decoder.

The focus of this paper is RNN-based captioning architectures and how they are affected by catastrophic forgetting. For more details on image captioning we refer to recent surveys .

Continual learning of recurrent networks. A fixed expansion layer technique was proposed to mitigate forgetting in RNNs in . A dedicated network layer that exploits sparse coding of RNN hidden state is used to reduce the overlap of pattern representations. In this method the network grows with each new task. A Net2Net technique was used for expanding the RNN in . The method uses GEM for training on a new task, but has several shortcomings: model weights continue to grow and it must retain previous task data in the memory.

Experiments on four synthetic datasets were conducted in to investigate forgetting in LSTM networks. The authors concluded that the LSTM topology has no influence on forgetting. This observations motivated us to take a close look to continual image captioning where the network architecture is more complex and an LSTM is used as a output decoder.

Continual LSTMs for transient tasks

We first describe our image captioning model and some details of LSTM networks. Then we describe how to apply classical continual learning approaches to LSTM networks.

We use a captioning model similar to Neural Image Captioning (NIC) . It is an encoder-decoder network that “translates” an image into a natural language description. It is trained end-to-end, directly maximizing the probability of correct sequential generation:

where s=[s1,…sN]s=[s_{1},\ldots s_{N}] is the target sentence for image II, θ\theta are the model parameters.

The decoder is an LSTM network in which words s1,…,sn−1s_{1},\ldots,s_{n-1} are encoded in the hidden state hnh_{n} and a linear classifier is used to predict the next word at time step nn:

where SS is a word embedding matrix, sns_{n} is the nn-th word of the ground-truth sentence for image II, CC is a linear classifier, and VV is the visual projection matrix that projects image features from the CNN encoder into the embedding space at time n=0n=0.

The LSTM network is defined by the following equations (for which we omit the bias terms): in\displaystyle i_{n} =\displaystyle= σ(Wixxn+Wihhn−1)\displaystyle\sigma(W_{ix}x_{n}+W_{ih}h_{n-1}) (6) on\displaystyle o_{n} =\displaystyle= σ(Woxxn+Wohhn−1)\displaystyle\sigma(W_{ox}x_{n}+W_{oh}h_{n-1}) (7) fn\displaystyle f_{n} =\displaystyle= σ(Wfxxn+Wfhhn−1)\displaystyle\sigma(W_{fx}x_{n}+W_{fh}h_{n-1}) (8) gn\displaystyle g_{n} =\displaystyle= tanh⁡(Wgxxn+Wghhn−1)\displaystyle\tanh(W_{gx}x_{n}+W_{gh}h_{n-1}) (9) hn\displaystyle h_{n} =\displaystyle= on⊙cn\displaystyle o_{n}\odot c_{n} (10) cn\displaystyle c_{n} =\displaystyle= fn⊙cn−1+in⊙gn\displaystyle f_{n}\odot c_{n-1}+i_{n}\odot g_{n} (11)

where ⊙\odot is the Hadamard (element-wise) product, σ\sigma the logistic function, cc the LSTM cell state. The WW matrices are the trainable LSTM parameters related to input xx and hidden state hh, for each gate ii, ff, oo, gg. The loss used to train the network is the sum of the negative log likelihood of the correct word at each step:

Inference. During training we perform teacher forcing using nn-th word of the target sentence as input to predict word n+1n+1. At inference time, since we have no target caption, we use the word predicted by the model at the previous step arg⁡max⁡ pn\arg\max\ p_{n} as input to the word embedding matrix SS.

2 Continual learning of recurrent models

Normally catastrophic forgetting is highlighted in continual learning benchmarks by defining tasks that are mutually disjoint in the classes they contain (i.e. no class belongs to more than one task). For sequential problems like image captioning, however, this is not so easy: sequential learners must classify words at each decoding step, and a large vocabulary of common words are needed for any practical captioning task.

Incremental model. Our models are trained on sequences of captioning tasks, each having different vocabularies. For this reason any captioning model must be able to enlarge its vocabulary. When a new task arrives we add a new column for each new word in the classifier and word embedding matrices. The recurrent network remains untouched because the embedding projects inputs into the same space. The basic approach to adapt to the new task is to fine-tune the network over the new training set. To manage the different classes (words) of each task we have two possibilities: (1) Use different classifier and word embedding matrices for each task; or (2) Use a common, growing classifier and a common, growing word embedding matrix.

The first option has the advantage that each task can benefit from ad hoc weights for the task, potentially initializing from the previous task for the common words. However, it also increases decoder network size consistently with each new task. The second option has the opposite advantage of keeping the dimension of the network bounded, sharing weights for all common words. Because of the nature of the captioning problem, many words will be shared and duplicating both word embedding matrix and classifier for all the common words seems wasteful. Thus we adopt the second alternative. With this approach, the key trick is to deactivate classifier weights for words not present in the current task vocabulary.

We use θ^t\hat{\theta}^{t} to denote optimal weights learned for task tt on dataset DtD_{t}. After training on task tt, we create a new model for task t+1t+1 with expanded weights for classifier and word embedding matrices. We use weights from θ^t\hat{\theta}^{t} to initialize the shared weights of the new model.

3 Recurrent continual learning baselines

We describe how to adapt two common continual learning approaches, one based on weight regularization and the other on knowledge distillation. We will use these as baselines in our comparison.

Weight regularization. A common method to prevent catastrophic forgetting is to apply regularization to important model weights before proceeding to learn a new task . Such methods can be directly applied to recurrent models with little effort. We choose Elastic Weight Consolidation (EWC) as a regularization-based baseline. The key idea of EWC is to limit change to model parameters vital to previously-learned tasks by applying a quadratic penalty to them depending on their importance. Parameter importance is estimated using a diagonal approximation of the Fisher Information Matrix. The additional loss function we minimize when learning task tt is:

where θ^t−1\hat{\theta}^{t-1} are the estimated model parameters for the previous task, θt\theta^{t} are the model parameters at the current task tt, L(x,S)\mathcal{L}(x,S) is the standard loss used for fine-tuning the network on task tt, ii indexes the model parameters shared between tasks tt and t−1t-1, Fit−1F^{t-1}_{i} is the ii-th element of a diagonal approximation of the Fisher Information Matrix for model after training on task t−1t-1, and λ\lambda weights the importance of the previous task. We apply Eq. 13 to all trainable weights. Due to the transient nature of words across tasks, we do not expect weight regularization to be optimal since some words are shared and regularization limits the plasticity needed to adjust to a new task.

Recurrent Learning without Forgetting. We also apply a knowledge distillation approach inspired by Learning without Forgetting (LwF) on the LSTM decoder network to prevent catastrophic forgetting. The model after training task t−1t-1 is used as a teacher network when fine-tuning on task tt. The aim is to let the new network freely learn how to classify new words appearing in task tt while keeping stable the predicted probabilities for words from previous tasks.

To do this, at each step nn of the decoder network the previous decoder is also fed with the data coming from the new task tt. Note that the input to the LSTM at each step nn is the embedding of the n−n-th word in the target caption, and the same embedding is given as input to both teacher and student networks – i.e. the student network’s embedding of word nn is also used as input for the teacher, while each network uses its own hidden state hn−1h_{n-1} and cell state cn−1c_{n-1} to decode the next word. At each decoding step, the output probabilities pn+1t,<tp^{t,<t}_{n+1} from the student network LSTM corresponding to words present in tasks 1,…,t−11,\dots,t-1 are compared with the those predicted by the teacher network, pn+1t−1,<tp^{t-1,<t}_{n+1}. A distillation loss ensures that the student network does not deviate from the teacher:

where γ(⋅)\gamma(\cdot) rescales a probability vector pp -with temperature parameter TT. This loss is combined with the LSTM training loss (see Eq. 12). Note that differently from , we do not fine-tune the classifier of the old network because we use a single, incremental word classifier.

Attention for continual learning of transient tasks

Inspired by the Hard Attention to the Task (HAT) method , we developed an attention-based technique applicable to recurrent networks. We name it Recurrent Attention to Transient Tasks (RATT), since it is specifically designed for recurrent networks with task transience. The key idea is to use an attention mechanism to allocate a portion of the activations of each layer to a specific task tt. An overview of RATT is provided in figure 1.

Attention masks. The number of neurons used for a task is limited by two task-conditioned attention masks: embedding attention axt∈a_{x}^{t}\in and hidden state attention aht∈a_{h}^{t}\in. These are computed with a sigmoid activation σ\sigma and a positive scaling factor ss according to:

where tt is a one-hot task vector, and AxA_{x} and AhA_{h} are embedding matrices. Next to the two attention mask, we have a vocabulary mask asta_{s}^{t} which is a binary mask identifying the words of the vocabulary used in task tt: as,it=1a^{t}_{s,i}=1 if word ii is part of the vocabulary of task tt and is zero otherwise. The forward pass (see Eqs.2 and 5) of the network is modulated with the attention masks according to:

Attention masks act as an inhibitor when their value is near . The main idea is to learn attention masks during training, and as such learn a limited set of neurons for each task. Neurons used in previous tasks can still be used in subsequent ones, however the weights which were important for previous tasks have reduced plasticity (depending on the amount of attention to for previous tasks).

Training. For training we define the cumulative forward mask as:

ah<ta_{h}^{<t} and as<ta_{s}^{<t} are similarly defined. We now define the following backward masks which have the dimensionality of the weight matrices of the network and are used to selectively backpropagate the gradient to the LSTM layers:

Note that we use ah,ia_{h,i} refer to the i-th element of vector aha_{h}, etc. The backpropagation with learning rate λ\lambda is then done according to

The only difference from standard backpropagation are the backward matrices BB which prevents the gradient from changing those weights that were highly attended in previous tasks. The backpropagation updates to the other matrices in Eqs. 6-9 are similar (see Suppl. Materials).

Other than we also define backward masks for the word embedding matrix SS, the linear classifier CC, and the image-projection matrix VV:

and the corresponding backpropagation updates:

The backward mask BVtB^{t}_{V} modulates the backpropagation to the image features. Since we do not define a mask on the output of the fixed image encoder, this is only defined by ax<ta_{x}^{<t}.

Linearly annealing the scaling parameter ss, used in Eq. 15, during training (like ) was found to be beneficial. We apply s=1smax+(smax−1smax)b−1B−1s=\frac{1}{s_{max}}+\left(s_{max}-\frac{1}{s_{max}}\right)\frac{b-1}{B-1} where bb is the batch index and BB is the total number of batches for the epoch. We used smax=2000s_{max}=2000 and smax=400s_{max}=400 for experiments on Flickr30k and MS-COCO, respectively.

The loss used to promote low network usage and to keep some neurons available for future tasks is:

This loss is combined with Eq. 12 for training. The loss encourages attention to only a few new neurons. However, tasks can attend to previously attended neurons without any penalty. This encourages forward transfer during training. If the attention masks are binary, the system would not suffer from any forgetting, however it would lose its backward transfer ability.

Differently than , when computing BStB_{S}^{t} we take into account the recurrency of the network, considering the classifier CC to be the previous layer of SS. In addition, our output masks asa_{s} allow for overlap to model the transient nature of the output vocabularies, whereas only considers non-overlapping classes for the various tasks.

Experimental results

All experiments use the same architecture: for the encoder network we used ResNet151 pre-trained on ImageNet . Note that the image encoder is frozen and is not trained during continual learning, as is common in many image captioning systems. The decoder consists of the word embedding matrix SS that projects the input words into a 256-dimensional space, an LSTM cell with hidden size 512 that takes the word (or image feature for the first step) embeddings as input, and a final fully connected layer CC that take as input the hidden state hnh_{n} at each LSTM step nn and outputs a probability distribution pn+1p_{n+1} over the ∣Vt∣|V^{t}| words in the vocabulary for current task tt.

We applied all techniques on the Flickr30K and MS-COCO captioning datasets (see next section for task splits). All experiments were conducted using PyTorch, networks were trained using the Adam optimizer, all hyperparameters were tuned over validation sets. Batch size, learning rate and max-decode length for evaluation were set, respectively, to 128, 4e-4, and 26 for MS-COCO, and 32, 1e-4 and 40 for Flickr30k. These differences are due to the size of the training set and by the average caption lengths in the two datasets.

Inference at test time is task-aware for all methods. For EWC and LwF this means that we consider only the word classifier outputs corresponding to the correct task, and for RATT that we use the fixed output masks for the correct task. All metrics where computed using the nlg-eval toolkit . Models where trained for a fixed number of epochs and the best model according to BLEU-4 performance on the validation set were chosen for each task. When proceeding to the next task, the best model from the previous task were used as a starting point.

For our experiments we use two different captioning datasets: MS-COCO and Flickr30k . We split MS-COCO into tasks using a disjoint visual categories procedure. For this we defined five tasks based on disjoint MS-COCO super-categories containing related classes (transport, animals, sports, food and interior). For Flickr30K we instead used an incremental visual categories procedure. Using the visual entities, phrase types, and splits from we identified four tasks: scene, animals, vehicles, instruments. In this approach the first task contains a set of visual concepts that can also be appear in future tasks.

Some statistics on number of images and vocabulary size for each task are given in table 1 for both datasets. See the supplementary material for a detailed breakdown of classes appearing in each task and more details on these dataset splits. MS-COCO does not provide a test set, so we randomly selected half of the validation set images and used them for testing only. Since images have at least five captions, we used the first five captions for each image as the target.

2 Ablation study

We conducted a preliminary study on our split of MS-COCO to evaluate the impact of our proposed Recurrent Attention to Transient Tasks (RATT) approach. In this experiment we progressively introduce the attention masks described in section 4. We start with the basic captioning model with no forgetting mitigation, and so is equivalent to fine-tuning. Then we introduce the mask on hidden state hnh_{n} of the LSTM (along with the corresponding backward mask), and then the constant binary mask on the classifier that depends on the words of the current task, then the visual and word embedding masks, and finally the combination of all masks.

In figure 2 we plot the BLEU-4 performance of these configurations for each training epoch and each of the five MS-COCO tasks. Note that for later tasks the performance on early epochs (i.e. before encountering the task) is noisy as expected – we are evaluating performance on future tasks. These results clearly show that applying the mask to LSTM decreases forgetting in the early epochs when learning a new task. However, performance continues to decrease and in some tasks the result is similar to fine-tuning. Even if the LSTM is forced to remember how to manage hidden states for previous tasks, the other parts of the network suffer from catastrophic forgetting. Adding the classifier mask improves the situation, but the main contribution comes from applying the mask to the embedding. Applying all masks we obtain zero or nearly-zero forgetting. This depends on the smaxs_{max} value used during training: in these experiments we use smax=400s_{max}=400, which results in zero forgetting of previous tasks. We also conducted an ablation study on the smaxs_{max} parameter. From the results in figure 3 we can see that higher smaxs_{max} values improve old task performance, and sufficiently high values completely avoid forgetting. Using moderate values, however, can be helpful to increase performance in later tasks. See supplementary material for additional ablations.

3 Results on MS-COCO

In table 2 we report the performance of a fine-tuning baseline with no forgetting mitigation (FT), EWC, LwF, and RATT on our splits for the MS-COCO captioning dataset. The forgetting percentage is computed by taking the BLEU-4 score for each model after training on the last task and dividing it by the BLEU-4 score at the end of the training of each individual task. From the results we see that all techniques consistently improve performance on previous tasks when compared to the FT baseline. Despite the simplicity of EWC, the improvement over fine-tuning is clear, but it struggles to learn a good model for the last task. LwF instead shows the opposite behavior: it is more capable of learning the last task, but forgetting is more noticeable. RATT achieves zero forgetting on MS-COCO, although at the cost of some performance on the final task. This is to be expected, though, as our approach deliberately and progressively limits network capacity to prevent forgetting of old tasks. Qualitative results on MS-COCO are provided in figure 4.

4 Results on Flickr30k

In table 3 we report performance of a fine-tuning baseline with no forgetting mitigation (FT), EWC, LwF, and RATT on our Flickr30k task splits. Because these splits are based on incremental visual categories, it does not reflect a classical continual-learning setup that enforce disjoint categories to maximize catastrophic forgetting: not only there are common words that share the same meaning between different tasks, but some of the visual categories in early tasks are also present in future ones. For this reason, learning how to describe task t=1t=1 also implies learning at least how to partially describe future tasks, so forward and backward transfer is significant.

Despite this, we see that all approaches increase performance on old tasks (when compared to FT) while retaining good performance on the last one. Note that both RATT and LwF result in negative forgetting: in these cases the training of a new task results in backward transfer that increases performance on an old one. EWC improvement is marginal, and LwF behaves a bit better and seems more capable of exploiting backward transfer. RATT backward transfer is instead limited by the choice of a high smaxs_{max}, which however guarantees nearly zero forgetting.

5 Human evaluation experiments

We performed an evaluation based on human quality judgments using 200 images (40 from each task) from the MS-COCO test splits. We generated captions with RATT, EWC, and LwF after training on the last task and then presented ten users with an image and RATT and baseline captions in random order. Users were asked (using forced choice) to select which caption best represents the image content. A similar evaluation was performed for the Flickr30k dataset with twelve users. The percentage of users who chose RATT over the baseline are given in the table 4. These results on MS-COCO dataset confirm that RATT is superior on all tasks, while on Flickr30k there is some uncertainty on the first task, especially when comparing RATT with LwF. Note that for the last task of each dataset there is no forgetting, so it is expected that baselines and RATT perform similarly.

Conclusions

In this paper we proposed a technique for continual learning of image captioning networks based on Recurrent Attention to Transient Tasks (RATT). Our approach is motivated by a feature of image captioning not shared with other continual learning problems: tasks are composed of transient classes (words) that can be shared across tasks. We also showed how to adapt Elastic Weight Consolidation and Learning without Forgetting, two representative approaches of continual learning, to the recurrent image captioning networks. We proposed task splits for the MS-COCO and Flickr30k image captioning datasets, and our experimental evaluation confirms the need of recurrent task attention in order to mitigate forgetting in continual learning with sequential, transient tasks. RATT is capable of zero forgetting at the expense of plasticity and backward transfer: the ability to adapt to new tasks is limited by the number of free neurons and it is difficult to exploit knowledge from future tasks to better predict older ones. The focus of this work is on how a simple encoder-decoder image captioning model forget, which limits the quality of captions when comparing with current state-of-the-art. As future work, we are interested in applying the developed method in more complex captioning systems.

Broader impact

Automatic image captioning has applications in image indexing, Content-Based Image Retrieval (CBIR), and industries like commerce, education, digital libraries, and web searching. Social media platforms could use it to directly generate descriptions of multimedia content. Image captioning systems can support peoples with disabilities, providing an access to the visual content unreachable before by changing its representation. Continual learning contrasts with the joint-training paradigm in common use today. It has the advantage that it can better protect privacy concerns since, once learned, the data not need to be retained. Furthermore, it is more efficient since networks continue learning and are not initialized from scratch every time a new task arrives. This ability is crucial for development of systems like virtual personal assistants, where the adaption to new tasks and environments is fundamental, especially where multi-modal communication channels are ubiquitous. Finally, the algorithm considered in this paper will reflect the biases present in the dataset. Therefore, special care should be taken when applying this technique to applications where possible biases in the dataset might result in biased outcomes towards minority and/or under-represented groups in the data.

Acknowledgments and Disclosure of Funding

We acknowledge the support from Huawei Kirin Solution. We thank NVIDIA Corporation for donating the Titan XP GPU that was used to conduct the experiments. We also acknowledge the project PID2019-104174GB-I00 of Ministry of Science of Spain and the ARS01_00421 Project "PON IDEHA - Innovazioni per l’elaborazione dei dati nel settore del Patrimonio Culturale" of the Italian Ministry of Education. Our acknowledged partners funded this project.

References