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

Minyoung HuhBrian CheungPulkit AgrawalPhillip Isola

article2023ICML120 citations

Identifies internal codebook covariate shift as the fundamental cause of index collapse in vector-quantized networks and introduces an affine re-parameterization alongside alternating optimization to stabilize training across vision and generative architectures.

Listen

Vector-quantized neural networks convert continuous data representations into discrete codes, making them foundational for modern compressed image generation, speech processing, and decision-making systems. However, training these models is notoriously brittle and unstable due to an issue known as index collapse, where the network abandons most of its available discrete codes early in optimization. Historically, practitioners have relied on heuristic fixes like randomly resetting unused codes rather than solving the underlying mathematical failure. The article investigates the root causes of this optimization instability and develops principled techniques to stabilize vector-quantized model training.

The research reveals that training instability stems from internal codebook covariate shift—a severe distributional divergence between the encoder features and the discrete codebook. Because the standard training objective uses an asymmetric distance calculation, unselected codes receive zero gradient updates and are permanently dropped. This divergence causes the straight-through estimator (the mathematical shortcut used to bypass non-differentiable code selection) to produce highly biased, inaccurate gradient updates. To resolve this, the authors propose three core interventions: an affine re-parameterization that shares global mean and variance parameters across all code vectors, an alternating optimization routine that updates code assignments before updating the network weights, and a synchronized commitment loss that accounts for immediate parameter updates rather than lagging behind.

The proposed techniques deliver consistent performance gains across both visual classification and image generation tasks using standard network architectures like AlexNet, ResNet-18, and Vision Transformers. On the ImageNet-100 classification benchmark, applying these optimizations improved absolute classification accuracy by 6.9 to 10.7 percentage points across architectures while maintaining high codebook utilization. For generative modeling on CelebA and CIFAR-10 datasets, the approach significantly improved image reconstruction and generation quality. For example, in a transformer-based generative test on CelebA, the method improved the standard image generation quality score (FID) from 90.4 down to 74.8, indicating much higher visual fidelity.

These results establish that index collapse is an optimization and gradient estimation error rather than an unavoidable property of discrete neural networks. The findings provide substantial practical value: they remove the need for memory-heavy sampling alternatives and reduce reliance on fragile heuristic resets, allowing models to train faster with minimal computational overhead (as low as 1.05 times standard training for fused passes). Teams developing discrete representation pipelines should immediately adopt the shared affine re-parameterization and synchronized loss updates, as they require minimal code modifications while substantially boosting robustness.

Users should note certain operational boundaries identified in the article. Performance remains sensitive to standard hyperparameter choices, such as warmup learning rate schedules and specific batch sizes, because sparse activations can still degrade discrete code utilization. While confidence in the experimental improvements is high across the tested image classification and generation benchmarks, the authors note that hyperparameter scales for affine updates may still require slight tuning depending on the specific model architecture.

arXiv: 2305.08842
  • Paper: Neural Discrete Representation Learning, Aäron van den Oord et al. (2017). This seminal paper introduces the Vector Quantised-Variational AutoEncoder (VQ-VAE) and the straight-through estimator for discrete latent representations, establishing the foundational architecture and optimization objective that the source diagnoses and repairs.
  • Paper: Taming Transformers for High-Resolution Image Synthesis, Patrick Esser et al. (2020). This work pairs vector-quantized codebooks with generative transformers to synthesize high-resolution images, providing the primary discrete representation framework that the source seeks to stabilize against codebook collapse.
  • Paper: Generating Diverse High-Fidelity Images with VQ-VAE-2, Ali Razavi et al. (2019). This paper extends discrete representation learning to multi-scale hierarchical codebooks, demonstrating the architectural benefits and training sensitivities of large-scale vector-quantized generative models.
  • Paper: Nonuniform-to-Uniform Quantization: Towards Accurate Quantization via Generalized Straight-Through Estimation, Zechun Liu et al. (2022). This study analyzes the mathematical biases and gradient approximations inherent to the straight-through estimator, offering critical background on STE failure modes in quantized neural networks.
Cover for Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks

Abstract

This work examines the challenges of training neural networks using vector quantization using straight-through estimation. We find that a primary cause of training instability is the discrepancy between the model embedding and the code-vector distribution. We identify the factors that contribute to this issue, including the codebook gradient sparsity and the asymmetric nature of the commitment loss, which leads to misaligned code-vector assignments. We propose to address this issue via affine re-parameterization of the code vectors. Additionally, we introduce an alternating optimization to reduce the gradient error introduced by the straight-through estimation. Moreover, we propose an improvement to the commitment loss to ensure better alignment between the codebook representation and the model embedding. These optimization methods improve the mathematical approximation of the straight-through estimation and, ultimately, the model performance. We demonstrate the effectiveness of our methods on several common model architectures, such as AlexNet, ResNet, and ViT, across various tasks, including image classification and generative modeling.

Project page: minyoungg.github.io/vqtorch

Table of Contents

  • 1. Introduction and related works
  • Post distribution shift
  • 2. Preliminaries
  • Deep neural networks
  • Vector-quantized networks
  • Updating codebook using EMA
  • 3. On the trainability of VQ networks
  • Stochastic sampling
  • Repeated K-means
  • Replacement policy
  • 3.1. Commitment loss is an asymmetric loss
  • 3.2. Gradient estimation gap
  • 4. Improved techniques for VQNs
  • 4.1. Minimizing internal codebook covariate shift with shared affine parameterization
  • 4.2. Alternated optimization
  • 4.3. One-step behind or in synchronous step?
  • 5. Results
  • 5.1. Classification
  • 5.2. Generative modeling
  • 5.3. Warmup and normalization can be helpful
  • 5.4. Ablation on alternating optimization
  • 5.5. Further reducing sparsity in VQNs
  • 6. Conclusion
  • 7. Acknowledgement
  • References
  • A. Appendix
  • A.1. Training details
  • VQ configuration and implementation
  • Classification
  • AlexNet quantization
  • ResNet18 quantization
  • ViT quantization
  • Generative modeling
  • Baseline
  • A.2. EMA and commitment loss
  • A.4. Warmup improves perplexity and performance
  • A.3. Gradient estimation gap
  • A.5. Sensitivity of VQ models
  • A.6. Index Collapse
  • A.7. The effect of VQ initialization
  • A.8. Generation with MaskGIT
  • A.9. Affine reparameterization on toy setting
  • A.10. Discussion and intuition of the commitment loss
  • A.11. Alternating optimization ablation
  • A.12. Affine parameterization using EMA

Knowls

  1. Knowl 1 — Straight-through vector quantization and commitment objective

    model/method

    A vector-quantized network encodes an input xx as ze=F(x)z_e=F(x), selects the nearest vector from a codebook C={c1,…,cm}C=\{c_1,\ldots,c_m\}, and decodes the selected vector zqz_q with GG. For a distance function dd, the quantizer is

    zq=ck,k=arg⁡min⁡j∈{1,…,m}d(ze,cj).z_q=c_k,\qquad k=\arg\min_{j\in\{1,\ldots,m\}}d(z_e,c_j).

    The network output is y^=G(zq)\hat y=G(z_q), and training minimizes a task loss together with a commitment loss:

    min⁡F,G,h  E(x,y)∼D[Ltask(G(h(F(x))),y)+αLcmt(ze,zq)],\min_{F,G,h}\;\mathbb{E}_{(x,y)\sim D}\left[L_{\mathrm{task}}\bigl(G(h(F(x))),y\bigr)+\alpha L_{\mathrm{cmt}}(z_e,z_q)\right],

    where hh is the quantizer, α>0\alpha>0 weights the commitment term, and yy is the target. With stop-gradient denoted by sg⁡(⋅)\operatorname{sg}(\cdot), the commitment loss used by the paper is

    Lcmt(ze,zq)=(1−β)d(ze,sg⁡(zq))+βd(sg⁡(ze),zq),L_{\mathrm{cmt}}(z_e,z_q)=(1-\beta)d\bigl(z_e,\operatorname{sg}(z_q)\bigr)+\beta d\bigl(\operatorname{sg}(z_e),z_q\bigr),

    where β∈[0,1]\beta\in[0,1] controls the relative movement of the embedding and codebook. The nondifferentiable nearest-neighbor operation is bypassed with a straight-through estimator (STE), which substitutes the derivative of the quantized embedding with respect to the encoder embedding by the identity. For the squared Euclidean distance d(a,b)=12∥a−b∥22d(a,b)=\tfrac12\|a-b\|_2^2, the paper gives α=10\alpha=10 and β=0.9\beta=0.9 as a general starting point; its experiments use α=5\alpha=5 and generally β∈[0.9,0.995]\beta\in[0.9,0.995]. When β=1\beta=1, SGD on the squared commitment loss produces the same code-vector update as an exponential moving average (EMA) with EMA decay equal to the SGD learning rate.

  2. Knowl 2 — Internal codebook covariate shift causes index collapse

    theoretical result

    The paper identifies the primary failure mechanism in straight-through vector-quantized networks as divergence between the encoder-embedding distribution PzP_z and the codebook distribution CzC_z. For a set or distribution of embeddings PzP_z and code vectors CzC_z, the commitment objective can be expressed as the average nearest-code divergence

    D(Pz,Cz)=1∣Pz∣∑zi∈Pzmin⁡cj∈Czd(zi,cj).D(P_z,C_z)=\frac{1}{|P_z|}\sum_{z_i\in P_z}\min_{c_j\in C_z}d(z_i,c_j).

    This objective is asymmetric: it averages over embedding points and minimizes with respect to code vectors. Only the subset Qz⊆CzQ_z\subseteq C_z that is selected by the nearest-neighbor assignments receives gradients. Code vectors in Cz∖QzC_z\setminus Q_z receive no update, so the selected subset exhibits mode-seeking behavior and tends to remain selected while inactive codes remain inactive. Stochastic encoder updates and nonstationary representations can therefore turn an initially well-aligned codebook into a bifurcated, underutilized codebook.

    The paper calls this mismatch internal codebook covariate shift. It is aggravated by the fact that ordinary network parameters receive dense gradients whereas code vectors receive sparse assignment-dependent gradients. In a ResNet18/ImageNet100 visualization initialized with K-means, more than 95%95\% of code vectors became unselected after a few iterations, while the embedding, selected-code, and full-codebook distributions separated. This distributional divergence precedes index collapse, in which the encoder eventually predicts only a small fraction of the available codes.

  3. Knowl 3 — Gradient estimation gap links quantization error to STE reliability

    theoretical result

    Let FF be the encoder, GG the decoder, ze=F(x)z_e=F(x) the continuous embedding, and zq=ze+ϵz_q=z_e+\epsilon the quantized embedding, where ϵ\epsilon is the quantization residual. The paper defines the straight-through gradient gap for an input xx as

    Δgap(F∣h)=∥∂Ltask(G(ze))∂F(x)−∂Ltask(G(ze+ϵ))∂F(x)∥.\Delta_{\mathrm{gap}}(F\mid h)=\left\|\frac{\partial L_{\mathrm{task}}(G(z_e))}{\partial F(x)}-\frac{\partial L_{\mathrm{task}}(G(z_e+\epsilon))}{\partial F(x)}\right\|.

    The first term is the gradient for the corresponding non-quantized network and the second is the gradient produced when the quantized representation is used. If ϵ=0\epsilon=0, then Δgap=0\Delta_{\mathrm{gap}}=0 and the STE update agrees with the non-quantized update. The paper states that when the decoder GG is KK-Lipschitz smooth, the estimation error is proportionally bounded by the quantization error, with scale K d(ze,zq)K\,d(z_e,z_q). Thus, reducing quantization error either at initialization or throughout training makes the STE more faithful.

    The gap is not a reliable standalone measure after collapse: once the encoder has learned to predict only a few surviving code vectors, the quantization error can become small and the gap can return toward zero even though codebook utilization and model quality are poor.

  4. Knowl 4 — Shared affine reparameterization aligns embeddings and code vectors

    model/method

    To counter internal codebook covariate shift, the paper reparameterizes every code vector using shared affine parameters:

    c(i)=cmean+cstd ∗ csignal(i), c^{(i)}=c_{\mathrm{mean}}+c_{\mathrm{std}}\,*\,c^{(i)}_{\mathrm{signal}},

    where csignal(i)c^{(i)}_{\mathrm{signal}} is the original code vector, cmeanc_{\mathrm{mean}} and cstdc_{\mathrm{std}} are learnable vectors with the same dimension as a code vector, and ∗* denotes elementwise multiplication. The shared parameters can instead be estimated from exponential moving averages of embedding and codebook statistics.

    The affine transformation gives every code vector an indirect gradient through the shared parameters, including code vectors that were not selected by the nearest-neighbor assignment. It can therefore shift and rescale the entire codebook toward the current embedding distribution instead of waiting for sparse per-code updates. Under a Gaussian modeling assumption, matching means and variances is equivalent to minimizing the corresponding KL divergence; the paper does not require that assumption for the basic parameterization. In the experiments, the learnable-affine variant was generally used because it was easier to implement and more stable. The distribution visualizations on ResNet18/ImageNet100 show that affine reparameterization avoids the standard codebook bifurcation and produces closer alignment between PzP_z, QzQ_z, and CzC_z.

  5. Knowl 5 — Alternating optimization reduces quantization-induced model updates

    algorithm

    The paper separates codebook assignment optimization from task-model optimization to keep the quantization error small while using the STE. For a data distribution pdatap_{\mathrm{data}}, the alternating objectives are

    min⁡h  E(x,y)∼pdata[Lcmt(h(F(x)),F(x))]\min_h\;\mathbb{E}_{(x,y)\sim p_{\mathrm{data}}}\left[L_{\mathrm{cmt}}\bigl(h(F(x)),F(x)\bigr)\right]

    followed by

    min⁡F,G  E(x,y)∼pdata[Ltask(G(h(F(x))),y)].\min_{F,G}\;\mathbb{E}_{(x,y)\sim p_{\mathrm{data}}}\left[L_{\mathrm{task}}\bigl(G(h(F(x))),y\bigr)\right].

    Here the inner step updates the quantizer/codebook using the current encoder representation, and the outer step updates the encoder and decoder using the newly improved assignments and the task loss. The procedure is analogous to online K-means with a nonlinear encoder and decoder: the codebook is first brought toward the current embedding distribution, then the model is optimized with a smaller quantization error.

    Input: encoder FF, decoder GG, quantizer hh, training minibatches, inner-step count rr, outer-step count ss
    Repeat for each training update:
        For rr inner iterations:
            Compute embeddings ze=F(x)z_e=F(x) on the current minibatch
            Update codebook assignments and code vectors using only the commitment loss
        For ss outer iterations:
            Recompute or use the updated quantized embeddings zq=h(F(x))z_q=h(F(x))
            Update FF and GG using the task loss and the straight-through estimator
    Return: the trained encoder, decoder, and quantizer

    The authors found that full convergence of the inner problem is unnecessary. A single inner iteration works well in practice, and one to two inner iterations suffice when all proposed techniques are combined. In an isolated ablation, eight inner iterations gave the best reported accuracy improvement, while increasing the number of outer iterations did not help.

  6. Knowl 6 — Synchronous codebook updates remove the historical EMA delay

    model/method

    The usual commitment-loss or EMA update uses the previous encoder representation:

    zq(t+1)←(1−η)zq(t)+ηze(t), z_q^{(t+1)}\leftarrow(1-\eta)z_q^{(t)}+\eta z_e^{(t)},

    where tt is the training step and η\eta is the codebook update rate. Because zq(t)z_q^{(t)} summarizes historical embeddings, the task model receives a delayed codebook representation. The synchronized alternative uses the just-updated embedding:

    zq(t+1)←(1−η)zq(t)+ηze(t+1). z_q^{(t+1)}\leftarrow(1-\eta)z_q^{(t)}+\eta z_e^{(t+1)}.

    The paper implements this behavior by modifying the straight-through path as

    z_q = z_e + (z_q - z_e).detach() + nu * (z_q - z_q.detach())
    

    where nu is a scalar controlling the additional task-gradient contribution through the codebook representation. The reported classification settings use ν=2\nu=2 for AlexNet, ν=0.2\nu=0.2 for ResNet18, and ν=0.01\nu=0.01 for ViT-Tiny. The synchronized rule reduces the delay between encoder and codebook updates and improves optimization when combined with affine reparameterization and alternating optimization.

  7. Knowl 7 — Classification performance improves across CNN and ViT architectures

    data/table

    The paper evaluates its optimization methods on ImageNet100 classification using K-means codebook initialization, 1024 codes, and the AlexNet, ResNet18, and ViT-Tiny architectures. “Affine,” “Sync.,” “Alt.,” and “Replace” denote affine reparameterization, synchronized updates, alternating optimization, and least-recently-used replacement, respectively. The table reports test accuracy and codebook perplexity; higher perplexity generally indicates more uniform code assignments, although excessive perplexity can also indicate redundancy.

    Could not parse LaTeX table

    Affine reparameterization provides the largest single improvement in each architecture. Synchronized and alternating optimization each add further gains, and the complete combination gives the best accuracy for all three models. The authors also report that ℓ2\ell_2 normalization hurt classification, plausibly because it removes embedding magnitude information used by magnitude-sensitive objectives such as softmax cross-entropy.

  8. Knowl 8 — The combined method improves image reconstruction and code usage

    data/table

    The paper evaluates VQVAE-style generative reconstruction on CIFAR10 and CelebA with a common backbone, MSE reconstruction loss, and no perceptual or discriminator loss. Metrics are computed on the test sets after hyperparameter tuning: MSE is reported in units of 10−310^{-3} and LPIPS in units of 10−110^{-1}, both lower-is-better; perplexity is higher-is-better as an indicator of code usage. “OPT” denotes the proposed alternating optimization, while “replace” is the LRU dead-code replacement policy.

    Could not parse LaTeX table

    The strongest reconstruction results arise when affine reparameterization and alternating optimization are combined with dead-code replacement; adding ℓ2\ell_2 normalization gives the best CIFAR10 result but slightly worsens CelebA MSE relative to the same method without normalization. The improvements demonstrate that better codebook alignment and less biased optimization translate into both higher code usage and better reconstruction.

  9. Knowl 9 — Warmup and activation probability govern codebook stability

    empirical result

    The paper finds that a small learning rate at the beginning of training allows code vectors to follow the changing embedding distribution. Cosine learning-rate decay with linear warmup improves both codebook perplexity and classification accuracy; without warmup, the ViT model collapses in the reported comparison. The authors interpret this as evidence that large initial encoder updates can create codebook-embedding misalignment before sparse codebook updates can recover.

    Under an i.i.d. code-selection approximation, the probability that a particular code vector cic_i is selected at least kk times for an image VQN is

    p(ci activates at least k times)=1−∑j=0k−1(bhw/2npoolj)(1C)j(1−1C)bhw/2npool−j,p(c_i\text{ activates at least }k\text{ times})=1-\sum_{j=0}^{k-1}\binom{bhw/2^{n_{\mathrm{pool}}}}{j}\left(\frac{1}{C}\right)^j\left(1-\frac{1}{C}\right)^{bhw/2^{n_{\mathrm{pool}}}-j},

    where bb is batch size, hh and ww are image height and width, npooln_{\mathrm{pool}} is the number of 2×22\times2 pooling layers, and CC is codebook size. Consequently, larger batches and images increase code-selection opportunities, whereas larger codebooks and more pooling reduce them. The paper’s sensitivity plots show that low activation ratios are associated with substantially worse performance; reducing the image size from 256×256256\times256 to 128×128128\times128 caused an over-20%20\% performance reduction in one experiment.

    These findings also qualify the proposed methods: affine reparameterization does not guarantee full code utilization, and normalization-based stabilization can reduce model expressivity. Architectural choices that determine the number of code selections remain important sources of sparsity.

  10. Knowl 10 — MaskGIT generation benefits from affine and alternating optimization

    empirical result

    The paper extends the generative evaluation to MaskGIT on CelebA. To fit the experiment on a 24 GB GPU, it uses a reduced-capacity setup: a 32-channel autoencoder rather than 128 channels, an 8-block transformer rather than 24 blocks, and a code resolution factor of 8. FID is computed on generated CelebA images, with lower values indicating better generation.

    Could not parse LaTeX table

    The proposed affine reparameterization plus alternating optimization improves FID by 15.615.6 relative to the reduced-capacity baseline and by 4.94.9 relative to the best individual stabilization variant, LRU replacement. The training curves on page 8 show lower reconstruction-FID and generation FID for the proposed method than for the baseline and the alternatives.

Coverage note — The prior-method survey, full optimizer and architecture configuration tables, detailed EMA-statistics variant, and proof-only commitment-loss bounds were omitted because they are either background, implementation detail, or intermediate analysis rather than top-ranked load-bearing contributions.

References

  1. 1.Asano, Y. M., Rupprecht, C., and Vedaldi, A. Self-labelling via simultaneous clustering and representation learning. In International Conference on Learning Representations, 2020.
  2. 2.Banerjee, A., Merugu, S., Dhillon, I. S., Ghosh, J., and Lafferty, J. Clustering with bregman divergences. Journal of machine learning research, 6(10), 2005.
  3. 3.Bengio, Y., Léonard, N., and Courville, A. Estimating or propagating gradients through stochastic neurons for conditional computation. arXiv preprint arXiv:1308.3432, 2013.
  4. 4.Caron, M., Bojanowski, P., Joulin, A., and Douze, M. Deep clustering for unsupervised learning of visual features. In Proceedings of the European conference on computer vision (ECCV), pp. 132–149, 2018.
  5. 5.Caron, M., Misra, I., Mairal, J., Goyal, P., Bojanowski, P., and Joulin, A. Unsupervised learning of visual features by contrasting cluster assignments. Advances in Neural Information Processing Systems, 33:9912–9924, 2020.
  6. 6.Chang, H., Zhang, H., Jiang, L., Liu, C., and Freeman, W. T. Maskgit: Masked generative image transformer. In Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition, pp. 11315–11325, 2022.
  7. 7.Chang, H., Zhang, H., Barber, J., Maschinot, A., Lezama, J., Jiang, L., Yang, M.-H., Murphy, K., Freeman, W. T., Rubinstein, M., et al. Muse: Text-to-image generation via masked generative transformers. arXiv preprint arXiv:2301.00704, 2023.
  8. 8.Chen, X., Hsieh, C.-J., and Gong, B. When vision transformers outperform resnets without pre-training or strong data augmentations. In International Conference on Learning Representations, 2022.
  9. 9.Chung, Y.-A., Tang, H., and Glass, J. Vector-quantized autoregressive predictive coding. In Interspeech, 2020.
  10. 10.Dhariwal, P., Jun, H., Payne, C., Kim, J. W., Radford, A., and Sutskever, I. Jukebox: A generative model for music. arXiv preprint arXiv:2005.00341, 2020.
  11. 11.Esser, P., Rombach, R., and Ommer, B. Taming transformers for high-resolution image synthesis. In Proceedings of the IEEE Conference on Computer Vision and Pattern Recognition, 2020.
  12. 12.Ghosh, D. Kl divergence for machine learning. https://dibyaghosh.com/blog/probability/kldivergence.html, 2018.
  13. 13.Goyal, P., Dollár, P., Girshick, R., Noordhuis, P., Wesolowski, L., Kyrola, A., Tulloch, A., Jia, Y., and He, K. Accurate, large minibatch sgd: Training imagenet in 1 hour. arXiv preprint arXiv:1706.02677, 2017.
  14. 14.Gray, R. Vector quantization. IEEE Assp Magazine, 1(2): 4–29, 1984.
  15. 15.He, K., Zhang, X., Ren, S., and Sun, J. Deep residual learning for image recognition. In Proceedings of the IEEE conference on computer vision and pattern recognition, pp. 770–778, 2016.
  16. 16.He, K., Chen, X., Xie, S., Li, Y., Dollár, P., and Girshick, R. Masked autoencoders are scalable vision learners. In Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition, pp. 16000–16009, 2022.
  17. 17.Heusel, M., Ramsauer, H., Unterthiner, T., Nessler, B., and Hochreiter, S. Gans trained by a two time-scale update rule converge to a local nash equilibrium. Advances in neural information processing systems, 30, 2017.
  18. 18.Ioffe, S. and Szegedy, C. Batch normalization: Accelerating deep network training by reducing internal covariate shift. In International conference on machine learning, pp. 448–456. PMLR, 2015.
  19. 19.Jang, E., Gu, S., and Poole, B. Categorical reparameterization with gumbel-softmax. In International Conference on Learning Representations, 2017.
  20. 20.Kaiser, L., Bengio, S., Roy, A., Vaswani, A., Parmar, N., Uszkoreit, J., and Shazeer, N. Fast decoding in sequence models using discrete latent variables. In International Conference on Machine Learning, pp. 2390–2399. PMLR, 2018.
  21. 21.Karpathy, A. deep-vector-quantization. https://github.com/karpathy/deep-vector-quantization, 2021.
  22. 22.Kohonen, T. Improved versions of learning vector quantization. In 1990 ijcnn international joint conference on Neural networks, pp. 545–550. IEEE, 1990.
  23. 23.Krizhevsky, A., Hinton, G., et al. Learning multiple layers of features from tiny images. 2009.
  24. 24.Kurz, G., Pfaff, F., and Hanebeck, U. D. Kullback-leibler divergence and moment matching for hyperspherical probability distributions. In 2016 19th International Conference on Information Fusion (FUSION), pp. 2087–2094. IEEE, 2016.
  25. 25.Łancucki, A., Chorowski, J., Sanchez, G., Marxer, R., Chen, N., Dolfing, H. J., Khurana, S., Alumäe, T., and Laurent, A. Robust training of vector quantized bottleneck models. In 2020 International Joint Conference on Neural Networks (IJCNN), pp. 1–7. IEEE, 2020.
  26. 26.Lee, D., Kim, C., Kim, S., Cho, M., and Han, W.-S. Autoregressive image generation using residual quantization. In Proceedings of the IEEE conference on computer vision and pattern recognition, 2022.
  27. 27.Liu, Z., Luo, P., Wang, X., and Tang, X. Deep learning face attributes in the wild. In Proceedings of International Conference on Computer Vision (ICCV), December 2015.
  28. 28.Loshchilov, I. and Hutter, F. Sgdr: Stochastic gradient descent with warm restarts. In International Conference on Learning Representations, 2017.
  29. 29.Ozair, S., Li, Y., Razavi, A., Antonoglou, I., Van Den Oord, A., and Vinyals, O. Vector quantized models for planning. In International Conference on Machine Learning, pp. 8302–8313. PMLR, 2021.
  30. 30.Paszke, A., Gross, S., Massa, F., Lerer, A., Bradbury, J., Chanan, G., Killeen, T., Lin, Z., Gimelshein, N., Antiga, L., Desmaison, A., Kopf, A., Yang, E., DeVito, Z., Raison, M., Tejani, A., Chilamkurthy, S., Steiner, B., Fang, L., Bai, J., and Chintala, S. Pytorch: An imperative style, high-performance deep learning library. In Wallach, H., Larochelle, H., Beygelzimer, A., d Alché-Buc, F., Fox, E., and Garnett, R. (eds.), Advances in Neural Information Processing Systems 32, pp. 8024–8035. Curran Associates, Inc., 2019.
  31. 31.Ramesh, A., Pavlov, M., Goh, G., Gray, S., Voss, C., Radford, A., Chen, M., and Sutskever, I. Zero-shot text-to-image generation. In International Conference on Machine Learning, pp. 8821–8831. PMLR, 2021.
  32. 32.Roy, A., Vaswani, A., Neelakantan, A., and Parmar, N. Theory and experiments on vector quantized autoencoders. arXiv preprint arXiv:1805.11063, 2018.
  33. 33.Russakovsky, O., Deng, J., Su, H., Krause, J., Satheesh, S., Ma, S., Huang, Z., Karpathy, A., Khosla, A., Bernstein, M., et al. Imagenet large scale visual recognition challenge. International journal of computer vision, 2015.
  34. 34.Sønderby, C. K., Poole, B., and Mnih, A. Continuous relaxation training of discrete latent variable image models. In Beysian DeepLearning workshop, NIPS, volume 201, 2017.
  35. 35.Takida, Y., Shibuya, T., Liao, W., Lai, C.-H., Ohmura, J., Uesaka, T., Murata, N., Takahashi, S., Kumakura, T., and Mitsufuji, Y. Sq-vae: Variational bayes on discrete representation with self-annealed stochastic quantization. arXiv preprint arXiv:2205.07547, 2022.
  36. 36.Tian, Y., Krishnan, D., and Isola, P. Contrastive multiview coding. In European conference on computer vision, pp. 776–794. Springer, 2020.
  37. 37.Van Den Oord, A., Vinyals, O., et al. Neural discrete representation learning. Advances in neural information processing systems, 30, 2017.
  38. 38.Williams, W., Ringer, S., Ash, T., MacLeod, D., Dougherty, J., and Hughes, J. Hierarchical quantized autoencoders. Advances in Neural Information Processing Systems, 33: 4524–4535, 2020.
  39. 39.Yan, W., Zhang, Y., Abbeel, P., and Srinivas, A. Videogpt: Video generation using vq-vae and transformers. arXiv preprint arXiv:2104.10157, 2021.
  40. 40.Yu, J., Li, X., Koh, J. Y., Zhang, H., Pang, R., Qin, J., Ku, A., Xu, Y., Baldridge, J., and Wu, Y. Vector-quantized image modeling with improved vqgan. In International Conference on Learning Representations, 2022.
  41. 41.Zeghidour, N., Luebs, A., Omran, A., Skoglund, J., and Tagliasacchi, M. Soundstream: An end-to-end neural audio codec. IEEE/ACM Transactions on Audio, Speech, and Language Processing, 30:495–507, 2021.
  42. 42.Zhang, R., Isola, P., Efros, A. A., Shechtman, E., and Wang, O. The unreasonable effectiveness of deep features as a perceptual metric. In CVPR, 2018.

Citation

MLA
Huh, M., et al. “Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks”. International Conference on Machine Learning, vol. 202, 2023, pp. 14096–113, https://proceedings.mlr.press/v202/huh23a.html.
APA
Huh, M., Cheung, B., Agrawal, P., & Isola, P. (2023). Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks. International Conference on Machine Learning, 202, 14096–14113. https://proceedings.mlr.press/v202/huh23a.html
Chicago
Huh, M., B. Cheung, P. Agrawal, and P. Isola. 2023. “Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks”. International Conference on Machine Learning 202: 14096–113. https://proceedings.mlr.press/v202/huh23a.html.
Harvard
Huh, M. et al. (2023) “Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks”, International Conference on Machine Learning. PMLR, pp. 14096–14113. Available at: https://proceedings.mlr.press/v202/huh23a.html.
Vancouver
1. Huh M, Cheung B, Agrawal P, Isola P (2023) Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks. In: International Conference on Machine Learning. PMLR, pp 14096–14113

BibTeX

@InProceedings{pmlr-v202-huh23a,
  title = 	 {Straightening Out the Straight-Through Estimator: Overcoming Optimization Challenges in Vector Quantized Networks},
  author =       {Huh, Minyoung and Cheung, Brian and Agrawal, Pulkit and Isola, Phillip},
  booktitle = 	 {Proceedings of the 40th International Conference on Machine Learning},
  pages = 	 {14096--14113},
  year = 	 {2023},
  editor = 	 {Krause, Andreas and Brunskill, Emma and Cho, Kyunghyun and Engelhardt, Barbara and Sabato, Sivan and Scarlett, Jonathan},
  volume = 	 {202},
  series = 	 {Proceedings of Machine Learning Research},
  month = 	 {23--29 Jul},
  publisher =    {PMLR},
  pdf = 	 {https://proceedings.mlr.press/v202/huh23a/huh23a.pdf},
  url = 	 {https://proceedings.mlr.press/v202/huh23a.html},
  abstract = 	 {This work examines the challenges of training neural networks using vector quantization using straight-through estimation. We find that the main cause of training instability is the discrepancy between the model embedding and the code-vector distribution. We identify the factors that contribute to this issue, including the codebook gradient sparsity and the asymmetric nature of the commitment loss, which leads to misaligned code-vector assignments. We propose to address this issue via affine re-parameterization of the code vectors. Additionally, we introduce an alternating optimization to reduce the gradient error introduced by the straight-through estimation. Moreover, we propose an improvement to the commitment loss to ensure better alignment between the codebook representation and the model embedding. These optimization methods improve the mathematical approximation of the straight-through estimation and, ultimately, the model performance. We demonstrate the effectiveness of our methods on several common model architectures, such as AlexNet, ResNet, and ViT, across various tasks, including image classification and generative modeling.}
}
Metadata:DOI registry

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/