Skip to content
AI.info

Research

Learning What to Remember: Adaptive Probabilistic Memory Retention for Memory-Efficient Language Models

Learning What to Remember: Adaptive Probabilistic Memory Retention for Memory-Efficient Language Models Authors: S M Rafiuddin (Department of Computer Science, Oklahoma State University) and Muntaha N

Learning What to Remember: Adaptive Probabilistic Memory Retention for Memory-Efficient Language Models
arXiv
2510.08798
Published
2025-10-09
Authors
S M Rafiuddin, Muntaha Nujat Khan

AI summary

Learning What to Remember: Adaptive Probabilistic Memory Retention for Memory-Efficient Language Models

Authors: S M Rafiuddin (Department of Computer Science, Oklahoma State University) and Muntaha Nujat Khan (Department of English, Oklahoma State University) arXiv: 2510.08798v1 [cs.CL], 09 Oct 2025 — Accepted at EMNLP 2025 Findings as a short paper

Overview

Research area: Natural Language Processing — efficient long-context language modeling and memory-constrained transformer inference.

Technical level: Intermediate. The core idea (learn which tokens to keep) is intuitive, but the paper includes Lagrangian optimization, Hard-Concrete reparameterization, and several theoretical appendices.

Scope (one sentence): The paper proposes and evaluates "Adaptive Retention," a probabilistic, layer-wise token selection mechanism that learns which hidden states to keep under a strict global budget, and measures its accuracy, throughput, and memory behavior across six NLP benchmarks.

What This Paper Is About

Transformer attention scales quadratically with sequence length, O(n²), which limits how much text a model can process at once and drives up memory use. The paper's goal is to let a model decide, layer by layer, which token representations are worth keeping and which can be discarded, while respecting a hard global budget on the total number of retained tokens. Rather than changing how attention itself is computed, the method learns a per-token retention probability and keeps only the top-scoring tokens, so the same base transformer can operate with progressively fewer active tokens at greater depth.

Key Contributions

  1. A probabilistic formulation of memory retention. Retention is modeled with binary indicators z = (z₁, …, z_T) drawn from Bernoulli(p), where p_t = Pr[z_t = 1], and the objective minimizes expected task loss subject to the constraint that the expected number of retained tokens does not exceed a budget M.

  2. A lightweight, context-aware retention scorer. A gated scoring network combines the current token state h_t with a decayed running summary state m_{t−1} (with decay parameter γ), so retention decisions use both local and global context.

  3. A differentiable training procedure with a budget constraint. The method uses a Lagrange multiplier λ updated by projected gradient ascent, combined with a Hard-Concrete reparameterization that allows gradients to flow through the discrete sampling step. At inference, retention is made deterministic via a top-M rule: keep the tokens whose probabilities p_t meet or exceed the M-th largest value.

  4. Empirical breadth plus theory. Six benchmarks (SST-2, IMDb, ArXiv, QASPER, PubMed RCT, CUAD), a component-level ablation, throughput and memory measurements, and six appendices covering a slackness guarantee, unbiasedness of the gradient estimator, a variance bound, convergence of the alternating scheme, a duality-gap bound, and a Rademacher-complexity generalization bound.

Main Findings

  • Accuracy holds up under heavy token budgets. The paper reports that keeping only 30–50% of tokens preserves at least 95% of full-model performance while cutting peak memory by roughly 35–45% and improving throughput by up to about 1.8×.

  • At 50% retention, gaps to the dense model are small. Adaptive Retention scores 91.5 on SST-2 (vs. 92.1 dense), 94.1 on IMDb (vs. 94.8), 80.9 R1 on ArXiv (vs. 81.3), 43.8 EM / 65.0 F1 on QASPER (vs. 44.0/65.0), 42.0 R1 / 22.0 RL on PubMed (vs. 44.0/22.0), and 85.8 micro / 87.8 macro on CUAD (vs. 86.0/88.0).

  • At 30% retention, the same qualitative pattern holds. Adaptive Retention scores 89.2 on SST-2 (vs. 90.8 dense), 92.3 on IMDb (vs. 93.6), 79.5 on ArXiv (vs. 80.1), 39.8/63.0 on QASPER, 40.0/21.0 on PubMed, and 84.0/85.8 on CUAD. PubMed ROUGE-1 shows the largest gap the paper highlights, trailing dense by 2.0 points at each budget while matching ROUGE-L.

  • It beats heuristic pruning. At 50%, Adaptive Retention outperforms H2O by +2.5/+2.6/+2.4 points on SST-2/IMDb/ArXiv, +5.3 EM and +4.5 F1 on QASPER, +2.0 R1 and +2.5 RL on PubMed, and +3.8/+3.8 on CUAD micro/macro. At 30%, it beats H2O by +3.7/+3.4/+3.9 points on SST-2/IMDb/ArXiv, +3.3 EM/+4.5 F1 on QASPER, +2.0 R1/+2.5 RL on PubMed, and +4.0/+3.8 on CUAD. Gains over random pruning are larger, including +6.5 points on SST-2 and +7.2 on IMDb at 30%.

  • It is competitive with sparse-attention architectures. On ArXiv, Adaptive Retention slightly exceeds both Longformer and BigBird at both budgets (80.9/79.5 vs. 80.1/78.0 for Longformer and 80.7/79.1 for BigBird). On other tasks, the paper describes gaps as typically ≤1 point, occasionally up to 2 points.

  • Zero-shot LLM references trail. Prompted GPT-3.5, Llama 2, Llama 3, Falcon, Mistral, Gemma, and Phi4-Mini score below fine-tuned encoder baselines on these supervised evaluations; the paper explicitly labels these comparisons as not directly comparable.

  • Every component matters. Removing the variational (Hard-Concrete) relaxation drops SST-2/IMDb/ArXiv by 1.3/1.4/1.9 points at 50% and 1.4/1.8/2.7 at 30%, and cuts throughput from 1.80× to 1.60× (memory 7.5 GB to 7.2 GB). Disabling alternating optimization costs 0.7/0.9, 1.0/1.1, and 1.4/2.3 points. Fixing the Lagrange multiplier costs 0.5/0.5, 0.5/0.5, and 1.1/1.6 points, with throughput of 1.70–1.75× and 7.3–7.4 GB. Simple threshold-based pruning trails by 2.4–3.9 points and reaches only 1.40× at 7.1 GB.

  • Throughput and memory. At 30% retention on a single 12 GB GPU, Adaptive Retention processes a batch in 0.27 s versus 0.48 s for the Full Transformer (0.54 s vs. 0.96 s per 1,000 tokens), a 1.80× speedup. Longformer reaches 1.55× (0.31 s per batch) and BigBird 1.45× (0.33 s).

  • Retention decays with depth. On SST-2 with DistilBERT-base-uncased, the retained fraction falls from 45.2% at layer 1 to 38.5% at layer 6 under the "30% budget" setting. On ArXiv with Longformer-base-4096, it falls from 50.0% at layer 1 to 34.0% at layer 12 under the same setting. The paper itself flags these numbers as inconsistent with a 30% target, since the method defines per-layer retention as keeping the top M_l = ⌊ρT_l⌋ tokens with ρ ∈ {0.5, 0.3}, and states they "should be corrected or explicitly explained."

  • Budget choice is a practical lever. Moving from 50% to 30% brings roughly an additional 20–25% latency reduction and 0.4–0.8× extra throughput, with accuracy drops of about 1–2 points on short-text classification and 0.6–1.0 points on long-document tasks.

Methodology in Plain English

The team starts from a simple idea: inside a transformer, not every token carries useful information forward. If the model could learn to identify the tokens worth keeping at each layer, it could shrink the active sequence as it goes deeper, saving compute and memory in the later layers without touching the attention mechanism itself.

To make this learnable, the researchers attach a small scoring network to the transformer. For each token position, this scorer looks at the current hidden state and a running summary of everything seen so far (a decayed average controlled by γ), and outputs a probability that the token should be kept. During training, each token is kept or dropped by sampling from a Bernoulli distribution with that probability.

Two technical obstacles are solved. First, sampling is discrete and blocks gradients — so training uses a Hard-Concrete relaxation, a smooth approximation that passes gradients through while still behaving like a near-binary gate. Second, the total number of kept tokens must stay under a budget — so the constraint is folded into the loss using a Lagrange multiplier λ, which is nudged upward whenever the model overspends its budget and held at zero otherwise, with SGD on the model and projected gradient ascent on λ alternating each step.

At test time, sampling is replaced with a deterministic rule: compute each token's retention probability, find the M-th largest value, and keep exactly the tokens at or above that threshold. Because the retained set gets smaller in deeper layers, attention and downstream heads operate on progressively shorter sequences. The method fine-tunes DistilBERT-base-uncased (~66M) for SST-2, IMDb, and CUAD, and Longformer-base-4096 (~149M) for ArXiv, QASPER, and PubMed RCT, using AdamW with a learning rate of 3×10⁻⁵ and weight decay of 0.01, three epochs on the short-context datasets (batch size 32) and one epoch on the long-document datasets (batch size 16). The Hard-Concrete settings are β = 0.66, γ = −0.1, ζ = 1.1, with λ updated at step size η = 1×10⁻².

Why This Matters

Research impact. The paper reframes token pruning from a fixed heuristic into a learned, budget-constrained probabilistic decision, and shows that choosing which tokens to keep can approach the performance of architectures that change how attention is computed — without modifying the base attention or task heads, making it drop-in for standard encoders. The appendices add formal support (unbiased gradient estimates, a variance bound, convergence of the two-timescale scheme, and a generalization bound showing masked-class Rademacher complexity scales with M/T).

Real-world applications:

  • Long-document question answering over scientific papers (QASPER) and arXiv-length texts (average 5,000 tokens).
  • Biomedical literature summarization on PubMed RCT abstracts.
  • Legal contract review, such as CUAD clause classification on long contracts.
  • Processing full-length documents and movie reviews (IMDb averages 230 words, many exceeding 512 tokens) on memory-limited hardware.

Industry relevance. The measured efficiency gains — 0.27 s per batch and 1.80× throughput at 30% retention on a single 12 GB GPU, with peak memory reduced by roughly 35–45% — target exactly the constraint that limits serving long-context models in production. Because the method leaves base attention unchanged and the retention scorer adds only O(T) overhead (under 2% latency in the authors' profiles), it is a practical retrofit rather than an architectural rewrite.

Future Directions

  • Extending to autoregressive decoding. The paper explicitly leaves this open and sketches a design: causal caching that stores hidden states only for retained tokens, amortized per-step scoring with a top-M cache maintained via a min-heap priority queue, and bounded attention over a fixed-size memory bank. It notes that direct comparison against decoder-side KV-cache methods such as SnapKV and PyramidKV is non-trivial and treats them as complementary.
  • Scaling up. Results are on medium-scale backbones (DistilBERT, Longformer-base); behavior at billion-parameter scales remains undemonstrated.
  • Robustness to hyperparameters. Performance shows some sensitivity to the budget penalty and relaxation controls (β, γ, ζ). The paper reports robustness plateaus and provides defaults plus sweeps, but broader guidance for new deployments is still needed.
  • Reconciling the reported layer-wise retention numbers. The appendix figures labeled "30% Budget" (roughly 34–50%) do not match a 30% target, and the paper states they should be corrected or explicitly explained before the depth-wise trends are interpreted.

Target Audience

Researchers and engineers working on efficient transformer inference, long-context modeling, and model compression will get the most from this paper. It also suits practitioners who need to serve long-document NLP tasks under tight GPU memory budgets and want a drop-in alternative to sparse-attention architectures or fixed pruning heuristics. Readers without a background in variational methods or constrained optimization will still follow the main results, but the appendices assume comfort with Lagrangian duality and stochastic approximation theory.

Authors’ abstract

Transformer attention scales quadratically with sequence length O(n^2), limiting long-context use. We propose Adaptive Retention, a probabilistic, layer-wise token selection mechanism that learns which representations to keep under a strict global budget M. Retention is modeled with Bernoulli gates trained via a Hard-Concrete/variational relaxation and enforced with a simple top-M rule at inference, making the method differentiable and drop-in for standard encoders. Across classification, extractive QA, and long-document summarization, keeping only 30-50% of tokens preserves >= 95% of full-model performance while cutting peak memory by ~35-45% and improving throughput by up to ~1.8x. This architecture-agnostic approach delivers practical long-context efficiency without modifying base attention or task heads.

Read the original paper