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.
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