Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models

Hong LiuSang Michael XieZhiyuan LiTengyu Ma

article2023ICML98 citations

Reveals that language models with identical pre-training losses can achieve substantially different downstream performance, proving both theoretically and empirically that SGD's implicit bias toward flatter loss minima governs transferability.

Listen

In modern artificial intelligence, large language models are typically evaluated and monitored using pre-training loss, which measures how well a model predicts masked or next words on training data. Because pre-training loss generally correlates with how well a model performs on practical downstream applications, conventional practice assumes that achieving a lower or equal pre-training loss guarantees equivalent or better downstream utility. However, this assumption fails to explain why different models with identical pre-training losses often display markedly different downstream accuracy.

The article investigates why models with the same pre-training loss achieve different downstream performance and demonstrates that the implicit bias of training algorithms—specifically their tendency to favor flatter loss minima—serves as a primary driver of downstream transferability.

To evaluate this dynamic, the authors conducted controlled empirical experiments across synthetic languages (probabilistic context-free grammars and hidden Markov models) and real-world corpora (OpenWebText and BookCorpus) using transformer models ranging from 2 million to 950 million parameters. In these environments, the authors examined models operating in a saturation regime where pre-training loss matches the optimal theoretical bound. They tested the effects of continuing training past convergence, altering model scale, and varying optimization techniques, while theoretically analyzing stochastic gradient descent and mathematically proving feature learning behaviors in a synthetic Dyck language setting.

The findings reveal three primary insights. First, pre-training loss alone does not dictate downstream success: models with identical, optimal pre-training losses showed substantial differences in downstream accuracy depending on training duration, model size, and optimization techniques. Second, standard optimization algorithms inherently possess an implicit bias toward flatter minima (measured by the trace of the Hessian), and this flatness strongly correlates with superior downstream accuracy where pre-training loss ceases to be predictive. For example, continuing training after loss convergence, scaling model capacity, or adding explicit flatness regularization improved downstream accuracy by several percentage points while loss remained flat. Third, theoretical analyses confirmed that standard stochastic gradient descent naturally oscillates along optimal loss manifolds toward flatter configurations, and in formal language settings, only the flattest models learn generalizable structural representations rather than memorizing data.

These results demonstrate that downstream performance depends heavily on the geometry of the learned parameter space rather than empirical pre-training loss alone. Relying exclusively on pre-training loss creates operational risks, such as prematurely halting training or selecting models that fit pre-training data via brittle memorization. Consequently, explicit regularization methods that promote flatness, such as weight decay, dropout, or sharpness-aware minimization, provide direct transferability benefits across downstream tasks.

For technical leaders and practitioners, the article suggests incorporating flatness metrics into model evaluation pipelines rather than relying solely on validation loss. When training budgets permit, organizations should continue pre-training past initial loss convergence or apply explicit flatness regularization to maximize transferability. Future work should focus on extending these theoretical guarantees to adaptive optimizers like Adam and developing computationally lightweight flatness regularizers suitable for large-scale production training.

The findings are bounded by certain theoretical and empirical limitations. While the theoretical proofs strictly address stochastic gradient descent and simplified synthetic grammars, the practical training of large language models frequently relies on adaptive optimizers such as AdamW on complex web-scale text. Nonetheless, the high consistency between empirical results across real-world datasets and theoretical models provides strong confidence in the core conclusion that flatness-driven implicit bias is a fundamental factor in model transferability.

Liu et al (2023).pdf
  • Paper: Visualizing the Loss Landscape of Neural Nets, Hao Li et al. (2017). Its method for visualizing neural loss landscapes and relating curvature to generalization provides essential grounding for the source’s analysis of flat minima and Hessian trace.
  • Paper: Exploring Generalization in Deep Learning, Behnam Neyshabur et al. (2017). It develops the broader generalization debate around sharpness and its limitations, clarifying the theoretical context for the source’s flatness-based account.
  • Paper: Why Does Unsupervised Pre-training Help Deep Learning?, Dumitru Erhan et al. (2010). Its account of pre-training as a regularizer establishes the earlier explanation of why pre-training can improve generalization that the source revisits through optimization bias.
Cover for Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models

Table of Contents

  • 1. Introduction
  • 2. Implicit Bias Affects Downstream Accuracy
  • 2.1. Formulations
  • 2.2. Experimental Setup
  • 2.3. Role of Implicit Bias on Downstream Accuracy
  • 3. SGD Prefers Flatter Minima in Language Modeling
  • 4. The Correlation Between Flatness and Downstream Performance
  • 5. Flatness Regularization Identifies Transferable Models on Synthetic Language
  • 6. Related Work
  • 7. Conclusion
  • Acknowledgement
  • References
  • A. Details in Section 2
  • A.1. Generating Simplified Datasets
  • A.2. Real-world Datasets
  • A.3. Compute the True Conditional Probabilities from Generative Models
  • A.4. Models
  • A.5. Algorithms
  • A.6. Results on Other Downstream Tasks
  • A.7. Evaluation of Pre-training Loss with KL Divergence
  • B. Details in Section 4
  • B.1. Unbiased Estimate of the Trace of Hessian
  • B.2. Details in Figure 4
  • B.3. Details in Figure 5
  • B.4. Embedding of a Smaller Transformer into a Larger Transformer
  • B.4.1. The Base Case with MLPs
  • B.4.2. Real Transformers
  • B.4.3. Viewing a Small Transformer as a Special Case of a Large Transformer
  • C. Limitations
  • D. Practical Implications
  • E. Additional Related Work
  • F. Omitted Proofs in Section 3
  • G. Omitted Proofs in Section 5
  • G.1. Omitted Proofs of Theorem 5.1
  • G.2. The Existence of Random Feature Solutions

Knowls

  1. Knowl 1 — The saturation regime fixes pre-training loss but not representations

    definition

    For masked language modeling (MLM), let xx be a sentence, let tt be a uniformly sampled masked position, and let x−tx_{-t} denote the sentence with position tt replaced by the mask token. The model with parameters θ\theta outputs a probability distribution fθ(x−t)f_\theta(x_{-t}) over the token vocabulary, and xtx_t is the true token. A model is in the saturation regime when its output equals the true conditional distribution for every masked context, so its expected cross-entropy reaches the conditional entropy lower bound:

    L(θ)=Ex,t[−log⁡fθ(x−t)xt]=Ex,t[−log⁡Pr⁡(xt∣x−t)].L(\theta)=\mathbb E_{x,t}\left[-\log f_\theta(x_{-t})_{x_t}\right] =\mathbb E_{x,t}\left[-\log \Pr(x_t\mid x_{-t})\right].

    Here L(θ)L(\theta) is the population MLM loss and fθ(x−t)xtf_\theta(x_{-t})_{x_t} is the predicted probability assigned to the true token. Multiple parameter configurations can compute the same true conditional probabilities and therefore have identical minimal pre-training loss, while still having different contextual representations.

  2. Knowl 2 — Equal pre-training loss can conceal large differences in transfer

    empirical result

    Across controlled synthetic-language experiments and real-corpus experiments, downstream performance differed among models with nearly identical pre-training loss. The differences arose in three comparisons: checkpoints from the same training run after the loss had converged, models of different sizes, and models trained with different pre-training procedures or regularization. In the first two comparisons, models could share an architecture while differing in their training history; in the size comparison, the authors also constructed larger-architecture views of smaller models that preserve their functionality. A model that outputs the exact true MLM conditional probabilities (the paper's lookup-table comparison) also transferred worse than ordinarily trained models. These results show that matching the pre-training loss does not determine downstream performance. In particular, even equal and correct token probabilities do not force contextual representations to have equal linear-probe utility.

  3. Knowl 3 — Longer training and larger models improve downstream accuracy at matched loss

    empirical result

    The reported comparisons show downstream gains even when pre-training loss is already near its minimum. For models with matched pre-training loss, increasing model size improved linear-probe accuracy by 6.9% on PCFG-generated data, 4.5% on HMM-generated data, and 2.0% on OPT-generated data. On PCFG and HMM checkpoints, continuing training after pre-training loss convergence improved downstream accuracy by 1.6% and 4.0%, respectively. On OpenWebText, continuing a 25M-parameter model from 400K to 1,400K steps increased SST-2 accuracy by 1.25% while the trace of the pre-training-loss Hessian decreased by 0.67. For the PCFG size comparison, among models larger than 9M parameters the loss was nearly unchanged, while the trace of the Hessian decreased from 19.8 to 12.6 and linear-probe accuracy on Task B rose from 40.4% to 50.5%. To compare flatness across sizes, the authors embedded smaller transformers into larger architectures by replicating weights and adding residual blocks with zeroed parameters; the construction preserves the represented function.

  4. Knowl 4 — Pre-training algorithm and regularization change transfer at similar loss

    data/table

    The table compares downstream accuracy for models with similar validation pre-training loss, including the trace of the pre-training-loss Hessian where reported. The PCFG adversarial procedure was designed to impair downstream adaptation while retaining MLM performance; it produced worse accuracy and a larger Hessian trace than AdamW. On BookCorpus, removing dropout and weight decay reduced downstream scores and increased the trace. Adding sharpness-aware minimization (SAM) to the unregularized setup reduced the trace and recovered much of the accuracy decrease, despite a slightly higher pre-training loss. Accuracies are percentages; Hessian traces are reported as in the paper.

    Setting Method Pre-training loss Downstream accuracy Hessian trace
    PCFG AdamW 3.204 Task A: 89.9 ±\pm 0.3; Task B: 49.2 ±\pm 0.8; Task C: 55.7 ±\pm 0.6 8.01 ±\pm 0.73
    PCFG Adversarial 3.206 Task A: 83.1 ±\pm 0.6; Task B: 42.3 ±\pm 1.5; Task C: 50.2 ±\pm 1.0 19.34 ±\pm 0.92
    PCFG Lookup table 3.196 Task A: 71.2; Task B: 39.7 –
    BookCorpus AdamW, dropout and weight decay 1.85 QNLI: 84.4; SST-2: 90.8; RTE: 62.9 24.55 ±\pm 0.18
    BookCorpus No dropout or weight decay 1.85 QNLI: 83.0; SST-2: 89.5; RTE: 60.2 34.91 ±\pm 0.40
    BookCorpus No dropout or weight decay, plus SAM 1.89 QNLI: 84.4; SST-2: 90.1; RTE: 62.5 22.15 ±\pm 0.26
  5. Knowl 5 — Mini-batch SGD follows a projected descent on Hessian trace after reaching a minimum

    theoretical result

    Consider population MLM cross-entropy L(θ)L(\theta), with fresh sentence and masked-position samples at each SGD step. Assume LL is C4C^4-smooth, its global minimum is the entropy of the true conditional distribution, and its set of global minimizers Γ⊆Rd\Gamma\subseteq\mathbb R^d is a smooth C2C^2 manifold of dimension d−Md-M, with rank⁡(∇2L(θ))=M\operatorname{rank}(\nabla^2L(\theta))=M at every θ∈Γ\theta\in\Gamma. Initialize SGD at a global minimizer θ∈Γ\theta\in\Gamma. For batch size one, the limiting motion along the minimizer manifold is the solution of

    dθ^(t)=−14∇ΓTr⁡[∇2L(θ^(t))]dt,θ^(0)=θ.d\hat\theta(t)=-\frac14\nabla_\Gamma\operatorname{Tr}\left[\nabla^2L(\hat\theta(t))\right]dt, \qquad \hat\theta(0)=\theta.

    Here ∇Γ\nabla_\Gamma is the Euclidean gradient projected onto the tangent space of Γ\Gamma, and tt is the continuous time in the limiting equation. In particular, for any K>0K>0 for which this equation has a solution, the SGD iterate after K/η2K/\eta^2 steps converges in distribution to θ^(K)\hat\theta(K) as the learning rate η→0\eta\to0. Thus, in this limit, SGD locally moves among equal-loss global minimizers toward smaller Hessian trace. For batch size BB, the same result holds with coefficient 1/(4B)1/(4B) instead of 1/41/4.

  6. Knowl 6 — At an exact MLM solution, stochastic-gradient covariance equals the loss Hessian

    theoretical result

    At any parameter value θ\theta in the saturation regime, define the per-example score sθ(x,t)=∇θlog⁡fθ(x−t)xts_\theta(x,t)=\nabla_\theta\log f_\theta(x_{-t})_{x_t}, where xx is a sampled sentence and tt its masked position. The covariance of the stochastic gradient of the MLM cross-entropy equals the Hessian of the population pre-training loss:

    Cov⁡x,t(sθ(x,t))=∇2L(θ).\operatorname{Cov}_{x,t}\left(s_\theta(x,t)\right)=\nabla^2L(\theta).

    The expectation of the score is zero at a global minimizer, so the covariance is also Ex,t[sθ(x,t)sθ(x,t)⊤]\mathbb E_{x,t}[s_\theta(x,t)s_\theta(x,t)^\top]. The identity applies to the paper's well-specified conditional-probability model at a global minimizer. It links the persistent mini-batch noise at an MLM minimum to the curvature that appears in the SGD flatness result.

  7. Knowl 7 — Hessian trace tracks downstream performance after pre-training loss converges

    empirical result

    The paper measures flatness by Tr⁡(∇2L(θ))\operatorname{Tr}(\nabla^2L(\theta)), with a smaller trace indicating a flatter minimum. For PCFG- and HMM-generated data, SGD checkpoints showed a clear decrease in this trace after validation pre-training loss converged, while downstream linear-probe accuracy rose by the reported 1.6% on PCFG Task C and 4.0% on HMM Task-10. On OpenWebText, the 25M-parameter model's continued training from 400K to 1,400K steps coincided with a 0.67 reduction in trace and a 1.25% increase in SST-2 accuracy. In comparisons across model sizes, the PCFG models larger than 9M parameters had nearly equal pre-training loss, while the trace fell from 19.8 to 12.6 as Task B linear-probe accuracy increased from 40.4% to 50.5%. These observations support the paper's claim that trace of Hessian is more informative about downstream differences than pre-training loss when the loss is already near its minimum.

  8. Knowl 8 — A Dyck-language result links the flattest exact MLM solution to perfect transfer

    theoretical result

    Consider sequences of even length TT over opening and closing brackets. The pre-training distribution is uniform over valid balanced-bracket strings, with one uniformly selected position masked. Encode a token at position jj as eje_j or −ej-e_j in R2T\mathbb R^{2T} according to bracket type, and encode the masked position using a random sign on coordinate j+Tj+T. A single-layer attention mechanism forms h(x)=∑j=1Tajxjh(x)=\sum_{j=1}^T a_jx_j, where the learned attention weights aja_j are nonnegative and sum to one. An mm-unit ReLU layer and output vector compute f(x)=m−1u⊤σ(Vh(x))f(x)=m^{-1}u^\top\sigma(Vh(x)), with V∈Rm×2TV\in\mathbb R^{m\times 2T}, u∈Rmu\in\mathbb R^m, and σ(z)=max⁡(z,0)\sigma(z)=\max(z,0) applied coordinatewise. The MLM objective is squared loss for predicting the masked-token target.

    For downstream evaluation, draw each of the TT bracket tokens independently and uniformly, and use the target g∗(x)g^*(x) equal to the number of closing brackets minus the number of opening brackets. Given nn such downstream examples, choose the flattest exact-MLM solution by minimizing Tr⁡(∇ψ2L)+Tr⁡(∇u2L)\operatorname{Tr}(\nabla^2_\psi L)+\operatorname{Tr}(\nabla^2_u L) subject to zero pre-training loss, where ψ=(Q,K,V)\psi=(Q,K,V) contains the attention and hidden-layer parameters. Hold that representation fixed and choose the minimum-Euclidean-norm output vector u~\tilde u that fits the nn downstream examples exactly. For m≥2m\ge2 and T≥6T\ge6, the resulting model has zero population downstream squared loss with probability at least 1−2−n1-2^{-n} over the downstream sample. The result establishes perfect transfer under these conditions for the flattest exact pre-training solution, despite the existence of other exact MLM solutions that need not learn the transferable bracket-count feature.

  9. Knowl 9 — An unbiased Hessian-trace estimate uses conditional-token samples

    algorithm

    In the saturation regime, the trace of the pre-training-loss Hessian can be estimated without constructing the full parameter Jacobian. For a sampled masked context (x−t,t)(x_{-t},t), draw the true token xtx_t from the model's conditional distribution fθ(x−t)f_\theta(x_{-t}) and calculate the squared Euclidean norm of the score ∇θlog⁡fθ(x−t)xt\nabla_\theta\log f_\theta(x_{-t})_{x_t}. Averaging this quantity over sampled contexts and conditional-token draws is an unbiased estimate of the trace:

    Tr⁡(∇2L(θ))=Et,x−tExt∼fθ(x−t)[∥∇θlog⁡fθ(x−t)xt∥22].\operatorname{Tr}(\nabla^2L(\theta)) =\mathbb E_{t,x_{-t}}\mathbb E_{x_t\sim f_\theta(x_{-t})} \left[\left\|\nabla_\theta\log f_\theta(x_{-t})_{x_t}\right\|_2^2\right].

    The paper's experiments use 10,000 sampled masked contexts and 50 conditional-token draws for each context. The estimate is intended for models near the saturation regime, where the model conditional distribution approximates the true one.

  10. Knowl 10 — The theory does not yet cover Adam or general transfer settings

    limitation

    The implicit-bias theorem in the paper is for mini-batch SGD; it does not explain the implicit bias of Adam, which the authors note is commonly used for transformers. The formal result connecting flatness to transfer is also restricted to the simplified Dyck-language construction. The paper does not establish a general theorem that low Hessian trace improves downstream performance for realistic language models; it presents empirical correlations and identifies broader theoretical coverage as an open problem.

Coverage note — Detailed dataset-generation and training hyperparameters, the construction details for embedding smaller transformers, and the auxiliary Gaussian-random-feature existence bound are omitted because they support the reported comparisons or Dyck analysis rather than constitute separate top-level findings.

References

  1. 1.Amid, E. and Warmuth, M. K. Reparameterizing mirror descent as gradient descent. Advances in Neural Information Processing Systems, 33:8430–8439, 2020a.
  2. 2.Amid, E. and Warmuth, M. K. Winnowing with gradient descent. In Conference on Learning Theory, pp. 163–182. PMLR, 2020b.
  3. 3.Arora, S., Cohen, N., Hu, W., and Luo, Y. Implicit regularization in deep matrix factorization. In Advances in Neural Information Processing Systems, pp. 7411–7422, 2019a.
  4. 4.Arora, S., Khandeparkar, H., Khodak, M., Plevrakis, O., and Saunshi, N. A theoretical analysis of contrastive unsupervised representation learning. In International Conference on Machine Learning, 2019b.
  5. 5.Arora, S., Li, Z., and Panigrahi, A. Understanding gradient descent on the edge of stability in deep learning. In Chaudhuri, K., Jegelka, S., Song, L., Szepesvari, C., Niu, G., and Sabato, S. (eds.), Proceedings of the 39th International Conference on Machine Learning, volume 162 of Proceedings of Machine Learning Research, pp. 948–1024. PMLR, 17–23 Jul 2022. URL https://proceedings.mlr.press/v162/arora22a.html.
  6. 6.Azulay, S., Moroshko, E., Nacson, M. S., Woodworth, B. E., Srebro, N., Globerson, A., and Soudry, D. On the implicit bias of initialization shape: Beyond infinitesimal mirror descent. In International Conference on Machine Learning, pp. 468–477. PMLR, 2021.
  7. 7.Ba, J. L., Kiros, J. R., and Hinton, G. E. Layer normalization. arXiv preprint arXiv:1607.06450, 2016.
  8. 8.Bahri, D., Mobahi, H., and Tay, Y. Sharpness-aware minimization improves language model generalization. arXiv preprint arXiv:2110.08529, 2021.
  9. 9.Bai, Y. and Lee, J. D. Beyond linearization: On quadratic and higher-order approximation of wide neural networks. International Conference on Learning Representations (ICLR), 2020.
  10. 10.Blanc, G., Gupta, N., Valiant, G., and Valiant, P. Implicit regularization for deep neural networks driven by an ornstein-uhlenbeck like process. arXiv preprint arXiv:1904.09080, 2019.
  11. 11.Borkar, V. S. Stochastic approximation: a dynamical systems viewpoint, volume 48. Springer, 2009.
  12. 12.Brown, T., Mann, B., Ryder, N., Subbiah, M., Kaplan, J. D., Dhariwal, P., Neelakantan, A., Shyam, P., Sastry, G., Askell, A., et al. Language models are few-shot learners. Advances in neural information processing systems, 33:1877–1901, 2020.
  13. 13.Chomsky, N. Three models for the description of language. IRE Transactions on information theory, 2(3):113–124, 1956.
  14. 14.Choromanski, K., Likhosherstov, V., Dohan, D., Song, X., Gane, A., Sarlos, T., Hawkins, P., Davis, J., Mohiuddin, A., Kaiser, L., et al. Rethinking attention with performers. arXiv preprint arXiv:2009.14794, 2020.
  15. 15.Dai, Z., Lai, G., Yang, Y., and Le, Q. Funnel-transformer: Filtering out sequential redundancy for efficient language processing. Advances in neural information processing systems, 33:4271–4282, 2020.
  16. 16.Damian, A., Ma, T., and Lee, J. D. Label noise sgd provably prefers flat global minimizers. Advances in Neural Information Processing Systems, 34:27449–27461, 2021.
  17. 17.Devlin, J., Chang, M.-W., Lee, K., and Toutanova, K. Bert: Pre-training of deep bidirectional transformers for language understanding. arXiv preprint arXiv:1810.04805, 2018.
  18. 18.Dinh, L., Pascanu, R., Bengio, S., and Bengio, Y. Sharp minima can generalize for deep nets. In Proceedings of the 34th International Conference on Machine Learning-Volume 70, pp. 1019–1028. JMLR. org, 2017.
  19. 19.Dosovitskiy, A., Beyer, L., Kolesnikov, A., Weissenborn, D., Zhai, X., Unterthiner, T., Dehghani, M., Minderer, M., Heigold, G., Gelly, S., et al. An image is worth 16x16 words: Transformers for image recognition at scale. arXiv preprint arXiv:2010.11929, 2020.
  20. 20.Du, S. S., Hu, W., and Lee, J. D. Algorithmic regularization in learning deep homogeneous models: Layers are automatically balanced. In Advances in Neural Information Processing Systems, pp. 384–395, 2018.
  21. 21.Du, S. S., Hu, W., Kakade, S. M., Lee, J. D., and Lei, Q. Few-shot learning via learning the representation, provably. arXiv preprint arXiv:2002.09434, 2020.
  22. 22.Duchi, J. C. and Ruan, F. Stochastic methods for composite and weakly convex optimization problems. SIAM Journal on Optimization, 28(4):3229–3259, 2018.
  23. 23.Dziugaite, G. K. and Roy, D. M. Computing nonvacuous generalization bounds for deep (stochastic) neural networks with many more parameters than training data. arXiv preprint arXiv:1703.11008, 2017.
  24. 24.Fehrman, B., Gess, B., and Jentzen, A. Convergence rates for the stochastic gradient descent method for non-convex objective functions. Journal of Machine Learning Research, 21:136, 2020.
  25. 25.Foret, P., Kleiner, A., Mobahi, H., and Neyshabur, B. Sharpness-aware minimization for efficiently improving generalization. arXiv preprint arXiv:2010.01412, 2020.
  26. 26.Foret, P., Kleiner, A., Mobahi, H., and Neyshabur, B. Sharpness-aware minimization for efficiently improving generalization. In International Conference on Learning Representations, 2021.
  27. 27.Forney, G. D. The viterbi algorithm. Proceedings of the IEEE, 61(3):268–278, 1973.
  28. 28.Fradkin, P., Atanackovic, L., and Zhang, M. R. Robustness to adversarial gradients: A glimpse into the loss landscape of contrastive pre-training. In First Workshop on Pre-training: Perspectives, Pitfalls, and Paths Forward at ICML 2022, 2022. URL https://openreview.net/forum?id=-b3MEzI6N3.
  29. 29.Gokaslan, A., Cohen, V., Pavlick, E., and Tellex, S. Openwebtext corpus, 2019.
  30. 30.Gunasekar, S., Lee, J. D., Soudry, D., and Srebro, N. Implicit bias of gradient descent on linear convolutional networks. In Advances in Neural Information Processing Systems, pp. 9461–9471, 2018.
  31. 31.HaoChen, J. Z., Wei, C., Lee, J. D., and Ma, T. Shape matters: Understanding the implicit bias of the noise covariance. arXiv preprint arXiv:2006.08680, 2020.
  32. 32.HaoChen, J. Z., Wei, C., Gaidon, A., and Ma, T. Provable guarantees for self-supervised deep learning with spectral contrastive loss. arXiv preprint arXiv:2106.04156, 2021.
  33. 33.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.
  34. 34.Hernandez, D., Kaplan, J., Henighan, T., and McCandlish, S. Scaling laws for transfer. arXiv preprint arXiv:2102.01293, 2021.
  35. 35.Hewitt, J. and Manning, C. D. A structural probe for finding syntax in word representations. In Proceedings of the 2019 Conference of the North American Chapter of the Association for Computational Linguistics: Human Language Technologies, Volume 1 (Long and Short Papers), pp. 4129–4138, 2019.
  36. 36.Hochreiter, S. and Schmidhuber, J. Flat minima. Neural Computation, 9(1):1–42, 1997a.
  37. 37.Hochreiter, S. and Schmidhuber, J. Long short-term memory. Neural computation, 9(8):1735–1780, 1997b.
  38. 38.Htut, P. M., Phang, J., Bordia, S., and Bowman, S. R. Do attention heads in bert track syntactic dependencies? arXiv preprint arXiv:1911.12246, 2019.
  39. 39.Izsak, P., Berchansky, M., and Levy, O. How to train bert with an academic budget. arXiv preprint arXiv:2104.07705, 2021.
  40. 40.Jastrz˛ebski, S., Kenton, Z., Arpit, D., Ballas, N., Fischer, A., Bengio, Y., and Storkey, A. Three factors influencing minima in sgd. arXiv preprint arXiv:1711.04623, 2017.
  41. 41.Ji, Z. and Telgarsky, M. Gradient descent aligns the layers of deep linear networks. arXiv preprint arXiv:1810.02032, 2018.
  42. 42.Jiang, Y., Neyshabur, B., Mobahi, H., Krishnan, D., and Bengio, S. Fantastic generalization measures and where to find them. arXiv preprint arXiv:1912.02178, 2019.
  43. 43.Kaplan, J., McCandlish, S., Henighan, T., Brown, T. B., Chess, B., Child, R., Gray, S., Radford, A., Wu, J., and Amodei, D. Scaling laws for neural language models. arXiv preprint arXiv:2001.08361, 2020.
  44. 44.Keskar, N. S., Mudigere, D., Nocedal, J., Smelyanskiy, M., and Tang, P. T. P. On large-batch training for deep learning: Generalization gap and sharp minima. arXiv preprint arXiv:1609.04836, 2016.
  45. 45.Kim, Y., Dyer, C., and Rush, A. M. Compound probabilistic context-free grammars for grammar induction. arXiv preprint arXiv:1906.10225, 2019.
  46. 46.Kushner, H. and Yin, G. G. Stochastic approximation and recursive algorithms and applications, volume 35. Springer Science & Business Media, 2003.
  47. 47.Lan, Z., Chen, M., Goodman, S., Gimpel, K., Sharma, P., and Soricut, R. Albert: A lite bert for self-supervised learning of language representations. arXiv preprint arXiv:1909.11942, 2019.
  48. 48.Lari, K. and Young, S. J. The estimation of stochastic context-free grammars using the inside-outside algorithm. Computer speech & language, 4(1):35–56, 1990.
  49. 49.Lee, H. B., Lee, H., Na, D., Kim, S., Park, M., Yang, E., and Hwang, S. J. Learning to balance: Bayesian meta-learning for imbalanced and out-of-distribution tasks. In International Conference on Learning Representations, 2020.
  50. 50.Li, Q., Tai, C., and Weinan, E. Stochastic modified equations and adaptive stochastic gradient algorithms. In International Conference on Machine Learning, pp. 2101–2110. PMLR, 2017a.
  51. 51.Li, Q., Tai, C., and Weinan, E. Stochastic modified equations and dynamics of stochastic gradient algorithms i: Mathematical foundations. The Journal of Machine Learning Research, 20(1):1474–1520, 2019.
  52. 52.Li, Y., Ma, T., and Zhang, H. Algorithmic regularization in over-parameterized matrix sensing and neural networks with quadratic activations. arXiv preprint arXiv:1712.09203, pp. 2–47, 2017b.
  53. 53.Li, Z., Luo, Y., and Lyu, K. Towards resolving the implicit bias of gradient descent for matrix factorization: Greedy low-rank learning. arXiv preprint arXiv:2012.09839, 2020.
  54. 54.Li, Z., Wang, T., and Arora, S. What happens after sgd reaches zero loss?–a mathematical framework. arXiv preprint arXiv:2110.06914, 2021.
  55. 55.Li, Z., Wang, T., and Yu, D. Fast mixing of stochastic gradient descent with normalization and weight decay. In Oh, A. H., Agarwal, A., Belgrave, D., and Cho, K. (eds.), Advances in Neural Information Processing Systems, 2022. URL https://openreview.net/forum?id=sof8l4cki9.
  56. 56.Liu, H., HaoChen, J. Z., Wei, C., and Ma, T. Meta-learning transferable representations with a single target domain. arXiv preprint arXiv:2011.01418, 2020.
  57. 57.Liu, H., Dai, Z., So, D., and Le, Q. V. Pay attention to mlps. Advances in Neural Information Processing Systems, 34:9204–9215, 2021.
  58. 58.Lyu, K. and Li, J. Gradient descent maximizes the margin of homogeneous neural networks. arXiv preprint arXiv:1906.05890, 2019.
  59. 59.Lyu, K., Li, Z., Wang, R., and Arora, S. Gradient descent on two-layer nets: Margin maximization and simplicity bias. Advances in Neural Information Processing Systems, 34:12978–12991, 2021.
  60. 60.Lyu, K., Li, Z., and Arora, S. Understanding the generalization benefit of normalization layers: Sharpness reduction. arXiv preprint arXiv:2206.07085, 2022.
  61. 61.Mamou, J., Le, H., Del Rio, M., Stephenson, C., Tang, H., Kim, Y., and Chung, S. Emergence of separable manifolds in deep language representations. arXiv preprint arXiv:2006.01095, 2020.
  62. 62.Mandt, S., Hoffman, M. D., and Blei, D. M. Stochastic gradient descent as approximate bayesian inference. Journal of Machine Learning Research, 18:1–35, 2017.
  63. 63.Min, S., Lyu, X., Holtzman, A., Artetxe, M., Lewis, M., Hajishirzi, H., and Zettlemoyer, L. Rethinking the role of demonstrations: What makes in-context learning work? arXiv preprint arXiv:2202.12837, 2022.
  64. 64.Neyshabur, B., Bhojanapalli, S., McAllester, D., and Srebro, N. Exploring generalization in deep learning. In Advances in Neural Information Processing Systems, pp. 5947–5956, 2017.
  65. 65.Nivat, M. On some families of languages related to the dyck language. In Proceedings of the second annual ACM symposium on Theory of computing, pp. 221–225, 1970.
  66. 66.Norton, M. D. and Royset, J. O. Diametrical risk minimization: Theory and computations. Machine Learning, pp. 1–19, 2021.
  67. 67.Peters, M. E., Neumann, M., Zettlemoyer, L., and Yih, W.-t. Dissecting contextual word embeddings: Architecture and representation. arXiv preprint arXiv:1808.08949, 2018.
  68. 68.Radford, A., Wu, J., Child, R., Luan, D., Amodei, D., Sutskever, I., et al. Language models are unsupervised multitask learners. OpenAI blog, 1(8):9, 2019.
  69. 69.Raffel, C., Shazeer, N., Roberts, A., Lee, K., Narang, S., Matena, M., Zhou, Y., Li, W., and Liu, P. J. Exploring the limits of transfer learning with a unified text-to-text transformer. Journal of Machine Learning Research, 21:1–67, 2020.
  70. 70.Raghu, A., Lorraine, J., Kornblith, S., McDermott, M., and Duvenaud, D. K. Meta-learning to improve pre-training. Advances in Neural Information Processing Systems, 34:23231–23244, 2021.
  71. 71.Roark, B. and Bacchiani, M. Supervised and unsupervised pcfg adaptation to novel domains. In Proceedings of the 2003 Human Language Technology Conference of the North American Chapter of the Association for Computational Linguistics, pp. 205–212, 2003.
  72. 72.Saunshi, N., Malladi, S., and Arora, S. A mathematical exploration of why language models help solve downstream tasks. arXiv preprint arXiv:2010.03648, 2020.
  73. 73.Saunshi, N., Ash, J., Goel, S., Misra, D., Zhang, C., Arora, S., Kakade, S., and Krishnamurthy, A. Understanding contrastive learning requires incorporating inductive biases. arXiv preprint arXiv:2202.14037, 2022.
  74. 74.Soudry, D., Hoffer, E., Nacson, M. S., Gunasekar, S., and Srebro, N. The implicit bias of gradient descent on separable data. The Journal of Machine Learning Research, 19(1):2822–2878, 2018.
  75. 75.Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I., and Salakhutdinov, R. Dropout: a simple way to prevent neural networks from overfitting. The journal of machine learning research, 15(1):1929–1958, 2014.
  76. 76.Su, W., Boyd, S., and Candes, E. A differential equation for modeling nesterov’s accelerated gradient method: Theory and insights. In Advances in Neural Information Processing Systems, pp. 2510–2518, 2014.
  77. 77.Tay, Y., Dehghani, M., Rao, J., Fedus, W., Abnar, S., Chung, H. W., Narang, S., Yogatama, D., Vaswani, A., and Metzler, D. Scale efficiently: Insights from pre-training and fine-tuning transformers. arXiv preprint arXiv:2109.10686, 2021.
  78. 78.Tolstikhin, I. O., Houlsby, N., Kolesnikov, A., Beyer, L., Zhai, X., Unterthiner, T., Yung, J., Steiner, A., Keysers, D., Uszkoreit, J., et al. Mlp-mixer: An all-mlp architecture for vision. Advances in Neural Information Processing Systems, 34:24261–24272, 2021.
  79. 79.Vaskevicius, T., Kanade, V., and Rebeschini, P. Implicit regularization for optimal sparse recovery. In Advances in Neural Information Processing Systems, pp. 2968–2979, 2019.
  80. 80.Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, L., and Polosukhin, I. Attention is all you need. arXiv preprint arXiv:1706.03762, 2017.
  81. 81.Wang, A., Singh, A., Michael, J., Hill, F., Levy, O., and Bowman, S. R. Glue: A multi-task benchmark and analysis platform for natural language understanding. arXiv preprint arXiv:1804.07461, 2018.
  82. 82.Wang, S., Li, B. Z., Khabsa, M., Fang, H., and Ma, H. Linformer: Self-attention with linear complexity. arXiv preprint arXiv:2006.04768, 2020.
  83. 83.Wei, C. and Ma, T. Data-dependent sample complexity of deep neural networks via lipschitz augmentation. In Advances in Neural Information Processing Systems, pp. 9722–9733, 2019a.
  84. 84.Wei, C. and Ma, T. Improved sample complexities for deep networks and robust classification via an all-layer margin. arXiv preprint arXiv:1910.04284, 2019b.
  85. 85.Wei, C., Kakade, S., and Ma, T. The implicit and explicit regularization effects of dropout. arXiv preprint arXiv:2002.12915, 2020.
  86. 86.Wei, C., Chen, Y., and Ma, T. Statistically meaningful approximation: a case study on approximating turing machines with transformers, 2021a.
  87. 87.Wei, C., Xie, S. M., and Ma, T. Why do pretrained language models help in downstream tasks? an analysis of head and prompt tuning. arXiv preprint arXiv:2106.09226, 2021b.
  88. 88.Wei, J., Wang, X., Schuurmans, D., Bosma, M., Chi, E., Le, Q., and Zhou, D. Chain of thought prompting elicits reasoning in large language models. arXiv preprint arXiv:2201.11903, 2022.
  89. 89.Wen, K., Ma, T., and Li, Z. How does sharpness-aware minimization minimize sharpness? arXiv preprint arXiv:2211.05729, 2022.
  90. 90.Woodworth, B., Gunasekar, S., Lee, J. D., Moroshko, E., Savarese, P., Golan, I., Soudry, D., and Srebro, N. Kernel and rich regimes in overparametrized models. arXiv preprint arXiv:2002.09277, 2020.
  91. 91.Wu, D., Xia, S.-T., and Wang, Y. Adversarial weight perturbation helps robust generalization. Advances in Neural Information Processing Systems, 33:2958–2969, 2020.
  92. 92.Xie, S. M., Raghunathan, A., Liang, P., and Ma, T. An explanation of in-context learning as implicit bayesian inference. arXiv preprint arXiv:2111.02080, 2021.
  93. 93.Yang, Z., Dai, Z., Yang, Y., Carbonell, J., Salakhutdinov, R. R., and Le, Q. V. Xlnet: Generalized autoregressive pretraining for language understanding. Advances in neural information processing systems, 32, 2019.
  94. 94.Yun, C., Bhojanapalli, S., Rawat, A. S., Reddi, S. J., and Kumar, S. Are transformers universal approximators of sequence-to-sequence functions? arXiv preprint arXiv:1912.10077, 2019.
  95. 95.Yun, C., Krishnan, S., and Mobahi, H. A unifying view on implicit bias in training linear neural networks. arXiv preprint arXiv:2010.02501, 2020.
  96. 96.Zhang, J., Karimireddy, S. P., Veit, A., Kim, S., Reddi, S., Kumar, S., and Sra, S. Why are adaptive methods good for attention models? Advances in Neural Information Processing Systems, 33:15383–15393, 2020.
  97. 97.Zhang, S., Roller, S., Goyal, N., Artetxe, M., Chen, M., Chen, S., Dewan, C., Diab, M., Li, X., Lin, X. V., et al. Opt: Open pre-trained transformer language models. arXiv preprint arXiv:2205.01068, 2022a.
  98. 98.Zhang, T. and Hashimoto, T. On the inductive bias of masked language modeling: From statistical to syntactic dependencies. arXiv preprint arXiv:2104.05694, 2021.
  99. 99.Zhang, Y., Backurs, A., Bubeck, S., Eldan, R., Gunasekar, S., and Wagner, T. Unveiling transformers with lego: a synthetic reasoning task. arXiv preprint arXiv:2206.04301, 2022b.
  100. 100.Zhang, Z., Yang, J., Ji, X., and Du, S. S. Variance-aware confidence set: Variance-dependent bound for linear bandits and horizon-free bound for linear mixture mdp. arXiv preprint arXiv:2101.12745, 2021.
  101. 101.Zheng, Y., Zhang, R., and Mao, Y. Regularizing neural networks via adversarial model perturbation. In Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition, pp. 8156–8165, 2021.
  102. 102.Zhu, Y., Kiros, R., Zemel, R., Salakhutdinov, R., Urtasun, R., Torralba, A., and Fidler, S. Aligning books and movies: Towards story-like visual explanations by watching movies and reading books. In Proceedings of the IEEE international conference on computer vision, pp. 19–27, 2015.

Citation

MLA
Liu, H., et al. “Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models”. International Conference on Machine Learning, vol. 202, 2023, pp. 22188–214, https://proceedings.mlr.press/v202/liu23ao.html.
APA
Liu, H., Xie, S. M., Li, Z., & Ma, T. (2023). Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models. International Conference on Machine Learning, 202, 22188–22214. https://proceedings.mlr.press/v202/liu23ao.html
Chicago
Liu, H., S. M. Xie, Z. Li, and T. Ma. 2023. “Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models”. International Conference on Machine Learning 202: 22188–214. https://proceedings.mlr.press/v202/liu23ao.html.
Harvard
Liu, H. et al. (2023) “Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models”, International Conference on Machine Learning. PMLR, pp. 22188–22214. Available at: https://proceedings.mlr.press/v202/liu23ao.html.
Vancouver
1. Liu H, Xie SM, Li Z, Ma T (2023) Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models. In: International Conference on Machine Learning. PMLR, pp 22188–22214

BibTeX

@InProceedings{pmlr-v202-liu23ao,
  title = 	 {Same Pre-training Loss, Better Downstream: Implicit Bias Matters for Language Models},
  author =       {Liu, Hong and Xie, Sang Michael and Li, Zhiyuan and Ma, Tengyu},
  booktitle = 	 {Proceedings of the 40th International Conference on Machine Learning},
  pages = 	 {22188--22214},
  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/liu23ao/liu23ao.pdf},
  url = 	 {https://proceedings.mlr.press/v202/liu23ao.html},
  abstract = 	 {Language modeling on large-scale datasets improves performance of various downstream tasks. The validation pre-training loss is often used as the evaluation metric for language models since the pre-training loss tends to be well-correlated with downstream performance (which is itself hard to evaluate comprehensively). Contrary to the conventional wisdom, this paper shows that 1) pre-training loss cannot fully explain downstream performance and 2) flatness of the model is well-correlated with downstream performance where pre-training loss is not. We identify three ways to produce models with the same pre-training loss but different downstream performance: continue pre-training after convergence, increasing the model size, and changing the pre-training algorithms. These experiments demonstrate the existence of implicit bias of pre-training algorithms—among models with the same minimal pre-training loss, they implicitly prefer more transferable ones. Toward understanding this implicit bias, we prove that SGD with standard mini-batch noise implicitly prefers flatter minima of pre-training loss in language models, and empirically observe a strong correlation between flatness (measured by the trace of Hessian) and downstream performance among models with the same pre-training loss. We also prove in a synthetic language setting that among models with the minimal pre-training loss, the flattest model transfers to downstream tasks.}
}
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/