Fisher8: Stabilizing Neural Heteroscedastic Regression via Output-Layer Fisher Geometry

arXiv:2608.10374 · cs.LG · Submitted 2026-08-11 · Read on arXiv

Sumedh Vemuganti, Nickvash Kani

University of Illinois at Urbana-Champaign

cs.LG

Submitted: 2026-08-11

Updated: 2026-08-17

Code: https://github.com/FAIR-Universe/Cosmology_Challenge

License: http://creativecommons.org/licenses/by-nc-nd/4.0/

Importance score: 75/100

The gist: Fisher8: Stabilizing Neural Heteroscedastic Regression via Output-Layer Fisher Geometry Summary This paper addresses the instability of training neural networks to jointly predict mean and

Terminology

Summary

Fisher8: Stabilizing Neural Heteroscedastic Regression via Output-Layer Fisher Geometry

Summary

This paper addresses the instability of training neural networks to jointly predict mean and uncertainty estimates from noisy observations using the Gaussian negative log-likelihood (NLL) loss. The authors argue that prior stabilization efforts—such as gradient reweighting (β-NLL), stop-gradients with Newton steps (Faithful), and careful regularization tuning—are not fixing an unstable loss, but instead correcting a misdirected and mis-scaled traversal of the loss landscape. They propose Fisher8, an output-layer gradient correction that reorients and rescales updates using Fisher geometry rather than Euclidean geometry.

Core Problem and Motivation

The paper begins by noting that training a neural network to jointly predict mean and variance using the Gaussian NLL loss suffers from recurring instabilities: mean estimates degrade relative to unit-variance baselines, models inflate σ2 to reduce loss without improving fit, and training exhibits sharp phase transitions under varying regularization strengths. The authors state: If a diverse set of stabilizers is repeatedly required during training, perhaps these interventions are not fixing an unstable loss, but instead correcting a misdirected and mis-scaled traversal of the loss landscape.

Methodology

The paper derives Fisher8 as a natural-gradient correction. The key insight is that standard gradient descent implicitly uses a Euclidean constraint on parameter perturbations, forcing δθ = (δµ, δs) to lie on a Euclidean ball. However, this Euclidean distance does not reflect the information-theoretic difference between the induced Gaussian distributions. The authors reinterpret (µ, s) as coordinates on a manifold where distance is measured by KL divergence between the distributions they induce.

The Fisher information matrix (FIM) for the parameterization θ = (µ, s) with s = ln σ2 is derived as:

F(θ) = diag(e-s, ½)

The natural gradient update is:

δθ* = -η F(θ)-1 ∇θ lθ =: -η ∇θ nat lθ

This yields output-layer natural gradients:

∇µ nat lθ = e s ∇µ lθ

∇s nat lθ = 2 ∇s lθ

For batched updates, the paper introduces a normalization step: each point's natural gradient is normalized to unit L2 norm before backpropagation. The batch update rule is:

∇µ nat = e s ⊙ ∇µ ∈ R B

∇s nat = 2 ∇s ∈ R B

Θ ← Θ - η [∇µ nat ∂µ/∂Θ + ∇s nat ∂s/∂Θ]

Approximate Trust Radius

Fisher8 admits an approximate KL trust radius between successive predictive distributions. Under nested approximations, the batch KL divergence is bounded by:

KL ≤ ½e-min(s) η2 + ¼ η2

This provides a post hoc readout of how far the joint output distribution has moved after each update. The authors note that if the network was overconfident (small s despite large mean error), the distributional contribution from changing the mean is large, implying a substantial mean correction was applied. Conversely, the KL contribution of further inflating s is upper bounded, limiting the incentive to inflate uncertainty to mask poor mean fits.

Comparison to Prior Stabilizers

The paper shows that three independently proposed stabilizers converge on overlapping components of this geometric correction:

  • β-NLL's gradient reweighting (e βs∇µl) matches Fisher8's mean update when β=1

  • Faithful's Newton steps for the mean head correspond to the same e s∇µl update

  • Fisher8 adds a factor of 2 to the scale gradient (2∇sl) and provides KL control

The key differences are that Fisher8 introduces no data-dependent hyperparameters beyond learning rate, provides KL control, and does not require stop-gradients or gradient severing.

Experiments

1D High-Frequency Sinusoid: Fisher8 corrects the documented training failure on a canonical 1D benchmark with both constant and input-dependent noise. The paper introduces local KL variance as a diagnostic metric to assess when feature-space activity fails to induce distributional change. Results show Fisher8 converges to ground-truth mean with calibrated uncertainty bands within 100k updates, while Baseline-NLL remains stalled. The authors add an addendum to Seitzer et al.'s hypothesis: incremental gains in feature-space expressivity improve predictive quality only if they also produce sufficient mobility in the distributions induced by the predicted parameters θ = (µ, s).

Multidimensional Regression (UCI benchmarks): Using SGD with a fixed budget of 100 gradient steps, Fisher8 achieves the best RMSE, NLL, and ECE on nearly every benchmark at conservative learning rates (lr=0.005). The paper demonstrates two fingerprints of second-order methods: (1) more progress per step at conservative learning rates, and (2) increased sensitivity to learning-rate misspecification at large step sizes. At lr=0.005, Fisher8 achieves RMSE of 8.02±1.34 on Yacht vs. 10.09±2.72 for β-NLL, and NLL of 2.44±0.12 vs. 2.50±0.20. The paper notes: Fisher8 inherits, and is not immune to, the high sensitivity of second-order approaches to the choice of learning rate.

Weak Lensing Cosmology: On a dataset of convergence maps (κ ∈ R 1424×176) for inferring cosmological parameters (Ωm, S8), Fisher8 achieves the highest score (8.528±0.165 vs. 7.683±0.261 for Baseline-NLL), lowest RMSE, and best NLL on both parameters. Interestingly, β-NLL and Faithful perform worse than Baseline-NLL, which the authors hypothesize is due to interaction with systematics in the noise composition.

Uncertainty-Aware Representations (Rotated MNIST): Fisher8 demonstrates superior downstream feature-space utility. On input-dependent noise, Fisher8 achieves 82.12% downstream digit classification accuracy vs. 70.98% for β-NLL and 54.22% for Faithful. On input-independent noise, Fisher8 achieves 80.94% accuracy while other methods collapse to 35–52%. The paper notes: With class-independent noise there is no need to learn a feature space that also encodes digit identity, yet Fisher8 learns digit-aware features by default. Additionally, Fisher8 requires at least 1000× less weight regularization (λ=0.1 vs. λ=100 or 1000 for baselines), and disabling regularization entirely creates negligible changes in performance.

Key Contributions

  1. Introducing local KL variance as a diagnostic metric to assess when feature-space activity fails to induce distributional change.

  2. Deriving Fisher8, an output-layer natural-gradient correction with approximate trust-region control and no dataset-dependent hyperparameters beyond learning rate.

  3. Providing a unifying lens that relates independently proposed stabilizers to overlapping components of this geometrically motivated update rule.

Limitations and Future Work

The paper acknowledges Fisher8's sensitivity to learning rate as a limitation. Future work includes extension to other likelihood classes, comparison with methods incorporating shared network curvature, analysis of interaction between Adam's preconditioning and Fisher8's gradient reorientation, and expansion to dense prediction tasks like depth estimation.

Improvements for AI systems

Improvements to AI Systems:

  1. Stabilized Heteroscedastic Regression: Replace standard Gaussian NLL training with Fisher8's output-layer natural-gradient correction. This yields AI systems that reliably learn both predictive means and calibrated uncertainty estimates without requiring fragile stabilizers (β-NLL, stop-gradients, or heavy regularization). The system can train on noisy data with fewer phase transitions, less variance inflation, and no degradation of mean predictions.

  2. Information-Geometric Gradient Traversal: Implement Fisher8's update rule (∇µ nat = e s ∇µl, ∇s nat = 2∇sl) so that every gradient step moves the predictive distribution along a KL-optimal path rather than a Euclidean one. The improved system achieves more progress per training step at conservative learning rates (e.g., 8.02 RMSE vs. 10.09 for β-NLL on Yacht), enabling faster convergence with fewer epochs.

  3. Automatic Trust-Region Control: Use Fisher8's approximate KL bound (KL ≤ ½e-min(s)η2 + ¼η2) as a built-in safety mechanism. The system can monitor distributional movement per update, preventing overconfident mean corrections or excessive uncertainty inflation. This eliminates the need for manual regularization tuning—Fisher8 works with λ=0.1 where baselines require λ=100–1000, and even with zero regularization.

  4. Uncertainty-Aware Feature Learning: Apply Fisher8 to regression heads in multi-task or representation-learning settings. The system learns feature spaces that encode task-relevant structure (e.g., digit identity in Rotated MNIST) even when the noise is input-independent, achieving 80.94% downstream accuracy vs. 35–52% collapse for other methods. This yields representations that are more transferable and robust to distribution shift.

  5. Diagnostic Monitoring via Local KL Variance: Integrate the paper's local KL variance metric as a training-time diagnostic. The improved system can detect when feature-space updates fail to produce distributional change (stalled learning) and trigger corrective actions, such as adjusting learning rate or reinitializing layers, before convergence fails.

  6. Robust Performance on High-Dimensional Structured Data: Deploy Fisher8 on scientific inference tasks (e.g., cosmology parameter estimation from convergence maps). The improved system achieves the highest predictive scores (8.528 vs. 7.683 baseline) and lowest RMSE/NLL on both parameters, making it suitable for real-world physics applications where calibrated uncertainty is critical.

  7. Unified Stabilizer Replacement: Replace β-NLL, Faithful, and manual gradient reweighting with Fisher8's single, principled correction. The improved system requires no data-dependent hyperparameters beyond learning rate, simplifying model selection and reducing sensitivity to initialization or noise composition (where β-NLL and Faithful actually degrade performance).

  8. Learning-Rate Sensitivity Awareness: Since Fisher8 inherits second-order sensitivity, the improved system can be paired with learning-rate schedulers or warm-up strategies. This allows the AI to operate at conservative rates (lr=0.005) for stability while optionally increasing rates with adaptive schedules for faster convergence in less noisy regimes.

Abstract

Training neural networks to jointly predict mean and uncertainty estimates from noisy observations can be unstable, prompting a series of independent stabilization efforts. We argue that these interventions highlight a common underlying issue where gradient steps are poorly aligned with the geometry of the loss landscape. To better align updates with local curvature, we derive Fisher8, an output-layer gradient correction that reorients and rescales updates using Fisher geometry rather than Euclidean geometry. Unlike past stabilizers, Fisher8 introduces no data-dependent hyperparameters beyond learning rate and admits an approximate KL trust radius between successive predictive distributions. We show that prior stabilizers converge on overlapping components of this geometric correction. Across multidimensional regression and representation-learning tasks, Fisher8 obtains superior likelihood--error tradeoffs, predicts calibrated uncertainty estimates, and learns rich uncertainty-aware feature spaces.

Sources

Related papers