The cost of useful natural gradient updates
Organizations: School of Computing, The Australian National University
Abstract
What information is needed to turn a natural-gradient direction into a useful finite update? Under a population Kullback-Leibler (KL) budget, we call a step useful if it is feasible and loses at most a fraction of the best feasible gain along the direction. We construct a four-state exponential family whose laws share their initial gradient, scalar Fisher information and natural gradient, yet two laws have disjoint useful-step sets. With these quantities supplied exactly and the law otherwise known only through draws, the family's worst-case sample complexity is for small , where scales rare-state probabilities and is the failure probability. The budget is fixed and the optimal gain stays bounded away from zero, so the step length, not the direction, carries this cost. For succinctly described event-tilt models, returning a useful step is NP-hard even with the exact natural gradient and efficient exact sampling. Recovering the unit natural gradient to constant error is also NP-hard even in a two-parameter logistic family with Fisher condition number at most 3. We also give matching sample bounds for event tilts, sample bounds for damped Fisher solves and a population-KL certificate for affine classifiers. In frozen-feature classifier heads, stopping at a sampled KL boundary succeeds in about half of the trials, and a 10% KL margin raises joint success above 93% at a KL budget of 0.01. Thus, knowing where to move is not enough: how far to move can carry an update's entire cost.
Figures & tables
Appendix figures & tables10 assets
Supplementary material from the paper’s appendix.
Appendix
| (a) Four-state calibration | (b) Binary calibration | |||
| Plug-in (%) | Conservative (%) | Confidence (%) | ||
| 48.97 | 52.78 | 95.80 | ||
| 48.05 | 93.90 | 96.39 | ||
| 50.29 | 100.00 | 96.29 | ||
| 97.07 | ||||
| Upper success | Plug-in success | Upper max. rate | Plug-in max. rate | |
|---|---|---|---|---|
| 24.02–73.19 | 24.02–73.19 | 40.09 | 40.09 | |
| 40.87–89.31 | 40.87–89.50 | 19.97 | 19.97 | |
| 47.61–99.17 | 46.78–99.32 | 3.81 | 4.25 | |
| 50.78–100.00 | 48.93–100.00 | 0.00 | 0.05 | |
| 53.12–100.00 | 49.95–100.00 | 0.00 | 0.00 | |
| 55.71–100.00 | 48.10–100.00 | 0.00 | 0.00 |
| P | P+R | P+O | P+O+R | P excess | P+O+R excess | ||
|---|---|---|---|---|---|---|---|
| 44.14 | 44.14 | 44.14 | 44.14 | 4.603 | 4.457 | ||
| 48.97 | 51.56 | 50.39 | 52.78 | 1.071 | 0.935 | ||
| 49.76 | 59.47 | 55.37 | 64.36 | 0.263 | 0.152 | ||
| 48.05 | 83.01 | 70.61 | 93.90 | 0.068 | 0.010 | ||
| 50.93 | 100.00 | 99.02 | 100.00 | 0.016 | 0.000 | ||
| 50.29 | 100.00 | 100.00 | 100.00 | 0.004 | 0.000 |
| Euclidean pass | Fisher cosine | Finite gain | ||
|---|---|---|---|---|
| 0.25 | 4.0 | 30.21 | 31.08 | 5.1 |
| 1 | 18.3 | 77.31 | 78.30 | 41.1 |
| 4 | 15.6 | 97.80 | 98.10 | 96.2 |
| 16 | 31.9 | 99.46 | 99.54 | 100.0 |
| 64 | 62.8 | 99.88 | 99.89 | 100.0 |
| 256 | 91.6 | 99.97 | 99.97 | 100.0 |
| Direction | Calibration | Feas. | Gain | [95% CI] | |
|---|---|---|---|---|---|
| 16 | Population | Reference | 100.0 | 100.0 | 100.0 [99.8, 100.0] |
| 16 | Population | Confidence | 99.9 | 78.0 | 1.1 [0.7, 1.6] |
| 16 | Pseudoinverse | Reference | 100.0 | 99.5 | 100.0 [99.8, 100.0] |
| 16 | Pseudoinverse | Confidence | 99.9 | 76.9 | 0.5 [0.3, 1.0] |
| 64 | Population | Reference | 100.0 | 100.0 | 100.0 [99.8, 100.0] |
| 64 | Population | Confidence | 100.0 | 88.0 | 23.6 [21.8, 25.5] |
| Feas. | Gain | [95% CI] | ||||
|---|---|---|---|---|---|---|
| 1 | 4,096 | 16 | 33.7 | 100.0 | 77.0 | 0.8 [0.5, 1.3] |
| 1 | 4,096 | 64 | 61.2 | 100.0 | 87.5 | 19.2 [17.6, 21.0] |
| 1 | 4,096 | 256 | 91.8 | 100.0 | 93.5 | 98.5 [97.9, 98.9] |
| 1 | 65,536 | 16 | 33.2 | 99.9 | 76.9 | 0.5 [0.3, 1.0] |
| 1 | 65,536 | 64 | 60.8 | 100.0 | 87.6 | 20.8 [19.1, 22.7] |
| 1 | 65,536 | 256 | 91.6 | 100.0 | 93.5 | 98.7 [98.1, 99.1] |
| (a) Matched expected count per block: | ||||
| Successes | Gain (%) | Successes | Gain (%) | |
| 2 | 361 | 410 | ||
| 8 | 0 | 1 | ||
| 32 | 0 | 0 | ||
| 128 | 0 | 0 | ||
| CIFAR-10 / ResNet-18 | SST-2 / DistilBERT | |
|---|---|---|
| Reference inputs | 50,000 training inputs | 67,349 training inputs |
| Query split (accuracy) | Test (95.34%) | Validation (91.055%) |
| Selected errors | 32 of 466 | 32 of 78 |
| Feature dimension / free head parameters | 513 / 4,617 | 769 / 769 |
| Contrast (budget) | CIFAR-10 | SST-2 |
|---|---|---|
| Reference minus empirical calibration ( ) | [ ] | [ ] |
| Reference direction minus estimated direction ( ) | 2.32 [ ] | 0.13 [ ] |
| Estimated direction minus gradient direction ( ) | 156.97 [71.08,249.32] | 1717.06 [1401.41,2021.48] |
| Reference direction minus estimated direction ( ) | 4.76 [1.10,8.64] | 76.58 [52.35,101.32] |
| Joint success (%) [95% CI] | |||
|---|---|---|---|
| Dataset | Full budget | Margin | Margin gain (%) |
| CIFAR-10 | 48.59 | 99.53 | 99.63 |
| SST-2 | 50.94 | 93.75 | 99.79 |