Summary (Overview)

  • σTransfer is a method that enables uncertainty transfer from small to large neural networks under the Maximal Update Parametrization (μP), allowing prior precision selection and posterior-based decisions to be performed on a cheap proxy network and transferred zero-shot to a much larger target network.
  • The key insight is a width-normalized prior covariance rescaling that makes the prior kernel, posterior covariance, and selected prior precision stable as model width grows, with theoretical guarantees under explicit conditions.
  • Prior-precision transfer: The optimal Laplace prior precision λ is stable across width, so a precision chosen on a small network is near-optimal for a large one, eliminating the need for expensive target-side precision sweeps.
  • Decision transfer: Uncertainty-based decisions (active learning acquisitions, OOD detection, abstention) made from the small network's posterior are approximately those the large network would make, without constructing a target posterior at all.
  • Empirically, σTransfer achieves speedups up to ~5000× on MNIST (width 128→4096) with negligible NLL degradation (0.002), and a median ~2.3× speedup on public 1B→7B language models with mean target-NLL increase below 10−410^{-4} across ten tasks.

Introduction and Theoretical Foundation

Background and Motivation

The Laplace approximation for neural networks provides predictive uncertainty estimates and supports evidence-based model selection. However, its reliability depends critically on the prior precision λ, a hyperparameter controlling the scale of posterior uncertainty. Selecting λ requires repeated posterior evaluations across candidate precisions—a sweep that becomes prohibitively expensive as network width grows, especially for models with billions of parameters.

The Transfer Problem

Under standard parametrization (SP), the same precision can produce different uncertainty at different widths, making transfer unreliable. The Maximal Update Parametrization (μP) (Yang & Hu, 2021; Yang et al., 2021) rescales each layer's initialization, forward multiplier, and learning rate with width, stabilizing training dynamics and enabling hyperparameter transfer. However, applying Laplace approximation to a μP network is insufficient: the optimal prior precision at one width is still not optimal at another.

Theoretical Foundation

The paper builds on the Tensor Programs framework (Yang, 2019; 2020a; Yang & Hu, 2021; Yang & Littwin, 2023), which establishes large-width limits for neural networks, and on the linearized Laplace approximation (Immer et al., 2021b). The central theoretical progression is:

Kn→K∞⏟ prior - kernel stability   ⟹  mnλ→m∞λ,Cnλ→C∞λ⏟ posterior mean and covariance stability   ⟹  λn⋆→λ⋆⏟ prior - precision transfer , same decision ⏟ decision transfer .\underbrace {K _ {n} \to K _ {\infty}} _ {\text { prior - kernel stability }} \implies \underbrace {m _ {n} ^ {\lambda} \to m _ {\infty} ^ {\lambda} , C _ {n} ^ {\lambda} \to C _ {\infty} ^ {\lambda}} _ {\text { posterior mean and covariance stability }} \implies \underbrace {\lambda_ {n} ^ {\star} \to \lambda^ {\star}} _ {\text { prior - precision transfer }}, \quad \underbrace {\text { same decision }} _ {\text { decision transfer }}.

Methodology

Linearized Laplace Approximation

Under the standard Gaussian prior θ∣λ∼N(0,λ−1I)\theta \mid \lambda \sim \mathcal{N}(0, \lambda^{-1}I), the Laplace approximation gives:

θ∣D,λ∼N(θ^n,(Hn+λI)−1),Hn=Jn(D)⊤ΩnJn(D)\theta \mid D, \lambda \sim \mathcal {N} \left(\hat {\theta} _ {n}, (H _ {n} + \lambda I) ^ {- 1}\right), \quad H _ {n} = J _ {n} (D) ^ {\top} \Omega_ {n} J _ {n} (D)

where HnH_n is the generalized Gauss–Newton curvature, JnJ_n is the Jacobian of network outputs with respect to parameters, and Ωn\Omega_n is the output-space NLL curvature. The output posterior has mean and covariance:

mnλ(X)=fn(X;θ^n),Cnλ(X,X′)=Jn(X)(Hn+λI)−1Jn(X′)⊤m _ {n} ^ {\lambda} (X) = f _ {n} (X; \hat {\theta} _ {n}), \qquad C _ {n} ^ {\lambda} (X, X ^ {\prime}) = J _ {n} (X) (H _ {n} + \lambda I) ^ {- 1} J _ {n} (X ^ {\prime}) ^ {\top}

The σTransfer Prior Kernel

The key methodological contribution is constructing a width-normalized prior. Using rescaled parameter coordinates ϕn=Tn−1θn\phi_n = T_n^{-1}\theta_n where Tn=blockdiag(t1(n)I,t2(n)I,… )T_n = \text{blockdiag}(t_1(n)I, t_2(n)I, \dots) applies width-dependent scale factors to each parameter tensor, the prior kernel becomes:

Kn(X,X′)=Jϕ,n(X)Jϕ,n(X′)⊤=Jn(X)TnTn⊤Jn(X′)⊤=Jn(X)SnJn(X′)⊤K _ {n} (X, X ^ {\prime}) = J _ {\phi , n} (X) J _ {\phi , n} (X ^ {\prime}) ^ {\top} = J _ {n} (X) T _ {n} T _ {n} ^ {\top} J _ {n} (X ^ {\prime}) ^ {\top} = J _ {n} (X) S _ {n} J _ {n} (X ^ {\prime}) ^ {\top}

where Sn=TnTn⊤S_n = T_n T_n^\top. The posterior covariance in stored parameter coordinates becomes:

Cnλ(X,X′)=Jn(X)(Hn+λSn−1)−1Jn(X′)⊤C _ {n} ^ {\lambda} (X, X ^ {\prime}) = J _ {n} (X) \big (H _ {n} + \lambda S _ {n} ^ {- 1} \big) ^ {- 1} J _ {n} (X ^ {\prime}) ^ {\top}

Key practical insight: Within a parameter tensor with width factor t, the prior precision is the scalar λ/t2\lambda/t^2—no dense matrix is ever constructed. σTransfer adds no new hyperparameters.

Woodbury Identity Reformulation

Applying the Woodbury identity for a general Gaussian prior with covariance λ−1Sn\lambda^{-1}S_n:

Cnλ(X,X′)=Kn(X,X′)λ−Kn(X,D)λ2Ωn1/2(I+λ−1Ωn1/2Kn(D,D)Ωn1/2)−1Ωn1/2Kn(D,X′)C _ {n} ^ {\lambda} (X, X ^ {\prime}) = \frac {K _ {n} (X , X ^ {\prime})}{\lambda} - \frac {K _ {n} (X , D)}{\lambda^ {2}} \Omega_ {n} ^ {1 / 2} \bigl (I + \lambda^ {- 1} \Omega_ {n} ^ {1 / 2} K _ {n} (D, D) \Omega_ {n} ^ {1 / 2} \bigr) ^ {- 1} \Omega_ {n} ^ {1 / 2} K _ {n} (D, X ^ {\prime})

This shows the prior geometry SnS_n enters the posterior covariance only through the prior kernel KnK_n, so stabilizing KnK_n suffices for posterior stability.

Empirical Validation / Results

Regression (FreeSolv, ESOL, Lipophilicity)

Table 1: Prior-precision transfer on three regression datasets (proxy width 512, target width 4096, full-covariance last-layer Laplace, five paired seeds, mean ± s.d.). |ΔNLL| scaled by 10310^3; lower is better.

MethodESOL |Δλ|ESOL |ΔNLL|FreeSolv |Δλ|FreeSolv |ΔNLL|Lipophilicity |Δλ|Lipophilicity |ΔNLL|
SP + isotropic1.5±0.41.5 \pm 0.441.8±30.741.8 \pm 30.75.0±2.25.0 \pm 2.2452.8±429.5452.8 \pm 429.54.0±2.74.0 \pm 2.7140.2±119.7140.2 \pm 119.7
SP + metric1.5±0.51.5 \pm 0.58.40±6.598.40 \pm 6.591.2±1.41.2 \pm 1.49.65±10.959.65 \pm 10.952.1±1.22.1 \pm 1.23.39±1.933.39 \pm 1.93
μP + isotropic3.0±1.43.0 \pm 1.40.947±0.5660.947 \pm 0.5663.1±0.53.1 \pm 0.51.564±0.8251.564 \pm 0.8253.2±1.03.2 \pm 1.00.094±0.0240.094 \pm 0.024
σTransfer (ours)0.9±0.80.9 \pm 0.80.548±0.4120.548 \pm 0.4120.5±0.40.5 \pm 0.40.461±0.4300.461 \pm 0.4300.4±0.20.4 \pm 0.20.033±0.0190.033 \pm 0.019

σTransfer achieves the smallest selection gap (0.4–0.9 bits vs. 1.5–5.0 for SP + isotropic) and predictive transfer gaps orders of magnitude smaller.

Image Classification (MNIST, Fashion-MNIST)

Table 2: Prior-precision transfer across three last-layer Laplace approximations (proxy width 128, target width 4096, 10 paired seeds). |ΔNLL| scaled by 10210^2.

MethodMNIST LL-full |Δλ|MNIST LL-full |ΔNLL|MNIST LL-diag |ΔNLL|MNIST LL-KFAC |ΔNLL|FMNIST LL-full |ΔNLL|FMNIST LL-diag |ΔNLL|FMNIST LL-KFAC |ΔNLL|
SP + isotropic1.751.759.89.80.80.86.96.99.09.00.90.913.213.2
SP + metric3.253.259.69.614.814.810.710.739.039.067.267.233.333.3
μP + isotropic4.754.7515.715.720.620.616.516.534.934.947.247.230.530.5
σTransfer (ours)0.250.250.20.20.10.10.20.20.40.40.60.60.30.3

σTransfer has the smallest predictive transfer gap in every column, with worst case 0.6×10−20.6 \times 10^{-2} (FMNIST LL-diagonal) still beating every other method. For OOD detection, σTransfer achieves 0.950 AUROC at width 4096 (matching target-side tuning), while SP loses 0.072 and μP with isotropic prior loses 0.327.

Language Models

  • Own ladder (widths 128–2048): Under σTransfer, the evidence-selected precision converges with width; under SP it drifts ~0.65 bits per width doubling without saturating. The proxy's choice lands 0.65±0.340.65 \pm 0.34 bits from the target's own, degrading target test NLL by only 0.0014 on average (vs. 2.65±0.632.65 \pm 0.63 bits and 0.013 degradation under SP).
  • Public 1B→7B (Blake et al., 2025): Seven of ten tasks select the same precision as the target; eight have absolute target-NLL changes below 10−410^{-4}.
  • Active learning on AG News: σTransfer has the highest mean Spearman correlation between proxy and target rankings across all proxy widths and initial label budgets; target accuracy is lower by only 0.0014 with proxy-selected labels.

Computational Savings

Table 3: Computational savings from σTransfer

TaskDatasetTransferApprox.SpeedupDegradation
λ sweepMNIST128→4096LL full~5000×0.002 NLL
OOD-aware λMNIST128→4096all + diag94×0.001 AUROC
evidence searchLLM (ours)128→2048LL full~7×0.0014 NLL
evidence searchLLM (public)1B→7BLL full~2.3×0 NLL
EPIG acquisitionAG News128→2048LL full~3×0.0014 ACC
evidence searchPenDigits32→128all + full~22×-0.007 NLL
evidence searchletter32→128all + full~32×0.078 NLL

Theoretical and Practical Implications

Theoretical Guarantees

The paper establishes four main theorems:

  1. Theorem 1 (Prior-kernel stability): For scalar-output, fixed-depth, equal-width ReLU μP MLPs trained for a fixed finite number of width-matched μP steps, the σTransfer prior kernel converges jointly almost surely as n→∞n \to \infty.

  2. Theorem 2 (Posterior-covariance stability): Given prior-kernel convergence and width-stable likelihood curvature, the linearized-Laplace output covariance converges uniformly over any interval [λ−,λ+][\lambda_-, \lambda_+] with 0<λ−≤λ+<∞0 < \lambda_- \leq \lambda_+ < \infty.

  3. Theorem 3 (Prior-precision transfer): If posterior mean and covariance converge uniformly in λ, validation-NLL curves converge uniformly. The excess NLL from using the width-n proxy's precision on a target of width N satisfies:

NLL⁡N(λn⋆)−min⁡λ∈ΛNLL⁡N(λ)≤2sup⁡λ∈Λ∣NLL⁡n(λ)−NLL⁡N(λ)∣\operatorname{NLL} _ {N} \left(\lambda_ {n} ^ {\star}\right) - \min _ {\lambda \in \Lambda} \operatorname{NLL} _ {N} (\lambda) \leq 2 \sup _ {\lambda \in \Lambda} | \operatorname{NLL} _ {n} (\lambda) - \operatorname{NLL} _ {N} (\lambda) |
  1. Theorem 4 (Decision transfer): If action scores depend continuously on posterior moments, they converge uniformly. With a positive gap for the best action, two sufficiently wide networks choose the same action.

Practical Implications

  • Cost reduction: σTransfer moves the expensive precision search to a small proxy network, reducing uncertainty-sensitive task costs by up to several orders of magnitude.
  • No new hyperparameters: The method adds no hyperparameters beyond the standard Laplace prior precision λ.
  • Broad applicability: Results extend beyond MLPs to any width-indexed architecture where the width-normalized prior kernel, predictive mean, and likelihood curvature converge (e.g., last-layer Laplace reduces to convergence of the normalized final-feature Gram matrix).
  • Decision transfer applications: Enables small networks to make acquisition, OOD-detection, and abstention decisions for much larger networks, with the large network needing no Laplace posterior at all.

Conclusion

σTransfer enables a small (cheap) network to make uncertainty estimation and uncertainty-based decision-making cheaper for a much larger (expensive) network through two forms of transfer:

  1. Prior-precision transfer: The precision search is performed on the small network and transferred zero-shot to the larger network, eliminating target-side search.
  2. Decision transfer: The small network's uncertainty estimates directly select data, flag OOD inputs, or decide when the larger network should abstain.

Both forms are justified theoretically through a series of results showing accurate transfer across network widths, and empirically σTransfer preserves predictive performance while reducing costs by up to several orders of magnitude. Key limitations acknowledged include: σTransfer does not always win (e.g., FMNIST under last-layer diagonal curvature where SP + isotropic selects marginally closer precision), and the theoretical guarantees require specific conditions (finite-step training, fixed-depth architectures). Future directions include extending to more general architectures and training regimes, and exploring additional downstream applications of decision transfer.

Related papers