This paper presents a lightweight and versatile learned optimizer that dynamically recombines gradient history, represented as averages over disjoint time spans. The optimizer reduces the prediction space to one scalar coefficient per gradient average, shared by multiple parameters. Progressively averaging older gradients minimizes memory cost of long history, while keeping their contributions independently accessible. A 37k-parameter network trained in 0.87 GPU-hours generalizes zero-shot to unseen tasks, lowering validation loss by 9.1% and 0.4% on BERT-Tiny and GPT-Tiny, and improving test accuracy over Adam by 3.5 %p on a Vision Transformer and by 2.7 %p on average across nine graph models, with FLOPs overhead as low as 0.3%.
Figures & tables
Figure 1: Recombination of gradient history as a learned update. Recombination of history contributions −∑hαhgh with predicted coefficients αh can redirect the Adam trajectory away from ascent.
Figure 2: Dynamic recombination of multi-resolution gradient history. For each parameter group, the current gradient enters a multi-resolution history. A group feature encoder converts each stored average into a fixed-dimensional slot feature, and a joint temporal predictor assigns the coefficients used to recombine the history into the group update. Step index t is omitted in the figure as it represents a single optimization step.
Figure 3: Multi-resolution history with exclusive temporal ranges, tier capacities 1,2,4 , and scaling factor r=2 . At step 8, carry propagates through full slots S2 and S3 , evicting the previous average over g1,…,g4 beyond the final tier.
Figure 4: Temporal coefficient prediction. For each history slot, the stored gradient average and current group parameters are converted into histogram features and scalar summaries, then fused into a fixed-dimensional slot feature. Positional encoding identifies the represented range, and self-attention processes the slot features jointly. Softmax produces normalized temporal allocations at,h(b) , while a group scale converts them into the final coefficients αt,h(b) used in Equation 1 . Step index t is omitted in the figure as it represents a single optimization step.
Data
Model
Learned
Adam
SGD
SGDM
Language modeling Terminal validation loss ↓
Wikipedia–BookCorpus
BERT-Tiny (masked-token)
3.463 ∗
3.810
6.703
6.222
WikiText-103
GPT-Tiny (next-token)
4.250 ∗
4.268
5.921
4.686
Vision Test accuracy (%) ↑
CIFAR-100
ViT (lowest val. loss)
31.13
27.61
17.29
32.54
ViT (terminal)
31.74
27.95
27.25
31.22
Table 1: Overall comparison on tasks unseen during optimizer training. Language models report terminal validation loss. Vision reports test accuracy both at the state with the lowest validation loss and at the terminal state. Graph models report test accuracy from the state with the lowest validation loss within 100 updates on Cora and CiteSeer and 500 updates on PubMed. Bold and underlined values denote the best and second-best result per row. Graph results are mean ± standard deviation over three seeds.
Figure 5: The learned optimizer updates all tensors except the embedding and output-head tensors. Adam trains those tensors using the same hyperparameters as the all-Adam baseline. BERT-Tiny follows a 100,000-update protocol with a context transition at update 90,000, while GPT-Tiny is evaluated at 50,000 updates. The GPT magnifies the later trajectory.
Figure 6: FLOP-matched Cora training loss under full span. The horizontal axis expresses cumulative FLOPs in units of one Adam step, with a budget equal to 100 Adam updates. Lines show three-seed means, and bands show seed-wise ranges.
Figure 7: ViT/CIFAR-100 epoch-mean training loss and validation accuracy over 50 epochs under matched initialization, split, and minibatch order. Table 1 reports validation-selected and terminal test accuracy.
Target
Metric
Full span
63-step span
Adam
BERT-Tiny
Val. loss ↓
3.463 ∗
3.824 ∗
3.810
GPT-Tiny
Val. loss ↓
4.250 ∗
4.261 ∗
4.268
ViT/CIFAR-100
Test acc. (best val.) (%) ↑
31.13
30.81
27.61
Test acc. (terminal) (%) ↑
31.74
30.86
27.95
GNN arithmetic mean across 9 cells
Test acc. (%) ↑
76.44 ± 0.28
76.47 ± 0.24
73.78 ± 0.26
Table 2: Effect of history span. Language models report terminal validation loss. ViT reports validation-selected and terminal test accuracy, and GNNs report validation-selected test accuracy. The GNN row is the unweighted arithmetic mean across nine model–dataset comparisons, reported as mean ± sample standard deviation across three target seeds. Bold and underlined values mark the best and second-best result in each row.
Figure 8: Temporal dynamics on GAT/PubMed. (a) Both paths share the trajectory through t=32 , with Adam states reconstructed from the learned path. At step t=33 , Adam branches and reaches terminal training loss of 0.30 ( t=500 ), and 0.25 for learned recombination. (b,d) Red bars show the scale-normalized temporal scores and signed contribution fractions for the layer-1 source-attention group at branch ( t=33 ). Signed fractions indicate each slot’s directional contribution to the constructed update. S1 corresponds to the slot with newest gradient, and S6 with the oldest. Gold bars show the closest nonnegative coefficient allocation fitted to the norm-matched Adam direction. (c,e) The learned scores and signed contributions vary over updates 24–55. Dotted boxes mark t=33 , and hatching marks unpopulated slots.
Figure 9: Deployment cost relative to Adam in isolated NVIDIA B200 profiles using bfloat16 arithmetic. A value of 1.0 denotes the matched Adam baseline. FLOPs/step include forward, backward, the learned update, and any update routed to Adam. Time/step measures the same end-to-end boundary. BERT values weight its 128- and 512-token contexts by 0.9/0.1 . GNN values are geometric means over the nine model–dataset tasks, with min–max error bars. Appendix A.4 reports all profiles for the learned optimizer.
Appendix figures & tables15 assets
Supplementary material from the paper’s appendix.
Appendix
Feature
Definition
Group size
log10nb
Stored-average scale
log10RMSϵ(S)
Current-parameter scale
log10RMSϵ(θ)
Coordinate uniformity
∥S∥1/(nb∥S∥2,ϵ+ϵ)
Gradient–parameter alignment
⟨S,θ⟩/(∥S∥2,ϵ∥θ∥2,ϵ+ϵ)
Gradient-to-parameter scale
log10RMSϵ(S)−log10RMSϵ(θ)
Appendix
Table 3: Statistics computed for the stored gradient average in each populated history slot.
Model
Dataset
Architecture
Params.
MLP
Synthetic binary
64 inputs; two 64-unit ReLU layers; 2-class head
8,450
ConvNet
CIFAR-10
5×5 convolutions with 2, 4 channels; ReLU and 2×2 pooling; 16-unit ReLU; 10-class head
2,142
MicroResNet
Fashion-MNIST
8-channel 3×3 BN stem; residual blocks of width 8, 16; adaptive pooling; 10-class head
5,122
MicroViT
CIFAR-10
8×8 patches; width 16; 2 heads; one pre-normalized block; FF width 64; 10-class head
6,810
Appendix
Table 4: Optimizer-training tasks. Parameter counts include all trainable parameters.
Item
Value
Tasks per meta-window
All four tasks in Table 4
Model-training batch
128
Reset trajectories / steps per trajectory
4 / 160
Successful meta-update cap
64
TBPTT horizon / meta-update cadence
20 / 10 model-training updates
Meta-objective
Normalized cumulative held-out-loss change in Equation 2
Table 6: Language-model configurations used for optimizer evaluation.
Figure 10: FLOPs and time per training step, reported as ratios to Adam, by history span.
Figure 11: FLOPs and time per training step, reported as ratios to Adam, by language-model parameter assignment.
Model
Learned optimizer updates
Adam updates
Learned
Adam
BERT-Tiny
All parameters
—
3.889
3.815
All except embeddings/output head
Embeddings/output head
3.463
3.810
Embeddings/output head
All other parameters
3.484
3.811
GPT-Tiny
All parameters
—
4.299
4.268
All except embeddings/output head
Embeddings/output head
4.250
4.268
Embeddings/output head
All other parameters
4.309
4.268
Appendix
Table 7: Terminal validation loss by language-model parameter assignment. Each row reports the paired Adam baseline. The same trained policy is routed to the listed tensors without separate optimizer training for each assignment. Bold marks the lowest learned-optimizer loss within each model.
Figure 12: Validation loss when the learned optimizer updates all tensors or either complementary parameter partition. Adam updates the tensors not routed to learned recombination using the same hyperparameters as the all-Adam baseline. BERT changes context at update 90,000, and the GPT inset magnifies updates 30,000–50,000.
Figure 13: FLOP-aligned GNN training loss for the full-span learned optimizer and Adam. Rows identify datasets, and columns identify model architectures. The horizontal axis measures cumulative profiler-visible FLOPs in units of one Adam step. Lines show three-seed means, and bands show seedwise ranges. Cora and CiteSeer use budgets of 100 Adam steps, while PubMed uses 500 Adam steps. PubMed learned curves continue beyond the dotted 500-step boundary to show the complete 500-update trajectory.
Figure 14: Step-aligned GNN validation loss. Rows identify datasets, and columns identify model architectures. Curves show three-seed means, bands show seedwise ranges, and circles mark the minimum of each seedwise validation trajectory. Checkpoints are selected independently within the fixed 100-update Cora and CiteSeer budgets and 500-update PubMed budget. Test accuracy does not enter checkpoint selection. Figure 13 separately compares training progress at equal cumulative FLOPs.
Policy training
Deployment grouping
Deployment groups
Train-loss AUC ↓
Terminal
MicroLeNet/CIFAR-10 (terminal test loss ↓ )
Output slice
Output slice
15
1.820 ± 0.012
1.743 ± 0.055
Output slice
Input slice
14
1.823 ± 0.016
1.705 ± 0.054
Output slice
Tensor-wide
8
1.822 ± 0.014
1.681 ± 0.028
Output slice
Network-wide
1
1.831 ± 0.017
1.666 ± 0.018
Input slice
Input slice
14
1.873 ± 0.043
1.748 ± 0.034
Appendix
Table 8: Parameter-grouping ablation. Entries are mean ± sample standard deviation over three paired target seeds. The group count is measured at deployment. Output and input slices prioritize complete leading- and second-dimension slices, respectively. Tensor-wide and network-wide conditions form one group per parameter tensor and optimizee. Bold marks the best mean within each task and metric. Adam uses the same target initializations and data streams.
Data
Model
63-step span (%)
Full span (%)
Cora
GCN
80.77 ± 0.67
80.83 ± 0.45
GraphSAGE
79.53 ± 0.35
79.27 ± 0.50
GAT
80.77 ± 0.71
80.83 ± 0.51
CiteSeer
GCN
70.07 ± 0.35
70.07 ± 0.23
GraphSAGE
70.13 ± 0.31
69.50 ± 0.46
GAT
68.07 ± 1.56
67.60 ± 1.73
Appendix
Table 9: GNN history-depth ablation. Entries are test accuracy (%) from the lowest-validation-loss state, reported as mean ± sample standard deviation over three paired target seeds. The final row is the unweighted arithmetic mean of the nine model–dataset means. Bold marks the best mean, and underlining marks the second-best; tied means are both bold.
Figure 15: Full-trajectory groupwise temporal dynamics on GAT/PubMed. Rows show a 64-coordinate layer-1 source-attention group, a 256-coordinate layer-1 dense-weight group, and the 3-coordinate layer-2 bias. Columns show temporal contributions ah , raw slot RMS, and signed contribution fractions; the final coefficient is αh=sah . RMS colors are logarithmically spaced while their labels retain linear values; signed-fraction colors use a symmetric logarithmic scale. Hatching marks temporal ranges not yet populated.
Figure 16: Spatio-temporal extension through pyramid grouping. A partition configuration p=(a1,…,am) assigns subdivision depth aj to tensor axis j , and spatial tier k=∑jaj contains configurations with the same total subdivision depth. Each cell Gp,b denotes parameter group b within configuration p . The illustration uses two tensor axes and split factor s=2 through tier k=2 ; the evaluated realization retains the whole-tensor and two single-axis configurations with 16 splits per active axis. Within each group, the temporal policy recombines the stored gradient history.
Data
Model
Row
Pyramid
Adam
SGD
SGDM
Cora
GCN
81.70 ± 0.20
81.57 ± 0.12
79.27 ± 0.83
31.80 ± 10.01
74.93 ± 0.92
GraphSAGE
79.60 ± 0.44
80.00 ± 0.69
78.73 ± 0.32
24.10 ± 3.58
76.63 ± 0.15
GAT
81.97 ± 0.31
82.03 ± 0.50
78.27 ± 1.26
47.23 ± 8.00
78.97 ± 1.19
CiteSeer
GCN
71.03 ± 0.81
70.97 ± 0.67
66.20 ± 1.61
33.00 ± 6.22
70.57 ± 0.71
GraphSAGE
69.87 ± 0.87
70.13 ± 0.29
65.57 ± 1.36
34.67 ± 9.38
70.23 ± 0.38
GAT
69.00 ± 0.26
69.17 ± 0.32
66.20 ± 0.82
52.87 ± 1.42
70.37 ± 0.78
Appendix
Table 10: Full-trajectory test accuracy (%) for row and pyramid grouping on unseen graph tasks. Cora and CiteSeer use 100 updates and PubMed uses 500. Entries are mean ± sample standard deviation over three target seeds. Bold and underlined values mark the best and second-best mean in each row.
AdamW is a default optimizer for deep learning, but its moment states add two parameter-sized buffers to training memory, increasing the cost of large-scale pretraining. We propose Gefen, a memory-efficient optimizer that automatically shares second-moment estimates across parameter blocks and quantizes the first moment using a learned codebook. Gefen reduces AdamW's optimizer memory footprint by up to 8x while maintaining performance, saving 6.5 GiB per billion parameters. Prior work shares second moments across parameters grouped along the Hessian's block-diagonal structure, but relies on hand-specified architectural rules and leaves unexplained why such grouping works. We prove that large mixed Hessian entries constrain the ratio of squared gradients toward one, explaining why shared second moments are accurate when the squared gradients they pool are similar. The Hessian need not be computed: its block structure is inherited by squared gradients, allowing blocks to be found directly. Gefen therefore infers block structure from initial squared gradients, requiring no architecture-specific metadata or user-tuned hyperparameters beyond AdamW defaults. Gefen learns an exact histogram-based dynamic-programming quantization codebook and reuses the blocks for first-moment scaling. Across diverse pretraining experiments, Gefen achieves the lowest peak optimizer memory among compared methods that maintain AdamW-level performance. In single-machine and distributed training, the reduced footprint enables larger microbatches and substantially improves throughput over AdamW, making Gefen a drop-in replacement that can train larger models or use larger global batch sizes. We provide the complete Python implementation, including fused CUDA kernels at https://github.com/ndvbd/Gefen
Nadav Benedek, Tomer Koren, Ohad Fried
Reichman University · Tel Aviv University, Google Research
Learned optimization aims to improve upon hand-designed optimizers (e.g., Adam and Muon) by meta-learning small neural network optimizers over a distribution of tasks. While recent work has greatly advanced the architectural design and inductive biases of learned optimizers (LOs), their meta-training remains biased toward short-unroll learning on particular tasks, resulting in redundant computation and leaving LOs often unable to compete with hand-designed optimizers. We introduce Efficient Long-hOrizon (ELO) learning, an efficient meta-training algorithm that (1) reallocates wasted meta-training compute to longer failure regimes, achieving efficient long-horizon learning, and (2) enforces decoupled progressive expert supervision, providing stable meta-learning signals that additionally improve the generalization of LOs. Our empirical study evaluates ELO for meta-training both element-wise and matrix-based LOs. Across downstream language modeling (GPT-2-124M/350M on FineWeb) and image classification (ViT-B/16, ResNet-50 on ImageNet-1K) tasks, ELO substantially improves the long-unroll performance and out-of-distribution generalization of the base LOs. In particular, ELO-Celo2 consistently outperforms well-tuned AdamW across all evaluated tasks, while remaining competitive with Muon on language modeling. \textit{Notably, all ELO baselines require less than 7 H100 GPU-hours for meta-training.}
Xiaolong Huang, Benjamin Thérien, James Harrison +1
Mila - Quebec AI Institute · Concordia University · Université de Montréal +1
We argue that forgetting is not confined to continual learning but is a general optimization phenomenon: during standard training, dominant mini-batch gradients suppress rare but useful update directions, causing short-term forgetting at every step. When such knowledge is never revisited, these losses compound into long-term forgetting-the classical failure mode of continual learning. We introduce FOGO, a scalable optimizer that continuously detects and resolves gradient interference across both regimes. FOGO spectrally orthogonalizes momentum updates to prevent dominant directions from monopolizing optimization, then stores representative past directions in a compact codebook memory built on random projection, where pairwise distances are provably preserved in low-dimensional space. At each step, conflicts between the current update and stored directions are resolved via lightweight orthogonal correction and lifted back through a proximal step, with minimal overhead and no data storage. Across class-imbalanced classification, continual visual learning under domain and class shifts, continual fine-tuning of LLaVA-7B, and GPT-2 pretraining, FOGO consistently improves convergence and knowledge retention, outperforming Adam and Muon.
Toan Nguyen, Yang Liu, Trung Le +2
School of Computer Science and Engineering University of New South Wales Sydney, Australia · Department of Data Science & AI Monash University Melbourne, Australia · DEVCOM Army Research Laboratory USA