Learning Fair Representation via Distributional Contrastive Disentanglement

Changdae Oh, Heeji Won, Junhyuk So, Taero Kim, Yewon Kim, Hosik Choi, Kyungwoo Song

Introduction

Machine learning algorithms show the great success on many tasks, and they have been widely adopted in real-world applications. Most works of machine learning utilize a simple predictor at the test time, and the extracted features of a given dataset greatly influence the model prediction. So, the success of machine learning algorithms largely depends on the data representation that the model learned (Bengio et al., 2013). However, the input features of a given dataset might contain noise and unnecessary information for the given task, and learning a representation that covers the important characteristics of the given dataset while invariant to unwanted information is crucial.

Neural networks are known to have an advantage in representation learning. The learned representation from the neural net represents the characteristics of data with fewer dimensional vectors. Neural network based methods show a significant performance improvement on many domains including image (He et al., 2016), text (Kenton and Toutanova, 2019), and tabular (Xu et al., 2019). However, recent studies raise that the traditional neural networks have difficulty in achieving fairness and domain generalization (Sarhan et al., 2020; Arjovsky et al., 2019). Neural networks provide a rich representation, but the representation also absorbs the sensitive (private) information or spurious correlation of a given dataset. Due to this unwanted information in learned representations, the model can induce unfair results in decision making system or fail to provide the correct predictions in the distribution shift scenario.

There are many directions to remove the sensitive information and spurious correlation in the model. One of them, adversarial representation learning (ARL) methods (Xie et al., 2017; Roy and Boddeti, 2019) set two goals, (i) maximally retain salient information about a given target attribute, and (ii) minimize the information leakage about a given sensitive attribute. While these methods have shown compelling results when optimized successfully, their convergence instability hinders achieving the above goals and limits the wide use of ARL methods practically. Another promising direction is a disentangled representation learning (Creager et al., 2019; Sarhan et al., 2020) that separates the non-sensitive representation and sensitive representation. Their empirical success proves the effectiveness of disentanglement-based approaches to fair expression learning. However, most of the existing methods for disentangling (Kim and Mnih, 2018; Creager et al., 2019) also rely on the adversarial learning technique to approximate Total Correlation (Watanabe, 1960) using the density ratio trick.

The other direction is to learn the invariant representation of a given dataset. Invariant representation denotes the shared important features invariant across domains or environments. Recently, Arovsky et al. (Arjovsky et al., 2019) proposed Invariant Risk Minimization (IRM) that encourages invariant representation learning with bi-level optimization. The representation trained with IRM may have shared important features, but there is no guarantee that the representation does not hold any sensitive information. Besides, Group-DRO (Sagawa et al., 2020) is known to be effective for learning robust representation, but it is also not guaranteed about the existence of sensitive information. Although invariant learning methods (Arjovsky et al., 2019; Creager et al., 2021; Sagawa et al., 2020) show meaningful performance improvements on domain generalization tasks, we observe that the representation learned from IRM and Group-DRO still has a spurious correlation or sensitive information largely that poses a potential risk in real-world applications.

In this paper, we propose FarconVAE (FAir Representation via distributional CONtrastive Variational AutoEncoder), a new disentangling approach with contrastive learning instead of adversarial learning. First, we construct a pair of two instances with different sensitive information and the same non-sensitive information. Then, FarconVAE 1) minimizes the distance between non-sensitive representations, 2) maximizes the dissimilarity between sensitive representations, and 3) maximizes the dissimilarity between sensitive and non-sensitive representations of each paired instance with our distributional contrastive loss. Finally, in the latent space divided into sensitive and non-sensitive representation, FarconVAE makes a fair prediction by using only non-sensitive representation.

Besides, we adopt a new feature swap based objective for FarconVAE, swap-recon, to improve the disentanglement. Swap-recon replaces the non-sensitive representation of a given data instance with that of another paired data instance while leaving the sensitive representation untouched. By reconstructing from both the swapped latent and the original latent to the same original data instance, swap-recon further boosts the disentanglement. In Appendix C, we empirically validate that swap-recon improves the representation disentanglement quality.

To the best of our knowledge, this is the first study to improve both fairness and out-of-distribution generalization performance. FarconVAE provides a fair representation that only includes non-sensitive core information by disentangling the sensitive information. Besides, it also removes the spurious correlation in representation learned from IRM and Group-DRO, while maintaining the model accuracy.

Our contributions in this work are three-fold:

We propose a novel framework FarconVAE that learns disentangled invariant representation with contrastive loss to achieve algorithmic fairness and domain generalization.

We provide a new distributional contrastive loss for disentanglement motivated by the Gaussian and Student-t kernels.

The proposed method is theoretically analyzed and empirically demonstrated on a broad range of data types (tabular, image, and text) and tasks, including fairness and domain generalization.

Related Works

The primary objective of domain generalization is to show stable predictive performance even though the test distributions are different from the train distribution (Gulrajani and Lopez-Paz, 2020). Traditional machine learning algorithm fails domain generalization, and there have been many research works to handle the distribution shift in diverse ways, including Bayesian neural net (Maddox et al., 2019), data augmentation (Xu et al., 2020), robust optimization (Sagawa et al., 2020; Levy et al., 2020), and invariant learning (Arjovsky et al., 2019; Creager et al., 2021).

Recently, Group-Distributionally Robust Optimization (Group-DRO) (Sagawa et al., 2020) and Invariant Risk Minmimzation (IRM) (Arjovsky et al., 2019) have been considered promising ways for domain generalization. Group-DRO optimizes the models to focus on the worst-case group with strong regularization, and IRM encourages the model to focus on the shared essential features across the group. However, there is no theoretical guarantee that the Group-DRO and IRM alleviate the spurious correlation or sensitive information removal. Recent works point out that the performance improvements of Group-DRO and IRM are limited on a shifted test distribution (Gulrajani and Lopez-Paz, 2020).

2. Learning Fair Representation

Fair representation learning aims to learn a representation that can be used for making accurate predictions without bias from sensitive information. We can divide the various fairness works depending on whether adversarial learning is used, and we introduce the related works for fair representation learning in this subsection.

Generative adversarial learning (GAN) (Goodfellow et al., 2014) shows a significant improvement in density estimation, and it has been widely utilized in many tasks. For fair representation learning, density estimation of a given instance without sensitive information is necessary. Therefore, there have been many research works that adopt adversarial learning for fair representation (Xie et al., 2017; Madras et al., 2018; Zhang et al., 2018; Roy and Boddeti, 2019). Controllable Invariance (CI) (Xie et al., 2017) adopts adversarial min-max game to filter out detrimental features such as sensitive information. CI introduces three types of network; encoder, discriminator, predictor, and adversarial minimax game between encoder and discriminator encourages the representation is invariant to sensitive information. Maximum Entropy Adversarial Representation Learning (MaxEnt-ARL) (Roy and Boddeti, 2019) is another kind of adversarial method, and MaxEnt-ARL utilizes a non-zero-sum game adversarial formulation by adopting different objectives for generator and discriminator to overcome the sub-optimal problems. The order approach is Flexibly Fair Variational AutoEncoder (FFVAE) (Creager et al., 2019) based on disentangled representation learning. FFVAE learns the separated latent space for sensitive and non-sensitive information. Although the algorithm does not contain an explicit adversary, it also relies on adversarial learning to approximate the Total Correlation (TC) (Watanabe, 1960) penalty used for disentanglement.

However, adversarial learning is known to have convergence instability problems, and it might hinder learning the robust representation. We empirically observe that the previous adversarial learning-based fair algorithm has difficulty separating the sensitive information. To solve the problems, we propose a contrastive learning-based disentangled representation learning method, FarconVAE, instead of adversarial learning.

2.2. Fair Representation without Adversarial Learning

Recently, there have been other fair representation learning research works without adversarial learning (Zemel et al., 2013; Louizos et al., 2015; Cheng et al., 2021; Sarhan et al., 2020). Zemel et al. (Zemel et al., 2013) propose fair clustering methods with probabilistic mapping, and Variational Fair AutoEncoder (VFAE) (Louizos et al., 2015) adopts Maximum Mean Discrepancy (MMD) measure (Gretton et al., 2006) to penalize the posterior. Besides, Fair Filter (FairFil) (Cheng et al., 2021) removes the sensitive information that is inherent in the sentence embedding from a pretrained language model. To learn debiased embeddings, FairFil utilizes contrastive learning and mutual information estimator.

Orthogonal Disentangled Fair Representations (ODFR) (Sarhan et al., 2020) is a non-adversarial disentangle-based methods, and it introduces orthogonal priors to enforce an orthogonality constraint between sensitive and non-sensitive representation. qϕs(zs∣x)q_{\phi_{s}}(z_{s}|x) and qϕx(zx∣x)q_{\phi_{x}}(z_{x}|x) denote the posteriors of sensitive and non-sensitive representation parameterized by ϕs\phi_{s} and ϕx\phi_{x} respectively, and p(zx)p(z_{x}) and p(zs)p(z_{s}) denote the priors, where p(zx)=N(T,I)p(z_{x})=\mathcal{N}(^{T},\mathbf{I}) and p(zs)=N(T,I)p(z_{s})=\mathcal{N}(^{T},\mathbf{I}).

However, minimizing LODL_{OD} in Eq. 1 does not guarantee a robust fair representation. It is known that the vanilla KL-divergence between posterior and fixed prior might cause the posterior collapse, and the representation might not contain meaningful information of given data (He et al., 2018). As a result, ODFR adopts additional auxiliary components such as entropy loss. Besides, there are diverse ways to set orthogonal priors, and the performance of ODFR may largely depend on the choice of prior.

FarconVAE has a relationship with FairFil and ODFR, but there are three major differences. First, FarconVAE measures the divergence between two posteriors for disentanglement instead of the fixed prior in ODFR. Thus, FarconVAE is relatively free from posterior collapse problems and prior choice. Second, we adopt a new distributional contrastive loss motivated by Gaussian and Student-t kernel. It makes FarconVAE get highly disentangled fair representation without the need for auxiliary components such as discriminators and entropy loss. Third, in contrast to ODFR and FairFil, which have shown effectiveness in limited data types for only fair prediction tasks, our FarconVAE has validated on both fairness and domain generalization tasks across three representative data types.

3. Contrastive Learning

Recently, Contrastive learning has emerged as a new promising paradigm for self-supervised representation learning, showing powerful performance in a broad domain such as Vision (Chen et al., 2020), NLP (Gao et al., 2021), and Graph (Li et al., 2019). InfoNCE (van den Oord et al., 2019) is a widely used loss function for contrastive learning that effectively formulates the instance discrimination task. InfoNCE loss encourages the similarity between positive pairs and the dissimilarity between negative pairs. Representation space learned by contrastive learning has performed well in various downstream tasks such as classification, clustering, or semantic similarity evaluation. However, there is limited work (Weinberger et al., 2022) that handled the disentanglement with contrastive learning. This paper provides a new distributional contrastive learning method to obtain disentangled invariant representation and apply it to fairness and domain generalization tasks.

Methodology

In this section, we firstly introduce the basic notation and problem formulation in Section 3.1. In Sections 3.2, 3.3, we describe FarconVAE in detail, and cover the kernel motivated distributional contrastive loss that induces stable disentanglement, respectively. In Section 3.4, we provide a description for swap-recon, a new feature swap based regularization. Finally, we theoretically analyze our two specified contrastive losses in Section 3.5.

Let xx be an observed input feature, ss be its sensitive attribute, and yy be a target label. As described in (Arjovsky et al., 2019; Sarhan et al., 2020), the sensitive attribute ss is highly correlated with feature xx and spuriously correlated with label yy on many real-world datasets, even if it is essentially irrelevant. In this situation, We want to build a fair ML algorithm. Mathematically, fairness can be defined as pθ(y∣x)=pθ(y∣x,s)p_{\theta}(y|x)=p_{\theta}(y|x,s) (Sarhan et al., 2020). The Fairness definition shows that the model output needs to be independent of the sensitive attribute. As a result, a fair representation debiased w.r.t sensitive information is necessary to achieve fairness. Furthermore, performance degradation for target label prediction should be minimized. However, the existing methods for building fair algorithms have difficulty in learning robust representation satisfying both predictiveness and fairness, as stated in Section 2.

In this work, we propose FarconVAE, which aims to learn a function that maps each data (x,s,y)(x,s,y) to disentangled representation zsz_{s} and zxz_{x} on the two separated latent spaces Zs,Zx\mathcal{Z}_{s},\mathcal{Z}_{x} that involve different semantic meanings. Specifically, Zs\mathcal{Z}_{s} is a latent space to absorb and isolate the whole sensitive information from observation, and Zx\mathcal{Z}_{x} maximally preserves only non-sensitive information to predict the target accurately. In other words, zxz_{x} is a fair representation used for target prediction that has high predictiveness without relying on sensitive information.

2. FarconVAE Structure

Figure 1 represents the FarconVAE in terms of graphical model and neural net view. FarconVAE assumes that there are three observable variables, input feature xx, sensitive attribute ss, and target label yy. FarconVAE has one encoder pϕ(⋅)p_{\bm{\phi}}(\cdot) that splits into two heads for encoding zxz_{x} and zsz_{s}, and two decoders: pθy(⋅)p_{\theta{y}}(\cdot) for decoding yy and pθ\textbackslashθy(⋅)p_{\bm{\theta}\textbackslash\theta{y}}(\cdot) that splits into two heads for decoding xx and ss. The goal of FarconVAE is to capture two disentangled latent zsz_{s} and zxz_{x}, where zsz_{s} and zxz_{x} denote the sensitive and non-sensitive representation, respectively. With these two latent variables, The log marginal likelihood of the x,sx,s and yy with model parameters θ\bm{\theta} is as follows:

Direct optimization of marginal likelihood, Eq. 2, is intractable, so we utilize the variational inference (Jordan et al., 1999) to approximate the marginal likelihood. We introduce the variational distribution qq, and we adopt encoder qϕq_{\bm{\phi}} and decoder pθp_{\bm{\theta}} as neural networks parameterized by ϕ={ϕbody,ϕheadx,ϕheads}\bm{\phi}=\{\phi_{body},\phi_{head_{x}},\phi_{head_{s}}\} and θ={θbody,θheadx,θheads,θy}\bm{\theta}=\{\theta_{body},\theta_{head_{x}},\theta_{head_{s}},\theta_{y}\}, respectively. In the following all sections, we will notate {ϕbody∪ϕheadx}\{\phi_{body}\cup\phi_{head_{x}}\} as ϕx\phi_{x} and {ϕbody∪ϕheads}\{\phi_{body}\cup\phi_{head_{s}}\} as ϕs\phi_{s}, {θbody∪θheadx}\{\theta_{body}\cup\theta_{head_{x}}\} as θx\theta_{x}, and {θbody∪θheads}\{\theta_{body}\cup\theta_{head_{s}}\} as θs\theta_{s}.

For the disentanglement between zxz_{x} and zsz_{s}, the conditional independence between zxz_{x} and zsz_{s} given x,s,yx,s,y is necessary. Under the conditional independence assumption, we construct the variational objective (Jordan et al., 1999; Kingma and Welling, 2014), called evidence lower bound (ELBO), as follows:

LELBO\mathcal{L}_{ELBO} in Eq. D consists of three components. First, the KL divergence in Eq. 3 can be interpreted as a regularization term to prevent the variational posteriors qϕx(zx∣x,s,y)q_{\phi_{x}}(z_{x}|x,s,y) and qϕs(zs∣x,s,y)q_{\phi_{s}}(z_{s}|x,s,y) moving too far from their priors. Second, reconstruction terms, pθx(x∣zx,zs)p_{\theta_{x}}(x|z_{x},z_{s}) and pθs(s∣zx,zs)p_{\theta_{s}}(s|z_{x},z_{s}) encourage the latent representation zxz_{x} and zsz_{s} to preserve the salient information of xx and ss. Third, prediction term pθy(y∣zx)p_{\theta_{y}}(y|z_{x}) gives an ability to model to predict the target value, and it injects task-specific information to representation zxz_{x}.

By maximizing the ELBO in Eq. D, we can infer the parameters of distribution over the joint latent variables zsz_{s} and zxz_{x} that depend on x,sx,s, and yy. These two latent representations include the salient information of the given dataset. Different with zsz_{s}, zxz_{x} has capability to predict yy, and its difference contributes to disentanglement between zsz_{s} and zxz_{x}. However, FarconVAE can be improved in two ways. First, LELBO\mathcal{L}_{ELBO} just assumes the conditional independence between zxz_{x} and zsz_{s} given x,s,yx,s,y. Second, zxz_{x} still has sensitive information if the yy is spuriously correlated with ss. Therefore, we provide a new contrastive loss with FarconVAE and introduce it in Section 3.3.

2.2. Model in Detail

In this subsection, we provide the distributional assumption in FarconVAE. For the prior p(zs)p(z_{s}) and p(zx)p(z_{x}), we adopt standard Gaussian N(zs;0,I)\mathcal{N}(z_{s};0,\mathbf{I}) and N(zx;0,I)\mathcal{N}(z_{x};0,\mathbf{I}), respectively.

For encoder, we formulate the variational posterior distribution qϕ(⋅∣⋅)q_{\bm{\phi}}(\cdot|\cdot) as Gaussian distribution with diagonal covariance structure, parameterized by neural network. In the remaining sections, we will omit the parameter of posterior and refer to it as q(⋅∣⋅)q(\cdot|\cdot). In the formulas below, z⋅z_{\cdot} can be zxz_{x} or zsz_{s}.

For decoder, we adopt Gaussian distribution for continuous variables and Bernoulli distribution for discrete variables. Like above, the parameters of the distribution are parameterized by θ\bm{\theta}.

To reconstruct feature xx and sensitive attribute ss, FarconVAE uses both zxz_{x} and zsz_{s}, while for predicting target label yy it utilizes zxz_{x} only.

3. Contrastive Learning for Disentanglement

By adopting Gaussian or Student-t kernel for kernel function k(⋅)\textit{k}(\cdot), we specify our two Distributional Contrastive losses as below:

In Section 3.5, we present a theoretical analysis for these losses.

4. Swap-Reconstruction

There were similar methods to our swap-recon, and they showed promising results in disentangling (Mathieu et al., 2016) or debiasing (Kim et al., 2021a). However, this is the first study that proposes swap-recon under the contrastive learning framework. We utilize LSR\mathcal{L}_{SR}, swap-recon loss, with our main distributional contrastive loss LDC\mathcal{L}_{DC}.

FarconVAE is a general framework for disentangled fair representation learning. To encourage disentanglement, we use a distributional contrastive loss and swap-recon loss as shown in Figure 2. The overall objective becomes as follows:

5. Theoretical Analysis

This subsection presents the theoretical analysis of our distributional contrastive loss. The naive way to achieve contrastive disentangling is to use a simple reciprocal of divergence as a k(⋅)k(\cdot) in Eq. 5. However, setting k(⋅)k(\cdot) as a simple reciprocal might cause a numeric instability, as shown in Figure 3. To mitigate the numerical instability and to encourage disentanglement, we provide new kernel motivated contrastive learning. We formulate Gaussian Kernel motivated similarity exp(−Div(P∣∣Q))exp(-Div(P||Q)) and Student-t motivated similarity (1+Div(P∣∣Q))−1(1+Div(P||Q))^{-1} for distribution PP and QQ as shown in Eq. 7 and Eq. 6. If the variance of PP and QQ are the same, the Student-t kernel based loss returns a greater loss than the Gaussian kernel based loss, when the disentanglement is not enough, as shown in Figure 3 and Proposition 1. Therefore, Student-t kernel contrastive loss may enforce more rigorous disentanglement. In the following propositions, DivDiv denotes the KL divergence.

Assume that univariate random variables z1z_{1} and z2z_{2} follow Gaussian distribution N(μ1,σ2)\mathcal{N}(\mu_{1},\sigma^{2}), and N(μ2,σ2)\mathcal{N}(\mu_{2},\sigma^{2}), respectively. Then (1+Div(p(z1)∣∣p(z2)))−1≥exp⁡(−Div(p(z1)∣∣p(z2)))(1+Div(p(z_{1})||p(z_{2})))^{-1}\geq\exp(-Div(p(z_{1})||p(z_{2}))).

We can derive similar results when the mean of PP and QQ are the same, as shown in Proposition 2. These rigorous disentanglement properties of Student-t based methods are effective in the general case, but they may overfit when the amount of data is restricted.

Besides, this phenomenon can also occur when we handle noisy datasets. A noisy dataset, which contains corrupted labels as well as some outliers, might have a relatively large variance (An and Cho, 2015). Proposition 2 ii) denotes that the Student-tt based methods enforce rigorous disentangling even though the variance of QQ is large enough. Therefore, this strict enforcement of disentanglement can also lead to overfitting in a noisy dataset, too. We provide an empirical validation on noisy settings in Section 4.1.

Assume that univariate random variables z1z_{1} and z2z_{2} follow Gaussian distribution N(μ1,σ12)\mathcal{N}(\mu_{1},\sigma_{1}^{2}), and N(μ2,σ22)\mathcal{N}(\mu_{2},\sigma_{2}^{2}), respectively. Then, i) the global minimum of (1+Div(p(z1)∣∣p(z2)))−1−exp⁡(−Div(p(z1)∣∣p(z2)))(1+Div(p(z_{1})||p(z_{2})))^{-1}-\exp(-Div(p(z_{1})||p(z_{2}))) is zero if μ1=μ2\mu_{1}=\mu_{2}, and ii) lim⁡σ2→∞(1+Div(p(z1)∣∣p(z2)))−1−exp⁡(−Div(p(z1)∣∣p(z2)))>0\lim_{\sigma_{2}\to\infty}(1+Div(p(z_{1})||p(z_{2})))^{-1}-\exp(-Div(p(z_{1})||p(z_{2})))>0.

Experiments

To validate FarconVAE, we define three research questions in terms of disentangled fairness, debiasing pretrained large-scale models, and domain generalization. We answer each question in Sections 4.1, 4.2, and 4.3, respectivelyWe release the code at: https://github.com/changdaeoh/FarconVAE. See Appendix A for detail setup. RQ1) Disentangled Fairness: Do latent representations zxz_{x} contain non-sensitive important information only, excluding sensitive information, while maintaining the predictive performance? RQ2) Pretrained Models Debiasing: Can FarconVAE be utilized for debiasing a pretrained large-scale model such as BERT? RQ3) Domain Generalization: Does FarconVAE alleviate the spurious correlation as well as unfairness?

Table 1 denotes the performance on Adult, German, and Extended YaleB (Georghiades et al., 2001) datasets. Adult and German are tabular datasets (Asuncion and Newman, 2007), and their targets are binary about income and credit risk, respectively. Gender is a sensitive attribute for both datasets. The YaleB dataset is a visual dataset, and the target task is to classify the facial identity as irrelevant to the light condition that is regarded as a sensitive attribute. FarconVAE-G and FarconVAE-t denote the FarconVAE with Gaussian and Student-t kernel contrastive loss. We report the mean and standard deviation of 10 runs. To measure ss accuracy, we first train the FarconVAE and encode the entire dataset, and then train a linear classifier for ss on it. The representation that has similar ss accuracy with Random-Guess can be interpreted as fair.

In terms of yy and ss accuracy, our FarconVAEs significantly improve the performance over baselines. They show the highest yy accuracy while ss accuracy is the closest with the Random-Guessing. It denotes that zxz_{x} extracted by our model contains the non-sensitive core information while removing the sensitive information. Note that FarconVAE-t has poor performance on ss accuracy in German dataset, which is relatively small. Thus, the Student-t kernel based contrastive learning may overfit, while its Gaussian kernel based counterpart still works well. Besides, we visualize the learned representation of FarconVAE and ODFR. Figure 4 denotes that zxz_{x} of FarconVAE has important information to classify the target label yy, while successfully removing the sensitive information about ss.

To validate the robustness of our model, we perform the evaluation on noisy data settings. We construct a training set by corrupting the original sensitive attribute ss with another one proportional to rate ϵ\epsilon. Figure 5 represents the performance according to the noise rate ϵ\epsilon. FarconVAE-t and FarconVAE-G maintain their performance even though the ϵ\epsilon increases, while other models show performance degradation. When we compare FarconVAE-t and FarconVAE-G, FarconVAE-G shows better performance than FarconVAE-t. It corresponds with our theoretical analysis, as shown in Section 3.5.

2. Debiased Sentence Representation

We evaluate our model on a text dataset as well as a tabular and image dataset. It is known that the traditional word embedding or pretrained language model such as BERT (Kenton and Toutanova, 2019) has a harmful bias. We validate whether our model, FarconVAE can debias the sentence representation, pretrained from BERT or not. Following the same experiment and evaluation setting with (Cheng et al., 2021), when the input sentence contains a sensitive word, we replace it with the word with the opposite semantic meaning from the pre-defined sensitive word dictionary to construct a contrastive pair. We use the absolute SEAT effect size (Liang et al., 2020) as a measure of bias. Table 2 denotes the results for our model and baseline models including original BERT, Sent-D (Liang et al., 2020), and FairFil (Cheng et al., 2021). Like FairFil, our FarconVAE acts like a small filter that takes the BERT’s sentence representation as input and outputs the debiased representation. When we adopt FarconVAE on BERT, the degree of bias reduces from 0.354 to 0.057. Another column, ”BERT post SST-2”, denotes the fine-tuning performance with debiased representation. FarconVAE makes relatively poor classification, but it alleviates the bias largely from 0.291 to 0.103. The results show that FarconVAE can better remove bias in a large-scale language model than existing methods.

3. Domain Generalization

Domain Generalization requires the invariant representations (Arjovsky et al., 2019) or robust representation (Sagawa et al., 2020), and it is commonly known that IRM and Group-DRO improve the domain generalization by alleviating the spurious correlation. However, there is no guarantee that the representations learned by IRM and Group-DRO are free from spurious correlation. We observe that the IRM and Group-DRO improve the predictiveness of representation for yy over Empirical Risk Minimization (ERM), but the representations from IRM and Group-DRO still have sensitive information largely. Table 3 denotes the performance of yy accuracy and ss accuracy on cMNIST (Arjovsky et al., 2019) and WaterbirdsFor a consistent evaluation, we report average accuracy for yy and ss. It is different from the weighted average accuracy (for y) reported in (Sagawa et al., 2020) for Waterbirds dataset. (Wah et al., 2011) those are intentionally constructed to have a spurious correlation between yy and ss. On Waterbirds, we additionally report the worst y acc. which measures the worst accuracy among four groups distinguished by the (s,y)(s,y) combinations. We applied FarconVAE on top of the feature extractor trained with IRM (for cMNIST) or Group-DRO (for Waterbirds) method. Our method successfully disentangles and removes ss information from the learned representation by IRM or Group-DRO, so the FarconVAE is free from spurious correlation and significantly improves the yy accuracy.

Figure 6 denotes the visualization of learned embedding from ERM, IRM, and FarconVAE on cMNIST. As shown in the first column of Figure 6, the learned representation from ERM and IRM are easily divided by sensitive attribute ss. However, the representation learned from FarconVAE is not easily divided by sensitive attributes. From the second and third columns, we can check that FarconVAE predicts yy well, while ERM and IRM predict opposite yy entirely or partially, respectively. To the best of our knowledge, this is the first work that provides general disentangling methods that can be utilized in both domain generalization and fairness tasks.

Conclusion

Algorithmic fairness demands fair representation learning that captures the non-sensitive core information. Domain generalization tasks also have difficulty removing spurious features for domain-invariant learning. This paper proposes FarconVAE, a new contrast-based disentangling approach that removes sensitive information or spurious correlation from representation while maintaining non-sensitive core information. We provide a kernel motivated distributional contrastive loss that stably induces disentangled invariant representation and a swap-recon loss that enforces swap-consistent reconstruction for further enhancing disentanglement.

FarconVAE improves the out-of-distribution generalization as well as the fairness. Besides, we observe that the representation learned by IRM or Group-DRO still has a strong spurious correlation between the targets and sensitive attributes, and FarconVAE mitigates this correlation significantly. We provide extensive empirical results about fairness, pretrained model debiasing, and domain generalization on tabular, image, and text datasets. Moreover, we derive a theoretical result for our new contrastive loss.

Acknowledgements

This work was partly supported by Institute of Information & Communications Technology Planning & Evaluation (IITP) grant funded by the Korea government(MSIT) (No.2021-0-02067,Next generation AI for multi-purpose video search,50%) and the National Research Foundation of Korea(NRF) grant funded by the Korea government(MSIT). (No. 2021R1F1A1060117, Research and Application of Artificial Intelligence Algorithm for Removing Spurious Bias,50%)

References

Appendix A Experimental Setup

For fair classification, we consider three benchmark datasets previously used in (Sarhan et al., 2020). The Adult dataset contains 45,222 instances, each with 14 attributes. The target yy is a binary label of annual income more or less than 50,000andgenderissensitiveattribute50,000 and gender is sensitive attributes$. The German dataset has 1,000 instances, each with 20 attributes, and the target task is to classify bank account holders with good or bad credit risk. The sensitive attribute is gender again. Both of these are tabular datasets obtained from the UCI ML-repository (Asuncion and Newman, 2007). The Extended YaleB (Georghiades et al., 2001) is a visual dataset that contains the face images of 38 people under five different light conditions. The target task is to identify one of the 38 people for a given data instance, while the light condition is the sensitive attribute here.

A.2. Pretrained Model Debiasing

For the pretrained model debiasing task, we validate our method on BERT (Kenton and Toutanova, 2019). Specifically, we attach the FarconVAE-t on top of BERT’s layers, takes [CLS] token embedding for each sentence as model input xx and the sensitive words are regarded as ss. Following the setup (Cheng et al., 2021), we use the same corpora consisting of 183,060 sentences for training, and the sensitive attribute is mainly gender-related words.

A.3. Domain Generalization

For domain generalization task, we experimented with two image datasets cMNIST (Arjovsky et al., 2019) and Waterbirds (Sagawa et al., 2020). The cMNIST is a synthetic dataset for binary classification, which intentionally makes correlation between labels (digit) and sensitive attributes (color) of the train set. A model is evaluated with the test set, which has the opposite correlation to the train set. The Waterbirds dataset is constructed by combining bird photographs from the CUB dataset (Wah et al., 2011) with backgrounds from the Places dataset (Zhou et al., 2017). The target label is the bird’s breed (waterbird or landbird), and the sensitive attribute is the background (water or land). Like cMNIST, there is a spurious correlation between yy and ss in the train dataset.

Appendix B Implementation Details

We put yy as FarconVAE input to increase the predictiveness of the representation. But unlike the training phase, we generally do not have access to true labels in the testing phase. So we train a separate classifier predicting yy from xx in advance or together with FarconVAE’s training phase. It is also possible to use a well-fitted pretrained model on the given dataset. After we build the best classifier that maps xx to yy, we use the predicted label y^\hat{y} by the classifier for each test set instance as the input of FarconVAE. For Adult, German, Extended YaleB, and CMNIST, this classifier is a simple multi-layer perceptron (MLP). For Waterbirds dataset, the classifier is ResNet-50 (He et al., 2016). For debiasing task on BERT, when FarconVAE is trained on the unannotated corpus (results of the left side in Table 2), the target label does not exist. In this case, we input the constant y=0.5y=0.5 to FarconVAE. When FarconVAE is finetuned with BERT on the labeled corpus (results of the right side in Table 2), we also input the constant y=0.5y=0.5 for consistent training.

B.2. Model Configuration

In this subsection, we introduce the setting of model architecture, hyperparameters, and other experimental options. To learn more informative representation, we use ELBO of β\beta-VAE formulation (Higgins et al., 2016), which can control the intensity of KLD regularization instead of the basic ELBO in Section 3.2. So, our ELBO has a parameter β\beta.

In the sections below, α\alpha, β\beta, γ\gamma, LRLR, and WDWD denote the weight of loss terms LDCL_{DC}, LKLDL_{KLD}, LSRL_{SR}, learning rate, and weight decay hyperparameter, respectively. See Githubhttps://github.com/changdaeoh/FarconVAE for details not listed here.

On Adult and German datasets, we follow the setup in (Roy and Boddeti, 2019; Sarhan et al., 2020) for all possible configurations. So, the encoder qϕq_{\bm{\phi}} and decoder pθ\textbackslashθyp_{\bm{\theta}\textbackslash\theta_{y}} are both one hidden layer with 64 hidden units and the decoder pθyp_{\theta_{y}} (refered as target predictor in (Roy and Boddeti, 2019; Sarhan et al., 2020)) is linear logistic regression. The latent dimension of zxz_{x} are 15 in Adult and 5 in German. For Extended YaleB dataset, we use one linear layer as encoder, decoder pθ\textbackslashθyp_{\bm{\theta}\textbackslash\theta_{y}} and decoder pθyp_{\theta_{y}} each contain 100 hidden units. The latent dimension is also 100.

B.2.2. Pretrained Model Debiasing

We use one linear layer for all components of FarconVAE with 128 hidden units and 128 latent dim. For doing contrastive learning on unlabeled corpus, we flip the sensitive words of a given sentence to follow (Cheng et al., 2021). If a sentence does not contain any sensitive words, a constant tensor of 0.5 which has the same shape as the embedding was inputted to FarconVAE and LDCL_{DC} and LSRL_{SR} are not used. In the fine-tuning stage, most of the sentences do not have sensitive words, so we use the mean embedding of pre-defined sensitive words list as FarconVAE input. We use (α\alpha, β\beta, γ\gamma, LRLR, WDWD) = (1.0, 0.2, 1.0, 5e-4, 1e-4) for FarconVAE contrastive-training, and (α\alpha, β\beta, γ\gamma, LRBERTLR_{BERT}, LRFarconLR_{Farcon}, WDWD) = (1.0, 0.2, 0.0, 2e-5, 1e-4, 1e-2) for entire fine-tuning.

B.2.3. Domain Generalization

Like above, we attach the FarconVAE-t on top of IRM or Group-DRO feature extractors (fixed) and use the representation from them as FarconVAE’s input feature. Again, we use one linear layer for all components of the FarconVAE latent dim set to 100 with 75 and 100 hidden units for cMNIST and Waterbirds, respectively. For cMNIST, we use (α\alpha, β\beta, γ\gamma, LRLR, WDWD) = (1.0, 0.2, 0.0, 1e-3, 1e-4). For Waterbirds, we use (α\alpha, β\beta, γ\gamma, LRLR, WDWD) = (0.5, 0.2, 0.5, 7e-4, 1e-4) and we anneal β\beta from zero to 0.2 during the first 10% epochs.

Appendix C Ablation Study

In this section, we provide the ablation study for our proposed methods. Figure 7 indicates that our proposed algorithm, FarconVAE with distributional contrastive loss is effective in disentangling the latent representation space, and as a result, it induces a fair representation. Moreover, the disentanglement is further enhanced when swap-recon (LSR\mathcal{L}_{SR}) is added. MRG in the right panel of Figure 7 denotes similarity between the s accuracy of model and that of random guessing, same with Figure 5.

Appendix D Derivation of ELBO

We assume that the zxz_{x} and zsz_{s} are conditionally independent given x,s,yx,s,yTo satisfy the assumption, we adopt contrastive loss as shown in Sec. 3 of main paper, i.e., qϕ(zx,zs∣x,s,y)=qϕzx(zx∣x,s,y)qϕzs(zs∣x,s,y){q_{\bm{\phi}}(z_{x},z_{s}|x,s,y)=q_{\phi_{z_{x}}}(z_{x}|x,s,y)q_{\phi_{z_{s}}}(z_{s}|x,s,y)}. Then, the evidence lower bound of the log marginal likelihood is:

Appendix E Proof

For simplicity, we denote tt as a μ1−μ2\mu_{1}-\mu_{2}. The KL divergence between two Gaussian distributions, N(μ1,σ2)\mathcal{N}(\mu_{1},\sigma^{2}) and N(μ2,σ2)\mathcal{N}(\mu_{2},\sigma^{2}) are t22σ2\frac{t^{2}}{2\sigma^{2}}. Then, (1+Div(p(z1)∣∣p(z2)))−exp⁡(Div(p(z1)∣∣p(z2)))=t22σ2−exp⁡(t22σ2)+1(1+Div(p(z_{1})||p(z_{2})))-\exp(Div(p(z_{1})||p(z_{2})))=\frac{t^{2}}{2\sigma^{2}}-\exp{(\frac{t^{2}}{2\sigma^{2}})}+1. When μ1=μ2\mu_{1}=\mu_{2}, 1+t22σ2−exp⁡(t22σ2)=01+\frac{t^{2}}{2\sigma^{2}}-\exp{(\frac{t^{2}}{2\sigma^{2}})}=0, and ∂(1+Div(p(z1)∣∣p(z2)))−exp⁡(Div(p(z1)∣∣p(z2)))∂t=t−texp⁡(x2/(2σ2))σ2\frac{\partial(1+Div(p(z_{1})||p(z_{2})))-\exp(Div(p(z_{1})||p(z_{2})))}{\partial t}=\frac{t-t\exp(x^{2}/(2\sigma^{2}))}{\sigma^{2}}¡0. Therefore, (1+Div(p(z1)∣∣p(z2)))−1≥exp⁡(−Div(p(z1)∣∣p(z2)))(1+Div(p(z_{1})||p(z_{2})))^{-1}\geq\exp(-Div(p(z_{1})||p(z_{2}))), and the equality holds when t=0t=0, i.e., μ1=μ2\mu_{1}=\mu_{2}.

For simplicity, we denote tt as a σ2σ1\frac{\sigma_{2}}{\sigma_{1}}. i) The KL divergence between two Gaussian distributions, N(μ,σ12)\mathcal{N}(\mu,\sigma_{1}^{2}) and N(μ,σ22)\mathcal{N}(\mu,\sigma_{2}^{2}) are log⁡(t)+12t2−0.5\log(t)+\frac{1}{2t^{2}}-0.5. Then, f(t)=(1+Div(p(z1)∣∣p(z2)))−exp⁡(Div(p(z1)∣∣p(z2)))=0.5+12t2−exp⁡(−0.5+12t2)t+log⁡(t).f(t)=(1+Div(p(z_{1})||p(z_{2})))-\exp(Div(p(z_{1})||p(z_{2})))=0.5+\frac{1}{2t^{2}}-\exp(-0.5+\frac{1}{2t^{2}})t+\log(t). f(t)f(t) is differentiable for all t>0t>0, and the only critical point of ∂f(t)∂t=0\frac{\partial{f(t)}}{\partial{t}}=0 is at t=1t=1. The domain of f′(t)f^{\prime}(t) is t∈R:t>0{t\in R:t>0}, and f(t)f(t) is −∞-\infty when t=0+t=0^{+} and ∞\infty. Therefore, the global minimum of (1+Div(p(z1)∣∣p(z2)))−1−exp⁡(−Div(p(z1)∣∣p(z2)))(1+Div(p(z_{1})||p(z_{2})))^{-1}-\exp(-Div(p(z_{1})||p(z_{2}))) is zero ii) lim⁡σ2→∞(1+Div(p(z1)∣∣p(z2)))−exp⁡(Div(p(z1)∣∣p(z2)))<0\lim_{\sigma_{2}\rightarrow\infty}(1+Div(p(z_{1})||p(z_{2})))-\exp(Div(p(z_{1})||p(z_{2})))<0. Therefore, (1+Div(p(z1)∣∣p(z2)))−1−exp⁡(−Div(p(z1)∣∣p(z2)))>0(1+Div(p(z_{1})||p(z_{2})))^{-1}-\exp(-Div(p(z_{1})||p(z_{2})))>0 for sufficiently large σ2\sigma^{2}.