Fisher-Guided Submodular Data Selection for Continual Pre-Training of Large Language Models

Authors: Zhenghao Zhao¹, Gaowen Liu², Zhiling Lan¹, Yan Yan¹
¹University of Illinois Chicago, ²Cisco Research


Summary (Overview)

  • Core problem: Data selection in continual pre-training (CPT) of LLMs must simultaneously maximize target-domain adaptation while preventing catastrophic forgetting of pretrained capabilities—a dual objective that parameter-agnostic selection methods (perplexity, loss) cannot address.
  • Key insight: The authors empirically demonstrate that parameter-agnostic CPT causes the post-CPT Fisher diagonal to drift downward on high-Fisher coordinates (committed directions) while leaving low-Fisher coordinates underused—revealing a parameter-space mechanism for catastrophic forgetting.
  • Proposed method: A Fisher-guided selector that decomposes each candidate's gradient into an anchor component (ganc=F1/2g(x)g_{anc} = F^{1/2}g(x), measuring interference with retained knowledge) and a frontier component (gfr=F−1/2g(x)g_{fr} = F^{-1/2}g(x), measuring acquisition capacity in unconstrained subspaces), aggregated via a log-determinant submodular objective optimized with Sieve-Streaming.
  • Key results: On TinyLlama-1.1B and Llama-3.1-8B medical CPT, the method achieves the best adaptation-forgetting Pareto frontier. Critically, 1B selected tokens outperform replay trained on 10B tokens on both axes—a 10× token-efficiency advantage.
  • Contributions: (1) Fisher-based anchor/frontier gradient decomposition; (2) log-det submodular objective with streaming optimization; (3) empirical validation at 1.1B and 8B scale with strict Pareto dominance over baselines.

Introduction and Theoretical Foundation

Background and Motivation

Data selection is a central bottleneck in LLM training: web-scale corpora are noisy and token budgets are finite. In CPT, this becomes a forgetting-control problem—a poorly chosen target-domain corpus can overwrite capabilities encoded in the pretrained checkpoint. Existing approaches fall short:

  • Quality classifiers/rules (Peng et al., 2025; Calian et al., 2025) score cleanliness/topical fit only.
  • Per-sample loss/perplexity (Sorscher et al., 2022) treat the model as a black box returning a scalar.
  • Replay-based CPT mixes general-domain data back into training but pays for retention with many extra tokens.

None of these directly ask: how will training on a candidate move the model parameters?

Theoretical Foundation: Fisher Information Geometry

The Fisher information matrix of the pretrained model encodes, direction-by-direction, how strongly the predictive distribution depends on parameter perturbations:

F:=F(θ0)=Ex∼ppre[g(x;θ0)g(x;θ0)⊤]∈RP×PF := F(\theta_0) = \mathbb{E}_{x \sim p_{\mathrm{pre}}} \big[ g(x;\theta_0) g(x;\theta_0)^{\top} \big] \in \mathbb{R}^{P \times P}
  • High-Fisher eigenvectors: directions the model has committed to—perturbing along them moves the output distribution, so they carry retained capability.
  • Low-Fisher eigenvectors: directions the model holds no opinion about—free for new content.

Key empirical observation (Figure 1): Parameter-agnostic CPT produces substantial negative drift on high-Fisher coordinates (damage to committed directions) while low-Fisher coordinates remain underused. This asymmetry is the parameter-space signature of catastrophic forgetting.

Relationship to Prior Work

  • EWC (Kirkpatrick et al., 2017) uses diagonal Fisher as a regularizer during training; the authors invert this direction—Fisher decides which tokens enter the stream.
  • Coreset/submodular methods (Mirzasoleiman et al., 2020; Killamsetty et al., 2021a) target static loss over labeled pools; the authors lift submodular framing to token-budgeted, unlabeled language modeling.

Methodology

Problem Formulation

Given pretrained parameters θ0∈RP\theta_0 \in \mathbb{R}^P, a candidate pool D={x1,…,xN}\mathcal{D} = \{x_1, \ldots, x_N\} of unlabeled target-domain documents, and budget B≪NB \ll N, find:

S⋆∈arg⁡max⁡S⊆D,∣S∣=BU(S)S^{\star} \in \arg\max_{S \subseteq \mathcal{D}, |S| = B} U(S)

such that CPT on SS: (i) maximizes target-domain adaptation, (ii) bounds general-capability degradation, (iii) covers complementary update directions.

Fisher Decomposition

For each candidate, define two Fisher-reweighted views of its gradient g(x)=∇θℓ(x;θ0)g(x) = \nabla_\theta \ell(x;\theta_0):

ganc(x)=F1/2g(x)∈RP,gfr(x)=F−1/2g(x)∈RPg_{\mathrm{anc}}(x) = F^{1/2} g(x) \in \mathbb{R}^P, \qquad g_{\mathrm{fr}}(x) = F^{-1/2} g(x) \in \mathbb{R}^P
  • Anchor view (gancg_{anc}): amplifies high-Fisher components → exposes overlap with retained behavior.
  • Frontier view (gfrg_{fr}): amplifies low-Fisher components → exposes capacity for low-interference acquisition.

Submodular Objective

The two-geometry log-determinant objective:

f(S)=log⁡det⁡(IP+β∑x∈Sganc(x)ganc(x)⊤)+α⋅log⁡det⁡(IP+γ∑x∈Sgfr(x)gfr(x)⊤)f(S) = \log \det \left(I_P + \beta \sum_{x \in S} g_{\text{anc}}(x) g_{\text{anc}}(x)^{\top}\right) + \alpha \cdot \log \det \left(I_P + \gamma \sum_{x \in S} g_{\text{fr}}(x) g_{\text{fr}}(x)^{\top}\right)

where α≥0\alpha \geq 0 weights the frontier term, and β,γ>0\beta, \gamma > 0 rescale covariance sums. Each log-det term rewards diversity among selected Fisher-reweighted gradients: the determinant is large when gradients span many independent directions, small when collinear.

Properties: ff is monotone submodular with f(∅)=0f(\varnothing) = 0, so the cardinality-constrained problem S⋆∈arg⁡max⁡S⊆D,∣S∣≤Bf(S)S^{\star} \in \arg\max_{S \subseteq \mathcal{D}, |S| \leq B} f(S) admits the classical (1−1/e)(1 - 1/e) greedy approximation guarantee (Nemhauser et al., 1978).

Scalable Pipeline (Figure 2)

Four engineering approximations make this tractable at LLM scale:

  1. LoRA subspace: Gradients computed in LoRA adapter subspace (D≪PD \ll P), with a short warmup to obtain operating point θ⋆\theta^{\star}.
  2. Diagonal Fisher: Replace full D×DD \times D Fisher with diagonal Λ∈R>0D\Lambda \in \mathbb{R}^D_{>0} (EWC approximation), estimated from a held-out proxy set Dproxy\mathcal{D}_{proxy}:
ganc(x)=Λ⊙g(x),gfr(x)=g(x)⊘Λεg_{\mathrm{anc}}(x) = \sqrt{\Lambda} \odot g(x), \qquad g_{\mathrm{fr}}(x) = g(x) \oslash \sqrt{\Lambda_{\varepsilon}}

where Λε=max⁡{Λ,ε}\Lambda_{\varepsilon} = \max\{\Lambda, \varepsilon\} element-wise.

  1. TRAK-style random projection: RD→Rd\mathbb{R}^D \to \mathbb{R}^d with d≪Dd \ll D, preserving pairwise inner products via Johnson–Lindenstrauss guarantees.
  2. Sieve-Streaming maximization: Single pass over the stream with a geometric grid of thresholds {τi}\{\tau_i\}; each candidate is accepted into a sieve only if its marginal gain Δ(x)=Δanc(x)+αΔfr(x)\Delta(x) = \Delta_{anc}(x) + \alpha \Delta_{fr}(x) exceeds the threshold. Per-candidate cost is O(d2)O(d^2) time, O(∣T∣⋅d2)O(|\mathcal{T}| \cdot d^2) memory, independent of pool size.

Empirical Validation / Results

Experimental Setup

  • Backbones: TinyLlama-1.1B (primary), Llama-3.1-8B (scale-up validation)
  • Target domain: Medical (PMC full-text 38B tokens + PubMed abstracts 7B tokens)
  • Reference distribution: FineWeb + Llama3-SynE English (held-out general domain)
  • LoRA configuration: rank r=128r=128, α=512\alpha=512, applied to Q/K/V/O projections
  • Budgets: 4B tokens (TinyLlama), 10B tokens (Llama-3.1-8B)
  • Decontamination: PubMedQA/BioASQ source-ID blocklist, 13-gram overlap filter, MinHash near-duplicate removal (removes 2.68% of documents)

Main Results (Table 1)

TinyLlama-1.1B (adaptation gain A↑A \uparrow, forgetting Φ↓\Phi \downarrow):

MethodMed-PPL↓PubMedQA↑BioASQ↑A↑Φ↓
Base model9.564.672.9--
Random7.664.473.6+0.23.5
High-PPL7.065.574.6+1.35.2
Replay8.361.071.4-2.60.9
EWC7.864.272.6-0.41.4
Ours6.270.078.5+5.50.4

Llama-3.1-8B (selected metrics):

MethodMed-PPL↓MedQA↑MedMCQA↑PubMedQA↑A↑Φ↓
Base model6.560.656.477.2--
Random5.260.856.878.1+0.52.5
Replay5.758.554.575.8-1.80.6
Ours4.065.560.581.5+4.40.1

Key findings:

  • Ours achieves the best (A, Φ) Pareto frontier, with adaptation gain exceeding the strongest perplexity baseline by 4.2 and forgetting more than halved versus the next-best method.
  • Ours surpasses domain-specialized Galactica on medical QA by a wide margin while preserving general capabilities.
  • On Llama-3.1-8B, Ours is the only method whose general-suite scores remain robustly flat.

Ablation: Fisher Decomposition (Table 2)

VariantA↑Φ↓A−Φ↑
Ours (full, α=1)+5.50.4+5.1
No gancg_{anc} (drop anchor)+4.25.0-0.8
No gfrg_{fr} (drop frontier)+0.80.5+0.3
Identity Fisher (Λ=I)+1.52.5-1.0
Set-level on raw gg+2.02.00.0

Interpretation: Removing the anchor term inflates forgetting; removing the frontier term stalls adaptation; identity Fisher collapses to raw gradient norm; diversity alone (without Fisher geometry) is insufficient.

Ablation: Subset-Level vs. Top-k (Table 3)

Selection ruleA↑Φ↓mean cos↓log det Ĝ↑
Top-B by ∥ganc∥\|g_{anc}\|+1.80.60.458.5
Top-B by ∥gfr∥\|g_{fr}\|+3.54.00.2211.5
Top-B by δ0(x)\delta_0(x)+3.02.20.3010.0
Sieve-Streaming f(S)+5.50.40.1218.5

Top-k rules cannot suppress mutual redundancy; the log-det rule selects subsets that are simultaneously high-magnitude and mutually orthogonal.

Token Efficiency (Table 4)

Token budget1B (A↑, Φ↓)2B (A↑, Φ↓)4B (A↑, Φ↓)10B (A↑, Φ↓)
Replay-3.5, 1.2-2.9, 1.0-2.6, 0.9-2.0, 0.8
Ours+3.0, 0.6+4.5, 0.5+5.5, 0.4-

Ours at 1B tokens already exceeds Replay at 10B tokens on both A and Φ—a 10× token-efficiency advantage.

Training Dynamics Analysis

Replay-50% reaches the lowest training loss (~1.47) yet performs worse on (A, Φ) than Ours (loss ~1.80). Lower training loss is not a proxy for better domain adaptation: replay makes optimization easier by adding familiar general-domain text, but the resulting model lies below Ours on the Pareto frontier.


Theoretical and Practical Implications

Theoretical Significance

  1. Parameter-space mechanism for forgetting: The paper provides direct empirical evidence that catastrophic forgetting manifests as asymmetric Fisher drift—mass decreases on committed (high-Fisher) coordinates while uncommitted directions remain underused. This reframes forgetting as a geometric failure rather than purely an optimization pathology.

  2. Dual-geometry selection signal: The anchor/frontier decomposition converts forgetting control and knowledge acquisition into a single selection signal, rather than two post-hoc training heuristics. This is a principled unification of two objectives typically treated separately.

  3. Submodularity with streaming guarantees: The log-det objective inherits the (1−1/e)(1 - 1/e) greedy guarantee offline and the (12−ε)(\frac{1}{2} - \varepsilon) Sieve-Streaming guarantee online, with per-candidate cost independent of pool size—enabling principled selection at web scale.

Practical Implications

  1. 10× token efficiency: The ability to match or exceed replay-based forgetting control with 1/10 the tokens has direct cost implications for CPT pipelines, particularly relevant given the expense of large-scale training.

  2. Replay-free forgetting control: The method bounds forgetting at selection time, potentially reducing or eliminating the need for expensive general-domain replay mixtures.

  3. Scalability: The LoRA-subspace restriction, diagonal Fisher approximation, TRAK projection, and Sherman–Morrison rank-1 updates make the method practical for billion-parameter models with modest computational overhead.


Conclusion

Main Takeaways

The paper presents a Fisher-guided submodular data selector for CPT that intrinsically balances target-domain acquisition with preservation of pretraining capabilities at selection time. By decomposing candidate gradients into Fisher-weighted anchor and frontier components and aggregating via a log-det objective optimized with Sieve-Streaming, the method:

  • Strictly Pareto-dominates perplexity- and replay-based baselines across 1.1B and 8B LLMs
  • Achieves ≥10× token-efficiency gain versus forgetting-aware replay
  • Isolates Fisher geometry as the key driver of selection quality (via ablations)

Limitations

  • Bounded by diagonal Fisher approximation (ignores off-diagonal curvature couplings)
  • Relies on a warmup checkpoint θ⋆\theta^{\star} as a local surrogate for the pretrained operating point

Future Directions

  • Extend to richer curvature models (block-diagonal or Kronecker-factored Fisher)
  • Apply to multimodal streaming and instruction-tuning curation
  • Potential integration with model-merging and routing approaches for post-hoc forgetting mitigation

"Forgetting can be read as a Fisher-geometric failure of parameter-agnostic CPT, and candidates should be judged by where their gradients move the pretrained model."

Related papers