Skip to content
AI.info

Research

Beyond Sharpness: A Flatness Decomposition Framework for Efficient Continual Learning

Overview Research area: Continual learning (CL), with a focus on optimization and loss-landscape geometry — specifically sharpness-aware minimization and flat-minima methods. Technical level: Intermed

Beyond Sharpness: A Flatness Decomposition Framework for Efficient Continual Learning
arXiv
2601.07636
Published
2026-01-12
Authors
Yanan Chen, Tieliang Gong, Yunjiao Zhang, Wen Wen

AI summary

Overview

Research area: Continual learning (CL), with a focus on optimization and loss-landscape geometry — specifically sharpness-aware minimization and flat-minima methods.

Technical level: Intermediate. The paper assumes familiarity with gradient-based training, stochastic optimizers such as SGD, and the standard CL paradigms (replay-based, regularization-based, architecture/expansion-based). The decomposition itself is described with projections and exponential moving averages, but the core intuition is accessible.

Scope: The paper proposes FLAD, a flatness-decomposition framework that isolates the stochastic-noise component of sharpness-aware perturbations and applies it with a partial-epoch schedule to improve continual learning accuracy at low computational cost.

What This Paper Is About

Continual learning models must learn tasks one after another without forgetting earlier ones, and one promising remedy is to steer training toward flatter loss minima. Existing sharpness-aware methods, however, treat sharpness as a single undifferentiated signal and pay for it with expensive extra gradient computations. This paper asks which part of the sharpness perturbation actually helps, and shows it is the stochastic-noise component rather than the gradient-aligned one, then builds an efficient training scheme around that finding.

Key Contributions

  1. A decomposed flatness signal for CL. The authors propose FLAD, a continual learning framework that penalizes only the stochastic-noise part of sharpness-aware perturbations, aiming to escape sharp minima while preserving past knowledge.
  2. A new interpretation of sharpness-aware optimization. Perturbation directions are explicitly split into gradient-aligned and gradient-orthogonal (stochastic-noise) components, and the paper argues that the noise-aligned perturbations are what drive generalization in CL.
  3. A lightweight scheduling scheme. Because the method can be applied for only part of training, the authors show that partial application already yields substantial gains, enabling large reductions in training time.
  4. Broad empirical validation. FLAD is plugged into six CL baselines spanning replay-based, regularization-based and expansion-based methods across CIFAR-10, CIFAR-100 and Tiny-ImageNet, consistently outperforming both standard and sharpness-aware optimizers.

Main Findings

  • The noise component is the useful one. In first-order experiments comparing perturbation variants (GAM, Pre, Random, Full, Noise), the Noise variant — restricted to the stochastic-noise component — outperformed all other variants, while the Full variant (perturbation aligned with the full gradient component) degraded performance.
  • Noise perturbations align with high curvature early, then settle. Measuring Tr(HΣ), which captures alignment between the Hessian H and the gradient noise covariance Σ, the noise variants corresponding to FLAD-0th and FLAD-1st generally showed higher values than their standard counterparts early in training with stronger fluctuations, and the trend reversed later in training.
  • Flatter final solutions. Hessian eigenvalue distributions and trace during CL showed a pronounced reduction in both the top eigenvalue and the trace compared with vanilla SGD and C-Flat, and PyHessian visualizations showed the MEMO landscape becoming noticeably smoother and wider when integrated with FLAD.
  • Consistent gains across six CL methods. Applying FLAD improved every baseline. The reported "Average Return" column, described as the average boost of FLAD toward C-Flat in each row, was +2.18% (Replay), +1.24% (iCaRL), +1.90% (WA), +1.41% (FOSTER), +1.83% (MEMO) and +0.97% (PODNet).
  • Example head-to-head numbers. On CIFAR-10 with N=5, Replay reached AAA 61.84 ± 6.66 and Acc 41.68 ± 3.61; with SAM 60.91 ± 5.42 / 42.57 ± 4.32; with GAM 61.92 ± 5.52 / 43.07 ± 6.82; with C-Flat 60.46 ± 6.62 / 42.32 ± 2.85; with FLAD 62.60 ± 4.85 / 43.13 ± 4.32.
  • Fastest convergence among the optimizers compared. On CIFAR-100 with the Replay baseline, FLAD achieved the fastest convergence and the highest final accuracy among the optimizers tested.
  • Partial application works. Applying the optimizer for only 10–20% of epochs already yielded significant improvement over vanilla SGD, and in many cases matched or exceeded training with it throughout. Using it for 30% of epochs reduced computational overhead by at least 50% compared with full-method training.
  • Fewer epochs beat more epochs with other optimizers. Training with only 20 epochs using FLAD achieved higher accuracy than training with 50 epochs using C-Flat and than 200 epochs using other optimizers.
  • Both sharpness orders help. Ablations on γ (weight of the first-order term) showed the method outperformed the vanilla optimizer across a wide range of γ values, and models trained with ρ > 0 uniformly outperformed those without gradient ascent. Replacing each gradient term with its decomposed variant consistently outperformed the original across all 6 CL methods.
  • Convergence guarantee. Under the stated assumptions (twice-differentiable loss bounded by M obeying the triangle inequality, loss and its second-order gradient β-Lipschitz smooth, learning rate η ≤ 1/β, perturbation radius ρ ≤ 1/4β, with η_i^T = η/√i and ρ_i^T = ρ/⁴√i), Theorem 1 gives FLAD a convergence rate of O(log n^T / √n^T), where n^T is the total number of iterations for task T.

Methodology in Plain English

Sharpness-aware training normally works by perturbing the weights in the direction of the gradient before computing the update, which requires extra forward and backward passes. FLAD starts from the observation that a mini-batch gradient is really two things added together: a shared direction that reflects the overall trend of the full dataset, and a batch-specific random fluctuation around it.

The authors separate the two by projecting the mini-batch gradient onto the full-batch direction and taking the orthogonal remainder. Because computing the true full-batch gradient every step is prohibitive, they approximate it with an exponential moving average of past mini-batch gradients, and they fix the cosine-similarity coefficient used in the projection as a constant during training.

They then use only that orthogonal remainder as the perturbation direction, for both the zeroth-order (gradient-norm) and first-order (gradient-sharpness) notions of flatness. The first-order term involves a Hessian-vector product, which keeps time and memory manageable. The combined strategy costs 2 forward and 4 backward passes per iteration and can be dropped into standard optimizers.

For continual learning, they plug the optimizer into existing methods: replay-based and regularization-based approaches use the curvature-regularized objective directly, while expansion-based approaches apply it to newly added components while shared components stay frozen. Inference follows post-processing such as classifier calibration or gating. Finally, they show that running FLAD for only a fraction of epochs retains most of the benefit.

Why This Matters

Impact on research: The paper challenges the common treatment of sharpness regularization as a monolithic objective. By showing that the gradient-aligned part of the perturbation can be counterproductive (“Full” degraded performance) while the orthogonal noise part drives the gains, it offers a mechanistic explanation for why sharpness-aware methods work — and suggests that future flatness methods should decompose rather than simply penalize curvature. The analysis via Tr(HΣ) connects the CL setting to existing work on gradient noise covariance.

Real-world applications (the paper does not enumerate specific application domains; these follow from the class-incremental setting it targets):

  • Personalization on edge devices, where a model must absorb new user data streams over time without discarding earlier behavior.
  • Continuously updated recognition or monitoring systems that receive new categories without storing or revisiting all past raw data.
  • Robotics and autonomous systems that encounter new environments and must adapt without losing competence in prior ones.
  • Medical or industrial deployment where models are periodically retrained on incoming data and full retraining is too costly or data is no longer available.

Industry relevance: Continual learning methods that depend on storing raw exemplars or expanding architectures carry storage and complexity costs that grow with the number of tasks. FLAD adds no extra modules and no increase in model complexity, requires only 2 forward and 4 backward passes per iteration, and can be applied in as little as 10–20% of epochs. For teams with constrained training budgets, that combination of plug-and-play integration and reduced compute is the practical selling point.

Future Directions

  • Extending the decomposition beyond the two sharpness orders studied. The framework unifies zeroth- and first-order sharpness, but whether the same noise-isolation principle holds for other curvature measures is not addressed.
  • Making the fixed projection coefficient adaptive. The paper fixes σ as a constant during optimization for efficiency; whether scheduling or learning it would improve results is left open.
  • Characterizing when the gradient-aligned component hurts. The “Full” variant degraded performance in the reported first-order experiments, but the paper does not provide a general condition predicting when this conflict arises.
  • Broadening the empirical scope. Evaluation covers CIFAR-10, CIFAR-100 and Tiny-ImageNet with ResNet-18/ResNet-32 and 3 runs per configuration; extending to other architectures, longer task sequences, and settings where task identity is available at inference is a natural next step.

Target Audience

This paper is most useful to researchers and graduate students working on continual learning, optimization for non-stationary data, or sharpness-aware minimization, as well as practitioners who want to add flatness-based regularization to an existing CL pipeline without rewriting the training loop or paying for expensive second-order updates. Readers without a background in gradient-based optimization will find the loss-landscape and Hessian-based analysis harder to follow, though the decomposition idea itself is intuitive.

Authors’ abstract

Continual Learning (CL) aims to enable models to sequentially learn multiple tasks without forgetting previous knowledge. Recent studies have shown that optimizing towards flatter loss minima can improve model generalization. However, existing sharpness-aware methods for CL suffer from two key limitations: (1) they treat sharpness regularization as a unified signal without distinguishing the contributions of its components. and (2) they introduce substantial computational overhead that impedes practical deployment. To address these challenges, we propose FLAD, a novel optimization framework that decomposes sharpness-aware perturbations into gradient-aligned and stochastic-noise components, and show that retaining only the noise component promotes generalization. We further introduce a lightweight scheduling scheme that enables FLAD to maintain significant performance gains even under constrained training time. FLAD can be seamlessly integrated into various CL paradigms and consistently outperforms standard and sharpness-aware optimizers in diverse experimental settings, demonstrating its effectiveness and practicality in CL.

Read the original paper