Best of Both Worlds: Transferring Knowledge from Discriminative Learning to a Generative Visual Dialog Model

Jiasen Lu, Anitha Kannan, Jianwei Yang, Devi Parikh, Dhruv Batra

Introduction

One fundamental goal of artificial intelligence (AI) is the development of perceptually-grounded dialog agents – specifically, agents that can perceive or understand their environment (through vision, audio, or other sensors), and communicate their understanding with humans or other agents in natural language. Over the last few years, neural sequence models (e.g. ) have emerged as the dominant paradigm across a variety of setting and datasets – from text-only dialog to more recently, visual dialog , where an agent must answer a sequence of questions grounded in an image, requiring it to reason about both visual content and the dialog history.

The standard training paradigm for neural dialog models is maximum likelihood estimation (MLE) or equivalently, minimizing the cross-entropy (under the model) of a ‘ground-truth’ human response. Across a variety of domains, a recurring problem with MLE trained neural dialog models is that they tend to produce ‘safe’ generic responses, such as ‘Not sure’ or ‘I don’t know’ in text-only dialog , and ‘I can’t see’ or ‘I can’t tell’ in visual dialog . One reason for this emergent behavior is that the space of possible next utterances in a dialog is highly multi-modal (there are many possible paths a dialog may take in the future). In the face of such highly multi-modal output distributions, models ‘game’ MLE by latching on to the head of the distribution or the frequent responses, which by nature tend to be generic and widely applicable. Such safe generic responses break the flow of a dialog and tend to disengage the human conversing with the agent, ultimately rendering the agent useless. It is clear that novel training paradigms are needed; that is the focus of this paper.

One promising alternative to MLE training proposed by recent work is sequence-level training of neural sequence models, specifically, using reinforcement learning to optimize task-specific sequence metrics such as BLEU , ROUGE , CIDEr . Unfortunately, in the case of dialog, all existing automatic metrics correlate poorly with human judgment , which renders this alternative infeasible for dialog models.

In this paper, inspired by the success of adversarial training , we propose to train a generative visual dialog model (GG) to produce sequences that score highly under a discriminative visual dialog model (DD). A discriminative dialog model receives as input a candidate list of possible responses and learns to sort this list from the training dataset. The generative dialog model (GG) aims to produce a sequence that DD will rank the highest in the list, as shown in Fig. 1.

Note that while our proposed approach is inspired by adversarial training, there are a number of subtle but crucial differences over generative adversarial networks (GANs). Unlike traditional GANs, one novelty in our setup is that our discriminator receives a list of candidate responses and explicitly learns to reason about similarities and differences across candidates. In this process, DD learns a task-dependent perceptual similarity and learns to recognize multiple correct responses in the feature space. For example, as shown in Fig. 1 right, given the image, dialog history, and question ‘Do you see any bird?’, besides the ground-truth answer ‘No, I do not’, DD can also assign high scores to other options that are valid responses to the question, including the one generated by GG: ‘Not that I can see’. The interaction between responses is captured via the similarity between the learned embeddings. This similarity gives an additional signal that GG can leverage in addition to the MLE loss. In that sense, our proposed approach may be viewed as an instance of ‘knowledge transfer’ from DD to GG. We employ a metric-learning loss function and a self-attention answer encoding mechanism for DD that makes it particularly conducive to this knowledge transfer by encouraging perceptually meaningful similarities to emerge. This is especially fruitful since prior work has demonstrated that discriminative dialog models significantly outperform their generative counterparts, but are not as useful since they necessarily need a list of candidate responses to rank, which is only available in a dialog dataset, not in real conversations with a user. In that context, our work aims to achieve the best of both worlds – the practical usefulness of GG and the strong performance of DD – via this knowledge transfer.

Our primary technical contribution is an end-to-end trainable generative visual dialog model, where the generator receives gradients from the discriminator loss of the sequence sampled from GG. Note that this is challenging because the output of GG is a sequence of discrete symbols, which naïvely is not amenable to gradient-based training. We propose to leverage the recently proposed Gumbel-Softmax (GS) approximation to the discrete distribution – specifically, a Recurrent Neural Network (RNN) augmented with a sequence of GS samplers, which when coupled with the straight-through gradient estimator enables end-to-end differentiability.

Our results show that our ‘knowledge transfer’ approach is indeed successful. Specifically, our discriminator-trained GG outperforms the MLE-trained GG by 1.7% on recall@5 on the VisDial dataset, essentially improving over state-of-the-art by 2.43% recall@5 and 2.67% recall@10. Moreover, our generative model produces more diverse and informative responses (see Table 3).

As a side contribution specific to this application, we introduce a novel encoder for neural visual dialog models, which maintains two separate memory banks – one for visual memory (where do we look in the image?) and another for textual memory (what facts do we know from the dialog history?), and outperforms the encoders used in prior work.

Related Work

GANs for sequence generation. Generative Adversarial Networks (GANs) have shown to be effective models for a wide range of applications involving continuous variables (e.g. images) c.f . More recently, they have also been used for discrete output spaces such as language generation – e.g. image captioning , dialog generation , or text generation – by either viewing the generative model as a stochastic parametrized policy that is updated using REINFORCE with the discriminator providing the reward , or (closer to our approach) through continuous relaxation of discrete variables through Gumbel-Softmax to enable backpropagating the response from the discriminator .

There are a few subtle but significant differences w.r.t. to our application, motivation, and approach. In these prior works, both the discriminator and the generator are trained in tandem, and from scratch. The goal of the discriminator in those settings has primarily been to discriminate ‘fake’ samples (i.e. generator’s outputs) from ‘real’ samples (i.e. from training data). In contrast, we would like to transfer knowledge from the discriminator to the generator. We start with pre-trained DD and GG models suited for the task, and then transfer knowledge from DD to GG to further improve GG, while keeping DD fixed. As we show in our experiments, this procedure results in GG producing diverse samples that are close in the embedding space to the ground truth, due to perceptual similarity learned in DD. One can also draw connections between our work and Energy Based GAN (EBGAN) – without the adversarial training aspect. The “energy” in our case is a deep metric-learning based scoring mechanism, instantiated in the visual dialog application.

Modeling image and text attention. Models for tasks at the intersection of vision and language – e.g., image captioning , visual question answering , visual dialog – typically involve attention mechanisms. For image captioning, this may be attending to relevant regions in the image . For VQA, this may be attending to relevant image regions alone or co-attending to image regions and question words/phrases .

In the context of visual dialog, uses attention to identify utterances in the dialog history that may be useful for answering the current question. However, when modeling the image, the entire image embedding is used to obtain the answer. In contrast, our proposed encoder HCIAE (Section 4.1) localizes the region in the image that can help reliably answer the question. In particular, in addition to the history and the question guiding the image attention, our visual dialog encoder also reasons about the history when identifying relevant regions of the image. This allows the model to implicitly resolve co-references in the text and ground them back in the image.

Preliminaries: Visual Dialog

We begin by formally describing the visual dialog task setup as introduced by Das et al. . The machine learning task is as follows. A visual dialog model is given as input an image I\bm{I}, caption c\bm{c} describing the image, a dialog history till round t−1t-1, H=((Q1,A1)c⏟H0,(q1,a1)⏟H1,…,(qt−1,at−1)⏟Ht−1)\bm{H}=(\underbrace{\vphantom{(Q_{1},A_{1})}\bm{c}}_{H_{0}},\underbrace{(\bm{q}_{1},\bm{a}_{1})}_{H_{1}},\ldots,\underbrace{(\bm{q}_{t-1},\bm{a}_{t-1})}_{H_{t-1}}), and the followup question qt\bm{q}_{t} at round tt. The visual dialog agent needs to return a valid response to the question.

Given the problem setup, there are two broad classes of methods – generative and discriminative models. Generative models for visual dialog are trained by maximizing the log-likelihood of the ground truth answer sequence atgt∈At\bm{a}^{gt}_{t}\in\mathcal{A}_{t} given the encoded representation of the input (I,H,qt)\bm{I},\bm{H},\bm{q}_{t}).

On the other hand, discriminative models receive both an encoding of the input (I,H,qt)\bm{I},\bm{H},\bm{q}_{t}) and as additional input a list of 100 candidate answers At={at(1),…,at(100)}\mathcal{A}_{t}=\{\bm{a}^{(1)}_{t},\ldots,\bm{a}^{(100)}_{t}\}. These models effectively learn to sort the list. Thus, by design, they cannot be used at test time without a list of candidates available.

Approach: Backprop Through Discriminative Losses for Generative Training

In this section, we describe our approach to transfer knowledge from a discriminative visual dialog model (DD) to generative visual dialog model (GG). Fig. 1 (a) shows the overview of our approach. Given the input image I\bm{I}, dialog history H\bm{H}, and question qt\bm{q}_{t}, the encoder converts the inputs into a joint representation et\bm{e}_{t}. The generator GG takes et\bm{e}_{t} as input, and produces a distribution over answer sequences via a recurrent neural network (specifically an LSTM). At each word in the answer sequence, we use a Gumbel-Softmax sampler SS to sample the answer token from that distribution. The discriminator DD in it’s standard form takes et\bm{e}_{t}, ground-truth answer atgt\bm{a}_{t}^{gt} and N−1N-1 “negative” answers {at,i−}i=1N−1\{\bm{a}^{-}_{t,i}\}_{i=1}^{N-1} as input, and learns an embedding space such that similarity(et,f(atgt))>similarity(et,f(at,⋅−))\text{similarity}(\bm{e}_{t},f(\bm{a}^{gt}_{t}))>\text{similarity}(\bm{e}_{t},f(\bm{a}^{-}_{t,\cdot})), where f(⋅)f(\cdot) is the embedding function. When we enable the communication between DD and GG, we feed the sampled answer a^t\bm{\hat{a}}_{t} into discriminator, and optimize the generator GG to produce samples that get higher scores in DD’s metric space.

We now describe each component of our approach in detail.

An important characteristic in dialogs is the use of co-reference to avoid repeating entities that can be contextually resolved. In fact, in the VisDial dataset nearly all (98%) dialogs involve at least one pronoun. This means that for a model to correctly answer a question, it would require a reliable mechanism for co-reference resolution.

A common approach is to use an encoder architecture with an attention mechanism that implicitly performs co-reference resolution by identifying the portion of the dialog history that can help in answering the current question . while using a holistic representation for the image. Intuitively, one would also expect that the answer is also localized to regions in the image, and be consistent with the attended history.

With this motivation, we propose a novel encoder architecture (called HCIAE) shown in Fig. 2. Our encoder first uses the current question to attend to the exchanges in the history, and then use the question and attended history to attend to the image, so as to obtain the final encoding.

Specifically, we use the spatial image features V∈Rd×k\bm{V}\in\mathcal{R}^{d\times k} from a convolution layer of a CNN. qt\bm{q}_{t} is encoded with an LSTM to get a vector mtq∈Rd\bm{m}^{q}_{t}\in\mathcal{R}^{d}. Simultaneously, each previous round of history (H0,…,Ht−1)({H}_{0},\ldots,{H}_{t-1}) is encoded separately with another LSTM as Mth∈Rd×t\bm{M}^{h}_{t}\in\mathcal{R}^{d\times t}. Conditioned on the question embedding, the model attends to the history. The attended representation of the history and the question embedding are concatenated, and used as input to attend to the image:

where We∈Rd×3d\bm{W}_{e}\in\mathcal{R}^{d\times 3d} is weight parameters and [⋅][\cdot] is the concatenation operation.

2 Discriminator Loss

Discriminative visual dialog models produce a distribution over the candidate answer list At\mathcal{A}_{t} and maximize the log-likelihood of the correct option atgt\bm{a}_{t}^{gt}. The loss function for DD needs to be conducive for knowledge transfer. In particular, it needs to encourage perceptually meaningful similarities. Therefore, we use a metric-learning multi-class N-pair loss defined as:

where ff is an attention based LSTM encoder for the answer. This attention can help the discriminator better deal with paraphrases across answers. The attention weight is learnt through a 1-layer MLP over LSTM output at each time step. The N-pair loss objective encourages learning a space in which the ground truth answer is scored higher than other options, and at the same time, encourages options similar to ground truth answers to score better than dissimilar ones. This means that, unlike the multiclass logistic loss, the options that are correct but different from the correct option may not be overly penalized, and thus can be useful in providing a reliable signal to the generator. See Fig. 1 for an example. Follwing , we regularize the L2 norm of the embedding vectors to be small.

3 Discriminant Perceptual Loss and Knowledge Transfer from D𝐷D to G𝐺G

At a high-level, our approach for transferring knowledge from DD to GG is as follows: GG repeatedly queries DD with answers a^t\bm{\hat{a}}_{t} that it generates for an input embedding et\bm{e}_{t} to get feedback and update itself. In each such update, GG’s goal is to update its parameters to try and have a^t\bm{\hat{a}}_{t} score higher than the correct answer, atgt\bm{a}^{gt}_{t}, under DD’s learned embedding and scoring function. Formally, the perceptual loss that GG aims to optimize is given by:

where ff is the embedding function learned by the discriminator as in (4). Intuitively, updating generator parameters to minimize LG\mathcal{L}_{G} can be interpreted as learning to produce an answer sequence a^t\bm{\hat{a}}_{t} that ‘fools’ the discriminator into believing that this answer should score higher than the human response atgt\bm{a}_{t}^{gt} under the discriminator’s learned embedding f(⋅)f(\cdot) and scoring function.

While it is straightforward to sample an answer a^t\bm{\hat{a}}_{t} from the generator and perform a forward pass through the discriminator, naïvely, it is not possible to backpropagate the gradients to the generator parameters since sampling discrete symbols results in zero gradients w.r.t. the generator parameters. To overcome this, we leverage the recently introduced continuous relaxation of the categorical distribution – the Gumbel-softmax distribution or the Concrete distribution .

At an intuitive level, the Gumbel-Softmax (GS) approximation uses the so called ‘Gumbel-Max trick’ to reparametrize sampling from a categorical distribution and replaces argmax with softmax to obtain a continuous relaxation of the discrete random variable. Formally, let x\bm{x} denote a KK-ary categorical random variable with parameters denoted by (p1,…pK)(p_{1},\ldots p_{K}), or x∼Cat(p)\bm{x}\sim Cat(\bm{p}). Let \big{(}g_{i}\big{)}_{1}^{K} denote KK IID samples from the standard Gumbel distribution, gi∼F(g)=e−e−gg_{i}\sim F(g)=e^{-e^{-g}}. Now, a sample from the Concrete distribution can be produced via the following transformation:

where τ\tau is a temperature parameter that control how close samples y\bm{y} from this Concrete distribution approximate the one-hot encoding of the categorical variable x\bm{x}.

As illustrated in Fig. 1, we augment the LSTM in GG with a sequence of GS samplers. Specifically, at each position in the answer sequence, we use a GS sampler to sample an answer token from that conditional distribution. When coupled with the straight-through gradient estimator this enables end-to-end differentiability. Specifically, during the forward pass we discretize the GS samples into discrete samples, and in the backward pass use the continuous relaxation to compute gradients. In our experiments, we held the temperature parameter fixed at 0.5.

Experiments

Dataset and Setup. We evaluate our proposed approach on the VisDial dataset , which was collected by Das et al. by pairing two subjects on Amazon Mechanical Turk to chat about an image. One person was assigned the role of a ‘questioner’ and the other of ‘answerer’. One worker (the questioner) sees only a single line of text describing an image (caption from COCO ); the image remains hidden to the questioner. Their task is to ask questions about this hidden image to “imagine the scene better”. The second worker (the answerer) sees the image and caption and answers the questions. The two workers take turns asking and answering questions for 10 rounds. We perform experiments on VisDial v0.9 (the latest available release) containing 83k dialogs on COCO-train and 40k on COCO-val images, for a total of 1.2M dialog question-answer pairs. We split the 83k into 82k for train, 1k for val, and use the 40k as test, in a manner consistent with . The caption is considered to be the first round in the dialog history.

Evaluation Protocol. Following the evaluation protocol established in , we use a retrieval setting to evaluate the responses at each round in the dialog. Specifically, every question in VisDial is coupled with a list of 100 candidate answer options, which the models are asked to sort for evaluation purposes. DD uses its score to rank these answer options, and GG uses the log-likelihood of these options for ranking. Models are evaluated on standard retrieval metrics – (1) mean rank, (2) recall @k@k, and (3) mean reciprocal rank (MRR) – of the human response in the returned sorted list.

Pre-processing. We truncate captions/questions/answers longer than 24/16/8 words respectively. We then build a vocabulary of words that occur at least 5 times in train, resulting in 8964 words.

Training Details In our experiments, all 3 LSTMs are single layer with 512d512d hidden state. We use VGG-19 to get the representation of image. We first rescale the images to be 224×224224\times 224 pixels, and take the output of last pooling layer (512×7×7512\times 7\times 7) as image feature. We use the Adam optimizer with a base learning rate of 4e-4. We pre-train GG using standard MLE for 20 epochs, and DD with supervised training based on Eq (4) for 30 epochs. Following , we regularize the L2L^{2} norm of the embedding vectors to be small. Subsequently, we train GG with LG+αLMLE\mathcal{L}_{G}+\alpha\mathcal{L}_{MLE}, which is a combination of discriminative perceptual loss and MLE loss. We set α\alpha to be 0.5. We found that including LMLE\mathcal{L}_{MLE} (with teacher-forcing) is important for encouraging GG to generate grammatically correct responses.

Baselines. We compare our proposed techniques to the current state-of-art generative and discriminative models developed in . Specifically, introduced 3 encoding architectures – Late Fusion (LF), Hierarchical Recurrent Encoder (HRE), Memory Network (MN) – each trained with a generative (-G) and discriminative (-D) decoder. We compare to all 6 models.

Our approaches. We present a few variants of our approach to systematically study the individual contributions of our training procedure, novel encoder (HCIAE), self-attentive answer encoding (ATT), and metric-loss (NP).

HCIAE-G-MLE is a generative model with our proposed encoder trained under the MLE objective. Comparing this variant to the generative baselines from establishes the improvement due to our encoder (HCIAE).

HCIAE-G-DIS is a generative model with our proposed encoder trained under the mixed MLE and discriminator loss (knowledge transfer). This forms our best generative model. Comparing this model to HCIAE-G-MLE establishes the improvement due to our discriminative training.

HCIAE-D-MLE is a discriminative model with our proposed encoder, trained under the standard discriminative cross-entropy loss. The answer candidates are encoded using an LSTM (no attention). Comparing this variant to the discriminative baselines from establishes the improvement due to our encoder (HCIAE) in the discriminative setting.

HCIAE-D-NP is a discriminative model with our proposed encoder, trained under the n-pair discriminative loss (as described in Section 4.2). The answer candidates are encoded using an LSTM (no attention). Comparing this variant to HCIAE-D-MLE establishes the improvement due to the n-pair loss.

HCIAE-D-NP-ATT is a discriminative model with our proposed encoder, trained under the n-pair discriminative loss (as described in Section 4.2), and using the self-attentive answer encoding. Comparing this variant to HCIAE-D-NP establishes the improvement due to the self-attention mechanism while encoding the answers.

Results. Tables 5.1, 5.1 present results for all our models and baselines in generative and discriminative settings. The key observations are:

Main Results for HCIAE-G-DIS: Our final generative model with all ‘bells and whistles’, HCIAE-G-DIS, uniformly performs the best under all the metrics, outperforming the previous state-of-art model MN-G by 2.43% on R@5. This shows the importance of the knowledge transfer from the discriminator and the benefit from our encoder architecture.

Knowledge transfer vs. encoder for GG: To understand the relative importance of the proposed history conditioned image attentive encoder (HCIAE) and the knowledge transfer, we compared the performance of HCIAE-G-DIS with HCIAE-G-MLE, which uses our proposed encoder but without any feedback from the discriminator. This comparison highlights two points: first, HCIAE-G-MLE improves R@5 by 0.7% over the current state-of-art method (MN-D) confirming the benefits of our encoder. Secondly, and importantly, its performance is lower than HCIAE-G-DIS by 1.7% on R@5, confirming that the modifications to encoder alone will not be sufficient to gain improvements in answer generation; knowledge transfer from DD greatly improves GG.

Metric loss vs. self-attentive answer encoding: In the purely discriminative setting, our final discriminative model (HCIAE-D-NP-ATT) also beats the performance of the corresponding state-of-art models by 2.53% on R@5. The n-pair loss used in the discriminator is not only helpful for knowledge transfer but it also improves the performance of the discriminator by 0.85% on R@5 (compare HCIAE-D-NP to HCIAE-D-MLE). The improvements obtained by using the answer attention mechanism leads to an additional, albeit small, gains of 0.4% on R@5 to the discriminator performance (compare HCIAE-D-NP to HCIAE-D-NP-ATT).

2 Does updating discriminator help?

Recall that our model training happens as follows: we independently train the generative model HCIAE-G-MLE and the discriminative model HCIAE-D-NP-ATT. With HCIAE-G-MLE as the initialization, the generative model is updated based on the feedback from HCIAE-D-NP-ATT and this results in our final HCIAE-G-DIS.

We performed two further experiments to answer the following questions:

What happens if we continue training HCIAE-D-NP-ATT in an adversarial setting? In particular, we continue training by maximizing the score of the ground truth answer atgt\bm{a}^{gt}_{t} and minimizing the score of the generated answer a^t\bm{\hat{a}}_{t}, effectively setting up an adversarial training regime LD=−LG\mathcal{L}_{D}=-\mathcal{L}_{G}. The resulting discriminator HCIAE-GAN1 has significant drop in performance, as can be seen in Table. 4 (32.97% R@5). This is perhaps expected because HCIAE-GAN1 updates its parameters based on only two answers, the ground truth and the generated sample (which is likely to be similar to ground truth). This wrecks the structure that HCIAE-D-NP-ATT had previously learned by leveraging additional incorrect options.

What happens if we continue structure-preserving training of HCIAE-D-NP-ATT? In addition to providing HCIAE-D-NP-ATT samples from GG as fake answers, we also include incorrect options as negative answers so that the structure learned by the discriminator is preserved. HCIAE-D-NP-ATT continues to train under loss LD\mathcal{L}_{D}. In this case (HCIAE-GAN2 in Table. 4), we find that there is a small improvement in the performance of GG. The additional computational overhead to training the discriminator supersedes the performance improvement. Also note that HCIAE-D-NP-ATT itself gets worse at the dialog task.

One might wonder, why not train a GAN for visual dialog? Formulating the task in a GAN setting would involve GG and DD training in tandem with DD providing feedback as to whether a response that GG generates is real or fake. We found this to be a particularly unstable setting, for two main reasons: First, consider the case when the ground truth answer and the generated answers are the same. This happens for answers that are typically short or ‘cryptic’ (e.g. ‘yes’). In this case, DD can not train itself or provide feedback, as the answer is labeled both positive and negative. Second, in cases where the ground truth answer is descriptive but the generator provides a short answer, DD can quickly become powerful enough to discard generated samples as fake. In this case, DD is not able to provide any information to GG to get better at the task. Our experience suggests that the discriminator, if one were to consider a ‘GANs for visual dialog’ setting, can not merely be focused on differentiating fake from real. It needs to be able to score similarity between the ground truth and other answers. Such a scoring mechanism provides a more reliable feedback to GG. In fact, as we show in the previous two results, a pre-trained DD that captures this structure is the key ingredient in sharing knowledge with GG. The adversarial training of DD is not central.

3 Qualitative Comparison

In Table 3 we present a couple of qualitative examples that compares the responses generated by G-MLE and G-DIS. G-MLE predominantly produces ‘safe’ and less informative answers, such as ‘Yes’ and or ‘I can’t tell’. In contrast, our proposed model G-DIS does so less frequently, and often generates more diverse yet informative responses.

Conclusion

Generative models for (visual) dialog are typically trained with an MLE objective. As a result, they tend to latch on to safe and generic responses. Discriminative (or retrieval) models on the other hand have been shown to significantly outperform their generative counterparts. However, discriminative models can not be deployed as dialog agents with a real user where canned candidate responses are not available. In this work, we propose transferring knowledge from a powerful discriminative visual dialog model to a generative model. We leverage the Gumbel-Softmax (GS) approximation to the discrete distribution –specifically, a RNN augmented with a sequence of GS samplers, coupled with a ST gradient estimator for end-to-end differentiability. We also propose a novel visual dialog encoder that reasons about image-attention informed by the history of the dialog; and employ a metric learning loss along with a self-attentive answer encoding to enable the discriminator to learn meaningful structure in dialog responses. The result is a generative visual dialog model that significantly outperforms state-of-the-art.

References