Learning Disentangled Textual Representations via Statistical Measures of Similarity
Pierre Colombo, Guillaume Staerman, Nathan Noiry, Pablo Piantanida
Introduction
As natural language processing (NLP) systems are taken up in an ever wider array of sectors (e.g., legal system Dale (2019), insurance Ly et al. (2020), education Litman (2016), healthcare Basyal et al. (2020)), there are growing concerns about the harmful potential of bias in such systems Leidner and Plachouras (2017). Recently, a large body of research aims at analyzing, understanding and addressing bias in various applications of NLP including language modelling Liang et al. (2021), machine translation Stanovsky et al. (2019), toxicity detection Dixon et al. (2018) and classification Elazar and Goldberg (2018). In NLP, current systems often rely on learning continuous embedding of the input text. Thus, it is crucial to ensure that the learnt continuous representations do not exhibit bias that could cause representational harms Blodgett et al. (2020); Barocas et al. (2017), i.e., representations less favourable to specific social groups. One way to prevent the aforementioned phenomenon is to enforce disentangled representations, i.e., representations that are independent of a sensitive attribute (see Fig. 1 for a visualization of different degrees of disentangled representations).
Learning disentangled representations has received a growing interest as it has been shown to be useful for a wide variety of tasks (e.g., style transfer Fu et al. (2017), few shot learning Karn et al. (2021), fair classification Colombo et al. (2021d)). For text, the dominant approaches to learn such representations can be divided into two classes. The first one, relies on an adversary that is trained to recover the discrete sensitive attribute from the latent representation of the input Xie et al. (2017). However, as pointed out by Barrett et al. (2019), even though the adversary seems to do a perfect job during training, a fair amount of the sensitive information can be recovered from the latent representation when training a new adversary from scratch. The second line of research involves a regularizer that is a trainable surrogate of the mutual information (MI) (e.g., CLUB Cheng et al. (2020a), MIReny Colombo et al. (2021d), KNIFE Pichler et al. (2020), MINE Belghazi et al. (2018); Colombo et al. (2021b)) and achieves higher degrees of disentanglement. However, as highlighted by recent works McAllester and Stratos (2020); Song and Ermon (2019), these estimators are hard to use in practice and the optimization procedure (see App. D.4) involves several updates of the regularizer parameters at each update of the representation model. As a consequence, these procedures are both time consuming and involve extra hyperparameters (e.g., optimizer learning rates, architecture, number of updates of the nested loop) that need to be carefully selected which is often not such an easy task.
Contributions. In this work, we focus our attention on learning to disentangle textual representations from a discrete attribute. Our method relies on a novel family of regularizers based on discrepancy measures. We evaluate both the disentanglement and representation quality on fair text classification. Formally, our contribution is two-fold: (1) A novel formulation of the problem of learning disentangled representations. Different from previous works–either minimizing a surrogate of MI or training an adversary–we propose to minimize a statistical measure of similarity between the underlying probability distributions conditioned to the sensitive attributes. This novel formulation allows us to derive new regularizers with convenient properties: (i) not requiring additional learnable parameters; (ii) alleviating computation burden; and (iii) simplifying the optimization dynamic. (2) Applications and numerical results. We carefully evaluate our new framework on four different settings coming from two different datasets. We strengthen the experimental protocol of previous works Colombo et al. (2021d); Ravfogel et al. (2020) and test our approach both on randomly initialized encoder (using RNN-based encoder) and during fine-tuning of deep contextualized pretrained representationsPrevious works (e.g., Ravfogel et al. (2020)) do not fine-tune the pretrained encoder when testing their methods.. Our experiments are conducted on four different main/sensitive attribute pairs and involve the training of over deep neural networks. Our findings show that: (i) disentanglement methods behave differently when applied to randomly initialized or to deep contextualized pretrained encoder; and (ii) our framework offers a better accuracy/disentanglement trade-off than existing methods (i.e., relying on an adversary or on a MI estimator) while being faster and easier to train. Model, data and code are available at https://github.com/PierreColombo/TORNADO.
Related Work
where refers to the main classifier; to its learnable parameters; CE to the cross-entropy loss; R denotes the disentanglement regularizer; its parameters and controls the trade-off between disentanglement and success in the classification task. We next review the two main methods that currently exist for learning textual disentangled representations: adversarial-based and MI-based.
2 MI-Based Regularizers
To better protect sensitive information, the second class of methods involves direct mutual information minimization. MI lies at the heart of information theory and measures statistical dependencies between two random variables and and find many applications in machine learning Boudiaf et al. (2020b, a, 2021). The MI is a non-negative quantity that is 0 if and only if and are independent and is defined as follows:
3 Limitations of Existing Methods
The aforementioned methods involve the use of extra parameters (i.e., ) in the regularizer. As the regularizer computes a quantity based on the representation given by the encoder with parameter , any modification of requires an adaptation of the parameter of (i.e., ). In practice, this adaptation is performed using gradient descent-based algorithms and requires several gradient updates. Thus, a nested loop (see App. D.4) is needed. Additional optimization parameters and the nested loop both induce additional complexity and require a fine-tuning which makes these procedures hard to be used on large-scale datasets. To alleviate these issues, the next section describes a parameter-free framework to get rid of the parameter present in .
Proposed Method
This section describes our approach to learn disentangled representations. We first introduce the main idea and provide an algorithm to implement the general loss. We next describe the four similarity measures proposed in this approach.
The proposed statistical measures of similarity, detailed in Section 3.2, have explicit and simple formulas. It follows that the use of neural networks is no longer necessary in the regularizer term which reduces drastically the complexity of the resulting learning problem. The disentanglement can be controlled by selecting appropriately the measure SM. For the sake of place, the algorithm we propose to solve (3) is deferred to the App. B.
2 Measure of Similarity between Distributions
In this work, we choose to focus on four different (dis-) similarity functions ranging from the most popular in machine learning such as the Maximum Mean Discrepancy measure (MMD) and the Sinkhorn divergence (SD) to standard statistical discrepancies such as the Jeffrey divergence (J) and the Fisher-Rao distance (FR).
The MMD can be estimated with a quadratic computational complexity where is the sample size. In this paper, MMD is computed using the Gaussian kernel , where is the usual euclidean norm.
2.2 Sinkhorn Divergence.
The Wasserstein distance aims at comparing two probability distributions through the resolution of the Monge-Kantorovich mass transportation problem (see e.g. Villani (2003); Peyré and Cuturi (2019)):
with .
2.3 Fisher-Rao Distance.
where is the univariate Fisher-Rao detailed in the App. A.1 for the sake of space.
2.4 Jeffrey Divergence.
The Jeffrey divergence (J) is a symmetric version of the Kullback-Leibler (KL) divergence and measures the similarity between two probability distributions. Formally, it is defined as follow:
where is the trace of .
FR and J are computed under the multivariate Gaussian with diagonal covariance matrix assumption. In this case, the Sinkhorn approximation is not needed as (4) can be efficiently computed thanks to the following closed-form:
Quantities defined in this section are replaced by their empirical estimate. Due to space constraints, the formula are described in App. A.2.
Experimental Setting
In this section, we describe the datasets, metrics, encoder and baseline choices. Additional experimental details can be found in App. D. For fair comparison, all models were re-implemented.
To ensure backward comparison with previous works, we choose to rely on the DIAL Blodgett et al. (2016) and the PAN Rangel et al. (2014) datasets. For both, main task labels () and sensitive labels () are binary, balanced and splits follow Barrett et al. (2019). Random guessing is expected to achieve near 50% of accuracy. The DIAL corpus has been automatically built from tweets and the main task is either polarityPolarity or emotion have been widely studied in the NLP community Jalalzai et al. (2020); Colombo et al. (2019) or mention prediction. The sensitive attribute is related to race (i.e., non-Hispanic blacks and non-Hispanic whites) which is obtained using the author geo-location and the words used in the tweet. The PAN corpus is also composed of tweets and the main task is to predict a mention label. The sensitive attribute is obtained through a manual process and annotations contain the age and gender information from 436 Twitter users.
2 Metrics
For the choice of the evaluation metrics, we follow the experimental setting of Colombo et al. (2021d); Elazar and Goldberg (2018); Coavoux et al. (2018). To measure the success of the main task, we report the classification accuracy. To measure the degree of disentanglement of the latent representation we train from scratch an adversary to predict the sensitive labels from the latent representation. In this framework, a perfect model would achieve a high main task accuracy (i.e., near 100%) and a low (i.e., near 50%) accuracy as given by the adversary prediction on the sensitive labels. Following Colombo et al. (2021d), we also report the disentanglement dynamic following variations of and train a different model for each .
3 Models
Choice of the encoder. Previous works that aim at learning disentangled representations either focus on randomly initialized RNN-encoders Colombo et al. (2021d); Elazar and Goldberg (2018); Coavoux et al. (2018) or only use pretrained representations as a feature extractor Ravfogel et al. (2020). In this work, we choose to fine-tune BERT during training as we believe it to be a more realistic setting. Choice of the baseline models. We choose to compare our methods against adversarial training from Elazar and Goldberg (2018); Coavoux et al. (2018) (model named ADV) and the recently MI bound introduced in Colombo et al. (2021d) (named MI) which has been shown to be more controllable than previous MI-based estimators.
Numerical Results
In this section, we gather experimental results for fair classification task. We study our framework when working either with RNN or BERT encoders. The parameter (see (3)) controls the trade-off between success on the main task and disentanglement for all models.
General observations. Learning disentangled representations is made more challenging when and are tightly entangled. By comparing Fig. 2 and Fig. 3, we notice that the race label (main task) is easier to disentangled from the sentiment compared to the mention. Randomly initialized RNN encoders. To allow a fair comparison with previous works, we start by testing our framework with RNN encoders on the DIAL dataset. Results are depicted in Fig. 2. It is worth mentioning that we are able to observe a similar phenomenon that the one reported in Colombo et al. (2021d). More specifically, we observe: (i) the adversary degenerates for and does not allow to reach perfectly disentangled representations nor to control the desirable degree of disentanglement; (ii) the MI allows better control over the desirable degree of disentanglement and achieves better-disentangled representations at a reduced cost on the main task accuracy. Fig. 2 shows that the encoder trained using the statistical measures of similarity–both with and without the multivariate Gaussian assumption–are able to learn disentangled representations. We can also remark that our losses follow an expected behaviour: when increases, more weight is given to the regularizer, the sensitive task accuracy decreases, thus the representations are more disentangled according to the probing-classifier. Overall, we observe that the W regularizer is the best performer with optimal performance for on both attributes. On the other hand, we observe that FR and J divergence are useful to learn to disentangle the representations but disentangling using these similarity measures comes with a greater cost as compared to W. Both MMD and SD also perform wellFor both losses when we did not remark any consistent improvements. and are able to learn disentangled representations with little cost on the main task performance. However, on DIAL, they are not able to learn perfectly disentangled representations. Similar conclusions can be drawn on PAN and results are reported in App. C.1.
BERT encoder. Results of the experiment conducted with BERT encoder are reported in Fig. 3. As expected, we notice that on both tasks the main and the sensitive task accuracy for small values of is higher than when working with RNN encoders. When training a classifier without disentanglement constraints (i.e., case in (1)), which corresponds to the dash lines in Fig. 2 and Fig. 3, we observe that BERT encoder naturally preserves more sensitive information (i.e., measured by the accuracy of the adversary) than randomly initialized encoder. Contrarily to what is usually undertaken in previous works (e.g., Ravfogel et al. (2020)), we allow the gradient to flow in BERT encoder while preforming fine-tuning. We observe a different behavior when compared to previous experiments. Our losses under the Multivariate diagonal Gaussian assumption (i.e., W, J, FR ) can only disentangle the representations at a high cost on the main task (i.e., perfect disentanglement corresponds to performance on the main task close to a random classifier). When training the encoder with either SD or MMD, we are able to learn disentangled representations with a limited cost on the main task accuracy: achieves good disentanglement with less than 3% of loss in the main task accuracy. The methods allow little control over the degree of disentanglement and there is a steep transition between light protection with no loss on the main task accuracy and strong protection with discriminative features destruction.
Takeaways. Our new framework relying on statistical Measures of Similarity introduces powerful methods to learn disentangled representations. When working with randomly initialized RNN encoders to learn disentangled representation, we advise relying on W. Whereas in presence of pretrained encoders (i.e., BERT), we observe a very different behavior To the best of our knowledge, we are the first to report such a difference in behavior when disentangling attributes with pretrained representations. and recommend using SD.
2 Speed Gain and Parameter Reduction
We report in Table 2 the training time and the number of parameters of each method. The reduced number of parameters brought by our method is marginal, however getting rid of these parameters is crucial. Indeed, they require a nested loop and require a fined selection of the hyperparameters which complexify the global system dynamic. Takeaways. Contrarily to MI or Adversarial based regularizer that are difficult (or even prohibitive) to be implemented on large-scale datasets, our framework is simpler and consistently faster which makes it a better candidate when working with large-scale datasets.
Further Analysis
Results presented in Section 5.1 have shown a different behaviour for RNN and BERT based encoders and, for different measures of similarity. Here, we aim at understanding of this phenomena.
In the previous section, we examine the change of the measures during the training. Takeaways. When using a RNN encoder, the system is able to maximize the main task accuracy while jointly minimizing most of the similarity measures. For BERT where the model is more complex, for measures relying on the diagonal gaussian multivariate assumption either the disentanglement plateau (e.g., FR or J) or the system fails to learn discriminative features and perform poorly on the main task (e.g., W). When combined with BERT both SD and MMD can achieve high main task accuracy while protecting the sensitive attribute.
2 Correlation Analysis
Takeaways. Both ADV and MI poorly are correlated with the degree of disentanglement of the learned representations. We find this result not surprising at light of the findings of Xie et al. (2017) and Song and Ermon (2019). All our losses achieve high correlation () except for J in the mention task with both encoders, and the FR with BERT on the mention task that achieves medium/low correlation. We believe, that the high correlation showcases the validity of the proposed approaches.
Summary and Concluding Remarks
We have introduced a new framework for learning disentangled representations which is faster to train, easier to tune and achieves better results than adversarial or MI-based methods. Our experiments on the fair classification task show that for RNN encoders, our methods relying on the closed-form of similarity measures under a multivariate Gaussian assumption can achieve perfectly disentangled representations with little cost on the main tasks (e.g. using Wasserstein). On BERT representations, our experiments show that the Sinkhorn divergence should be preferred. It can achieves almost perfect disentanglement at little cost but allows for fewer control over the degree of disentanglement.
Acknowledgments
This work was also granted access to the HPC resources of IDRIS under the allocation 2021-AP010611665 as well as under the project 2021-101838 made by GENCI. This work has been supported by the project PSPC AIDA: 2019-PSPC-09 funded by BPI-France.
References
Appendix A Additional details on Statistical Measures of Similarity
It is the purpose of this part to recall additional details on similarity measures defined in the core paper.
A.2 Empirical versions of Statistical Measures of Similarity
Maximum Mean Discrepancy. The MMD is defined as:
where is the euclidean distance between and . We limit ourselves to the 1-Wasserstein for the sake of place. The Sinkhorn divergence is then:
It is worth noticing that a robust version of the Wasserstein distance can be found in Staerman et al. (2021a) (see also Staerman et al. (2021c)).
Fisher-Rao distance. The Fisher-Rao distance is defined as
where is defined as in Section A.1, and and are replaced by and the classical (univariate) unbiased mean and standard deviation estimators respectively.
Jeffrey divergence. Let and be the mean and the covariance matrix estimators of the samples and respectively. Jeffrey divergence–under the multivariate Gaussian assumption–boils down to:
Furthermore, under the multivariate Gaussian assumption, the Wasserstein distance writes as follows:
Appendix B Algorithm
The algorithm we propose to compute (3) involves a simple training loop and is described in Algorithm 1.
Appendix C Additional Results
In this section, we gather additional experimental results.
We report in Fig. 6 and Fig. 7 the results of the disentanglement analysis on the PAN dataset. RNN encoders. We can make the same observations that the one done on DIAL in Section 5.1. We observe that the W regularizer performs well and is among the most controllable loss. It is worth noting the good performance of the SD and MMD losses which both work well on the RNN encoder. BERT. For BERT encoder, we observe a similar steep transition than in Section 5.1 and we can draw similar conclusions. FR , W and J fail to disentangle BERT representation with little cost on the main task. SD and MMD achieve good results. Takeaways. When working with randomly initialized RNN encoders to learn disentangled representation we advise relying on W and when working with pretrained encoder we advise to rely on the SD.
C.2 On the Diagonal Gaussian Assumption
Our closed-form for the Fisher-Rao metric relies on the diagonal Gaussian assumption that we have also made for W and J for a fair comparison. In this experiment (see Fig. 8), we examine this assumption by evaluating the relative distance (using a -norm) between the empirical covariance matrix and a diagonal matrix. Takeaways. Interestingly, as increases, the empirical covariance matrix becomes closer to a diagonal matrix. For BERT, we observe that the W saturates and the distance for is higher than for RNN. This might be the result of the optimization problems identified in Fig. 4. Hence, we observe that our methods–when learning more disentangled representations–is that the covariance matrix becomes closer to a diagonal matrix.
Appendix D Experimental Details
In this section we gather the model details we used in our experiments. All models rely on the tokenizer based on Word Piece Schuster and Nakajima (2012) and is similar to the one used for BERT (i.e bert-base-uncased) and possess over tokens.
For the randomly initialized RNN encoder, we use a bidirectionnal GRU Chung et al. (2014) that is composed of 2 layers with an hidden dimension of 128. For activation, we use LeakyReLU Xu et al. (2015) and the classification head is composed of fully connected layers of input dimension 256. The learning rate of AdamW Loshchilov and Hutter (2017) is set to and the dropout Srivastava et al. (2014) is set to 0.2. The number of warmup steps Vaswani et al. (2017) is set to 1000.
For all 140 models, we train on NVIDIA-V100 with 32GB of RAM. Each model is trained for 30k steps and the model with the best disentanglement accuracy is selected based on the validation set. Each model takes around 5 hours to train. Evaluation requires to train and adversary composed of 3 hidden layers of input 128-128-128-2. The evaluation which involves the training of the probing classifier takes below 1 hour of GPU time. Overall, we train 6 different classifiers per model which correspond to 840 models.
For the BERT encoder, we add a classification head composed of one fully connected layer. We use a learning rate of for AdamW and he number of warmup steps is set to 1000.
For all the 140 models, we train on NVIDIA-V100 with 32GB of RAM. Each model is trained for 10k steps, which correspond to the convergence of the model and the model with the best disentanglement accuracy is selected based on the validation set. Each model takes approximately 3 hours to train. Evaluation requires to train and adversary composed of 3 hidden layers and involes LeakyRely and dropout rate of 0.1 of input 768-768-768-2. The evaluation which involves the training of the probing classifier takes below 1 hour of GPU time. Overall, we train 6 different classifiers per model which correspond to 840 models.
D.2 Negative Results
We briefly describe a few ideas that did not look promising in our experiments to help future research. Specifically,
We attempt to combine our work with MINE from Belghazi et al. (2018) and we observe high instability during the training.
We additionally used the clozed-form of MMD under a multivariate Gaussian assumption which lead to poor results (there was no protection against the classifier).
We also used the Hausdoff distance Serra (1998) which interpolates between the Iterative closest point Chetverikov et al. (2002) loss and a kernel distance as well as MMD with Laplacian kernel Kondor and Pan (2016). For both case, we ended with optimization issues and poor trade-offs.
D.3 Dataset Examples
For completness, we gather in this section examples of the DIAL and PAN corpus. Note that this samples have been randomly selected. We report in Table 4 same randomly sampled examples text from the DIAL corpus and order them based on the sensitive attribute race. The polarity label is obtained through emojis. The goal of the mention task is to predict if a tweet is conversational (i.e., contains a @mentions tokens)
We report in Tab. 5 examples from the PAN corpus. The age attribute is obtained through birth-date published on the user’s Linkedin profile whereas for the gender the authors rely on both the user’s name and photograph.
D.4 Related Work General Algorithm
For completeness we provide in Algorithm 2 the algorithm used for training adversarial or MI-based regularizers. It is worth noting that these baselines require extra learnable parameters that need to be tuned using a Nested Loop.
D.5 Future Work.
As future work we plan to disentangled more complex labels such as dialog acts Colombo et al. (2020, 2021a), emotions Witon et al. (2018) and linguistic phenomena such as disfluencies Dinkar et al. (2020) and other spoken language phenomenon Chapuis et al. (2020). Future research also include extending these losses to data augmentation Dhole et al. (2021); Colombo et al. (2021e) and sentence generation Colombo et al. (2021c, f) and study the trade-off using rankings Colombo et al. (2022) or anomaly detection Staerman et al. (2019, 2020, 2021b, 2022).
D.6 Libraries used.
For this project among the library we used we can cite:
Pytorch Paszke et al. (2017) for the GPU support.
Geomloss Feydy et al. (2019) for the SD and MMD. It can be found at https://www.kernel-operations.io/geomloss