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 (, measuring interference with retained knowledge) and a frontier component (, 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:
- 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 , a candidate pool of unlabeled target-domain documents, and budget , find:
such that CPT on : (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 :
- Anchor view (): amplifies high-Fisher components → exposes overlap with retained behavior.
- Frontier view (): amplifies low-Fisher components → exposes capacity for low-interference acquisition.
Submodular Objective
The two-geometry log-determinant objective:
where weights the frontier term, and 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: is monotone submodular with , so the cardinality-constrained problem admits the classical greedy approximation guarantee (Nemhauser et al., 1978).
Scalable Pipeline (Figure 2)
Four engineering approximations make this tractable at LLM scale:
- LoRA subspace: Gradients computed in LoRA adapter subspace (), with a short warmup to obtain operating point .
- Diagonal Fisher: Replace full Fisher with diagonal (EWC approximation), estimated from a held-out proxy set :
where element-wise.
- TRAK-style random projection: with , preserving pairwise inner products via Johnson–Lindenstrauss guarantees.
- Sieve-Streaming maximization: Single pass over the stream with a geometric grid of thresholds ; each candidate is accepted into a sieve only if its marginal gain exceeds the threshold. Per-candidate cost is time, 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 , , 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 , forgetting ):
| Method | Med-PPL↓ | PubMedQA↑ | BioASQ↑ | A↑ | Φ↓ |
|---|---|---|---|---|---|
| Base model | 9.5 | 64.6 | 72.9 | - | - |
| Random | 7.6 | 64.4 | 73.6 | +0.2 | 3.5 |
| High-PPL | 7.0 | 65.5 | 74.6 | +1.3 | 5.2 |
| Replay | 8.3 | 61.0 | 71.4 | -2.6 | 0.9 |
| EWC | 7.8 | 64.2 | 72.6 | -0.4 | 1.4 |
| Ours | 6.2 | 70.0 | 78.5 | +5.5 | 0.4 |
Llama-3.1-8B (selected metrics):
| Method | Med-PPL↓ | MedQA↑ | MedMCQA↑ | PubMedQA↑ | A↑ | Φ↓ |
|---|---|---|---|---|---|---|
| Base model | 6.5 | 60.6 | 56.4 | 77.2 | - | - |
| Random | 5.2 | 60.8 | 56.8 | 78.1 | +0.5 | 2.5 |
| Replay | 5.7 | 58.5 | 54.5 | 75.8 | -1.8 | 0.6 |
| Ours | 4.0 | 65.5 | 60.5 | 81.5 | +4.4 | 0.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)
| Variant | A↑ | Φ↓ | A−Φ↑ |
|---|---|---|---|
| Ours (full, α=1) | +5.5 | 0.4 | +5.1 |
| No (drop anchor) | +4.2 | 5.0 | -0.8 |
| No (drop frontier) | +0.8 | 0.5 | +0.3 |
| Identity Fisher (Λ=I) | +1.5 | 2.5 | -1.0 |
| Set-level on raw | +2.0 | 2.0 | 0.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 rule | A↑ | Φ↓ | mean cos↓ | log det Ĝ↑ |
|---|---|---|---|---|
| Top-B by | +1.8 | 0.6 | 0.45 | 8.5 |
| Top-B by | +3.5 | 4.0 | 0.22 | 11.5 |
| Top-B by | +3.0 | 2.2 | 0.30 | 10.0 |
| Sieve-Streaming f(S) | +5.5 | 0.4 | 0.12 | 18.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 budget | 1B (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
-
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.
-
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.
-
Submodularity with streaming guarantees: The log-det objective inherits the greedy guarantee offline and the Sieve-Streaming guarantee online, with per-candidate cost independent of pool size—enabling principled selection at web scale.
Practical Implications
-
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.
-
Replay-free forgetting control: The method bounds forgetting at selection time, potentially reducing or eliminating the need for expensive general-domain replay mixtures.
-
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 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
- On the Recall Scaling Laws in Mamba: A Theoretical and Mechanistic Study via Hashing
Mamba models learn linear hash functions for associative recall, requiring state memory scaling as ND = Θ(Nf log V) — linear in facts, logarithmic in vocabulary.
- iS-KV: Online Low-Rank KV Cache Compression via Block-Incremental SVD
iS-KV compresses KV caches via block-incremental SVD, jointly updating basis and coordinates to retain all tokens, achieving near-original accuracy at 4-7x compression, outperforming eviction methods.
- CATCH: A Controllable Analysis Testbed for Reward Hacking in Coding RL
CATCH is a controllable coding-RL testbed revealing that chain-of-thought monitors suppress reward hacking initially but erode as policies learn to mislead them with code comments.