Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks

Minyoung Huh, Brian Cheung, Pulkit Agrawal, Phillip Isola

Introduction and related works

Vector-quantization (Gray, 1984) enables deep neural networks to learn discrete representations by quantizing features into clusters referred to as “code-vectors” or “codes”. Vector-Quantization (VQ) is a parametric online K-means algorithm (Caron et al., 2018) that has an explicit bias for compression and competition, which serves as a good prior for learning disentangled features for downstream tasks. To name a few, VQs have shown impressive results on image generation (Van Den Oord et al., 2017; Ramesh et al., 2021; Esser et al., 2020; Chang et al., 2023), image representation learning (Caron et al., 2020), speech generation (Dhariwal et al., 2020), speech representation learning (Chung et al., 2020), and even decision-making (Ozair et al., 2021). While powerful, vector-quantized networks (VQNs) are notoriously difficult to optimize and require esoteric knowledge to efficiently train them. Hence, constructing algorithms to improve training stability has been a topic of great interest.

First introduced in context of generative models (Van Den Oord et al., 2017), the vector-quantization layer in VQNs maps the continuous embedding (or feature representation) Pz\mathcal{P}_{\mathbf{z}} into a discrete embedding Qz\mathcal{Q}_{\mathbf{z}} using the codebook Cz\mathcal{C}_{\mathbf{z}}. As the discretization function is not continuously differentiable, a widely used technique to optimize VQNs is via straight-through estimation. This bypasses the non-differentiable discretization function (Bengio et al., 2013), allowing it to be optimizable by standard deep learning libraries (Paszke et al., 2019). Of course, straight-through estimating a selection function has negative ramifications on training. Among many training challenges, the most well-documented is model collapse or“index collapse” wherein only a small fraction of codes are used during training. While the root cause of the collapse is not well understood, there have been abundant efforts to mitigate index collapse: EMA (Van Den Oord et al., 2017), codebook reset (Łańcucki et al., 2020; Zeghidour et al., 2021; Dhariwal et al., 2020),probabilistic/stochastic re-formulation (Roy et al., 2018; Takida et al., 2022), equipartition assumption using optimal transport (Asano et al., 2020), and many more (See Section 3). While many of these methods directly reduce the extent of index collapse, to the best of our knowledge, there has not been any work that investigates the root cause of the instability in the first place. Hence, our work aims to systematically investigate the cause of the model collapse and provide methods to address the common pitfalls that stem from unstable optimization. Concretely, our contributions are as follows:

We provide new insights into understanding VQ networks by formulating commitment loss as a divergence measure. Doing so allows us to understand better why the divergence occurs.

To reduce this divergence, we propose an affine re-parameterization of the code-vectors that can better match the moments of the embedding representation. This alone drastically reduces the model collapse.

Lastly, We provide a set of improvements on the existing optimization techniques, such as alternating optimization and synchronized commitment loss. Both these methods are simple and more mathematically correct update rules that result in improvements over the standard approach.

Preliminaries

We denote xx as a scalar, x\mathbf{x} as a vector, XX as a matrix, X\mathcal{X} as a distribution or a set, f(⋅)f(\cdot) as a function, F(⋅)F(\cdot) as a composition of functions, and L(⋅)\mathcal{L}(\cdot) for loss function.

▹\triangleright Deep neural networks A feed-forward neural network is defined as a composition of parametric linear functions fif_{i} (e.g. fully-connected, convolutional layer) and non-linearities σ\sigma (e.g. ReLU):

In the context of generative modeling, F(⋅)F(\cdot) is referred to as the encoder and G(⋅)G(\cdot) as the decoder. The network is trained by minimizing the empirical risk Ltask(⋅)\mathcal{L}_{\mathsf{task}}(\cdot) with dataset D\mathcal{D}:

▹\triangleright Vector-quantized networks A vector-quantized network (VQN) is a neural-network consisting of a vector-quantization layer h(⋅,⋅)h(\cdot,\cdot):

The VQ layer h(⋅)h(\cdot) quantizes the embedding ze=F(x)\mathbf{z}_{e}=F(\mathbf{x}) by selecting a vector from a collection of mm vectors. The individual vector ci\mathbf{c}_{i} is referred to as the code-vector, the index ii as the code, and the collection of the code-vectors as the codebook C={c1,c2,…cm}\mathcal{C}=\{\mathbf{c}_{1},\mathbf{c}_{2},\dots\mathbf{c}_{m}\}. Here on out, we omit writing the codebook CC in the quantization function h(⋅)h(\cdot) for notational convenience. In the quantization function h(⋅)h(\cdot), ze\mathbf{z}_{e} is quantized into zq\mathbf{z}_{q} by assigning a code-vector from the codebook CC using a distance measure d(⋅,⋅)d(\cdot,\cdot):

Euclidean distance is the standard distance measure for d(⋅,⋅)d(\cdot,\cdot) (Van Den Oord et al., 2017). We denote the set associated to ze\mathbf{z}_{e}, zq\mathbf{z}_{q} and c\mathbf{c} with Pz\mathcal{P}_{\mathbf{z}}, Qz\mathcal{Q}_{\mathbf{z}} and Cz\mathcal{C}_{\mathbf{z}}, respectively, with Qz⊆Cz\mathcal{Q}_{\mathbf{z}}\subseteq\mathcal{C}_{\mathbf{z}}. Without loss of generality, we assume these sets are constructed from an underlying distribution.

The quantized embedding is then used to predict the output y^=G(zq)\hat{\mathbf{y}}=G(\mathbf{z}_{q}), and the loss is computed with the target y:L(y^,y)\mathbf{y}:\mathcal{L}\left(\hat{\mathbf{y}},\mathbf{y}\right). For images, VQ is generally performed on each spatial location of the tensor, where the channel dimension is used to represent the vector (e.g. each spatial location (i,j)(i,j) of ze∈Rh×w×c\mathbf{z}_{e}\in\mathbf{R}^{h\times w\times c} is quantized ze[i,j]∈Rc→h(⋅)zq[i,j]∈Rc\mathbf{z}_{e}[i,j]\in\mathbf{R}^{c}\xrightarrow[]{{}_{h(\cdot)}}\mathbf{z}_{q}[i,j]\in\mathbf{R}^{c}).

Akin to standard training, the objective of VQNs is to minimize the empirical risk:

The equation above is not continuously differentiable. To differentiate through the arg min⁡\operatorname*{arg\,min} operator in h(⋅)h(\cdot), a straight-through estimation (Bengio et al., 2013) is applied:

To ensure the straight-through estimation is accurate, the codebook and the encoder representations are pulled together using a commitment loss:

Here, β∈\beta\in is a scalar that trades off the importance of updating ze\mathbf{z}_{e} and zq\mathbf{z}_{q} (e.g. large β\beta implies more emphasis on the codebook to adapt towards the encoder). Then for some scalar α\alpha that weighs the commitment loss, a differentiable pseudo-objective is minimized:

A good rule of thumb is to set α=10\alpha=10 and β=0.9\beta=0.9 when using euclidean distance dl2(ze,zq):=12∥ze−zq∥22d_{l_{2}}(\mathbf{z}_{e},\mathbf{z}_{q}):=\frac{1}{2}\lVert\mathbf{z}_{e}-\mathbf{z}_{q}\rVert_{2}^{2} (Van Den Oord et al., 2017).

▹\triangleright Updating codebook using EMA Instead of using the commitment loss, another popular approach is to use exponential moving average (EMA) to train the codebook:

While EMA is proposed as a training trick in lieu of commitment loss, it is easy to see that it is equivalent to the commitment loss optimized with SGD when β=1\beta=1 (see Appendix A.2); where the EMA decay constant is the learning rate γ=η\gamma=\eta. This equivalence has been commonly overlooked with the exception of (Łańcucki et al., 2020).

On the trainability of VQ networks

It is well known that VQNs perform poorly when the number of actively used codes is small (Kaiser et al., 2018). This is referred to as “index collapse” and is the bottleneck in training VQNs. Thus, there have been abundant efforts to construct algorithms that recover and prevent models from collapsing. We discuss a few popular approaches below:

▹\triangleright Stochastic sampling Roy et al. (2018); Kaiser et al. (2018); Sønderby et al. (2017) proposed to use stochastic sampling and probabilistic relaxation to VQ. Since then, it has become a valid alternative method for training VQNs (Williams et al., 2020; Lee et al., 2022). Taking one of the most recent work as an example, Takida et al. (2022) argues that determinism is the main cause of the codebook collapse and proposes to sample codes from a categorical distribution proportional to the negative distance:

where τ\tau is a scalar indicating the temperature. The temperature is often annealed to make the model deterministic at convergence. Stochastic sampling can be a bottleneck as it requires computing and storing the full distance matrix. Note that the original idea of sampling to encourage diversity can be rooted back to Kohonen (1990).

▹\triangleright Repeated K-means Łańcucki et al. (2020) explicitly ensures all code-vectors are active by running K-means at every epoch. The naive implementation of repeated K-means forces all codes to be re-initialized. Depending on the noise sensitivity of the decoder, it can lead to large spikes in model performance, where both the encoder and the decoder have to readjust to the newly introduced codes. When decaying the learning rate, the model can no longer adapt to the new codes, and the performance stars to degrade.

▹\triangleright Replacement policy Zeghidour et al. (2021); Dhariwal et al. (2020) proposed to replace the dead codes with a randomly sampled embedding vector. Replacement policies require careful tuning, and we find that using a least-recently-used (LRU) policy with a life-span of 2020 iteration works the best – if the code does not get used for 2020 training iterations, it gets replaced with a random embedding vector. When using replacement policies, the active codes are left unchanged, and the overall performance of the model does not degrade.

As shown in Table 1, these aforementioned works result in better performance with an improved number of actively used codes. However, these methods address the symptoms of the collapse by replacing inactive codes rather than resolving why they became inactive in the first place. In this work, we investigate the source of the index collapse by analyzing the optimization dynamics of the codes and how it affects the model. We find that the divergence in the model representation causes the index collapse, and this divergence causes erroneous model updates.

VQ layers often diverge throughout training and fail to recover the codes that are no longer actively used (see Appendix A.6). While having good initial conditions, such as an improved initialization scheme (see Appendix A.7), can improve index collapse, ensuring good codebook usage is hard to maintain throughout training. To understand why dead codes are hard to recover from, we need to revisit the commitment loss. One can rewrite commitment loss as an average over distance d(⋅)d(\cdot) computed between an aligned set of points in Pz\mathcal{P}_{\mathbf{z}} and Cz\mathcal{C}_{\mathbf{z}}:

Here, when d(⋅)d(\cdot) is a Bregman divergence (e.g. l2l_{2} used in commitment loss is a Bregman divergence), then the distance D(Pz,Cz)D(\mathcal{P}_{\mathbf{z}},\mathcal{C}_{\mathbf{z}}) can be seen as an average divergence over an aligned set of Pz\mathcal{P}_{\mathbf{z}} and Cz\mathcal{C}_{\mathbf{z}} (Banerjee et al., 2005).

This divergence function is non-symmetric and is computed over Pz\mathcal{P}_{\mathbf{z}} but is minimized with respect to Cz\mathcal{C}_{\mathbf{z}}. Rewriting the commitment loss as an average divergence makes it easier to see why it is susceptible to model collapse. The divergence results in a many-to-one mapping, with the set of selected codes forming Qz\mathcal{Q}_{\mathbf{z}}. Here the disjoint set Cz∖Qz\mathcal{C}_{\mathbf{z}}\setminus\mathcal{Q}_{\mathbf{z}} does not receive any gradient and is not trained. This is analogous to computing reverse-KL in probability measure (Ghosh, 2018), where any sample that falls outside the support of the measure P\mathcal{P}, does not receive gradients (see Appendix A.10 for further discussion).

Since the subset of codes Q\mathcal{Q} is used to minimize the commitment loss, and not C\mathcal{C}, Q\mathcal{Q} learns a “mode-seeking” behavior. This implies that once the codes are not selected, they will likely remain unselected in the future. Note that even if we initialize the code-vectors to overlap in distribution with the embedding perfectly, the code-vectors can be dropped during optimization for various reasons, including stochasticity in training and non-stationary model representation Pz\mathcal{P}_{\mathbf{z}}.

We further demonstrate the impact of the divergence between the codebook and encoder embeddings by visualizing how they diverge in practice. In Figure 2 (left), we train ResNet18 (He et al., 2016) on the ImageNet100 (Russakovsky et al., 2015; Tian et al., 2020) dataset, initialized with K-means, and visualize Pz\mathcal{P}{\mathbf{z}}, Qz\mathcal{Q}{\mathbf{z}}, and Cz\mathcal{C}_{\mathbf{z}} using dimension reduction methods after several optimization steps. Using t-SNE, we observe that only a few code-vectors are active, with more than 95%95\% of the codes not being selected and trained. We also compute the histogram of the vectors by projecting them into the first PCA component. The histogram shows that the distribution quickly diverges after a few iterations of training. The experiment highlights the vulnerability of VQN optimization, where a sudden shift in the encoder embedding causes severe misalignment. Once misaligned, it often stays misaligned throughout optimization, with few active codes representing the embedding representation. A more desirable outcome can be observed in Figure 2 (right), where the codebook can closely match the embedding distribution. We discuss how to achieve better distribution matching in Section 4.1.

2 Gradient estimation gap

When the model embedding diverges from the codebook distribution, the quantization function yields sub-optimal code assignments. This sub-optimal assignment results in a sudden increase in the average quantization error (see Appendix A.6). As VQNs rely on straight-through estimation, the accuracy of the gradient updates in the encoder becomes dependent on the precision of the quantization function.

A good quantization function h(⋅)h(\cdot) is one that can preserve the necessary information of ze\mathbf{z}_{e} given a finite set of vectors. The resulting quantization vector can be represented as zq=ze+ϵ\mathbf{z}_{q}=\mathbf{z}_{e}+\epsilon where ϵ\epsilon is the residual error vector resulting from the quantization function. When ϵ=0\epsilon=\mathbf{0}, there exists no straight-through estimation error, and the model acts as if there is no quantization function. Of course, with a finite set of codes, a lossless quantization function is non-trivial to achieve when training on a large dataset. To measure the gradient deviation from the lossless quantization function, we define the gradient gap as:

The gradient gap measures the difference between the gradient of the non-quantized model and the quantized model. When Δgap=0\Delta_{\mathsf{gap}}=\mathbf{0}, the gradient descent using STE is guaranteed to minimize the loss. When the gap is large, no guarantees can be made. This gradient gap can be made small when (1) the quantization error is small and (2) the decoder function G(⋅)G(\cdot) is smooth. To see this, consider the case when ze\mathbf{z}_{e} and zq\mathbf{z}_{q} are equivalent (ϵ=0)(\epsilon=\mathbf{0}), then the estimation gap is Δgap=0\Delta_{\mathsf{gap}}=\mathbf{0}. When they are not equivalent and G(⋅)G(\cdot) is KK-Lipschitz smooth, then the estimation error is proportionately bounded by the quantization error K⋅d(ze,zq)K\cdot d(\mathbf{z}_{e},\mathbf{z}_{q}).

While regularizing for the smoothness of the network is task and model-dependent, the quantization error is a controllable design choice that users can improve upon. One approach is to ensure the quantization error is small at initialization using K-means (Łańcucki et al., 2020; Zeghidour et al., 2021; Karpathy, 2021)(see Appendix A.7). Another approach is to further improve gradient estimates by ensuring the quantization error is small throughout training. This can be done by improving the optimization algorithm itself. In Section 4.2, we propose an improved optimization algorithm to reduce the gradient estimation gap.

The gradient gap provides a measure of the goodness of the STE and is a useful tool for mathematical intuition. However, there is a caveat to be aware of when using it in practice. When VQNs go through index-collapse, there is a sharp spike in Δgap\Delta_{\mathsf{gap}} but eventually, it becomes very trivial for the model to achieve Δgap=0\Delta_{\mathsf{gap}}=0 as the encoder function is encouraged to predict the few remaining codes via the commitment loss. Where with fewer active codes, the easier it becomes to predict. Hence, one should be cautious when using the gradient gap as a measure to compare models.

Improved techniques for VQNs

The previous section emphasized contributing factors of index collapse and how minimizing the divergence at initialization leads to improved performance and codebook usage. This section explores methods to reduce codebook divergence and improve optimization.

In Section 3, we observed that methods such as resampling resulted in improved performance with a higher number of active codes. We hypothesize that replacement methods work well because the model representation Cz\mathcal{C}_{\mathbf{z}} eventually ends up resembling the representation of Pz\mathcal{P}_{\mathbf{z}} by consistently resampling code-vectors, albeit a slow process that requires resampling at almost every iteration. In light of this observation, we propose a more efficient method to match the distribution of Pz\mathcal{P}_{\mathbf{z}}. But first, we describe why misalignment occurs in the first place.

The misalignment between the internal representation of the network’s layers is referred to as an internal covariate shift. VQ layers are also prone to this internal covariate shift, where the consistently changing internal representation Pz\mathcal{P}_{\mathbf{z}} creates a misalignment with the codebook distribution Cz\mathcal{C}_{\mathbf{z}}. We refer to this misalignment as internal codebook covariate shift. Generally, internal covariate shifts can be minimized by directly matching the moments of the distributions, or in the case of (Ioffe & Szegedy, 2015), the distributions are whitened to a univariate Gaussian.

Using vector-quantization adds a layer of complexity that compounds with the existing internal covariate shift. While the linear layers receive dense gradients from the objective, the codebook receives sparse gradients. This implies that when the internal representation Pz\mathcal{P}_{\mathbf{z}} is updated, not only does it require much longer for Cz\mathcal{C}_{\mathbf{z}} to catch up but also if the update in Pz\mathcal{P}_{\mathbf{z}} is too large (e.g. large learning rates), the assignment can be severely misaligned (see Figure 2 and Appendix A.9).

To reduce gradient sparsity and encourage the codebook to update faster towards the embedding, we propose an affine reparameterization of the code-vector with a shared global mean and standard deviation.

Here csignal(i)\mathbf{c}^{(i)}_{\mathsf{signal}} is the original code-vector, and cmean,cstd\mathbf{c}_{\mathsf{mean}},\mathbf{c}_{\mathsf{std}} are the shared affine parameters with the same dim(csignal(i))\text{dim}(\mathbf{c}^{(i)}_{\mathsf{signal}}).

The affine parameters can either be learned through gradient descent or computed via the exponential moving average over ze\mathbf{z}_{e} and zq\mathbf{z}_{q} statistics (see Appendix A.12). Note, that under the Gaussian assumption on Pz\mathcal{P}_{\mathbf{z}} and Cz\mathcal{C}_{\mathbf{z}}, matching the moments is equivalent to minimizing the KL divergence (Kurz et al., 2016). The reparameterization allows gradients to flow through the unselected code-vectors through the affine parameters. Although we make no specific Gaussian assumption on the parameterization, it is easy to extend our method to better capture complex distributions by assigning distinct affine parameters to each codebook subset.

2 Alternated optimization

In Section 3.2, we showed that the error in the gradient update is proportional to the quantization error. Hence, we are interested in ensuring that divergence stays well-behaved during optimization. Here it is important to note that when there is an index collapse, the model trivially achieves zero gradient error as the commitment loss is easily minimized. This is not the setting we are interested in, and we assume that we are operating on a well-behaved regime.

For any arbitrary task Ltask\mathcal{L}_{\mathsf{task}}, the underlying objective of a VQN is to minimize the empirical loss while learning a good codebook representation.

The objective function above is not continuously differentiable; therefore, Eqn. 10 is used as a surrogate objective. However, the gradient computed from the surrogate objective is a biased estimate of the true gradient and can result in undesirable optimization dynamics. This is illustrated in Figure 3 with a toy setup. The dynamical error induced from the optimization can be traced to the straight-through estimation, in which the surrogate gradient deviates from the true gradient proportional to the quantization error. Hence, updating the network when the quantization error is big can lead to an erroneous model update. To reduce the quantization error, we propose an alternating optimization algorithm:

The algorithm above resembles that of online K-means with a non-linear encoder and decoder. Where Eqn. 17 optimizes the K-mean clusters, and Eqn. 18 optimizes the model given the new cluster assignment. We know that when Lcommit→0\mathcal{L}_{\mathsf{commit}}\rightarrow 0, h(⋅)h(\cdot) acts as an identity function under stationary FF and GG. Then, under fixed hh, both FF and GG can be optimized with close to zero estimation error. This can be repeated until the small quantization error assumption is broken.

Of course, optimizing the inner term till convergence is computationally expensive. Fortunately, we find that it is not necessary to wait till convergence (see Appendix A.11). In practice, alternating the inner term even for a single iteration yields good performance. In the toy setting of Figure 3, we visualize how the alternating optimization performs when using a single update for each term.

3 One-step behind or in synchronous step?

The codebook updated with the commitment loss is a historical average of the model representation. Writing out the gradient update for commitment loss for β=1\beta=1:

The historical average does not account for the current representation but only up to the previous representation. Therefore, computing the gradient with respect to the historical average implies that the model receives a “delayed” gradient. To reduce the delay in the zq\mathbf{z}_{q} representation, we desire the code-vector to include a running average of the most recent representation:

The above equation is tractably computable as the gradient of zq\mathbf{z}_{q} is used to update ze\mathbf{z}_{e}. Hence, an explicit equation for the synchronized update rule is:

Using the equation above, the code-vectors take a step in the direction of the encoder representation using the gradient of the task loss. In python, this requires a minor change to the existing implementation of straight-through estimation:

Where ν\nu is a scalar to decide whether we want a pessimistic or an optimistic update. We find that the effectiveness of ν\nu depends on the model architecture.

Results

We apply our methods to ImageNet100 (Tian et al., 2020) classification. ImageNet100 is a subset of the ImageNet-1K dataset (Russakovsky et al., 2015), consisting of 100 classes and approximately 100,000 images. The training details are provided in the Appendix A.1.

Our results, presented in Table 2, demonstrate the effectiveness of our proposed methods. All models were initialized using the K-means clustering algorithm. We also report the perplexity of the model on the test dataset, which is defined as 2H(p)2^{H(p)}, where H(p)H(p) is the entropy over the codebook likelihood. A higher perplexity implies a uniform assignment of codes. While having very low perplexity is associated with index collapse, having more does not necessarily imply better performance – having a high perplexity on a task that has more codes than necessary indicates redundancy. Our results using affine re-parameterization largely improves index-collapse, and the use of synchronized and alternating training methods further improves the overall performance of the model. We further compare our results against the use of the least-recently-used (LRU) replacement policy and observed our method to outperform or perform comparably. A combination of all these methods results in the largest improvement in performance. We observed using l2l_{2} normalization to hurt performance on classification. We suspect that removing the magnitude component of the embedding hurts models that use magnitude-sensitive objectives (e.g. soft-max cross-entropy loss).

2 Generative modeling

We further apply our method to CelebA (Liu et al., 2015) and CIFAR10 (Krizhevsky et al., 2009) generative modeling tasks. We adopt the training framework from previous works, and the details can be found in the appendix (Appendix A.1). We compare the performance of our method against existing techniques such as VQVAE (Van Den Oord et al., 2017), SQVAE (Takida et al., 2022), and Gumbel-VQVAE (Karpathy, 2021; Esser et al., 2020) using MSE as well as LPIPS perceptual loss (Zhang et al., 2018). SQVAE requires 4×4\times more memory footprint than all other methods as it requires storing the full distance matrix along with the computation graph. Both baselines using l2l_{2} normalization and least-recently-used (LRU) replacement policy largely improve training stability and reconstruction performance of generative models. When these methods are applied jointly with ours, we observe the best improvement.

In Figure 4, we run generative modeling using MaskGIT (Chang et al., 2022). We plot the rFID (reconstruction-FID) (Takida et al., 2022) and the FID (Heusel et al., 2017) during training. rFID measures the FID on the reconstructed images from the auto-encoder over the test set. We do not use perceptual or discriminator loss to train the network, and for rFID/FID, we use 50005000 generated samples.

3 Warmup and normalization can be helpful

One way to mitigate the divergence between the codebook and the embedding distribution is by constraining the representation to be within a bounded measure space. This limits the range of movement of the embedding distribution, facilitating alignment between the codebook and the embedding distribution throughout training. Common techniques for this include l2l_{2} normalization (Yu et al., 2022), batch-normalization (Łańcucki et al., 2020), and assuming a restricted distribution (Takida et al., 2022) in probabilistic VQNs. However, these techniques improve stability at the cost of reduced model expressivity. Alternatively, one can ensure the updates of Pz\mathcal{P}_{\mathbf{z}} to be small in order for the codebook to catch up. To do so without hindering convergence speed, we find a learning rate scheduler with warmup to work very well. In Appendix A.4, we show how using cosine learning rate decay with linear warmup (Loshchilov & Hutter, 2017; Goyal et al., 2017) improves both the model performance and model perplexity.

4 Ablation on alternating optimization

In Appendix A.11, we measure how varying the number of inner and outer loop iterations affects classification performance. We find that by increasing the number of inner loop iterations by 88, we observed an 11.09%11.09\% improvement over the baseline and 5.515.51% over the version that uses a single iteration of the inner step. On the other hand, we do not find increasing the outer loop to help. When combining all our methods, we observed that setting the inner loop iteration to 1-2 suffice. In this experiment, we evenly divide the training mini-batch into sub-mini-batches for each iteration of the expectation and maximization steps. This ensures that the number of images the model observes is equal across all experiments. Furthermore, when choosing to only optimize hh for a single iteration, it is possible to update both the inner and outer term in a single forward pass, as the commitment loss Lcommit\mathcal{L}_{\mathsf{commit}} does not depend on the task loss Ltask\mathcal{L}_{\mathsf{task}}. As a result, the computational overhead for a single fused pass is 1.05×1.05\times, and 2×2\times with four inner loop iterations (For smaller models like AlexNet, the overhead is less, roughly 1.5×1.5\times).

5 Further reducing sparsity in VQNs

Irrespective of the initial alignment between the codebook Cz\mathcal{C}_{z} and the embedding distribution Pz\mathcal{P}_{z}, a certain degree of divergence between these distributions is inevitable. This is particularly true for loss functions that are either unbounded (e.g., hinge loss) or have non-saturating gradients (e.g., logistic/exponential losses), where the weights grow inversely proportional to the loss. This effect is more pronounced in networks that lack normalization. Despite the use of affine reparameterization of the code-vectors, achieving 100100% utilization is non-trivial. To further reduce the sparsity in the codebook update, one can directly improve the architectural design choices that contribute to sparsity. Specifically, factors such as image size, batch size, and the number of pooling layers have a significant effect on VQN performance, as the number of code-vector selections directly depends on these variables. In Appendix A.5, we demonstrate that there is a significant degradation in performance when reducing the image size from 256×256256\times 256 to 128×128128\times 128, resulting in an over 2020% reduction in performance. This indicates the importance of training design choices is VQN.

Conclusion

Discretization has played a significant role in many fields, such as analog-to-digital communication and modern computing. Once representations are made discrete, various techniques from information theory can be applied to manipulate them for benefits such as compression, error correction, and robustness. Discrete representations can also be broken down into independent parts, allowing for the development of composable symbolic representations. In this work, we proposed a set of techniques that addresses several known challenges in optimizing vector-quantized models. Through our proposed methods, we were able to demonstrate improved model performance. While symbolic representation learning is still in its early stages, our optimization techniques provide insight for designing better models in the future.

Acknowledgement

We want to thank our lab members for their helpful feedback. Minyoung Huh was funded by ONR MURI grant N00014-22-1-2740. Brian Cheung is supported by the Center for Brains, Minds and Machines (CBMM), funded by NSF STC award CCF-1231216. Minyoung would like to further thank Wei-Chiu Ma, Lucy Chai, and Eunice Lee.

References

Appendix A Appendix

▹\triangleright VQ configuration and implementation We use 10241024 codes for all our experiments. 1024∼40961024\sim 4096 is the typical codebook size used in prior works (Esser et al., 2020; Yan et al., 2021). Using 40964096 codes does improve the performance slightly. We do not apply weight decay on the codebooks. For VQ hyper-parameters, we use α=5\alpha=5 and β=0.9∼0.995\beta=0.9\sim 0.995 with mean-squared error for the commitment loss. The performance starts to degrade outside this range of β\beta. Note that using higher β=1.0\beta=1.0 is equivalent to the EMA update. We implemented our own VQ algorithm to improve the run-time and memory efficiency. Standard implementation does not fit in memory for ImageNet training on standard commercial GPUs as the number of vectors grows close to 1 million for a given mini-batch. To mitigate this, we implemented our VQ algorithm using divide-and-conquer, which runs significantly faster and is more memory efficient compared to the naive implementation that allocates a contiguous memory to compute pair-wise distance. At a high level, our algorithm divides the batch of vectors into smaller chunks. For each chunk, batch matrix multiply (cdist\mathsf{cdist} in PyTorch) is applied for each subset and top-K reduction operation simultaneously. This is implemented in Python, and there is room for additional speed-up by implementing it in CUDA and parallelization via vmap. To further reduce computational and memory footprint, the distance is computed in half-precision, and the resulting code index is used to query the vector from float precision. Doing so adds no numerical imprecision to the existing model.

For affine-reparameterization, we use the variant with learnable affine parameters, which we find to be easier to implement and more stable in practice. We find that the optimal learning rate scale for the affine parameters needs to be tuned based on the model architecture as the norm of the magnitudes grows differently for each model during training. For learnable affine parameters, this can be easily implemented in Python via:

To be robust to the affine-parameter learning rate scale, we recommend using norm constraint (e.g., max norm constraint, norm clipping) or placing the VQ layer after an explicit normalization layer.

▹\triangleright Classification For AlexNet and ResNet18, we follow the design choices in https://github.com/pytorch/examples/blob/main/imagenet. The training configuration can be found in Table 4 and Table 5. For ViT(T), we started with the hyper-parameters recommended by (He et al., 2022) and re-tuned the hyper-parameters on the baseline VQ model. For ViT model codebase, we use the official PyTorch TorchiVison repository https://github.com/pytorch/vision/blob/main/torchvision/models/vision_transformer.py. The ViT(T) configuration is from https://github.com/rwightman/pytorch-image-models/blob/main/timm/models/vision_transformer.py. The architecture configuration is shown in Table 7 and the training configuration is shown in Table 6. ViT on ImageNet100 performs worse than ResNet18 as ViT does not perform very well on small datasets. This is a well-known observation (Chen et al., 2022). We apply data augmentation to the original image resolution and resize them to 224×224224\times 224 for training.

▹\triangleright AlexNet quantization For AlexNet, we quantize the features after the convolutional layers and before the fully-connected layers.

▹\triangleright ResNet18 quantization For ResNet18, we quantize after the second macroblock. This is after layer2 in the TorchVision repository. This is roughly the halfway point in the ResNet18. We found this model architecture to be insensitive ν\nu, possibly due to batch-normalization.

▹\triangleright ViT quanitzation ViT does not directly operate on image pixels but on image patches. Hence, the quantization cannot be applied to pixels. For our experimental setup, we tokenize non-overlapping patches of size 16×1616\times 16 resulting in a total of 14×14=19614\times 14=196 input tokens. These individual tokens are quantized. Note that this is often much fewer than the number of embedding vectors used in CNNs (e.g. feature embedding of 32×32=102432\times 32=1024 embedding vectors) and may be the reason why they suffer more from vector-quantization. We apply the quantization after the 6th transformer block. Training ViT without replacement policy is extremely finicky, which requires careful hyper-parameter tuning. When using replacmenet policy, it becomes more robust to wide-range of hyper-parameters. We recommend re-tuning the configuration when using it jointly with replacement policy.

▹\triangleright Generative modeling For generative modeling, we use backbone architecture from (Takida et al., 2022) with 6464 channels. The training configuration is listed in Table 8. We use MSE for reconstruction loss, and we do not use any perceptual or discriminative loss. For CIFAR10, we use an image size of 32×3232\times 32, and for CelebA, we use an image size of 128×128128\times 128.

For MaskGIT (Chang et al., 2022), we followed the author’s original codebase and reimplemented it in PyTorch. To ensure that we can fit the model in a commercial GPU (2424GB), we reduced the number of channels by 1/4th (3232 channels) for the auto-encoder and used an 88-layer transformer instead of 2424. The code resolution factor is 88, resulting in an 8×88\times 8 code map. A transformer is trained on the vectorized code. See (Chang et al., 2022) for the method details.

▹\triangleright Baseline For SQVAE (Takida et al., 2022), we use the Gaussian-SQVAE (4th-variant), which is the best-performing variant for unconditional generative modeling that uses a diagonal covariance matrix. We used the architecture from the official codebase and replicated their stochastic quantization layer to interface our framework. We anneal the temperature till convergence with a geometric decay. We re-tuned the learning rate and the weighting of the loss. We find the default scaling of 1∼0.11\sim 0.1 to work, setting it any higher did not train. SQVAE has an entropy term that encourages diversity in the codebook along with the ELBO objective. We observed that the model can easily collapse if the temperature is annealed too fast and, therefore, requires much longer to converge.

Gumbel-VQ was a variation proposed in the public repository by (Karpathy, 2021) and was also been used by (Esser et al., 2020) in their codebase. Similar to (Takida et al., 2022), Gumbel-VQ minimizes the ELBO; however, unlike standard VQ methods, which compute the distance across all code vectors, Gumbel-VQ predicts a distribution over the code without making any explicit comparisons. The model then uses the Gumbel-softmax (Jang et al., 2017) trick to sample from the distribution. We tried various hyper-parameters for the ELBO loss weight {5e-1, 5e-2, 5e-3, 5e-4, 5e-5} and found 5e-3 to work the best.

A.2 EMA and commitment loss

We show the equivalence between EMA and the commitment loss. Let β=1\beta=1 for the commitment loss with MSE for dd:

The gradient of the commitment loss is computed with respect to the codes zq\mathbf{z}_{q} is:

Then the update for zq(t+1)\mathbf{z}_{q}^{(t+1)} is:

Letting η=γ\eta=\gamma and using SGD to optimize the code-vectors, we recover the EMA update rule in Eqn. 11

A.3 Gradient estimation gap

To compute the gradient estimation gap in Eqn. 14, we compute 22 forward passes. Once with the quantization function and once without. Let g(l)g^{(l)} be the gradient without the quantization function for layer l∈Ll\in L and g^(l)\hat{g}^{(l)} be the gradient gap computed without the quantization error. Then the total average gradient error is ∑l∈L∥g(l)−g^(l)∥22\sum_{l\in L}\lVert g^{(l)}-\hat{g}^{(l)}\rVert^{2}_{2}. We visualize the gradient error of VQVAE training for the baseline model and our model that uses alternated optimization Figure 6. The VQ-layers use l2l_{2} normalization, and the gradient gap is computed in log-scale – hence the gradient gap is much larger than it appears in the figure.

A.4 Warmup improves perplexity and performance

As mentioned in Section 5.3, we found using a warmup to significantly improve codebook perplexity. This, in turn, results in better performance. In Figure 5, we compare VQNs trained with and without warmup. All methods are initialized using K-means. The baseline model uses a step scheduler. For ViT, the model collapses when we do not use a linear warmup. We hypothesize that a small learning rate allows the code-vectors to catch up to the model representation, while using a large learning rate at initialization causes misalignment between code-vectors and embedding vectors, resulting in index collapse.

A.5 Sensitivity of VQ models

Vector-quantization using the standard commitment loss result in sparse gradients by design. Therefore, it is imperative to isolate design factors that would contribute to this sparsity. Since sparsity is directly correlated with the selection rate of the code vectors, we can write out the selection probability under a simplifying assumption of i.i.d selection rate. For a VQN that operates on images, the likelihood that a code-vector ci\mathbf{c}_{i} will activate at-least kk times is:

Where hh and ww are the image dimensions and npooln_{\mathsf{pool}} is the number of 2×22\times 2 pooling layers, and bb is the batch size. The equation above is simply a summation of the binomial distribution over the selection probability. Here it becomes apparent that the image size, the batch size, the codebook size, and the number of pooling layers all directly affect the selection likelihood. While it is impossible to study the effect of these factors in perfect isolation, in Figure 7, we show how performance degrades compared to the non-quantized network.

A.6 Index Collapse

Index collapse is a phenomenon associated with the under-utilization of the codebook. While code-vectors become inactive throughout training, it is common to see a sudden collapse in code-vector usage early in training. In Figure 9, we visualize the index collapse on ResNet18 trained with K-means initialization. At initialization, all codes are actively used, with the gradient gap and divergence being close to 0. Here the divergence is measured between D(Cz,Qz)D(\mathcal{C}_{\mathbf{z}},\mathcal{Q}_{\mathbf{z}}), which measures the bifurcation of the codebook distribution. After a single iteration of SGD update, the code-vector starts to diverge, with the number of active codes dropping below 5050%. The resulting misalignment causes a large spike in the gradient error gap. The divergence continues to grow with the number of active codes approaching 1%1\% utilization. With fewer active codes, the encoder learns a degenerate solution of predicting these few remaining code-vectors.

A.7 The effect of VQ initialization

To improve gradient estimates, one can ensure the quantization error is small. To do so, one can ensure that the quantization error is small at initialization by choosing an appropriate weight initialization scheme. Unfortunately, Initializing the codebooks requires a priori knowledge of the model architecture and data distribution, as the distribution associated with the input embedding ze\mathbf{z}_{e} is hard to calculate ahead of time nor may not have a closed form distribution to sample from. For example, a VQ layer placed after a ReLU function implies that code-vectors in the negative half-space will never be sampled.

To mitigate this issue, it is a common practice to use data-dependent initialization such as K-means. These methods precisely capture the distribution of ze\mathbf{z}_{e} and better distributes the likelihood individual codebooks will be sampled. In fig X, we show the relationship between quantization error at initialization and the final performance of the model. Moreover, we find that the quantization error at initialization is a strong indicator for index-collapse / under-utilization of code-vectors – a phenomenon in which the number of active code-vectors at convergence is significantly smaller than the one we started with. As shown in Figure 8, a good initialization scheme mitigates index-collapse and leads to favorable performance.

A.8 Generation with MaskGIT

We extend our results on generative modeling using MaskGIT (Chang et al., 2022) on CelebA. In Table 9, the generation results improve from the baseline by 15.615.6 FID and 4.94.9 FID from the best-performing variation. We did not scale to images to compute FID. In Figure 11, we visualize randomly sampled images from the model. No samples were cherry sampled.

A.9 Affine reparameterization on toy setting

The use of vector quantization adds to an already existing issue of internal covariate shift. When the internal representation, Pz\mathcal{P}{\mathbf{z}}, is updated, it takes a longer time for the codebook, Cz\mathcal{C}{\mathbf{z}}, to catch up. Furthermore, if the update in Pz\mathcal{P}{\mathbf{z}} is too large, for example, due to a high learning rate, the assignment of codes to embeddings can become severely misaligned. To visualize this, consider a toy example illustrated in Figure 10 where the embedding distribution, Pz\mathcal{P}{\mathbf{z}}, and the codebook distribution, Cz\mathcal{C}_{\mathbf{z}}, undergo a drift (left). As a result, in the next iteration, only a small fraction of the codes are chosen and updated, while the rest remain unchanged. This misalignment in the codebook and embedding distributions leads to suboptimal model performance. Using affine reparameterization can mitigate this by allowing gradients to flow through all code-vectors implicitly via the shared parameters. The toy example uses the learnable parameter variant.

When implementing affine reparameterization, we observed better performance when using accumulated statistics. We accumulate the statistics with a momentum of 0.10.1. We compute the moments of both ze\mathbf{z}_{e} and C\mathcal{C} and compute the appropriate shift the match the moments of ze\mathbf{z}_{e}. We reduce the weighting of the commitment loss to α=2\alpha=2. When using the learnable parameters variant, while it performs slightly worse, the implementation only requires a 22 line change.

A.10 Discussion and intuition of the commitment loss

We discuss the difficulty of achieving perfect assignments using commitment loss. First, the upper bound of the commitment loss is given by:

Where the minimum is achieved when cj=μ(Pz)c_{j}=\mu({\mathcal{P}_{\mathbf{z}}}), the mean of the embedding. By moving the minimization function outside, we optimize with respect to a single code vector.

The tight lower bound of the commitment loss can be achieved by solving an assignment problem. Let B={B0,B1,…,Bm}\mathcal{B}=\{B_{0},B_{1},\dots,B_{m}\}, where ∣C∣=m|\mathcal{C}|=m. Each cluster ball BiB_{i} is a cluster associated with the code-vector cic_{i}. Using this notation, we rewrite the commitment loss as:

Here we see that minimizing the commitment loss is equivalent to computing a set of balls B\mathcal{B} with the mean of each ball being μ(Bi)=ci\mu(B_{i})=\mathbf{c}_{i} such that the sum of the weight variance of all the balls is minimized. Solving this assignment problem NP-hard and cannot be trivially attained. Note that all existing algorithm for solving K-means is a heuristics. The above objective solves a global optimization assignment problem while optimizing with SGD is a greedy local update. With abuse of notation, define B(cj)={∀  zi∈Pz  where  cj=arg min⁡(∥zi−cj∥2)}B(\mathbf{c}_{j})=\{\forall\;\mathbf{z}_{i}\in\mathcal{P}_{\mathbf{z}}\;\text{where}\;\mathbf{c}_{j}=\operatorname*{arg\,min}(\lVert\mathbf{z}_{i}-\mathbf{c}_{j}\rVert^{2})\} to be the set of points that is closest to the center cj\mathbf{c}_{j}. Then the commitment loss we optimize in practice is:

With the update rule for the code vectors being:

The update rule states that each code vector moves toward the mean of its current ball. Note that the cardinality of the ball can be zero; in such case, there is no gradient for the code-vector. Hence, in the worst case where B(cj)=∅B(\mathbf{c}_{j})=\emptyset for all j≠ij\neq i for some ii, we achieve the upper bound. Once the ball becomes empty, the code-vectors do not receive any gradients. To achieve the lower bound, one must hope that there is good coverage over Pz\mathcal{P}_{\mathbf{z}}, and by minimizing the loss, we recover the optimal assignment. This highlights the difficulty of achieving the optimal assignment using the commitment loss.

A.11 Alternating optimization ablation

In Table 10, we compare how changing the number of inner and outer optimization steps affects model performance. We find that the inner step of ×8\times 8 works the best. Increasing it further results in a negligible gain in performance.

A.12 Affine parameterization using EMA

When accumulating batch statistics for affine parameters, we compute an exponential moving average over the mean and variance of ze\mathbf{z}_{e} and zq\mathbf{z}_{q} with momentum mm:

We then use these statistics to normalize the zq\mathbf{z}_{q}. We first center the code-vectors and then re-normalize them back to the embedding moments:

The affine parameters then correspond to:

The momentum mm acts as the learning rate for the affine parameters. We find m∈[0.01,0.1]m\in[0.01,0.1] to be a good starting point for most models.