Prune and schedule trials¶
Goal: stop trials that are clearly losing, and spend the saved budget on ones that are not.
Most tuning budget goes on configurations that were already behind after a tenth of their work. A scheduler reads a trial's intermediate reports and decides, while it is still running, whether to let it continue. atune calls the decision points reports and the deciders schedulers; Optuna calls the same idea pruning.
You need two things: an objective that reports as it goes, and a scheduler attached to the study.
Report intermediate values¶
An objective that never reports cannot be pruned — there is nothing for the scheduler to look at. Reporting is one call per step against the curve being trained:
/// The loss of configuration `x` after `step` resource units.
///
/// `x + LATE_PENALTY / step` — decaying towards `x`, so more resource is always
/// better and a pruned trial records a *worse* number than a survivor (see the
/// module header). `f64::from` on a `u32` is a lossless widening, matching the
/// Python arm's `float(step)`.
fn learning_curve(x: f64, step: u32) -> f64 {
x + LATE_PENALTY / f64::from(step)
}
def learning_curve(x: float, step: int) -> float:
"""The loss of configuration ``x`` after ``step`` resource units.
``x + LATE_PENALTY / step`` — decaying towards ``x``, so more resource is
always better and a pruned trial records a *worse* number than a survivor
(see the module docstring). ``float(step)`` is a lossless widening, matching
the Rust arm's ``f64::from(step)``.
"""
return x + LATE_PENALTY / float(step)
def objective(trial: atune.Trial) -> float:
"""Trains one configuration, reporting its loss at every step."""
x = trial.suggest_float("x", LOW, HIGH)
for step in range(MIN_RESOURCE, MAX_RESOURCE + 1):
# `report` raises `atune.Pruned` when the scheduler prunes; letting it
# propagate is the idiomatic path. The trial is then recorded `pruned`
# and keeps the intermediate it last reported — which is why the curve
# has to *fall* (see the module docstring).
trial.report(learning_curve(x, step), step)
# A trial that survives the ladder is worth its converged loss.
return learning_curve(x, MAX_RESOURCE)
In Rust the call is ctx.report(step, &values)? — it sits inside the study
listing below — and when the scheduler says stop, that call returns the
pruned sentinel that ?
propagates. In Python the same moment is trial.report(value, step) raising
atune.Pruned, as the listing above shows. Either way the objective returns
early and the study records the trial as pruned with its last intermediate
value as its objective value. That last detail matters more than it looks —
see the curve.
The step you report at is a resource: an epoch, a batch, a fold, a simulated second. It must increase, and it must mean the same thing on every trial, because comparing it across trials is the whole mechanism.
Attach a scheduler¶
let study = Study::builder()
.parallelism(THREADS)
.budget(Budget::trials(TRIALS))
// Successive halving over a 1..27 ladder: rungs at 1, 3 and 9.
.scheduler(Arc::new(AshaPruner::new(
u64::from(MIN_RESOURCE),
u64::from(MAX_RESOURCE),
REDUCTION_FACTOR,
)?))
.sampler(Arc::new(Tpe::new()))
.create(StudyConfig::new(STUDY).with_seed(SEED))?;
study.optimize(|ctx| {
let x = ctx.suggest_f64("x", LOW..=HIGH, Scale::Linear)?;
for step in MIN_RESOURCE..=MAX_RESOURCE {
// `report` returns `Error::TrialPruned` when the scheduler prunes;
// letting `?` propagate it is the idiomatic path. The trial is then
// recorded `Pruned` and keeps the intermediate it last reported —
// which is why the curve has to *fall* (see the module header).
ctx.report(u64::from(step), &[learning_curve(x, step)])?;
}
// A trial that survives the ladder is worth its converged loss.
Ok(learning_curve(x, MAX_RESOURCE).into())
})?;
study = atune.create_study(
direction="minimize",
# Successive halving over a 1..27 ladder: rungs at 1, 3 and 9.
scheduler=atune.schedulers.Asha(
MIN_RESOURCE,
MAX_RESOURCE,
reduction_factor=REDUCTION_FACTOR,
),
sampler=atune.samplers.Tpe(),
seed=SEED,
name=STUDY,
)
study.optimize(objective, n_trials=TRIALS, n_jobs=JOBS)
AshaPruner::new(1, 27, 3) is asynchronous successive halving over a resource
ladder from 1 to 27 with reduction factor 3: at each rung a trial must be in the
top third of the trials that have reached that rung, or it stops. Asynchronous
means a trial is judged against whoever has already arrived, so no worker waits
for a cohort to fill.
The other schedulers differ in what they compare, not in how they attach:
| Scheduler | Stops a trial when |
|---|---|
MedianPruner |
it is worse than the running median at the same step |
AshaPruner |
it falls outside the top fraction at a rung of the ladder |
HyperbandPruner |
it loses within a bracket, across several ladder configurations |
WilcoxonPruner |
a signed-rank test says it is behind — for noisy objectives |
Patient |
it has not improved for a given number of reports |
Pbt and FreezeThaw implement the same seam but do something else with it: they
fork and revive trials rather than only stopping them. See
population-based training.
The curve must improve with the resource¶
This is the one thing that will silently ruin a pruning setup, and it is worth stating plainly because this project got it wrong first.
A pruned trial's objective value is its last intermediate report. So if the metric you report gets worse as the resource grows — a loss reported as "accumulated so far", a score reported as a running total — then a trial killed at the first rung records a small number while one that ran to the end records a large one. The pruned trials look like the winners, and any sampler that learns from the study's history is being told to chase exactly what the pruner rejected.
Measured here, on an earlier version of the example above: with a rising curve the best trial in the study was one pruned at step 1, and TPE — fed those records — converged on the far end of the range from the true optimum. The fix is not a special case in the pruner. Report a metric that improves with the resource: a loss that decays, an accuracy that climbs. If your natural metric does not, report its best-so-far instead.
Check what actually happened¶
Pruning is easy to configure and easy to configure into a no-op.
let view = study.view()?;
let total = view.len();
let pruned = view
.all()
.filter(|trial| trial.state == TrialState::Pruned)
.count();
let complete = view
.all()
.filter(|trial| trial.state == TrialState::Complete)
.count();
states = [trial.state for trial in study.trials]
n_pruned = sum(1 for state in states if state == "pruned")
n_complete = sum(1 for state in states if state == "complete")
The example prints those counts in its human-readable summary (the
trials : … complete, … pruned line), and that count is what tells you
whether the pruner ran at all.
Zero pruned trials means the scheduler never fired — usually a ladder no trial
reached, or reports at steps the ladder does not contain. Almost everything
pruned usually means a ladder too aggressive for how noisy the objective is,
and WilcoxonPruner is the answer to that rather than a gentler ladder.
From the command line, atune trials shows each trial's state, so pruned and
complete trials are distinguishable without writing code — see the
CLI reference.
What this costs you¶
Pruning trades a correct answer for a cheaper one, and the trade is real.
- A pruned trial is a guess. Successive halving assumes the ranking at a small resource predicts the ranking at a large one. When that fails — and it fails for configurations that need a warm-up before they look good — the pruner discards the eventual winner.
- Pruning interacts with the sampler. Every prune changes the history the sampler fits its model to, so the same seed with a different ladder is a different study. Expected, not a bug — but it makes the ladder part of your configuration, and it belongs in whatever you record about a run.
- Parallel pruning is history-dependent. Which trials have reached a rung depends on the interleaving, so a study with a pruner and several workers is not reproducible trial-for-trial the way a stateless-sampler study is. Determinism says exactly what still holds.
Next¶
- Samplers and schedulers — the seam, and how the two interact.
- Population-based training — a scheduler that forks winners instead of stopping losers.
- Error reference —
TrialPrunedandTrialPausedare?-able sentinels, not failures.