Disentanglement via Latent Quantization

Kyle Hsu, Will Dorrell, James C. R. Whittington, Jiajun Wu, Chelsea Finn

Introduction

Our increasing reliance on black-box methods for processing high-dimensional data underscores the importance of developing techniques for learning human-interpretable representations. To name but a few possible benefits, such representations could foster more informed human decision-making , facilitate efficient model debugging and improvement , and streamline auditing and regulation . In this context, disentangled representation learning serves as a worthwhile scaffolding: loosely speaking, its goal is for a model to tease apart a dataset’s underlying sources of variation and represent them independently of one another.

Accomplishing this, however, has proven difficult. Conceptually, the field lacks a formal problem statement that resolves fundamental ambiguities without overly restrictive assumptions (reviewed in Section 6); methodologically, evaluation metrics have been found to be sensitive to hyperparameters, ad hoc, and/or sample inefficient ; and empirically, there remains a need for an inductive bias that enables consistently good performance in the purely unsupervised setting.

In this work, we answer the call for a better inductive bias for disentanglement. Our solution is motivated by observing that many datasets of interest are generated from their sources in a compositional manner, which entails a neatly organized source space (Figure 1). This distinguishing property of realistic generative processes applies to real world physics as well as human approximations thereof (e.g., rendering). Hence, our broad strategy to uncover the true underlying sources is to bias the model towards encoding to and decoding from a similarly structured latent space.

We manifest this inductive bias by drawing from two common ideas in the machine learning literature: discrete representations and model regularization. Specifically, we propose (i) quantizing a model’s latent representation into learnable discrete values with a separate scalar codebook per dimension and (ii) applying strong regularization via an unusually high weight decay . Intuitively, forcing the model to use a small number of scalar values to combinatorially construct many latent codes encourages it to assign a consistent meaning to each value, an outcome that weight decay explicitly incentivizes by regularizing towards parsimony.

A side benefit of models with quantized latents is that they sidestep one issue that has hindered evaluation in previous works: they enable the use of simpler, more robust distribution estimation techniques for discrete variables. As a further methodological contribution, we present InfoMEC, a new set of metrics for the modularity, explicitness, and compactness of (both continuous and discrete) representations that is cohesively grounded in information theory and fixes other well-established shortcomings in existing disentanglement metrics.

We demonstrate the broad applicability of latent quantization by adding it to both basic data-reconstructing (vanilla autoencoder) and latent-reconstructing (InfoGAN) generative models. Together with regularization, this is sufficient to dramatically improve the modularity and explicitness of the learned representations of a representative suite of four disentangled representation learning datasets with image observations and ground-truth source evaluations . In particular, our quantized-latent autoencoder (QLAE, pronounced like clay) consistently outperforms strong methods from prior work without compromising data reconstruction. We think of QLAE as a minimalist implementation of a combinatorial representation in neural networks, suggesting that our recipe of latent quantization and regularization could be broadly useful to other areas of machine learning.Code for models and InfoMEC metrics: https://github.com/kylehkhsu/latent_quantization.

Preliminaries

In order to properly contextualize our proposed inductive bias, methodological contributions, and experiments, we first devote some attention to explaining the problem of disentangled representation learning. We also discuss prior disentangled representation learning methods we build upon.

We begin by considering the standard data-generating model of nonlinear independent components analysis (ICA) , a problem very related to but more conceptually precise than disentanglement:

where s=(s1,…,sns)\mathbf{s}=(\mathbf{s}_{1},\dots,\mathbf{s}_{n_{s}}) comprise the ns{n_{s}} mutually independent source variables (sources); x\mathbf{x} is the observed data variable; and g:S→Xg:\mathcal{{S}}\to\mathcal{{X}} is the nonlinear data-generating function. The nonlinear ICA problem is to recover the underlying sources given a dataset D\mathcal{D} of samples from this model. Specifically, the solution should include an approximate inverting function g^−1:X→Z\hat{{g}}^{-1}:\mathcal{{X}}\to\mathcal{{Z}} such that, assuming latent variables (latents) z=(z1,…,zns)\mathbf{z}=(\mathbf{z}_{1},\dots,\mathbf{z}_{{n_{s}}}) are matched correctly to the sources, each source is perfectly determined by its corresponding latent. A typical technical phrasing is that g^−1∘g\hat{{g}}^{-1}\circ g should be the composition of a permutation and a dimension-wise invertible function.

As stated, this problem is nonidentifiable (or underspecified). Given D\mathcal{D}, one may find many sets of independent latents (and associated nonlinear generators) that fit the data despite being non-trivially different from the true generative sources . As such, reliably recovering the true sources from the data is impossible. Much recent work in nonlinear ICA has focused on proposing additional problem assumptions so as to provably pare down the possibilities to a unique solution. These theoretical assumptions can then be transcribed into architectural choices or regularization terms. Such approaches have shown promise in increasing our understanding of what assumptions are required to disentangle; we review these works in Section 6.

While identifiability is conceptually appealing, achieving it under sufficiently generalizable assumptions that apply to non-toy datasets has proven hard. The field of disentangled representation learning has taken a more pragmatic approach, focusing on empirically evaluating the recovery of each dataset’s designated source set. Given this empirical focus, robust performance metrics are vital. Unfortunately, there is a plethora of approaches in use, and subtle yet impactful issues arise even in the most common choices . To address these concerns, in Section 4 we propose new metrics for three existing, complementary notions of disentanglement . They measure the following three properties: modularity—the extent to which sources are encoded into disjoint sets of latents; explicitness—how simply the latents encode each source; and compactness—the extent to which latents encode information about disjoint sets of sources. We frame these three in a cohesive information-theoretic framework that we name InfoMEC.

The disentangled representation learning problem statement considered in this work is as follows. Given a dataset of paired source-data samples {(s,x=g(s))}\{(s,x=g(s))\} from the nonlinear ICA model (1), learn an encoder g^−1:X→Z\hat{{g}}^{-1}:\mathcal{{X}}\to\mathcal{{Z}} and decoder g^:Z→X\hat{g}:\mathcal{{Z}}\to\mathcal{{X}} solely using the data {x}\{x\} such that (i) the InfoMEC as estimated from samples {(s,z=g^−1∘g(s))}\{(s,z=\hat{{g}}^{-1}\circ g(s))\} from the joint source-latent distribution is high, while (ii) maintaining an acceptable level of reconstruction error between xx and g^∘g^−1(x)\hat{g}\circ\hat{{g}}^{-1}(x).

2 Autoencoding and InfoGAN as Data and Latent Reconstruction

We will apply our proposed latent quantization scheme to two foundational approaches for disentangled representation learning: vanilla autoencoders (AEs) and information-maximizing generative adversarial networks (InfoGANs) . Here, we provide a brief overview of the two and defer complete implementation details to Appendices A and C. Both approaches involve learning an encoder g^−1\hat{{g}}^{-1} and decoder g^\hat{g}. An autoencoder takes a datapoint x∈Xx\in\mathcal{{X}} as input and produces a reconstruction g^∘g^−1(x)∈X\hat{g}\circ\hat{{g}}^{-1}(x)\in\mathcal{{X}} that is optimized to match the input:

An InfoGAN instead takes a latent code z∈Zz\in\mathcal{{Z}} as input. The decoder (aka generator) maps zz to the data space, and from this the encoder produces a reconstruction of the latent:

Unlike the data reconstruction loss, this is clearly insufficient for learning as the dataset D\mathcal{D} isn’t even used. InfoGAN can be thought of as grounding latent reconstruction by making the marginal distribution of generated datapoints, g^(z)\hat{g}(\mathbf{z}), indistinguishable from the empirical data distribution. A concrete measure of this is provided by an additional binary classifier (aka discriminator) or value model (aka critic) trained alongside but in opposition to the decoder. While InfoGAN was originally motivated as maximizing a variational lower bound on the mutual information between the latent and the generated data, we find the above interpretation to be unifying.

Latent Quantization

Our goal is to encourage our model to disentangle by biasing it towards using an organized latent space. Why would this mitigate the nonidentifiability of nonlinear ICA? Our key motivation is that generative processes for realistic data are compositional and hence necessarily use highly organized source spaces. We discuss connections to related works in Section 6.

A lesson from nonlinear ICA is that, given a flexible enough model, data can be mapped to and from latent spaces in many convoluted ways. We motivate the use of strong model regularization with the conjecture that, of all the possible mappings from organized latent space to data, the most parsimonious will be the true generative model or something close enough to it. We operationalize this by using a high weight decay on both the encoder and decoder networks. Ablation studies (Section 5) show that both quantized latents and weight decay are necessary to disentangle well.

To train a quantized-latent model, we use the straight-through gradient estimator and co-opt the quantization and commitment losses proposed for vector quantization:

The straight-through gradient estimator facilitates the flow of gradients through the nondifferentiable quantization step. Lquantize\mathcal{L}_{\text{quantize}} pulls the discrete values constituting zz onesidedly towards the pre-quantized continuous output of the encoder. This is needed to optimize V⁡\operatorname{V}, since straight-through gradient estimation disconnects V⁡\operatorname{V} from the computation graph. Conversely, the commitment loss prevents the pre-quantized representation, which does see gradients from downstream computation, from straying too far from the codes. While this is a significant failure mode for vector quantization, we find that the use of scalars instead of high-dimensional vectors alleviates this issue, allowing us to drastically downweight the quantization and commitment losses while maintaining training stability. This gives the model much-needed flexibility to reorganize the discrete latent space. Finally, while using a shared global codebook like in vector quantization is certainly feasible, we find it better to maintain dimension-specific codebooks to enable the stable optimization of each individual value. Algorithm 1 contains pseudocode for latent quantization and computing the quantization and commitment losses. Appendix A presents pseudocode for training a quantized-latent autoencoder (QLAE) in Algorithm 2 and a quantized-latent InfoWGAN-GP in Algorithm 3.

InfoMEC: Information-Theoretic Metrics for Disentanglement

In this section, we derive InfoMEC, metrics for modularity, explicitness, and compactness, building upon and otherwise taking inspiration from several prior works . We take care to motivate our design decisions, and while we do not expect this to be the final word on disentanglement metrics, we hope our presentation enables others to clearly understand InfoMEC and propose further improvements.

Nonlinear ICA asks for the latents to recover the sources up to a permutation and dimension-wise invertible transformation. The mutual information between an individual source and latent,

is a granular measure of the extent to which they are deterministic functions of each other. Unlike other measures such as correlation (used in MCC ), LASSO weights (used in linear DCI ), or linear predictive accuracy (used in SAP ), mutual information takes into account arbitrary nonlinear dependence between its two arguments, making it invariant within the nonlinear ICA equivalence class for any candidate solution.

When both arguments are discrete, estimating the mutual information is simple via the empirical joint distribution, but if either is continuous, estimation becomes non-trivial. Previous works bin a continuous variable and pretend it is discrete , but this is sensitive to the binning strategy . Instead, for evaluating continuous latents, we choose the celebrated kk-nearest neighbor based KSG estimator , in particular a variant designed to handle a mix of discrete and continuous arguments. We use k=3k=3. See Appendix B for experimental vignettes demonstrating the severe sensitivity of binning-based estimation to the binning strategy (Figure 5) and the robustness of KSG-based estimation to kk (Figure 6). We remark that latent quantization enables reliable evaluation using the discrete-discrete estimator.

To facilitate aggregation, we desire a normalization to the interval $.Tothisend,notethattheidentity. To this end, note that the identityI(\mathbf{s}_{i};\mathbf{z}_{j})=H(\mathbf{s}_{i})-H(\mathbf{s}_{i}\mid\mathbf{z}_{j})andthenonnegativityofentropyimplyand the nonnegativity of entropy implyI(\mathbf{s}_{i};\mathbf{z}_{j})\leq H(\mathbf{s}_{i})$ for discrete sources. Following , we define a normalized mutual information as

We prefer this normalization scheme over others since i) it is the proportion of a source’s entropy reduced by conditioning on a latent and thus scales consistently to $foranymodel,andii)itavoidsthescale−dependent(andpossiblynegative)differentialentropyofacontinuouslatent.Wegatherallevaluationsoffor any model, and ii) it avoids the scale-dependent (and possibly negative) differential entropy of a continuous latent. We gather all evaluations of\operatorname{NMI}(\mathbf{s}_{i},\mathbf{z}_{j})intoa2−dimensionalarrayinto a 2-dimensional array\operatorname{NMI}\in^{{n_{s}}\times{n_{z}}}.Weremoveinactivelatents(columnsof. We remove inactive latents (columns of\operatorname{NMI}),whicharethosewithzerorange(overtheevaluationsample)fordiscretelatents.Forcontinuouslatents,zeroistoostrict,soweheuristicallydefinethethresholdtobe), which are those with zero range (over the evaluation sample) for discrete latents. For continuous latents, zero is too strict, so we heuristically define the threshold to be\nicefrac{{1}}{{20}},appliedafterdividingtherangesbytheirmaximum.SeeFigure3forexamplesof, applied after dividing the ranges by their maximum. See Figure 3 for examples of\operatorname{NMI}^{\top}$.

Modularity is the extent to which sources are separated into disjoint sets of latents. Perfect modularity occurs when each latent is informative of only one source, i.e. when every column of NMI⁡\operatorname{NMI} has only one nonzero element. This has been measured as the gap between the two largest entries in a column, or the ratio of the largest entry in the column to the column sum. We prefer the ratio since the gap is agnostic to the smallest ns−2{n_{s}}-2 values in the column, but these values matter and should influence the measure . Since the possible range of values for this ratio is [\nicefrac1ns,1][\nicefrac{{1}}{{{n_{s}}}},1], we re-normalize to $$. Finally, we define InfoModularity (InfoM) as the average of this quantity over latents:

Compactness complements modularity; it is the extent to which latents only contain information about disjoint sets of sources. We therefore define InfoCompactness (InfoC) analogously to InfoM, but considering rows of NMI⁡\operatorname{NMI} instead of columns, and averaging over sources instead of latents, etc.:

We advocate for this terminology since previous names such as “mutual information gap” and “mutual information ratio” are ambiguous, and indeed the former of these works considered solely compactness and the latter solely modularity, with neither mentioning the distinction. We remark that when nz>ns{n_{z}}>{n_{s}} (after pruning inactive latents), it is impossible to achieve both perfect modularity and perfect compactness. Of the two, modularity should be prioritized and indeed has been referred to as disentanglement itself .

2 Explicitness

Modularity and compactness are measured in terms of mutual information, so they are agnostic to how this information is encoded. Our third metric, explicitness, measures the extent to which the relationship between the sources and latents is simple (e.g., linear ). Since previous explicitness metrics have been rather ad hoc, we propose a formalism using the framework of predictive V\mathcal{V}-information, a generalization of mutual information that specifies an allowable function class, denoted V\mathcal{V}, for the computation of information . We first estimate the predictive V\mathcal{V}-information of each source si\mathbf{s}_{i} given all latents z\mathbf{z}:

This requires estimating the predictive conditional V\mathcal{V}-entropy

and the marginal V\mathcal{V}-entropy of the source

where ∅\varnothing is an uninformative constant. The predictive conditional V\mathcal{V}-entropy measures how well a source, si\mathbf{s}_{i}, can be predicted by mapping the latents, z\mathbf{z}, through a function in function class V\mathcal{V}. Note that this estimation uses the best in-sample negative log likelihood. We choose V\mathcal{V} to be the space of linear models (though one could pick V\mathcal{V} to fit particular needs) and so use logistic regression (linear regression) for discrete (continuous) sources. We use no regularization. We compute the marginal V\mathcal{V}-entropy HV(si∣∅)H_{\mathcal{V}}(\mathbf{s}_{i}\mid\varnothing) in the same way, but substituting a universal constant for all inputs. We propose a simple normalization analogous to the one done for NMI⁡\operatorname{NMI}:

which can be interpreted as the relative reduction in the V\mathcal{V}-entropy of a source achieved by knowing the latents, and is in $:forclassificationnegativeloglikelihoodisthecross−entropy,andforregressionweleveragePropositions1.3and1.5fromXuetal.toarguethat: for classification negative log likelihood is the cross-entropy, and for regression we leverage Propositions 1.3 and 1.5 from Xu et al. to argue that\operatorname{NMI}_{\mathcal{V}}(\mathbf{z}\to\mathbf{s}_{i})=R^{2}$, the coefficient of determination. We can now compute explicitness as:

3 Summary and Comparison to Nonlinear DCI

We have derived three metrics for evaluating the modularity, explicitness, and compactness of a representation. Each metric has a straightforward information-theoretic interpretation and all share a range of $.Weorderthemindecreasingimportanceandcollectivelyrefertothemas. We order them in decreasing importance and collectively refer to them as\text{InfoMEC}{}:=(\text{InfoM}{},\text{InfoE}{},\text{InfoC}{})$.

Nonlinear DCI , a widely used three-pronged framework that measures similar disentanglement properties, suffers several practical drawbacks in comparison to InfoMEC from being defined in terms of relative counts of decision tree splits: determining this requires considering all latents jointly while fitting p(si∣z)p\left(\mathbf{s}_{i}\mid\mathbf{z}\right). This results in a cumbersome computational footprint that is exacerbated by highly sensitive hyperparameters such as tree depth , the tuning of which has even seen omission in prior work . These drawbacks worsen with increased latent space dimensionality. In contrast, InfoMEC avoids these issues as it isolates InfoM and InfoC from the choice of predictive function class and only computes pairwise interactions between individual sources and latents. (InfoE also fits p(si∣z)p\left(\mathbf{s}_{i}\mid\mathbf{z}\right), but does so with function classes of severely limited capacity for which fitting procedures scale well.) See Appendix B for experimental vignettes demonstrating the hyperparameter sensitivity of nonlinear DCI (Figure 7) and the robustness of InfoMEC (Figure 6).

Experiments

Experimental design. We design our experiments to answer the following questions: Does latent quantization improve disentanglement? How does it compare against the strongest known methods that operate under the same assumptions? And, finally, which of our design choices were critical? We benchmark on four established datasets: Shapes3D , MPI3D , Falcor3D , and Isaac3D . Each consists of RGB image observations generated (near-)noiselessly from categorical or discretized numerical sources. Shapes3D is toyish, but the others are chosen for their difficulty . In particular, we use the complex shapes variant of MPI3D collected on a real world robotics apparatus. See Appendix C.1 for further dataset details. Aside from baseline AE and InfoGAN (specifically InfoWGAN-GP, a Wasserstein GAN with gradient penalty ), we compare to β\beta-VAE , β\beta-TCVAE , and BioAE , the strongest methods from prior work that obey our problem assumptions and make design decisions mutually exclusive with latent quantization. We also compare to VQ-VAE with d=64d=64 and nv=512{n_{v}}=512 (Figure 2). We ablate weight decay, scalar codebooks, and dimension-specific codebooks from QLAE and weight decay from QLInfoWGAN-GP. We quantify modularity, explicitness, and compactness using both InfoMEC and nonlinear DCI. We qualitatively inspect representations via decoded latent interventions (Figure 4 and Appendix D).

Select experimental details. The choice of decoder architecture is known to impose inductive biases relevant for disentanglement . We use an expressive architecture (Appendix C.3) based on StyleGAN for all methods and datasets. We downsample the observations to 64×6464\times 64 (if necessary). We follow prior work in considering a statistical learning problem rather than a machine learning one: we train on the entire dataset then evaluate on 1000010000 i.i.d. samples. We fix the number of latents in all methods to twice the number of sources. For quantized-latent models, we fix nv=10{n_{v}}=10 discrete values per codebook. We tune one key regularization hyperparameter per method per dataset with a thorough sweep (Table 11, Appendix C.2). We use the best performing configurations over 2 seeds and rerun with 5 more seeds. Despite our modest list of methods and ablations, just the last stage took over 1000 GPU-hours.

DCI:=(D  I  C)↑\text{DCI}:=(\text{D}\;\text{I}\;\text{C})\uparrow QLAE (ours) (0.59(\mathbf{0.59} 0.95\mathbf{0.95} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.47}}) (0.81(\mathbf{0.81} 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.61}}) (0.36(\mathbf{0.36} 0.85\mathbf{0.85} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.36}}) (0.50(\mathbf{0.50} 0.96\mathbf{0.96} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.38}}) (0.69(\mathbf{0.69} 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.54}}) QLAE w/ global codebook (0.52(0.52 0.93\mathbf{0.93} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.41}) (0.83(\mathbf{0.83} 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.60}}) (0.36(\mathbf{0.36} 0.86\mathbf{0.86} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.34}}) (0.36(0.36 0.900.90 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.26}) (0.53(0.53 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.43}) QLAE w/o weight decay (0.49(0.49 0.94\mathbf{0.94} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.40}) (0.63(0.63 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.47}) (0.30(\mathbf{0.30} 0.84\mathbf{0.84} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.29}}) (0.46(\mathbf{0.46} 0.97\mathbf{0.97} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.36}}) (0.58(0.58 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.48}}) VQ-VAE w/ weight decay (0.43(0.43 0.840.84 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.37}) (0.74(0.74 0.99\mathbf{0.99} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.57}) (0.22(0.22 0.680.68 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.20}) (0.41(0.41 0.850.85 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.32}) (0.34(0.34 0.850.85 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.39}) QLInfoWGAN-GP (ours) (0.26(\mathbf{0.26} 0.77\mathbf{0.77} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.26}}) (0.38(\mathbf{0.38} 0.85\mathbf{0.85} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.29}}) (0.24(\mathbf{0.24} 0.71\mathbf{0.71} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.25}}) (0.20(\mathbf{0.20} 0.73\mathbf{0.73} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.24}}) (0.24(\mathbf{0.24} 0.79\mathbf{0.79} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.25}}) QLInfoWGAN-GP w/o w.d. (0.19(0.19 0.730.73 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.19}) (0.16(0.16 0.710.71 {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.13}) (0.28(\mathbf{0.28} 0.74\mathbf{0.74} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.23}}) (0.14(0.14 0.72\mathbf{0.72} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}0.17}) (0.20(\mathbf{0.20} 0.77\mathbf{0.77} {\color[rgb]{.5,.5,.5}\definecolor[named]{pgfstrokecolor}{rgb}{.5,.5,.5}\pgfsys@color@gray@stroke{.5}\pgfsys@color@gray@fill{.5}\mathbf{0.23}})

Ablations on latent space design and regularization. We observe that ablating dimension-specific codebooks, weight decay, and scalar codebooks from QLAE each causes a significant drop in InfoM and D (Table 5), verifying the importance of these design decisions. The effect of ablating weight decay from QLInfoWGAN-GP is less pronounced. We note that while VQ-VAE w/ weight decay performs somewhat closely to QLAE in terms of InfoMEC, this is only because we use the categorical codes (as opposed to the high-dimensional vector representation) for evaluation. In addition, the scalar codebook design of QLAE enables meaningful interpolation between discrete values, whereas this is not supported by vector quantization.

Nonlinear ICA and disentangled representation learning. There is a long history of trying to build interpretable representations that separately represent the sources of variation in a dataset. This goes back to classic work on (linear) ICA , and has been known in deep learning as disentanglement . Without further assumptions, nonlinear ICA and its relatives in disentangled representation learning are provably underspecified . Approaches to resolve this indeterminacy include labeling a small number of datapoints or showing pairs of datapoints in which only one or a few sources differ . More in line with our work are approaches that assume additional structure in the data generative process and disentangle by ensuring the representation reflects this structure . Assumptions and methods include: factorized priors , biologically inspired activity constraints , sparse source variation over time , structurally sparse source to pixel influence , geometric assumptions on the source to image mapping , sparse underlying causal graphs between sources , piecewise linearity , and hierarchical generation . Many of these ideas are in principle compatible with latent quantization, and we leave discovery of fruitful combinations to future work.

Factorized latent spaces. VAEs that specify an isotropic Gaussian latent prior regularize the marginalized variational distribution (aka aggregate posterior) towards being a factorized distribution . Roth et al. bias latents to have pairwise factorized support via a Hausdorff set distance regularization. For linear ICA, Whittington et al. prove that regularizing latents to be nonnegative and energy-minimizing results in them having factorized support. Differently from all of these works, latent quantization imbues a model with factorized structure by construction, instead of relying on the optimization of regularized objectives to manifest this structure. The favorable disentanglement that QLAE yields over β\beta-TCVAE and BioAE suggests that this strategy is more effective.

Discrete representation learning. Oord et al. first demonstrated the feasibility of discrete neural representation learning at scale, and their techniques have since been broadly applied, e.g., to videos , audio , and anomaly detection . The following works design discrete representations similarly to how we do, though for purposes other than unsupervised disentanglement. Several works use one scalar codebook per latent dimension to achieve high efficiency in retrieval . Kobayashi et al. disentangle normal and abnormal features in medical images into separate vector codebooks via pixel-space supervision. Liu et al. and Träuble et al. use multiple codebooks with separately parameterized key and value vectors and investigate the effect of discretization in systematic generalization and continual learning, respectively.

We have proposed to use latent quantization and model regularization to impose an inductive bias towards disentanglement that enables our models to outperform strong prior methods. Ablations verify that our main design decisions are critical. We have also synthesized previously proposed ideas for evaluation into InfoMEC, three information-theoretic disentanglement metrics that rectify or sidestep key drawbacks in existing approaches.

While our results are promising, one concern might be that we have overfit our inductive bias to existing disentanglement benchmarks, in which, just like our model, the sources are discrete and the generative process is (near-)noiseless. Our experiments have already demonstrated the ability of latent quantization to represent sources that have more values (up to 40) than the per-dimension codebook size (fixed to 10) via allocating multiple latent dimensions. Future work should strive to construct disentanglement benchmarks that better reflect realistic conditions, e.g. continuous sources.

Beyond the intuitions and connections to related works we have presented, we do not understand why our method performs as well as it does. It may be fruitful to tackle this empirically, e.g. by probing how QLAE distributes data around its latent space, and how weight decay changes this. Achieving satisfactory understanding would enable the field to better position latent quantization within the ongoing body of work that aims to develop generalizable conditions for successful disentanglement.

Lastly, we hope this method, its future versions, and other methods the field develops are able to deliver on the original motivation for disentangled representation learning—to learn human-interpretable representations in complex, real-world situations, and to leverage the interpretability to empower human decision-making. This will require methods that can disentangle out-of-distribution data samples, that work for generic data types, and that can learn compositionally from sparse interactions with data. We suspect that latent quantization may have a role to play in these directions.

We gratefully acknowledge the developers of open-source software packages that facilitated this research: NumPy , JAX , Equinox , matplotlib , seaborn , and scikit-learn . We also thank Evan Liu, Kaylee Burns, Karsten Roth, Anirudh Goyal, and Cian Eastwood for feedback on previous drafts.

KH was funded by a Sequoia Capital Stanford Graduate Fellowship. Part of WD’s work on this project happened during a visit to Stanford funded by the Bogue Fellowship. JCRW was funded by a Henry Wellcome Post-doctoral Fellowship (222817/Z/21/Z). This work was also in part supported by the Stanford Institute for Human-Centered AI (HAI), NSF RI #2211258, Air Force Office of Scientific Research (AFOSR) YIP FA9550-23-1-0127, and ONR MURI N00014-22-1-2740.

Appendix A Quantized-Latent Models

This section contains pseudocode for latent quantization and for training QLAE and QLInfoWGAN-GP.

Appendix B Disentanglement Metrics Vignettes

Appendix C Experiment Details

This section contains details on the experiments conducted in this work.

C.2 Hyperparameters

This section specifies fixed and tuned hyperparameters for all methods considered.

C.3 Network Architectures

Inspired by recent works showing how well-designed decoder architectures can facilitate inductive biases relevant for disentanglement , we use an expressive architecture for all results presented in this work.

Encoder. We use a simple feedforward convolutional encoder network. Each convolutional block consists of two resolution-preserving convolutional layers (kernel size 33, stride 11) and one downsampling convolutional layer (kernel size 44, stride 22) at a consistent width (number of channels). Each convolution operation is followed by a leaky ReLU (slope 0.30.3), then instance normalization. There are four such blocks with widths 3232, 6464, 128128, and 256256. The 256×4×4256\times 4\times 4 output is then flattened. Two dense layers each of width 256256 with ReLU activation (and no normalization) follow. The final operation is an affine projection to the latent layer.

Decoder. We use a decoder architecture based on StyleGAN . The latent code zz of shape nz{n_{z}} passes through two dense layers each of width 256256 with ReLU activation (and no normalization); call the output of this ww. We directly parameterize a starting input feature map of shape 256×4×4256\times 4\times 4 with all entries initialized to 0.10.1. Each style-decoding layer consists of processing an input feature map via a transposed convolution followed by a leaky ReLU (slope 0.30.3) and then an adaptive instance normalization (AdaIN): ww undergoes an affine projection to a scale and bias scalar for each channel of the output feature map, and the output feature map is instance normalized then affinely transformed by spatially broadcasting the scales and biases. Similar to the encoder, style-decoding layers are grouped into blocks, with each block consisting of two resolution-preserving style-decoding layers (kernel size 33, stride 11) and one upsampling style-decoding layer (kernel size 44, stride 22) at a consistent width. There are four such blocks with widths 256256, 128128, 6464, and 3232. The final operation is a 1×11\times 1 convolution of width 33 to yield an output of shape 3×64×643\times 64\times 64.

C.4 Negative Results

We tried a number of additional modifications to our basic methods, QLAE and QLInfoWGAN-GP, beyond the ablations presented in the main paper. Here is a list of those that only marginally helped, didn’t help, or made things worse:

Keeping the original VQ-VAE hyperparameter settings for λquantize\lambda_{\text{quantize}} (11) and λcommit\lambda_{\text{commit}} (0.250.25).

A step function schedule for the weight decay in AdamW: for the first half of optimization, and a specified value for the second half.

Linearly annealing λquantize\lambda_{\text{quantize}} and/or λcommit\lambda_{\text{commit}} from up to a specified value.

(QLInfoWGAN-GP-specific) Updating V⁡\operatorname{V} with the encoder instead of with the decoder.

C.5 Nonlinear DCI Evaluation

For each p(si∣z)p\left(\mathbf{s}_{i}\mid\mathbf{z}\right) learning problem, we train a random forest classifier with 100100 trees and the information gain splitting criterion. We use a 0.90.9/0.10.1 train/test split of the evaluation sample. We tune the maximum tree depth hyperparameter on held-out accuracy (see Figure 7 for an example of why this is important). We compute D and C using the most predictive model’s relative feature importances, following their definitions . We use held-out accuracy for I.

Appendix D Qualitative Results

Appendix E Quantitative Results

This section contains unabridged results for the experiments in the main text. Intervals denote 95% confidence intervals of the mean estimated assuming a tt-distribution. Bolded intervals overlap with the interval with highest endpoint in the column. AE and InfoGAN variants are presented and bolded separately as all AEs are filtered for near-perfect data reconstruction, whereas InfoGANs are generally more lossy.