43 Commits

Author SHA1 Message Date
Manuel Raimann 147a64708d chore: release v0.7.0 2026-02-11 19:05:52 +01:00
Manuel Raimann b36137bc96 feat: add benchmark suite with criterion and standard test functions
Add micro-benchmarks for TPE, Random, and Grid samplers with varying
history sizes, end-to-end optimization benchmarks on Sphere and
Rosenbrock, and a convergence tracking example that outputs CSV data.
Includes 6 standard test functions (Sphere, Rosenbrock, Rastrigin,
Ackley, Branin, Hartmann6) with correctness tests at known optima.
2026-02-11 19:05:33 +01:00
Manuel Raimann f3941a4b3e chore: remove unused proc-macro2 dependency from optimizer-derive 2026-02-11 18:57:30 +01:00
Manuel Raimann 0cae3b4227 feat: add parameter importance via Spearman rank correlation
Compute per-parameter importance scores by measuring the absolute
Spearman rank correlation between each parameter's values and the
objective across completed trials. Scores are normalized to sum to 1.0
and returned sorted by descending importance.
2026-02-11 18:56:19 +01:00
Manuel Raimann 09480a973b feat: add constraint handling with feasibility-aware trial ranking
Add constraint support so that optimization problems with constraints
(e.g., "model size < 100MB") prefer feasible solutions. Convention:
constraint value <= 0.0 means feasible.

- Add constraint_values field to Trial with set_constraints/getter
- Add constraints field to CompletedTrial with is_feasible() method
- Propagate constraints through complete_trial/prune_trial
- Make best_trial() and top_trials() constraint-aware: feasible trials
  rank above infeasible, infeasible ranked by total violation
2026-02-11 18:49:08 +01:00
Manuel Raimann aa9734587a feat: add BOHB sampler for budget-aware Bayesian optimization
BOHB combines TPE's model-guided sampling with Hyperband's budget-aware
evaluation by conditioning the TPE model on trials at specific budget
levels, producing better-calibrated parameter proposals.
2026-02-11 18:42:09 +01:00
Manuel Raimann a1dd6d393c feat: add CMA-ES sampler behind cma-es feature flag
Implement CMA-ES (Covariance Matrix Adaptation Evolution Strategy) as a
new sampler for continuous optimization. Uses nalgebra for matrix
operations and implements the full Hansen 2016 algorithm: mean update,
evolution paths, covariance matrix adaptation, and cumulative step-size
adaptation.

Key features:
- Builder API with configurable sigma0, population_size, and seed
- Three-phase state machine: discovery, active sampling, generation update
- Rejection sampling with bounds clipping fallback
- Log-scale and step parameter support via internal-space mapping
- Categorical parameters handled via random sampling fallback
- Numerical robustness: eigenvalue clamping, symmetry enforcement
2026-02-11 18:31:12 +01:00
Manuel Raimann a53ba4f472 feat: add Sobol quasi-random sampler behind sobol feature flag
Add SobolSampler using Owen-scrambled Sobol sequences (sobol_burley
crate) for better uniform coverage of the parameter space compared to
random sampling. Supports all distribution types including log-scale
and stepped parameters.
2026-02-11 17:59:48 +01:00
Manuel Raimann 9b4a8321b5 feat: implement IntoIterator for &Study and add iter() method
Enables idiomatic `for trial in &study` iteration over completed trials.
2026-02-11 17:54:09 +01:00
Manuel Raimann db0314a1d1 feat: add optimize_with_retries() for automatic retry of failed trials 2026-02-11 17:51:36 +01:00
Manuel Raimann df5fcedecf feat: add tracing integration behind tracing feature flag
Add optional structured logging via the tracing crate:
- info-level spans on all optimize* methods (direction, n_trials/duration)
- info events for trial completed, new best value, trial pruned
- debug events for trial failed and parameter sampled
- Zero overhead when the feature is disabled
2026-02-11 17:47:13 +01:00
Manuel Raimann 5061df9876 feat: add optimize_with_checkpoint() and atomic save writes
Adds a convenience method that combines optimize_with_callback + save
to automatically checkpoint every N trials. Also makes save() use
atomic writes (write to temp file, then rename) to prevent corrupt
files on crash.
2026-02-11 17:37:57 +01:00
Manuel Raimann 5fe0a75f78 feat: add serde serialization support behind serde feature flag
Add Serialize/Deserialize derives to all public data types (ParamValue,
Distribution, Direction, TrialState, ParamId, AttrValue, CompletedTrial)
gated behind a `serde` feature flag. Introduce StudySnapshot struct and
Study::save()/Study::load() for persisting and restoring study state as
human-readable JSON.
2026-02-11 17:33:48 +01:00
Manuel Raimann dcf3403261 feat: add Study::summary() and Display impl
Add a human-readable summary method and Display trait for Study<V>
where V: Display. The summary shows optimization direction, trial
counts (with complete/pruned breakdown), best value, and best
parameters with their labels.

Also derive PartialOrd + Ord on ParamId for deterministic parameter
ordering in summary output.
2026-02-11 17:27:43 +01:00
Manuel Raimann bb79d98027 chore: release v0.6.0 2026-02-11 17:24:10 +01:00
Manuel Raimann f4b0631178 feat: add enqueue_trial() for pre-specified parameter evaluation
Add Study::enqueue() to push specific parameter configurations onto a
FIFO queue. The next call to ask()/create_trial() or the next iteration
in optimize() pops the front entry and injects it into the trial so
that suggest_param() returns the pre-filled value instead of sampling.

Parameters missing from an enqueued map fall back to normal sampling,
and once the queue is drained regular sampler-driven trials resume.
2026-02-11 17:23:25 +01:00
Manuel Raimann 52f3c074dc feat: add trial user attributes for logging and analysis
Add AttrValue enum (Float, Int, String, Bool) with From impls and
set_user_attr/user_attr/user_attrs methods on both Trial and
CompletedTrial. Attributes set during optimization are propagated
through complete_trial and prune_trial. AttrValue is re-exported
at crate root and in the prelude.
2026-02-11 17:18:17 +01:00
Manuel Raimann f67188e3a7 feat: add ask() and tell() methods for ask-and-tell interface 2026-02-11 17:11:38 +01:00
Manuel Raimann 72da883add feat: add Study::top_trials(n) for retrieving best N trials
Returns the top N completed trials sorted by objective value,
respecting the study's optimization direction. Pruned and failed
trials are excluded.
2026-02-11 17:08:10 +01:00
Manuel Raimann bc278d83a3 feat: add timeout-based optimization methods (optimize_until)
Add duration-based variants of all optimization methods that run trials
until a wall-clock deadline rather than for a fixed trial count:

- optimize_until(duration, objective)
- optimize_until_with_callback(duration, objective, callback)
- optimize_until_async(duration, objective)
- optimize_until_parallel(duration, concurrency, objective)
2026-02-11 17:06:08 +01:00
Manuel Raimann c586650df4 feat: add From<RangeInclusive> for FloatParam and IntParam 2026-02-11 17:02:18 +01:00
Manuel Raimann cc78a92ed6 feat: add Study::minimize() and Study::maximize() constructor shortcuts 2026-02-11 17:00:35 +01:00
Manuel Raimann 12c35e7cb4 feat: add WilcoxonPruner for statistics-based trial pruning 2026-02-11 16:59:05 +01:00
Manuel Raimann a885d10436 feat: add HyperbandPruner for multi-bracket trial pruning 2026-02-11 16:52:59 +01:00
Manuel Raimann 41150057d6 feat: add SuccessiveHalvingPruner for SHA-based trial pruning 2026-02-11 16:47:57 +01:00
Manuel Raimann a522fb2f34 feat: add PatientPruner for patience-based trial pruning
Wraps any inner Pruner and requires it to recommend pruning for N
consecutive steps before actually pruning. Useful for noisy
intermediate values where a single bad step shouldn't trigger pruning.
2026-02-11 16:40:00 +01:00
Manuel Raimann ec04757740 feat: add PercentilePruner for configurable percentile-based trial pruning
Generalizes MedianPruner with a configurable percentile threshold.
MedianPruner now reuses the shared compute_percentile function at 50%.
2026-02-11 16:36:57 +01:00
Manuel Raimann 432c74b927 feat: add MedianPruner for statistics-based trial pruning
Prunes trials whose intermediate values are worse than the median
of completed trials at the same step. Supports configurable warmup
steps and minimum trial count before pruning activates.
2026-02-11 16:32:03 +01:00
Manuel Raimann 8199ba6b43 feat: add ThresholdPruner for fixed-bound trial pruning 2026-02-11 16:28:22 +01:00
Manuel Raimann abe15bdd7b feat: add TrialPruned error variant and Pruned trial state
- Add `Pruned` variant to `TrialState`
- Add `Error::TrialPruned` variant and standalone `TrialPruned` struct
  with `From<TrialPruned> for Error` for ergonomic `?` usage
- Add `state` field to `CompletedTrial` (defaults to `Complete`)
- Add `Study::prune_trial()` and `Study::n_pruned_trials()`
- `optimize()` and `optimize_with_callback()` detect `TrialPruned`
  errors via Any downcasting and record pruned trials instead of
  failing them
- `best_trial()` / `best_value()` now filter to only `Complete` trials
- Re-export `TrialPruned` from crate root and prelude
2026-02-11 16:24:57 +01:00
Manuel Raimann 4d8af3242b feat: add intermediate values and pruner integration to Trial
- Add report(step, value) and should_prune() methods to Trial
- Store intermediate_values in Trial, copied to CompletedTrial on completion
- Pass pruner through trial factory so trials can consult it
- Rebuild trial factory when pruner is changed via set_pruner()
2026-02-11 16:08:25 +01:00
Manuel Raimann 2651e61b2c feat: add Pruner trait and NopPruner default implementation
Introduce the pruning extension point for the trial pruning system.
The Pruner trait allows deciding whether to stop a trial early based
on intermediate values. NopPruner (never prunes) is the default.
2026-02-11 15:46:56 +01:00
Manuel Raimann d88280c152 chore: release v0.5.1 2026-02-10 09:30:58 +01:00
Manuel Raimann 0e8fc82201 refactor: update random number generation to use rand::make_rng() and upgrade rand to version 0.10 2026-02-10 09:30:39 +01:00
Manuel Raimann d7ad3f60db chore: release v0.5.0 2026-02-06 18:55:06 +01:00
Manuel Raimann 5dc81fa0ab Implement Parameters API
- Add `.name()` builder method on all 5 parameter types for custom labels
- Add `CompletedTrial::get(&param)` for typed parameter access
- Add `Display` impl on `ParamValue`
- Add prelude module at `optimizer::prelude::*`
- Shadow `_with_sampler` methods on `Study<f64>` so `optimize()` auto-uses
  the configured sampler; deprecate `_with_sampler` variants
- Use runtime `Any` downcasting with `trial_factory` to avoid E0592
- Update all examples and tests to use the new API

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-06 18:54:55 +01:00
Manuel Raimann 55e50e6afb Implement Parameters API 2026-02-06 17:15:47 +01:00
Manuel Raimann 2fef1ab3c3 docs: add missing derive feature flag to README
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-06 16:51:13 +01:00
Manuel Raimann aa1718fed7 docs: update quick start example to use typed Parameter API
The old example used the removed string-based trial.suggest_float()
API. Updated to use FloatParam::new() with param.suggest(trial).

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-06 16:50:59 +01:00
Manuel Raimann f9227df96e docs: update features list with boolean, enum, and derive support
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-06 16:50:43 +01:00
Manuel Raimann 02ec9498c4 refactor: reorganize sampler module and update imports 2026-02-04 13:14:22 +01:00
Manuel Raimann 183a78c09e chore: remove log dependency from Cargo.toml 2026-02-02 17:33:03 +01:00
Manuel Raimann ce817c1f38 feat: add 'consts' to default extended words in typos configuration 2026-02-02 17:32:20 +01:00
54 changed files with 12197 additions and 2459 deletions
+32 -3
View File
@@ -1,6 +1,9 @@
[workspace]
members = ["optimizer-derive"]
[package]
name = "optimizer"
version = "0.4.0"
version = "0.7.0"
edition = "2024"
rust-version = "1.88"
license = "MIT"
@@ -13,18 +16,39 @@ categories = ["algorithm", "science", "data-structures"]
readme = "README.md"
[dependencies]
log = "0.4"
rand = "0.9"
rand = "0.10"
thiserror = "2"
parking_lot = "0.12"
tokio = { version = "1", features = ["sync", "rt-multi-thread"], optional = true }
optimizer-derive = { version = "0.1.0", path = "optimizer-derive", optional = true }
serde = { version = "1", features = ["derive"], optional = true }
serde_json = { version = "1", optional = true }
tracing = { version = "0.1", optional = true }
sobol_burley = { version = "0.5", optional = true }
nalgebra = { version = "0.33", optional = true }
[features]
default = []
async = ["dep:tokio"]
derive = ["dep:optimizer-derive"]
serde = ["dep:serde", "dep:serde_json"]
tracing = ["dep:tracing"]
sobol = ["dep:sobol_burley"]
cma-es = ["dep:nalgebra"]
[dev-dependencies]
tokio = { version = "1", features = ["rt-multi-thread", "macros", "time"] }
optimizer-derive = { version = "0.1.0", path = "optimizer-derive" }
serde_json = "1"
criterion = { version = "0.8", features = ["html_reports"] }
[[bench]]
name = "samplers"
harness = false
[[bench]]
name = "optimization"
harness = false
[[example]]
name = "async_api_optimization"
@@ -34,3 +58,8 @@ required-features = ["async"]
[[example]]
name = "ml_hyperparameter_tuning"
path = "examples/ml_hyperparameter_tuning.rs"
[[example]]
name = "parameter_api"
path = "examples/parameter_api.rs"
required-features = ["derive"]
+16 -4
View File
@@ -13,28 +13,39 @@ A Rust library for black-box optimization with multiple sampling strategies.
- **Random Search** - Simple random sampling for baseline comparisons
- **TPE (Tree-Parzen Estimator)** - Bayesian optimization for efficient search
- **Grid Search** - Exhaustive search over a specified parameter grid
- Float, integer, and categorical parameter types
- Float, integer, categorical, boolean, and enum parameter types
- Log-scale and stepped parameter sampling
- Sync and async optimization with parallel trial evaluation
- `#[derive(Categorical)]` for enum parameters
## Quick Start
```rust
use optimizer::{Direction, Study};
use optimizer::parameter::{FloatParam, Parameter};
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, Study};
// Create a study with TPE sampler
let sampler = TpeSampler::builder().seed(42).build().unwrap();
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
// Define parameter search space
let x_param = FloatParam::new(-10.0, 10.0);
// Optimize x^2 for 20 trials
study
.optimize_with_sampler(20, |trial| {
let x = trial.suggest_float("x", -10.0, 10.0)?;
let x = x_param.suggest(trial)?;
Ok::<_, optimizer::Error>(x * x)
})
.unwrap();
// Get the best result
let best = study.best_trial().unwrap();
println!("Best value: {} at x={:?}", best.value, best.params);
println!("Best value: {}", best.value);
for (id, label) in &best.param_labels {
println!(" {}: {:?}", label, best.params[id]);
}
```
## Samplers
@@ -134,6 +145,7 @@ let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
## Feature Flags
- `async` - Enable async optimization methods (requires tokio)
- `derive` - Enable `#[derive(Categorical)]` for enum parameters
## Documentation
+112
View File
@@ -0,0 +1,112 @@
#[allow(dead_code)]
mod test_functions;
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
use optimizer::parameter::Parameter;
use optimizer::sampler::random::RandomSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{FloatParam, Study};
fn make_params(dims: usize) -> Vec<FloatParam> {
(0..dims)
.map(|i| FloatParam::new(-5.0, 5.0).name(format!("x{i}")))
.collect()
}
fn bench_tpe_sphere(c: &mut Criterion) {
let mut group = c.benchmark_group("tpe_sphere");
group.sample_size(10);
for dims in [2, 10, 50] {
let params = make_params(dims);
group.bench_with_input(BenchmarkId::new("dims", dims), &params, |b, params| {
b.iter(|| {
let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap());
study
.optimize(100, |trial| {
let x: Vec<f64> = params
.iter()
.map(|p| p.suggest(trial))
.collect::<Result<_, _>>()
.unwrap();
Ok::<_, optimizer::Error>(test_functions::sphere(&x))
})
.unwrap();
});
});
}
group.finish();
}
fn bench_tpe_rosenbrock(c: &mut Criterion) {
let mut group = c.benchmark_group("tpe_rosenbrock");
group.sample_size(10);
for dims in [2, 10] {
let params = make_params(dims);
group.bench_with_input(BenchmarkId::new("dims", dims), &params, |b, params| {
b.iter(|| {
let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap());
study
.optimize(100, |trial| {
let x: Vec<f64> = params
.iter()
.map(|p| p.suggest(trial))
.collect::<Result<_, _>>()
.unwrap();
Ok::<_, optimizer::Error>(test_functions::rosenbrock(&x))
})
.unwrap();
});
});
}
group.finish();
}
fn bench_random_vs_tpe(c: &mut Criterion) {
let mut group = c.benchmark_group("random_vs_tpe");
group.sample_size(10);
let params = make_params(5);
group.bench_function("random_5d", |b| {
b.iter(|| {
let study = Study::minimize(RandomSampler::with_seed(42));
study
.optimize(100, |trial| {
let x: Vec<f64> = params
.iter()
.map(|p| p.suggest(trial))
.collect::<Result<_, _>>()
.unwrap();
Ok::<_, optimizer::Error>(test_functions::sphere(&x))
})
.unwrap();
});
});
group.bench_function("tpe_5d", |b| {
b.iter(|| {
let study = Study::minimize(TpeSampler::builder().seed(42).build().unwrap());
study
.optimize(100, |trial| {
let x: Vec<f64> = params
.iter()
.map(|p| p.suggest(trial))
.collect::<Result<_, _>>()
.unwrap();
Ok::<_, optimizer::Error>(test_functions::sphere(&x))
})
.unwrap();
});
});
group.finish();
}
criterion_group!(
benches,
bench_tpe_sphere,
bench_tpe_rosenbrock,
bench_random_vs_tpe
);
criterion_main!(benches);
+118
View File
@@ -0,0 +1,118 @@
use std::collections::HashMap;
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
use optimizer::parameter::Parameter;
use optimizer::sampler::Sampler;
use optimizer::sampler::grid::GridSearchSampler;
use optimizer::sampler::random::RandomSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{CompletedTrial, FloatParam};
/// Build a synthetic history of `n` completed trials over `dims` float parameters.
fn build_history(n: usize, dims: usize) -> Vec<CompletedTrial<f64>> {
let params: Vec<FloatParam> = (0..dims)
.map(|i| FloatParam::new(-5.0, 5.0).name(format!("x{i}")))
.collect();
let mut history = Vec::with_capacity(n);
let sampler = RandomSampler::with_seed(42);
for trial_id in 0..n {
let id = trial_id as u64;
let mut param_values = HashMap::new();
let mut distributions = HashMap::new();
let mut param_labels = HashMap::new();
for p in &params {
let dist = p.distribution();
let val = sampler.sample(&dist, id, &history);
param_values.insert(p.id(), val);
distributions.insert(p.id(), dist);
param_labels.insert(p.id(), p.label());
}
// Use sphere function value as objective
let value: f64 = param_values
.values()
.map(|v| {
let optimizer::ParamValue::Float(f) = v else {
unreachable!()
};
f * f
})
.sum();
history.push(CompletedTrial::new(
id,
param_values,
distributions,
param_labels,
value,
));
}
history
}
fn bench_tpe_sample(c: &mut Criterion) {
let mut group = c.benchmark_group("tpe_sample");
let dist = FloatParam::new(-5.0, 5.0).distribution();
let tpe = TpeSampler::builder().seed(42).build().unwrap();
for history_size in [10, 100, 1000] {
let history = build_history(history_size, 2);
group.bench_with_input(
BenchmarkId::new("history", history_size),
&history,
|b, history| {
b.iter(|| tpe.sample(&dist, history.len() as u64, history));
},
);
}
group.finish();
}
fn bench_random_sample(c: &mut Criterion) {
let mut group = c.benchmark_group("random_sample");
let dist = FloatParam::new(-5.0, 5.0).distribution();
let sampler = RandomSampler::with_seed(42);
for history_size in [10, 100, 1000] {
let history = build_history(history_size, 2);
group.bench_with_input(
BenchmarkId::new("history", history_size),
&history,
|b, history| {
b.iter(|| sampler.sample(&dist, history.len() as u64, history));
},
);
}
group.finish();
}
fn bench_grid_sample(c: &mut Criterion) {
let mut group = c.benchmark_group("grid_sample");
let dist = FloatParam::new(-5.0, 5.0).distribution();
let history: Vec<CompletedTrial<f64>> = Vec::new();
for grid_points in [5, 10, 50] {
group.bench_with_input(
BenchmarkId::new("points", grid_points),
&grid_points,
|b, _| {
b.iter(|| {
// Fresh sampler each iteration since grid tracks used points
let sampler = GridSearchSampler::builder()
.n_points_per_param(grid_points)
.build();
sampler.sample(&dist, 0, &history)
});
},
);
}
group.finish();
}
criterion_group!(
benches,
bench_tpe_sample,
bench_random_sample,
bench_grid_sample
);
criterion_main!(benches);
+84
View File
@@ -0,0 +1,84 @@
//! Standard optimization test functions for benchmarking.
/// Sphere function: unimodal, convex. Global minimum f(0,...,0) = 0.
pub fn sphere(x: &[f64]) -> f64 {
x.iter().map(|xi| xi * xi).sum()
}
/// Rosenbrock function: narrow valley. Global minimum f(1,...,1) = 0.
pub fn rosenbrock(x: &[f64]) -> f64 {
x.windows(2)
.map(|w| 100.0 * (w[1] - w[0] * w[0]).powi(2) + (1.0 - w[0]).powi(2))
.sum()
}
/// Rastrigin function: highly multimodal. Global minimum f(0,...,0) = 0.
pub fn rastrigin(x: &[f64]) -> f64 {
let n = x.len() as f64;
10.0 * n
+ x.iter()
.map(|xi| xi * xi - 10.0 * (2.0 * std::f64::consts::PI * xi).cos())
.sum::<f64>()
}
/// Ackley function: nearly flat with a deep well. Global minimum f(0,...,0) = 0.
pub fn ackley(x: &[f64]) -> f64 {
let n = x.len() as f64;
let sum_sq: f64 = x.iter().map(|xi| xi * xi).sum();
let sum_cos: f64 = x
.iter()
.map(|xi| (2.0 * std::f64::consts::PI * xi).cos())
.sum();
-20.0 * (-0.2 * (sum_sq / n).sqrt()).exp() - (sum_cos / n).exp() + 20.0 + std::f64::consts::E
}
/// Branin function (2D only). Three global minima with f* ≈ 0.397887.
///
/// # Panics
///
/// Panics if `x` does not have exactly 2 elements.
pub fn branin(x: &[f64]) -> f64 {
assert!(x.len() == 2, "Branin requires exactly 2 dimensions");
let (x1, x2) = (x[0], x[1]);
let pi = std::f64::consts::PI;
let a = 1.0;
let b = 5.1 / (4.0 * pi * pi);
let c = 5.0 / pi;
let r = 6.0;
let s = 10.0;
let t = 1.0 / (8.0 * pi);
a * (x2 - b * x1 * x1 + c * x1 - r).powi(2) + s * (1.0 - t) * x1.cos() + s
}
/// Hartmann 6D function. Global minimum f* ≈ -3.3224.
///
/// # Panics
///
/// Panics if `x` does not have exactly 6 elements.
pub fn hartmann6(x: &[f64]) -> f64 {
assert!(x.len() == 6, "Hartmann6 requires exactly 6 dimensions");
let alpha = [1.0, 1.2, 3.0, 3.2];
let a_matrix = [
[10.0, 3.0, 17.0, 3.5, 1.7, 8.0],
[0.05, 10.0, 17.0, 0.1, 8.0, 14.0],
[3.0, 3.5, 1.7, 10.0, 17.0, 8.0],
[17.0, 8.0, 0.05, 10.0, 0.1, 14.0],
];
let p_matrix = [
[0.1312, 0.1696, 0.5569, 0.0124, 0.8283, 0.5886],
[0.2329, 0.4135, 0.8307, 0.3736, 0.1004, 0.9991],
[0.2348, 0.1451, 0.3522, 0.2883, 0.3047, 0.6650],
[0.4047, 0.8828, 0.8732, 0.5743, 0.1091, 0.0381],
];
let mut result = 0.0;
for i in 0..4 {
let mut inner = 0.0;
for (j, xj) in x.iter().enumerate() {
inner += a_matrix[i][j] * (xj - p_matrix[i][j]).powi(2);
}
result -= alpha[i] * (-inner).exp();
}
result
}
+125 -74
View File
@@ -6,7 +6,7 @@
//!
//! # Key Concepts Demonstrated
//!
//! - Async optimization with `optimize_parallel_with_sampler`
//! - Async optimization with `optimize_parallel`
//! - Running multiple trials concurrently for faster optimization
//! - Boolean and categorical parameter types
//! - Measuring speedup from parallelism
@@ -26,8 +26,7 @@
use std::time::{Duration, Instant};
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, ParamValue, Study, Trial};
use optimizer::prelude::*;
// ============================================================================
// Configuration: Service parameters we want to tune
@@ -133,32 +132,27 @@ async fn evaluate_service(config: &ServiceConfig) -> f64 {
/// 3. Return both the Trial and the result value as a tuple
///
/// This ownership pattern allows the trial to be used across await points.
async fn objective(mut trial: Trial) -> optimizer::Result<(Trial, f64)> {
// Sample configuration parameters
// Stepped integers: only sample multiples of the step value
let cache_size_mb = trial.suggest_int_step("cache_size_mb", 64, 1024, 64)?;
let connection_pool_size = trial.suggest_int_step("connection_pool_size", 10, 200, 10)?;
let request_timeout_ms = trial.suggest_int_step("request_timeout_ms", 1000, 10000, 500)?;
// Regular integer
let retry_count = trial.suggest_int("retry_count", 0, 5)?;
// Log-scale integer: good for parameters like batch sizes
// that might vary from 1 to 256
let batch_size = trial.suggest_int_log("batch_size", 1, 256)?;
// Regular integer for compression level
let compression_level = trial.suggest_int("compression_level", 0, 9)?;
// Boolean: internally uses categorical with [false, true]
let use_http2 = trial.suggest_bool("use_http2")?;
// Categorical: choose from a list of options
let load_balancing = trial.suggest_categorical(
"load_balancing",
&["round_robin", "least_connections", "random", "ip_hash"],
)?;
#[allow(clippy::too_many_arguments)]
async fn objective(
mut trial: Trial,
cache_size_mb_param: &IntParam,
connection_pool_size_param: &IntParam,
request_timeout_ms_param: &IntParam,
retry_count_param: &IntParam,
batch_size_param: &IntParam,
compression_level_param: &IntParam,
use_http2_param: &BoolParam,
load_balancing_param: &CategoricalParam<&str>,
) -> optimizer::Result<(Trial, f64)> {
// Sample configuration parameters using parameter definitions
let cache_size_mb = cache_size_mb_param.suggest(&mut trial)?;
let connection_pool_size = connection_pool_size_param.suggest(&mut trial)?;
let request_timeout_ms = request_timeout_ms_param.suggest(&mut trial)?;
let retry_count = retry_count_param.suggest(&mut trial)?;
let batch_size = batch_size_param.suggest(&mut trial)?;
let compression_level = compression_level_param.suggest(&mut trial)?;
let use_http2 = use_http2_param.suggest(&mut trial)?;
let load_balancing = load_balancing_param.suggest(&mut trial)?;
// Build configuration
let config = ServiceConfig {
@@ -183,24 +177,6 @@ async fn objective(mut trial: Trial) -> optimizer::Result<(Trial, f64)> {
// Helper Functions
// ============================================================================
/// Formats a parameter value for display.
fn format_param(name: &str, value: &ParamValue) -> String {
match (name, value) {
(_, ParamValue::Float(v)) => format!("{v:.4}"),
(_, ParamValue::Int(v)) => format!("{v}"),
("use_http2", ParamValue::Categorical(idx)) => {
if *idx == 1 { "true" } else { "false" }.to_string()
}
("load_balancing", ParamValue::Categorical(idx)) => {
["round_robin", "least_connections", "random", "ip_hash"]
.get(*idx)
.unwrap_or(&"unknown")
.to_string()
}
(_, ParamValue::Categorical(idx)) => format!("category_{idx}"),
}
}
/// Prints the results of the optimization.
fn print_results(study: &Study<f64>, elapsed: Duration, n_trials: usize) {
println!("\n{}", "=".repeat(60));
@@ -219,31 +195,46 @@ fn print_results(study: &Study<f64>, elapsed: Duration, n_trials: usize) {
}
/// Prints the best configuration found.
fn print_best_config(study: &Study<f64>) -> optimizer::Result<()> {
#[allow(clippy::too_many_arguments)]
fn print_best_config(
study: &Study<f64>,
cache_size_mb_param: &IntParam,
connection_pool_size_param: &IntParam,
request_timeout_ms_param: &IntParam,
retry_count_param: &IntParam,
batch_size_param: &IntParam,
compression_level_param: &IntParam,
use_http2_param: &BoolParam,
load_balancing_param: &CategoricalParam<&str>,
) -> optimizer::Result<()> {
let best = study.best_trial()?;
println!("\nBest configuration found:");
println!(" Score: {:.6}", best.value);
println!("\n Parameters:");
// Print parameters in a logical order
let param_order = [
"cache_size_mb",
"connection_pool_size",
"request_timeout_ms",
"retry_count",
"batch_size",
"compression_level",
"use_http2",
"load_balancing",
];
for name in param_order {
if let Some(value) = best.params.get(name) {
let display = format_param(name, value);
println!(" {name}: {display}");
}
}
println!(
" cache_size_mb: {}",
best.get(cache_size_mb_param).unwrap()
);
println!(
" connection_pool_size: {}",
best.get(connection_pool_size_param).unwrap()
);
println!(
" request_timeout_ms: {}",
best.get(request_timeout_ms_param).unwrap()
);
println!(" retry_count: {}", best.get(retry_count_param).unwrap());
println!(" batch_size: {}", best.get(batch_size_param).unwrap());
println!(
" compression_level: {}",
best.get(compression_level_param).unwrap()
);
println!(" use_http2: {}", best.get(use_http2_param).unwrap());
println!(
" load_balancing: {}",
best.get(load_balancing_param).unwrap()
);
Ok(())
}
@@ -284,7 +275,35 @@ async fn main() -> optimizer::Result<()> {
// Step 2: Create a study to minimize the score
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
// Step 3: Configure optimization
// Step 3: Define parameter search spaces
let cache_size_mb_param = IntParam::new(64, 1024).name("cache_size_mb").step(64);
let connection_pool_size_param = IntParam::new(10, 200).name("connection_pool_size").step(10);
let request_timeout_ms_param = IntParam::new(1000, 10000)
.name("request_timeout_ms")
.step(500);
let retry_count_param = IntParam::new(0, 5).name("retry_count");
let batch_size_param = IntParam::new(1, 256).name("batch_size").log_scale();
let compression_level_param = IntParam::new(0, 9).name("compression_level");
let use_http2_param = BoolParam::new().name("use_http2");
let load_balancing_param = CategoricalParam::new(vec![
"round_robin",
"least_connections",
"random",
"ip_hash",
])
.name("load_balancing");
// Clone params for use after the closure moves them
let cache_size_mb_p = cache_size_mb_param.clone();
let connection_pool_size_p = connection_pool_size_param.clone();
let request_timeout_ms_p = request_timeout_ms_param.clone();
let retry_count_p = retry_count_param.clone();
let batch_size_p = batch_size_param.clone();
let compression_level_p = compression_level_param.clone();
let use_http2_p = use_http2_param.clone();
let load_balancing_p = load_balancing_param.clone();
// Step 4: Configure optimization
let n_trials = 40;
let concurrency = 4; // Run 4 trials in parallel
@@ -292,25 +311,57 @@ async fn main() -> optimizer::Result<()> {
let start = Instant::now();
// Step 4: Run parallel async optimization
// Step 5: Run parallel async optimization
//
// optimize_parallel_with_sampler:
// optimize_parallel:
// - Runs up to `concurrency` trials simultaneously
// - Each trial calls the objective function
// - Uses a semaphore to limit concurrent evaluations
// - Collects results as trials complete
//
// The "_with_sampler" suffix means the TPE sampler gets access to
// trial history for informed sampling.
// The sampler gets access to trial history for informed sampling.
study
.optimize_parallel_with_sampler(n_trials, concurrency, objective)
.optimize_parallel(n_trials, concurrency, move |trial| {
let cache_size_mb_param = cache_size_mb_param.clone();
let connection_pool_size_param = connection_pool_size_param.clone();
let request_timeout_ms_param = request_timeout_ms_param.clone();
let retry_count_param = retry_count_param.clone();
let batch_size_param = batch_size_param.clone();
let compression_level_param = compression_level_param.clone();
let use_http2_param = use_http2_param.clone();
let load_balancing_param = load_balancing_param.clone();
async move {
objective(
trial,
&cache_size_mb_param,
&connection_pool_size_param,
&request_timeout_ms_param,
&retry_count_param,
&batch_size_param,
&compression_level_param,
&use_http2_param,
&load_balancing_param,
)
.await
}
})
.await?;
let elapsed = start.elapsed();
// Step 5: Print results
print_results(&study, elapsed, n_trials);
print_best_config(&study)?;
print_best_config(
&study,
&cache_size_mb_p,
&connection_pool_size_p,
&request_timeout_ms_p,
&retry_count_p,
&batch_size_p,
&compression_level_p,
&use_http2_p,
&load_balancing_p,
)?;
print_top_trials(&study, 5);
Ok(())
+124
View File
@@ -0,0 +1,124 @@
use std::ops::ControlFlow;
use std::time::Instant;
use optimizer::parameter::Parameter;
use optimizer::sampler::random::RandomSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{FloatParam, Study};
/// Standard optimization test functions.
mod functions {
pub fn sphere(x: &[f64]) -> f64 {
x.iter().map(|xi| xi * xi).sum()
}
pub fn rosenbrock(x: &[f64]) -> f64 {
x.windows(2)
.map(|w| 100.0 * (w[1] - w[0] * w[0]).powi(2) + (1.0 - w[0]).powi(2))
.sum()
}
pub fn rastrigin(x: &[f64]) -> f64 {
let n = x.len() as f64;
10.0 * n
+ x.iter()
.map(|xi| xi * xi - 10.0 * (2.0 * std::f64::consts::PI * xi).cos())
.sum::<f64>()
}
}
fn run_convergence(
name: &str,
sampler_name: &str,
study: Study<f64>,
params: &[FloatParam],
objective: fn(&[f64]) -> f64,
n_trials: usize,
) {
let start = Instant::now();
study
.optimize_with_callback(
n_trials,
|trial| {
let x: Vec<f64> = params
.iter()
.map(|p| p.suggest(trial))
.collect::<Result<_, _>>()
.unwrap();
Ok::<_, optimizer::Error>(objective(&x))
},
|study, _trial| {
let elapsed = start.elapsed().as_millis();
let best = study.best_value().unwrap();
let n = study.n_trials();
println!("{n},{best},{elapsed},{sampler_name},{name}");
ControlFlow::Continue(())
},
)
.unwrap();
}
fn main() {
println!("trial,best_value,elapsed_ms,sampler,function");
let dims = 5;
let params: Vec<FloatParam> = (0..dims)
.map(|i| FloatParam::new(-5.0, 5.0).name(format!("x{i}")))
.collect();
let n_trials = 200;
// Sphere: Random vs TPE
run_convergence(
"sphere_5d",
"random",
Study::minimize(RandomSampler::with_seed(1)),
&params,
functions::sphere,
n_trials,
);
run_convergence(
"sphere_5d",
"tpe",
Study::minimize(TpeSampler::builder().seed(1).build().unwrap()),
&params,
functions::sphere,
n_trials,
);
// Rosenbrock: Random vs TPE
run_convergence(
"rosenbrock_5d",
"random",
Study::minimize(RandomSampler::with_seed(2)),
&params,
functions::rosenbrock,
n_trials,
);
run_convergence(
"rosenbrock_5d",
"tpe",
Study::minimize(TpeSampler::builder().seed(2).build().unwrap()),
&params,
functions::rosenbrock,
n_trials,
);
// Rastrigin: Random vs TPE
run_convergence(
"rastrigin_5d",
"random",
Study::minimize(RandomSampler::with_seed(3)),
&params,
functions::rastrigin,
n_trials,
);
run_convergence(
"rastrigin_5d",
"tpe",
Study::minimize(TpeSampler::builder().seed(3).build().unwrap()),
&params,
functions::rastrigin,
n_trials,
);
}
+87 -81
View File
@@ -23,9 +23,7 @@
use std::ops::ControlFlow;
use optimizer::sampler::CompletedTrial;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, ParamValue, Study, Trial};
use optimizer::prelude::*;
// ============================================================================
// Configuration: Hyperparameters we want to tune
@@ -87,34 +85,32 @@ fn evaluate_model(config: &ModelConfig) -> f64 {
/// The objective function that the optimizer calls for each trial.
///
/// This function:
/// 1. Uses `trial.suggest_*()` methods to sample hyperparameter values
/// 2. Builds a model configuration from those values
/// 1. Uses parameter definitions passed as arguments
/// 2. Builds a model configuration from the suggested values
/// 3. Evaluates the model and returns the loss
///
/// The optimizer learns from the results to suggest better parameters
/// in future trials.
fn objective(trial: &mut Trial) -> optimizer::Result<f64> {
// Sample hyperparameters using different strategies:
// Log-scale: Good for parameters spanning multiple orders of magnitude
// The learning rate might be 0.001, 0.01, or 0.1 - log-scale samples evenly across these
let learning_rate = trial.suggest_float_log("learning_rate", 0.001, 0.3)?;
// Regular integer: Uniformly samples from the range [3, 12]
let max_depth = trial.suggest_int("max_depth", 3, 12)?;
// Stepped integer: Only samples multiples of 50 (50, 100, 150, ..., 500)
// Useful when you only want to test specific values
let n_estimators = trial.suggest_int_step("n_estimators", 50, 500, 50)?;
// Regular float: Uniformly samples from [0.5, 1.0]
let subsample = trial.suggest_float("subsample", 0.5, 1.0)?;
let colsample_bytree = trial.suggest_float("colsample_bytree", 0.5, 1.0)?;
// More parameters
let min_child_weight = trial.suggest_int("min_child_weight", 1, 10)?;
let reg_alpha = trial.suggest_float_log("reg_alpha", 1e-3, 10.0)?;
let reg_lambda = trial.suggest_float_log("reg_lambda", 1e-3, 10.0)?;
#[allow(clippy::too_many_arguments)]
fn objective(
trial: &mut Trial,
learning_rate_param: &FloatParam,
max_depth_param: &IntParam,
n_estimators_param: &IntParam,
subsample_param: &FloatParam,
colsample_bytree_param: &FloatParam,
min_child_weight_param: &IntParam,
reg_alpha_param: &FloatParam,
reg_lambda_param: &FloatParam,
) -> optimizer::Result<f64> {
let learning_rate = learning_rate_param.suggest(trial)?;
let max_depth = max_depth_param.suggest(trial)?;
let n_estimators = n_estimators_param.suggest(trial)?;
let subsample = subsample_param.suggest(trial)?;
let colsample_bytree = colsample_bytree_param.suggest(trial)?;
let min_child_weight = min_child_weight_param.suggest(trial)?;
let reg_alpha = reg_alpha_param.suggest(trial)?;
let reg_lambda = reg_lambda_param.suggest(trial)?;
// Build configuration and evaluate
let config = ModelConfig {
@@ -148,35 +144,12 @@ fn objective(trial: &mut Trial) -> optimizer::Result<f64> {
/// Return `ControlFlow::Continue(())` to keep optimizing.
/// Return `ControlFlow::Break(())` to stop early.
fn on_trial_complete(study: &Study<f64>, trial: &CompletedTrial<f64>) -> ControlFlow<()> {
// Helper to extract parameter values
let get_float = |name: &str| -> f64 {
match trial.params.get(name) {
Some(ParamValue::Float(v)) => *v,
_ => 0.0,
}
};
let get_int = |name: &str| -> i64 {
match trial.params.get(name) {
Some(ParamValue::Int(v)) => *v,
_ => 0,
}
};
// Print progress
println!(
"{:>5} {:>10.5} {:>10} {:>12} {:>10.3} {:>12.3} {:>8} {:>10.4} {:>10.4} {:>12.6}",
study.n_trials(),
get_float("learning_rate"),
get_int("max_depth"),
get_int("n_estimators"),
get_float("subsample"),
get_float("colsample_bytree"),
get_int("min_child_weight"),
get_float("reg_alpha"),
get_float("reg_lambda"),
trial.value,
);
// Print trial number and objective value
print!("{:>5} ", study.n_trials());
for value in trial.params.values() {
print!("{value:>12} ");
}
println!("{:>12.6}", trial.value);
// Early stopping: if we find an excellent solution, stop early
if trial.value < 0.16 {
@@ -218,29 +191,47 @@ fn main() -> optimizer::Result<()> {
// Print header
println!("Starting hyperparameter optimization...\n");
println!(
"{:>5} {:>10} {:>10} {:>12} {:>10} {:>12} {:>8} {:>10} {:>10} {:>12}",
"Trial",
"LR",
"MaxDepth",
"Estimators",
"Subsample",
"ColSample",
"MinCW",
"Alpha",
"Lambda",
"Loss"
"{:>5} {:>12} (parameters...) {:>12}",
"Trial", "Params", "Loss"
);
println!("{}", "-".repeat(110));
println!("{}", "-".repeat(60));
// Step 3: Run optimization
// Step 3: Define parameter search spaces
let learning_rate_param = FloatParam::new(0.001, 0.3)
.name("learning_rate")
.log_scale();
let max_depth_param = IntParam::new(3, 12).name("max_depth");
let n_estimators_param = IntParam::new(50, 500).name("n_estimators").step(50);
let subsample_param = FloatParam::new(0.5, 1.0).name("subsample");
let colsample_bytree_param = FloatParam::new(0.5, 1.0).name("colsample_bytree");
let min_child_weight_param = IntParam::new(1, 10).name("min_child_weight");
let reg_alpha_param = FloatParam::new(1e-3, 10.0).name("reg_alpha").log_scale();
let reg_lambda_param = FloatParam::new(1e-3, 10.0).name("reg_lambda").log_scale();
// Step 4: Run optimization
//
// optimize_with_callback_sampler runs the objective function for up to
// optimize_with_callback runs the objective function for up to
// n_trials iterations. After each trial, it calls the callback.
// The "_sampler" suffix means the TPE sampler gets access to trial
// history for informed sampling.
// The sampler gets access to trial history for informed sampling.
let n_trials = 50;
study.optimize_with_callback_sampler(n_trials, objective, on_trial_complete)?;
study.optimize_with_callback(
n_trials,
|trial| {
objective(
trial,
&learning_rate_param,
&max_depth_param,
&n_estimators_param,
&subsample_param,
&colsample_bytree_param,
&min_child_weight_param,
&reg_alpha_param,
&reg_lambda_param,
)
},
on_trial_complete,
)?;
// Step 4: Get the best result
println!("\n{}", "=".repeat(110));
@@ -251,14 +242,29 @@ fn main() -> optimizer::Result<()> {
println!("\nBest trial:");
println!(" Loss: {:.6}", best.value);
println!(" Parameters:");
for (name, value) in &best.params {
match value {
ParamValue::Float(v) => println!(" {name}: {v:.6}"),
ParamValue::Int(v) => println!(" {name}: {v}"),
ParamValue::Categorical(v) => println!(" {name}: category {v}"),
}
}
println!(
" learning_rate: {:.6}",
best.get(&learning_rate_param).unwrap()
);
println!(" max_depth: {}", best.get(&max_depth_param).unwrap());
println!(
" n_estimators: {}",
best.get(&n_estimators_param).unwrap()
);
println!(" subsample: {:.6}", best.get(&subsample_param).unwrap());
println!(
" colsample_bytree: {:.6}",
best.get(&colsample_bytree_param).unwrap()
);
println!(
" min_child_weight: {}",
best.get(&min_child_weight_param).unwrap()
);
println!(" reg_alpha: {:.6}", best.get(&reg_alpha_param).unwrap());
println!(
" reg_lambda: {:.6}",
best.get(&reg_lambda_param).unwrap()
);
// Step 5: Use the best parameters (in a real app)
//
+56
View File
@@ -0,0 +1,56 @@
use optimizer::prelude::*;
use optimizer_derive::Categorical;
#[derive(Clone, Debug, Categorical)]
enum Activation {
Relu,
Sigmoid,
Tanh,
}
fn main() {
let study: Study<f64> = Study::new(Direction::Minimize);
// Define parameters outside the objective function
let lr_param = FloatParam::new(1e-5, 1e-1).name("lr").log_scale();
let n_layers_param = IntParam::new(1, 5).name("n_layers");
let units_param = IntParam::new(32, 512).name("units").step(32);
let optimizer_param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]).name("optimizer");
let activation_param = EnumParam::<Activation>::new().name("activation");
let batch_size_param = IntParam::new(16, 256).name("batch_size").log_scale();
let use_dropout_param = BoolParam::new().name("use_dropout");
study
.optimize(20, |trial| {
let lr = lr_param.suggest(trial)?;
let n_layers = n_layers_param.suggest(trial)?;
let units = units_param.suggest(trial)?;
let optimizer = optimizer_param.suggest(trial)?;
let use_dropout = use_dropout_param.suggest(trial)?;
let activation = activation_param.suggest(trial)?;
let batch_size = batch_size_param.suggest(trial)?;
// Simulate a loss function
let loss = lr * (n_layers as f64) + (units as f64) * 0.001
- if use_dropout { 0.1 } else { 0.0 };
println!(
"Trial {}: lr={lr:.6}, layers={n_layers}, units={units}, opt={optimizer}, \
dropout={use_dropout}, activation={activation:?}, batch={batch_size} -> loss={loss:.4}",
trial.id()
);
Ok::<_, Error>(loss)
})
.unwrap();
let best = study.best_trial().unwrap();
println!("\nBest trial: value={:.4}", best.value);
println!(" lr: {:.6}", best.get(&lr_param).unwrap());
println!(" n_layers: {}", best.get(&n_layers_param).unwrap());
println!(" units: {}", best.get(&units_param).unwrap());
println!(" optimizer: {}", best.get(&optimizer_param).unwrap());
println!(" activation: {:?}", best.get(&activation_param).unwrap());
println!(" batch_size: {}", best.get(&batch_size_param).unwrap());
println!(" use_dropout: {}", best.get(&use_dropout_param).unwrap());
}
+15
View File
@@ -0,0 +1,15 @@
[package]
name = "optimizer-derive"
version = "0.1.0"
edition = "2024"
rust-version = "1.88"
license = "MIT"
description = "Derive macros for the optimizer crate"
repository = "https://github.com/raimannma/rust-optimizer"
[lib]
proc-macro = true
[dependencies]
syn = { version = "2", features = ["full"] }
quote = "1"
+71
View File
@@ -0,0 +1,71 @@
use proc_macro::TokenStream;
use quote::quote;
use syn::{Data, DeriveInput, Fields, parse_macro_input};
/// Derive macro for the `Categorical` trait on fieldless enums.
///
/// Generates an implementation of `optimizer::Categorical` that maps
/// enum variants to/from sequential indices.
///
/// # Example
///
/// ```ignore
/// use optimizer::Categorical;
///
/// #[derive(Clone, Categorical)]
/// enum Color {
/// Red,
/// Green,
/// Blue,
/// }
/// ```
#[proc_macro_derive(Categorical)]
pub fn derive_categorical(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let Data::Enum(data_enum) = &input.data else {
return syn::Error::new_spanned(&input, "Categorical can only be derived for enums")
.to_compile_error()
.into();
};
// Validate all variants are fieldless
for variant in &data_enum.variants {
if !matches!(variant.fields, Fields::Unit) {
return syn::Error::new_spanned(
variant,
"Categorical can only be derived for enums with unit variants (no fields)",
)
.to_compile_error()
.into();
}
}
let n_choices = data_enum.variants.len();
let variant_names: Vec<_> = data_enum.variants.iter().map(|v| &v.ident).collect();
let indices: Vec<usize> = (0..n_choices).collect();
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let expanded = quote! {
impl #impl_generics optimizer::Categorical for #name #ty_generics #where_clause {
const N_CHOICES: usize = #n_choices;
fn from_index(index: usize) -> Self {
match index {
#(#indices => #name::#variant_names,)*
_ => panic!("invalid index {} for {} with {} variants", index, stringify!(#name), #n_choices),
}
}
fn to_index(&self) -> usize {
match self {
#(#name::#variant_names => #indices,)*
}
}
}
};
expanded.into()
}
+4
View File
@@ -2,6 +2,7 @@
/// Distribution for floating-point parameters.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct FloatDistribution {
/// Lower bound (inclusive).
pub low: f64,
@@ -15,6 +16,7 @@ pub struct FloatDistribution {
/// Distribution for integer parameters.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct IntDistribution {
/// Lower bound (inclusive).
pub low: i64,
@@ -28,6 +30,7 @@ pub struct IntDistribution {
/// Distribution for categorical parameters.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CategoricalDistribution {
/// Number of choices available.
pub n_choices: usize,
@@ -35,6 +38,7 @@ pub struct CategoricalDistribution {
/// Enum wrapping all parameter distribution types.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum Distribution {
/// A floating-point distribution.
Float(FloatDistribution),
+34
View File
@@ -72,6 +72,10 @@ pub enum Error {
got: usize,
},
/// Returned when a trial is pruned (stopped early by the objective function).
#[error("trial was pruned")]
TrialPruned,
/// Returned when an internal invariant is violated.
#[error("internal error: {0}")]
Internal(&'static str),
@@ -83,3 +87,33 @@ pub enum Error {
}
pub type Result<T> = core::result::Result<T, Error>;
/// Convenience type for signalling a pruned trial from an objective function.
///
/// Implements `Into<Error>` so it can be used with `?` in objectives that
/// return `Result<V, Error>`.
///
/// # Examples
///
/// ```
/// use optimizer::{Error, TrialPruned};
///
/// fn objective_that_prunes() -> Result<f64, Error> {
/// // ... some computation ...
/// Err(TrialPruned)?
/// }
/// ```
#[derive(Debug)]
pub struct TrialPruned;
impl core::fmt::Display for TrialPruned {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "trial was pruned")
}
}
impl From<TrialPruned> for Error {
fn from(_: TrialPruned) -> Self {
Error::TrialPruned
}
}
+93
View File
@@ -0,0 +1,93 @@
//! Parameter importance via Spearman rank correlation.
/// Assign average ranks to a slice of `f64` values (handles ties).
#[allow(clippy::cast_precision_loss, clippy::float_cmp)]
pub(crate) fn rank(values: &[f64]) -> Vec<f64> {
let n = values.len();
let mut indexed: Vec<(usize, f64)> = values.iter().copied().enumerate().collect();
indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(core::cmp::Ordering::Equal));
let mut ranks = vec![0.0; n];
let mut i = 0;
while i < n {
// Find the run of tied values.
let mut j = i + 1;
while j < n && indexed[j].1 == indexed[i].1 {
j += 1;
}
// Average rank for the tie group (1-based ranks).
let avg = (i + 1..=j).sum::<usize>() as f64 / (j - i) as f64;
for item in &indexed[i..j] {
ranks[item.0] = avg;
}
i = j;
}
ranks
}
/// Pearson correlation coefficient on two equal-length slices.
#[allow(clippy::cast_precision_loss)]
fn pearson(x: &[f64], y: &[f64]) -> f64 {
let n = x.len() as f64;
let mean_x = x.iter().sum::<f64>() / n;
let mean_y = y.iter().sum::<f64>() / n;
let mut cov = 0.0;
let mut var_x = 0.0;
let mut var_y = 0.0;
for (xi, yi) in x.iter().zip(y.iter()) {
let dx = xi - mean_x;
let dy = yi - mean_y;
cov += dx * dy;
var_x += dx * dx;
var_y += dy * dy;
}
let denom = (var_x * var_y).sqrt();
if denom == 0.0 { 0.0 } else { cov / denom }
}
/// Spearman rank correlation (Pearson on ranks).
pub(crate) fn spearman(x: &[f64], y: &[f64]) -> f64 {
pearson(&rank(x), &rank(y))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rank_no_ties() {
let ranks = rank(&[30.0, 10.0, 20.0]);
assert_eq!(ranks, vec![3.0, 1.0, 2.0]);
}
#[test]
fn rank_with_ties() {
let ranks = rank(&[10.0, 20.0, 20.0, 30.0]);
assert_eq!(ranks, vec![1.0, 2.5, 2.5, 4.0]);
}
#[test]
fn perfect_positive_correlation() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![2.0, 4.0, 6.0, 8.0, 10.0];
let r = spearman(&x, &y);
assert!((r - 1.0).abs() < 1e-10);
}
#[test]
fn perfect_negative_correlation() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![10.0, 8.0, 6.0, 4.0, 2.0];
let r = spearman(&x, &y);
assert!((r + 1.0).abs() < 1e-10);
}
#[test]
fn zero_variance_returns_zero() {
let x = vec![5.0, 5.0, 5.0];
let y = vec![1.0, 2.0, 3.0];
assert!(spearman(&x, &y).abs() < f64::EPSILON);
}
}
+1 -1
View File
@@ -5,7 +5,7 @@
//! parameter independently, the multivariate KDE models the joint distribution
//! to better capture correlations between parameters.
use rand::Rng;
use rand::{Rng, RngExt};
use crate::error::{Error, Result};
+1 -1
View File
@@ -3,7 +3,7 @@
//! This module provides a Gaussian kernel density estimator used by the TPE
//! sampler to model probability distributions over good and bad trial regions.
use rand::Rng;
use rand::{Rng, RngExt};
use crate::error::{Error, Result};
+134 -28
View File
@@ -17,6 +17,9 @@
//! - **Random Search** - Simple random sampling for baseline comparisons
//! - **TPE (Tree-Parzen Estimator)** - Bayesian optimization for efficient search
//! - **Grid Search** - Exhaustive search over a specified parameter grid
//! - **Sobol (QMC)** - Quasi-random sampling for better space coverage (requires `sobol` feature)
//! - **CMA-ES** - Covariance Matrix Adaptation Evolution Strategy for continuous optimization (requires `cma-es` feature)
//! - **BOHB** - Bayesian Optimization + `HyperBand` for budget-aware TPE sampling
//!
//! Additional features include:
//!
@@ -28,24 +31,26 @@
//! # Quick Start
//!
//! ```
//! use optimizer::sampler::tpe::TpeSampler;
//! use optimizer::{Direction, Study};
//! use optimizer::prelude::*;
//!
//! // Create a study with TPE sampler
//! let sampler = TpeSampler::builder().seed(42).build().unwrap();
//! let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
//!
//! // Define parameter search space
//! let x = FloatParam::new(-10.0, 10.0).name("x");
//!
//! // Optimize x^2 for 20 trials
//! study
//! .optimize_with_sampler(20, |trial| {
//! let x = trial.suggest_float("x", -10.0, 10.0)?;
//! Ok::<_, optimizer::Error>(x * x)
//! .optimize(20, |trial| {
//! let x_val = x.suggest(trial)?;
//! Ok::<_, Error>(x_val * x_val)
//! })
//! .unwrap();
//!
//! // Get the best result
//! let best = study.best_trial().unwrap();
//! println!("Best value: {} at x={:?}", best.value, best.params);
//! println!("x = {}", best.get(&x).unwrap());
//! ```
//!
//! # Creating a Study
@@ -69,29 +74,35 @@
//!
//! # Suggesting Parameters
//!
//! Within the objective function, use [`Trial`] to suggest parameter values:
//! Within the objective function, use parameter types to suggest values:
//!
//! ```
//! use optimizer::parameter::{BoolParam, CategoricalParam, FloatParam, IntParam, Parameter};
//! use optimizer::{Direction, Study};
//!
//! let study: Study<f64> = Study::new(Direction::Minimize);
//!
//! // Define parameter search spaces
//! let x_param = FloatParam::new(0.0, 1.0);
//! let lr_param = FloatParam::new(1e-5, 1e-1).log_scale();
//! let step_param = FloatParam::new(0.0, 1.0).step(0.1);
//! let n_param = IntParam::new(1, 10);
//! let batch_param = IntParam::new(16, 256).log_scale();
//! let units_param = IntParam::new(32, 512).step(32);
//! let flag_param = BoolParam::new();
//! let optimizer_param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]);
//!
//! study
//! .optimize(10, |trial| {
//! // Float parameters
//! let x = trial.suggest_float("x", 0.0, 1.0)?;
//! let lr = trial.suggest_float_log("learning_rate", 1e-5, 1e-1)?;
//! let step = trial.suggest_float_step("step", 0.0, 1.0, 0.1)?;
//! let x = x_param.suggest(trial)?;
//! let lr = lr_param.suggest(trial)?;
//! let step = step_param.suggest(trial)?;
//! let n = n_param.suggest(trial)?;
//! let batch = batch_param.suggest(trial)?;
//! let units = units_param.suggest(trial)?;
//! let flag = flag_param.suggest(trial)?;
//! let optimizer = optimizer_param.suggest(trial)?;
//!
//! // Integer parameters
//! let n = trial.suggest_int("n_layers", 1, 10)?;
//! let batch = trial.suggest_int_log("batch_size", 16, 256)?;
//! let units = trial.suggest_int_step("units", 32, 512, 32)?;
//!
//! // Categorical parameters
//! let optimizer = trial.suggest_categorical("optimizer", &["sgd", "adam", "rmsprop"])?;
//!
//! // Return objective value
//! Ok::<_, optimizer::Error>(x * n as f64)
//! })
//! .unwrap();
@@ -147,35 +158,130 @@
//!
//! ```ignore
//! use optimizer::{Study, Direction};
//! use optimizer::parameter::{FloatParam, Parameter};
//!
//! let x_param = FloatParam::new(0.0, 1.0);
//!
//! // Sequential async
//! study.optimize_async(10, |mut trial| async move {
//! let x = trial.suggest_float("x", 0.0, 1.0)?;
//! Ok((trial, x * x))
//! study.optimize_async(10, |mut trial| {
//! let x_param = x_param.clone();
//! async move {
//! let x = x_param.suggest(&mut trial)?;
//! Ok((trial, x * x))
//! }
//! }).await?;
//!
//! // Parallel with bounded concurrency
//! study.optimize_parallel(10, 4, |mut trial| async move {
//! let x = trial.suggest_float("x", 0.0, 1.0)?;
//! Ok((trial, x * x))
//! study.optimize_parallel(10, 4, |mut trial| {
//! let x_param = x_param.clone();
//! async move {
//! let x = x_param.suggest(&mut trial)?;
//! Ok((trial, x * x))
//! }
//! }).await?;
//! ```
//!
//! # Feature Flags
//!
//! - `async`: Enable async optimization methods (requires tokio)
//! - `derive`: Enable `#[derive(Categorical)]` for enum parameters
//! - `serde`: Enable `Serialize`/`Deserialize` on public types and `Study::save()`/`Study::load()`
//! - `sobol`: Enable the Sobol quasi-random sampler for better space coverage
//! - `cma-es`: Enable the CMA-ES sampler for continuous optimization
//! - `tracing`: Emit structured log events via the [`tracing`](https://docs.rs/tracing) crate at key optimization points
/// Emit a `tracing::info!` event when the `tracing` feature is enabled.
/// No-op otherwise.
#[cfg(feature = "tracing")]
macro_rules! trace_info {
($($arg:tt)*) => { tracing::info!($($arg)*) };
}
#[cfg(not(feature = "tracing"))]
macro_rules! trace_info {
($($arg:tt)*) => {};
}
/// Emit a `tracing::debug!` event when the `tracing` feature is enabled.
/// No-op otherwise.
#[cfg(feature = "tracing")]
macro_rules! trace_debug {
($($arg:tt)*) => { tracing::debug!($($arg)*) };
}
#[cfg(not(feature = "tracing"))]
macro_rules! trace_debug {
($($arg:tt)*) => {};
}
mod distribution;
mod error;
mod importance;
mod kde;
mod param;
pub mod parameter;
pub mod pruner;
pub mod sampler;
mod study;
mod trial;
mod types;
pub use error::{Error, Result};
pub use error::{Error, Result, TrialPruned};
#[cfg(feature = "derive")]
pub use optimizer_derive::Categorical;
pub use param::ParamValue;
pub use parameter::{
BoolParam, Categorical, CategoricalParam, EnumParam, FloatParam, IntParam, ParamId, Parameter,
};
pub use pruner::{
HyperbandPruner, MedianPruner, NopPruner, PatientPruner, PercentilePruner, Pruner,
SuccessiveHalvingPruner, ThresholdPruner, WilcoxonPruner,
};
pub use sampler::CompletedTrial;
pub use sampler::bohb::BohbSampler;
#[cfg(feature = "cma-es")]
pub use sampler::cma_es::CmaEsSampler;
pub use sampler::grid::GridSearchSampler;
pub use sampler::random::RandomSampler;
#[cfg(feature = "sobol")]
pub use sampler::sobol::SobolSampler;
pub use sampler::tpe::TpeSampler;
pub use study::Study;
pub use trial::{SuggestableRange, Trial};
#[cfg(feature = "serde")]
pub use study::StudySnapshot;
pub use trial::{AttrValue, Trial};
pub use types::{Direction, TrialState};
/// Convenient wildcard import for the most common types.
///
/// ```
/// use optimizer::prelude::*;
/// ```
pub mod prelude {
#[cfg(feature = "derive")]
pub use optimizer_derive::Categorical as DeriveCategory;
pub use crate::error::{Error, Result, TrialPruned};
pub use crate::param::ParamValue;
pub use crate::parameter::{
BoolParam, Categorical, CategoricalParam, EnumParam, FloatParam, IntParam, Parameter,
};
pub use crate::pruner::{
HyperbandPruner, MedianPruner, NopPruner, PatientPruner, PercentilePruner, Pruner,
SuccessiveHalvingPruner, ThresholdPruner,
};
pub use crate::sampler::CompletedTrial;
pub use crate::sampler::bohb::BohbSampler;
#[cfg(feature = "cma-es")]
pub use crate::sampler::cma_es::CmaEsSampler;
pub use crate::sampler::grid::GridSearchSampler;
pub use crate::sampler::random::RandomSampler;
#[cfg(feature = "sobol")]
pub use crate::sampler::sobol::SobolSampler;
pub use crate::sampler::tpe::TpeSampler;
pub use crate::study::Study;
#[cfg(feature = "serde")]
pub use crate::study::StudySnapshot;
pub use crate::trial::{AttrValue, Trial};
pub use crate::types::Direction;
}
+11
View File
@@ -6,6 +6,7 @@
/// For categorical parameters, the `Categorical` variant stores
/// the index into the choices array.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum ParamValue {
/// A floating-point parameter value.
Float(f64),
@@ -14,3 +15,13 @@ pub enum ParamValue {
/// A categorical parameter value, stored as an index into the choices array.
Categorical(usize),
}
impl core::fmt::Display for ParamValue {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Float(v) => write!(f, "{v}"),
Self::Int(v) => write!(f, "{v}"),
Self::Categorical(v) => write!(f, "category({v})"),
}
}
}
+1039
View File
File diff suppressed because it is too large Load Diff
+559
View File
@@ -0,0 +1,559 @@
use core::sync::atomic::{AtomicU64, Ordering};
use std::collections::HashMap;
use std::sync::Mutex;
use super::Pruner;
use crate::sampler::CompletedTrial;
use crate::types::{Direction, TrialState};
/// Hyperband pruner that manages multiple Successive Halving brackets.
///
/// Hyperband addresses SHA's sensitivity to the `min_resource` choice by
/// running multiple brackets, each with a different tradeoff between the
/// number of configurations and the starting budget:
///
/// - Bracket 0: many trials, very small starting budget (aggressive pruning)
/// - Bracket 1: fewer trials, larger starting budget (moderate pruning)
/// - ...
/// - Bracket `s_max`: few trials, full budget (no pruning)
///
/// Trials are assigned to brackets in round-robin fashion. Each bracket
/// runs SHA with its own `min_resource` and rung schedule.
///
/// # Examples
///
/// ```
/// use optimizer::Direction;
/// use optimizer::pruner::HyperbandPruner;
///
/// let pruner = HyperbandPruner::new()
/// .min_resource(1)
/// .max_resource(81)
/// .reduction_factor(3)
/// .direction(Direction::Minimize);
/// ```
pub struct HyperbandPruner {
min_resource: u64,
max_resource: u64,
reduction_factor: u64,
direction: Direction,
/// Tracks which bracket each trial belongs to.
trial_brackets: Mutex<HashMap<u64, usize>>,
/// Counter for round-robin bracket assignment.
next_bracket: AtomicU64,
}
impl HyperbandPruner {
/// Create a new `HyperbandPruner` with default parameters.
///
/// Defaults: `min_resource=1`, `max_resource=81`, `reduction_factor=3`,
/// `direction=Minimize`.
#[must_use]
pub fn new() -> Self {
Self {
min_resource: 1,
max_resource: 81,
reduction_factor: 3,
direction: Direction::Minimize,
trial_brackets: Mutex::new(HashMap::new()),
next_bracket: AtomicU64::new(0),
}
}
/// Set the minimum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn min_resource(mut self, r: u64) -> Self {
assert!(r > 0, "min_resource must be > 0, got {r}");
self.min_resource = r;
self
}
/// Set the maximum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn max_resource(mut self, r: u64) -> Self {
assert!(r > 0, "max_resource must be > 0, got {r}");
self.max_resource = r;
self
}
/// Set the reduction factor (eta). At each rung, the top 1/eta trials survive.
///
/// # Panics
///
/// Panics if `eta` is less than 2.
#[must_use]
pub fn reduction_factor(mut self, eta: u64) -> Self {
assert!(eta >= 2, "reduction_factor must be >= 2, got {eta}");
self.reduction_factor = eta;
self
}
/// Set the optimization direction.
#[must_use]
pub fn direction(mut self, d: Direction) -> Self {
self.direction = d;
self
}
/// Compute `s_max = floor(log(max_resource / min_resource) / log(eta))`.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn s_max(&self) -> u64 {
let eta = self.reduction_factor as f64;
let ratio = self.max_resource as f64 / self.min_resource as f64;
(ratio.ln() / eta.ln()).floor() as u64
}
/// Compute the rung steps for a given bracket `s`.
///
/// For bracket `s`, the starting resource is `max_resource / eta^(s_max - s)`,
/// and rungs are spaced at powers of eta from there up to `max_resource`.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn rung_steps_for_bracket(&self, bracket: usize) -> Vec<u64> {
let s_max = self.s_max();
let eta = self.reduction_factor as f64;
// Starting resource for this bracket
let exponent = s_max.saturating_sub(bracket as u64);
let min_resource_bracket =
(self.max_resource as f64 / eta.powi(exponent as i32)).ceil() as u64;
let mut steps = Vec::new();
let mut rung: u32 = 0;
while let Some(power) = self.reduction_factor.checked_pow(rung) {
let step = min_resource_bracket.saturating_mul(power);
if step > self.max_resource {
break;
}
steps.push(step);
rung += 1;
}
steps
}
/// Assign a trial to a bracket (round-robin) and return the bracket index.
#[allow(clippy::cast_possible_truncation)]
fn assign_bracket(&self, trial_id: u64) -> usize {
let n_brackets = (self.s_max() + 1) as usize;
let mut map = self.trial_brackets.lock().expect("lock poisoned");
*map.entry(trial_id).or_insert_with(|| {
let idx = self.next_bracket.fetch_add(1, Ordering::Relaxed);
(idx as usize) % n_brackets
})
}
}
impl Default for HyperbandPruner {
fn default() -> Self {
Self::new()
}
}
#[allow(clippy::cast_precision_loss)]
impl Pruner for HyperbandPruner {
fn should_prune(
&self,
trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
let bracket = self.assign_bracket(trial_id);
let rungs = self.rung_steps_for_bracket(bracket);
// Find the highest rung step <= current step
let Some(&rung_step) = rungs.iter().rev().find(|&&r| r <= step) else {
return false;
};
// Never prune at the last rung (full budget)
if rung_step >= self.max_resource {
return false;
}
// Get the current trial's value at this rung step
let current_value =
if let Some(&(_, v)) = intermediate_values.iter().find(|(s, _)| *s == rung_step) {
v
} else if let Some(&(_, v)) = intermediate_values
.iter()
.rev()
.find(|(s, _)| *s <= rung_step)
{
v
} else {
return false;
};
self.is_pruned_at_rung(current_value, rung_step, bracket, completed_trials)
}
}
impl HyperbandPruner {
/// Determine whether a trial should be pruned at the given rung within its bracket.
///
/// Only compares against other trials in the same bracket.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn is_pruned_at_rung(
&self,
current_value: f64,
rung_step: u64,
bracket: usize,
completed_trials: &[CompletedTrial],
) -> bool {
let eta = self.reduction_factor as usize;
// Collect values at this rung step from trials in the same bracket
let map = self.trial_brackets.lock().expect("lock poisoned");
let mut values_at_rung: Vec<f64> = completed_trials
.iter()
.filter(|t| t.state == TrialState::Complete || t.state == TrialState::Pruned)
.filter(|t| map.get(&t.id).copied() == Some(bracket))
.filter_map(|t| {
t.intermediate_values
.iter()
.find(|(s, _)| *s == rung_step)
.map(|(_, v)| *v)
})
.collect();
drop(map);
// Need at least eta trials to make a meaningful comparison
if values_at_rung.len() < eta {
return false;
}
values_at_rung.push(current_value);
values_at_rung
.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(core::cmp::Ordering::Equal));
if self.direction == Direction::Maximize {
values_at_rung.reverse();
}
let n_keep = (values_at_rung.len() as f64 / eta as f64).ceil() as usize;
let threshold_idx = n_keep.max(1) - 1;
let threshold = values_at_rung[threshold_idx];
match self.direction {
Direction::Minimize => current_value > threshold,
Direction::Maximize => current_value < threshold,
}
}
}
#[cfg(test)]
#[allow(clippy::cast_precision_loss)]
mod tests {
use super::*;
fn make_trial(id: u64, values: &[(u64, f64)]) -> CompletedTrial {
use std::collections::HashMap;
use crate::parameter::ParamId;
CompletedTrial::with_intermediate_values(
id,
HashMap::<ParamId, crate::ParamValue>::new(),
HashMap::new(),
HashMap::new(),
0.0,
values.to_vec(),
HashMap::new(),
)
}
fn make_pruned_trial(id: u64, values: &[(u64, f64)]) -> CompletedTrial {
let mut t = make_trial(id, values);
t.state = TrialState::Pruned;
t
}
#[test]
fn s_max_default() {
let pruner = HyperbandPruner::new();
// s_max = floor(ln(81/1) / ln(3)) = floor(4.0) = 4
assert_eq!(pruner.s_max(), 4);
}
#[test]
fn s_max_custom() {
let pruner = HyperbandPruner::new()
.min_resource(1)
.max_resource(16)
.reduction_factor(2);
// s_max = floor(ln(16) / ln(2)) = floor(4.0) = 4
assert_eq!(pruner.s_max(), 4);
}
#[test]
fn bracket_count() {
let pruner = HyperbandPruner::new();
// s_max=4, so brackets 0..=4 → 5 brackets
assert_eq!(pruner.s_max() + 1, 5);
}
#[test]
fn rung_steps_bracket_0_default() {
let pruner = HyperbandPruner::new();
// Bracket 0: min_resource_bracket = ceil(81 / 3^4) = ceil(81/81) = 1
// Rungs: 1, 3, 9, 27, 81
assert_eq!(pruner.rung_steps_for_bracket(0), vec![1, 3, 9, 27, 81]);
}
#[test]
fn rung_steps_bracket_2_default() {
let pruner = HyperbandPruner::new();
// Bracket 2: min_resource_bracket = ceil(81 / 3^(4-2)) = ceil(81/9) = 9
// Rungs: 9, 27, 81
assert_eq!(pruner.rung_steps_for_bracket(2), vec![9, 27, 81]);
}
#[test]
fn rung_steps_bracket_4_default() {
let pruner = HyperbandPruner::new();
// Bracket 4 (s_max): min_resource_bracket = ceil(81 / 3^0) = 81
// Rungs: 81 only (no pruning, full budget)
assert_eq!(pruner.rung_steps_for_bracket(4), vec![81]);
}
#[test]
fn rung_steps_eta2() {
let pruner = HyperbandPruner::new()
.min_resource(1)
.max_resource(16)
.reduction_factor(2);
// s_max = 4
// Bracket 0: min=ceil(16/2^4)=1, rungs: 1,2,4,8,16
assert_eq!(pruner.rung_steps_for_bracket(0), vec![1, 2, 4, 8, 16]);
// Bracket 2: min=ceil(16/2^2)=4, rungs: 4,8,16
assert_eq!(pruner.rung_steps_for_bracket(2), vec![4, 8, 16]);
// Bracket 4: min=16, rungs: 16
assert_eq!(pruner.rung_steps_for_bracket(4), vec![16]);
}
#[test]
fn round_robin_bracket_assignment() {
let pruner = HyperbandPruner::new(); // 5 brackets (0..=4)
// Trials get assigned in round-robin: 0→0, 1→1, 2→2, 3→3, 4→4, 5→0, ...
assert_eq!(pruner.assign_bracket(100), 0);
assert_eq!(pruner.assign_bracket(101), 1);
assert_eq!(pruner.assign_bracket(102), 2);
assert_eq!(pruner.assign_bracket(103), 3);
assert_eq!(pruner.assign_bracket(104), 4);
assert_eq!(pruner.assign_bracket(105), 0); // wraps around
// Repeated calls for same trial return same bracket
assert_eq!(pruner.assign_bracket(100), 0);
assert_eq!(pruner.assign_bracket(103), 3);
}
#[test]
fn no_prune_before_first_rung() {
let pruner = HyperbandPruner::new().direction(Direction::Minimize);
// Assign trial 0 to bracket 0 (rungs: 1, 3, 9, 27, 81)
pruner.assign_bracket(0);
// Register completed trials in bracket 0
let mut completed = Vec::new();
for i in 1..=9 {
pruner.assign_bracket(i);
completed.push(make_trial(i, &[(1, i as f64)]));
}
// Trial at step 0 (before rung 1) → don't prune
assert!(!pruner.should_prune(0, 0, &[(0, 100.0)], &completed));
}
#[test]
fn no_prune_at_max_resource() {
let pruner = HyperbandPruner::new().direction(Direction::Minimize);
// Put all trials in bracket 0
let mut completed = Vec::new();
for i in 0..9 {
pruner.assign_bracket(i);
completed.push(make_trial(i, &[(81, (i + 1) as f64)]));
}
let trial_id = 9;
pruner.assign_bracket(trial_id);
// At max_resource (81), never prune
assert!(!pruner.should_prune(trial_id, 81, &[(81, 100.0)], &completed));
}
#[test]
fn prune_worst_in_bracket_minimize() {
let pruner = HyperbandPruner::new().direction(Direction::Minimize);
// Force all trials into bracket 0 by assigning sequentially
// With 5 brackets, trials 0,5,10,... go to bracket 0
let bracket_0_ids: Vec<u64> = (0..5).map(|i| i * 5).collect();
// Assign all 25 trial IDs to fill brackets
for i in 0..25 {
pruner.assign_bracket(i);
}
// Create 9 completed trials in bracket 0 at rung step=1
let completed: Vec<_> = bracket_0_ids
.iter()
.take(3)
.enumerate()
.map(|(idx, &id)| make_trial(id, &[(1, (idx + 1) as f64)]))
.collect();
// Trial 25 → bracket 0 (25 % 5 == 0)
let test_id = 25;
pruner.assign_bracket(test_id);
assert_eq!(pruner.assign_bracket(test_id), 0);
// 3 completed + 1 current = 4. eta=3. ceil(4/3)=2. Threshold = 2.0
// Value 2.0 → keep
assert!(!pruner.should_prune(test_id, 1, &[(1, 2.0)], &completed));
// Value 3.0 → prune
assert!(pruner.should_prune(test_id, 1, &[(1, 3.0)], &completed));
}
#[test]
fn prune_worst_in_bracket_maximize() {
let pruner = HyperbandPruner::new().direction(Direction::Maximize);
// Assign trials so they end up in bracket 0
for i in 0..25 {
pruner.assign_bracket(i);
}
let completed: Vec<_> = [0u64, 5, 10]
.iter()
.enumerate()
.map(|(idx, &id)| make_trial(id, &[(1, (idx + 1) as f64)]))
.collect();
let test_id = 25;
pruner.assign_bracket(test_id);
// For maximize, best = highest. Values: 1,2,3 + current
// Value 2.0 → keep (threshold = 2.0 when sorted desc: 3,2,current,1)
assert!(!pruner.should_prune(test_id, 1, &[(1, 2.0)], &completed));
// Value 1.0 → prune
assert!(pruner.should_prune(test_id, 1, &[(1, 0.5)], &completed));
}
#[test]
fn different_brackets_have_different_aggressiveness() {
let pruner = HyperbandPruner::new()
.min_resource(1)
.max_resource(81)
.reduction_factor(3)
.direction(Direction::Minimize);
let rungs_0 = pruner.rung_steps_for_bracket(0);
let rungs_2 = pruner.rung_steps_for_bracket(2);
let rungs_4 = pruner.rung_steps_for_bracket(4);
// Bracket 0 has the most rungs (most aggressive)
assert!(rungs_0.len() > rungs_2.len());
// Bracket 4 has just 1 rung (no pruning)
assert_eq!(rungs_4.len(), 1);
// Bracket 0 starts earliest
assert!(rungs_0[0] < rungs_2[0]);
}
#[test]
fn trials_in_different_brackets_independent() {
let pruner = HyperbandPruner::new().direction(Direction::Minimize);
// Assign trials: bracket 0 gets IDs 0,5,10,15,20
for i in 0..25 {
pruner.assign_bracket(i);
}
// Bracket 0 trials: bad values at rung step=1
let bracket_0_trials: Vec<_> = [0u64, 5, 10]
.iter()
.map(|&id| make_trial(id, &[(1, 100.0)]))
.collect();
// Bracket 1 trials: good values at rung step=1
let bracket_1_trials: Vec<_> = [1u64, 6, 11]
.iter()
.map(|&id| make_trial(id, &[(1, 1.0)]))
.collect();
let mut all_trials = bracket_0_trials;
all_trials.extend(bracket_1_trials);
// A new bracket-0 trial with value 50 should be compared against
// bracket-0 peers (100,100,100), not bracket-1 peers (1,1,1)
let test_id = 25; // bracket 0
pruner.assign_bracket(test_id);
// 3 peers at 100.0 + current at 50.0. ceil(4/3)=2. Sorted: 50,100,100,100. Threshold=100.0
// Value 50.0 < 100.0 → keep
assert!(!pruner.should_prune(test_id, 1, &[(1, 50.0)], &all_trials));
}
#[test]
fn includes_pruned_trials() {
let pruner = HyperbandPruner::new().direction(Direction::Minimize);
for i in 0..25 {
pruner.assign_bracket(i);
}
let completed = vec![
make_trial(0, &[(1, 1.0)]),
make_pruned_trial(5, &[(1, 8.0)]),
make_pruned_trial(10, &[(1, 9.0)]),
];
let test_id = 25;
pruner.assign_bracket(test_id);
// Values: 1.0, 8.0, 9.0 + current. eta=3.
// Value 1.0 → keep
assert!(!pruner.should_prune(test_id, 1, &[(1, 1.0)], &completed));
// Value 5.0 → prune (sorted: 1,5,8,9 → keep ceil(4/3)=2 → threshold=5.0, 5.0 not > 5.0 → keep)
assert!(!pruner.should_prune(test_id, 1, &[(1, 5.0)], &completed));
// Value 6.0 → prune (sorted: 1,6,8,9 → threshold=6.0, 6.0 not > 6.0 → keep)
assert!(!pruner.should_prune(test_id, 1, &[(1, 6.0)], &completed));
// Value 9.5 → prune
assert!(pruner.should_prune(test_id, 1, &[(1, 9.5)], &completed));
}
#[test]
#[should_panic(expected = "min_resource must be > 0")]
fn rejects_zero_min_resource() {
let _ = HyperbandPruner::new().min_resource(0);
}
#[test]
#[should_panic(expected = "max_resource must be > 0")]
fn rejects_zero_max_resource() {
let _ = HyperbandPruner::new().max_resource(0);
}
#[test]
#[should_panic(expected = "reduction_factor must be >= 2")]
fn rejects_reduction_factor_one() {
let _ = HyperbandPruner::new().reduction_factor(1);
}
}
+127
View File
@@ -0,0 +1,127 @@
use super::Pruner;
use super::percentile::compute_percentile;
use crate::sampler::CompletedTrial;
use crate::types::{Direction, TrialState};
/// Prune trials that are performing worse than the median of completed trials
/// at the same step.
///
/// This is the most commonly used pruner. It compares the current trial's
/// intermediate value at each step with the median of all completed trials'
/// values at that same step.
///
/// Equivalent to `PercentilePruner::new(50.0, direction)`.
///
/// # Examples
///
/// ```
/// use optimizer::Direction;
/// use optimizer::pruner::MedianPruner;
///
/// // Prune trials worse than median when minimizing, after 5 warmup steps
/// let pruner = MedianPruner::new(Direction::Minimize)
/// .n_warmup_steps(5)
/// .n_min_trials(3);
/// ```
pub struct MedianPruner {
/// The optimization direction.
direction: Direction,
/// Don't prune in the first N steps (let the trial warm up).
n_warmup_steps: u64,
/// Require at least N completed trials before pruning.
n_min_trials: usize,
}
impl MedianPruner {
/// Create a new `MedianPruner` for the given optimization direction.
///
/// By default, `n_warmup_steps` is 0 and `n_min_trials` is 1.
#[must_use]
pub fn new(direction: Direction) -> Self {
Self {
direction,
n_warmup_steps: 0,
n_min_trials: 1,
}
}
/// Set the number of warmup steps. No pruning occurs before this step.
#[must_use]
pub fn n_warmup_steps(mut self, n: u64) -> Self {
self.n_warmup_steps = n;
self
}
/// Set the minimum number of completed trials required before pruning.
#[must_use]
pub fn n_min_trials(mut self, n: usize) -> Self {
self.n_min_trials = n;
self
}
}
impl Pruner for MedianPruner {
fn should_prune(
&self,
_trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
// 1. Don't prune during warmup
if step < self.n_warmup_steps {
return false;
}
// Get the current trial's latest value
let Some(&(_, current_value)) = intermediate_values.last() else {
return false;
};
// 2. Collect values at this step from completed (non-pruned) trials
let mut values_at_step: Vec<f64> = completed_trials
.iter()
.filter(|t| t.state == TrialState::Complete)
.filter_map(|t| {
t.intermediate_values
.iter()
.find(|(s, _)| *s == step)
.map(|(_, v)| *v)
})
.collect();
// 3. Not enough trials
if values_at_step.len() < self.n_min_trials {
return false;
}
// 4. Compute median (50th percentile)
let median = compute_percentile(&mut values_at_step, 50.0);
// 5. Compare against median based on direction
match self.direction {
Direction::Minimize => current_value > median,
Direction::Maximize => current_value < median,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn compute_median_odd() {
assert!((compute_percentile(&mut [3.0, 1.0, 2.0], 50.0) - 2.0).abs() < f64::EPSILON);
}
#[test]
fn compute_median_even() {
assert!((compute_percentile(&mut [4.0, 1.0, 3.0, 2.0], 50.0) - 2.5).abs() < f64::EPSILON);
}
#[test]
fn compute_median_single() {
assert!((compute_percentile(&mut [5.0], 50.0) - 5.0).abs() < f64::EPSILON);
}
}
+74
View File
@@ -0,0 +1,74 @@
//! Pruner trait and implementations for trial pruning.
//!
//! Pruners decide whether to stop (prune) a trial early based on its
//! intermediate values compared to other trials. This is useful for
//! discarding unpromising trials before they complete, saving compute.
mod hyperband;
mod median;
mod nop;
mod patient;
pub(crate) mod percentile;
mod successive_halving;
mod threshold;
mod wilcoxon;
pub use hyperband::HyperbandPruner;
pub use median::MedianPruner;
pub use nop::NopPruner;
pub use patient::PatientPruner;
pub use percentile::PercentilePruner;
pub use successive_halving::SuccessiveHalvingPruner;
pub use threshold::ThresholdPruner;
pub use wilcoxon::WilcoxonPruner;
use crate::sampler::CompletedTrial;
/// Trait for pluggable trial pruning strategies.
///
/// Pruners are consulted after each intermediate value is reported to
/// decide whether the trial should be stopped early. The trait requires
/// `Send + Sync` to support concurrent and async optimization.
///
/// # Implementing a custom pruner
///
/// ```
/// use optimizer::pruner::Pruner;
/// use optimizer::sampler::CompletedTrial;
///
/// struct MyPruner {
/// threshold: f64,
/// }
///
/// impl Pruner for MyPruner {
/// fn should_prune(
/// &self,
/// _trial_id: u64,
/// _step: u64,
/// intermediate_values: &[(u64, f64)],
/// _completed_trials: &[CompletedTrial],
/// ) -> bool {
/// // Prune if the latest value exceeds the threshold
/// intermediate_values
/// .last()
/// .is_some_and(|&(_, v)| v > self.threshold)
/// }
/// }
/// ```
pub trait Pruner: Send + Sync {
/// Decide whether to prune a trial at the given step.
///
/// # Arguments
///
/// * `trial_id` - The current trial's ID.
/// * `step` - The step at which the intermediate value was reported.
/// * `intermediate_values` - All `(step, value)` pairs reported so far for this trial.
/// * `completed_trials` - History of all completed trials (for comparison).
fn should_prune(
&self,
trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool;
}
+17
View File
@@ -0,0 +1,17 @@
use super::Pruner;
use crate::sampler::CompletedTrial;
/// A pruner that never prunes. This is the default when no pruner is configured.
pub struct NopPruner;
impl Pruner for NopPruner {
fn should_prune(
&self,
_trial_id: u64,
_step: u64,
_intermediate_values: &[(u64, f64)],
_completed_trials: &[CompletedTrial],
) -> bool {
false
}
}
+160
View File
@@ -0,0 +1,160 @@
use std::collections::HashMap;
use std::sync::Mutex;
use super::Pruner;
use crate::sampler::CompletedTrial;
/// Wraps another pruner and adds a patience window.
///
/// The inner pruner must recommend pruning for `patience` consecutive
/// steps before this pruner actually prunes the trial. This is useful
/// to prevent premature pruning when intermediate values are noisy.
///
/// # Examples
///
/// ```
/// use optimizer::pruner::{PatientPruner, ThresholdPruner};
///
/// // Only prune after the threshold pruner recommends pruning 3 times in a row
/// let inner = ThresholdPruner::new().upper(100.0);
/// let pruner = PatientPruner::new(inner, 3);
/// ```
pub struct PatientPruner {
inner: Box<dyn Pruner>,
patience: u64,
/// Track consecutive prune recommendations per trial.
consecutive_counts: Mutex<HashMap<u64, u64>>,
}
impl PatientPruner {
/// Create a new `PatientPruner` wrapping the given inner pruner.
///
/// The inner pruner must recommend pruning for `patience` consecutive
/// calls before this pruner returns `true`.
pub fn new(inner: impl Pruner + 'static, patience: u64) -> Self {
Self {
inner: Box::new(inner),
patience,
consecutive_counts: Mutex::new(HashMap::new()),
}
}
}
impl Pruner for PatientPruner {
fn should_prune(
&self,
trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
let inner_says_prune =
self.inner
.should_prune(trial_id, step, intermediate_values, completed_trials);
let mut counts = self.consecutive_counts.lock().expect("lock poisoned");
let count = counts.entry(trial_id).or_insert(0);
if inner_says_prune {
*count += 1;
*count >= self.patience
} else {
*count = 0;
false
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pruner::ThresholdPruner;
/// A test pruner that always returns the given value.
struct ConstPruner(bool);
impl Pruner for ConstPruner {
fn should_prune(
&self,
_trial_id: u64,
_step: u64,
_intermediate_values: &[(u64, f64)],
_completed_trials: &[CompletedTrial],
) -> bool {
self.0
}
}
/// A pruner that returns values from a sequence.
struct SequencePruner(Mutex<Vec<bool>>);
impl Pruner for SequencePruner {
fn should_prune(
&self,
_trial_id: u64,
_step: u64,
_intermediate_values: &[(u64, f64)],
_completed_trials: &[CompletedTrial],
) -> bool {
self.0.lock().expect("lock poisoned").remove(0)
}
}
fn call(pruner: &PatientPruner, trial_id: u64, step: u64) -> bool {
pruner.should_prune(trial_id, step, &[(step, 0.0)], &[])
}
#[test]
fn patience_1_behaves_like_inner() {
let pruner = PatientPruner::new(ConstPruner(true), 1);
assert!(call(&pruner, 0, 0));
assert!(call(&pruner, 0, 1));
let pruner = PatientPruner::new(ConstPruner(false), 1);
assert!(!call(&pruner, 0, 0));
assert!(!call(&pruner, 0, 1));
}
#[test]
fn patience_3_requires_consecutive_recommendations() {
let pruner = PatientPruner::new(ConstPruner(true), 3);
assert!(!call(&pruner, 0, 0)); // count=1
assert!(!call(&pruner, 0, 1)); // count=2
assert!(call(&pruner, 0, 2)); // count=3 → prune
}
#[test]
fn counter_resets_on_no_prune() {
// Sequence: prune, prune, no-prune, prune, prune, prune
let seq = vec![true, true, false, true, true, true];
let pruner = PatientPruner::new(SequencePruner(Mutex::new(seq)), 3);
assert!(!call(&pruner, 0, 0)); // count=1
assert!(!call(&pruner, 0, 1)); // count=2
assert!(!call(&pruner, 0, 2)); // reset → count=0
assert!(!call(&pruner, 0, 3)); // count=1
assert!(!call(&pruner, 0, 4)); // count=2
assert!(call(&pruner, 0, 5)); // count=3 → prune
}
#[test]
fn independent_per_trial() {
let pruner = PatientPruner::new(ConstPruner(true), 2);
assert!(!call(&pruner, 0, 0)); // trial 0: count=1
assert!(!call(&pruner, 1, 0)); // trial 1: count=1
assert!(call(&pruner, 0, 1)); // trial 0: count=2 → prune
assert!(!call(&pruner, 2, 0)); // trial 2: count=1
assert!(call(&pruner, 1, 1)); // trial 1: count=2 → prune
}
#[test]
fn works_with_threshold_pruner() {
let inner = ThresholdPruner::new().upper(10.0);
let pruner = PatientPruner::new(inner, 2);
// Value below threshold → inner says no
assert!(!pruner.should_prune(0, 0, &[(0, 5.0)], &[]));
// Value above threshold → inner says yes, count=1
assert!(!pruner.should_prune(0, 1, &[(0, 5.0), (1, 15.0)], &[]));
// Value above threshold again → count=2 → prune
assert!(pruner.should_prune(0, 2, &[(0, 5.0), (1, 15.0), (2, 20.0)], &[]));
}
}
+294
View File
@@ -0,0 +1,294 @@
use super::Pruner;
use crate::sampler::CompletedTrial;
use crate::types::{Direction, TrialState};
/// Prune trials that are not in the top `percentile`% of completed trials
/// at the same training step.
///
/// `PercentilePruner::new(50.0, direction)` is equivalent to `MedianPruner`.
/// `PercentilePruner::new(25.0, direction)` keeps only the top 25% of trials.
///
/// # Examples
///
/// ```
/// use optimizer::Direction;
/// use optimizer::pruner::PercentilePruner;
///
/// // Keep only the top 25% of trials (aggressive pruning)
/// let pruner = PercentilePruner::new(25.0, Direction::Minimize)
/// .n_warmup_steps(5)
/// .n_min_trials(3);
/// ```
pub struct PercentilePruner {
/// Keep trials in the top `percentile`%. Range: (0.0, 100.0).
percentile: f64,
/// Don't prune in the first N steps (let the trial warm up).
n_warmup_steps: u64,
/// Require at least N completed trials before pruning.
n_min_trials: usize,
/// The optimization direction.
direction: Direction,
}
impl PercentilePruner {
/// Create a new `PercentilePruner` for the given percentile and direction.
///
/// The `percentile` value must be in `(0.0, 100.0)`.
/// A percentile of 50.0 is equivalent to median pruning.
///
/// # Panics
///
/// Panics if `percentile` is not in `(0.0, 100.0)`.
#[must_use]
pub fn new(percentile: f64, direction: Direction) -> Self {
assert!(
percentile > 0.0 && percentile < 100.0,
"percentile must be in (0.0, 100.0), got {percentile}"
);
Self {
percentile,
n_warmup_steps: 0,
n_min_trials: 1,
direction,
}
}
/// Set the number of warmup steps. No pruning occurs before this step.
#[must_use]
pub fn n_warmup_steps(mut self, n: u64) -> Self {
self.n_warmup_steps = n;
self
}
/// Set the minimum number of completed trials required before pruning.
#[must_use]
pub fn n_min_trials(mut self, n: usize) -> Self {
self.n_min_trials = n;
self
}
}
impl Pruner for PercentilePruner {
fn should_prune(
&self,
_trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
// 1. Don't prune during warmup
if step < self.n_warmup_steps {
return false;
}
// Get the current trial's latest value
let Some(&(_, current_value)) = intermediate_values.last() else {
return false;
};
// 2. Collect values at this step from completed (non-pruned) trials
let mut values_at_step: Vec<f64> = completed_trials
.iter()
.filter(|t| t.state == TrialState::Complete)
.filter_map(|t| {
t.intermediate_values
.iter()
.find(|(s, _)| *s == step)
.map(|(_, v)| *v)
})
.collect();
// 3. Not enough trials
if values_at_step.len() < self.n_min_trials {
return false;
}
// 4. Compute percentile threshold
let threshold = compute_percentile(&mut values_at_step, self.percentile);
// 5. Compare against threshold based on direction
match self.direction {
Direction::Minimize => current_value > threshold,
Direction::Maximize => current_value < threshold,
}
}
}
/// Compute the given percentile of a non-empty slice. Sorts the slice in place.
///
/// Uses linear interpolation between the two nearest ranks.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
pub(crate) fn compute_percentile(values: &mut [f64], percentile: f64) -> f64 {
values.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(core::cmp::Ordering::Equal));
let len = values.len();
if len == 1 {
return values[0];
}
// Rank in [0, len-1] range
let rank = percentile / 100.0 * (len - 1) as f64;
let lower = rank.floor() as usize;
let upper = rank.ceil() as usize;
if lower == upper {
values[lower]
} else {
let frac = rank - lower as f64;
values[lower] * (1.0 - frac) + values[upper] * frac
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn compute_percentile_median_odd() {
// Percentile 50 on odd-length slice = median
let val = compute_percentile(&mut [3.0, 1.0, 2.0], 50.0);
assert!((val - 2.0).abs() < f64::EPSILON);
}
#[test]
fn compute_percentile_median_even() {
// Percentile 50 on even-length slice = median (interpolated)
let val = compute_percentile(&mut [4.0, 1.0, 3.0, 2.0], 50.0);
assert!((val - 2.5).abs() < f64::EPSILON);
}
#[test]
fn compute_percentile_25() {
// [1.0, 2.0, 3.0, 4.0], rank = 0.25 * 3 = 0.75
// interpolate: 1.0 * 0.25 + 2.0 * 0.75 = 1.75
let val = compute_percentile(&mut [4.0, 1.0, 3.0, 2.0], 25.0);
assert!((val - 1.75).abs() < f64::EPSILON);
}
#[test]
fn compute_percentile_75() {
// [1.0, 2.0, 3.0, 4.0], rank = 0.75 * 3 = 2.25
// interpolate: 3.0 * 0.75 + 4.0 * 0.25 = 3.25
let val = compute_percentile(&mut [4.0, 1.0, 3.0, 2.0], 75.0);
assert!((val - 3.25).abs() < f64::EPSILON);
}
#[test]
fn compute_percentile_single() {
let val = compute_percentile(&mut [5.0], 50.0);
assert!((val - 5.0).abs() < f64::EPSILON);
}
#[test]
#[should_panic(expected = "percentile must be in (0.0, 100.0)")]
fn new_rejects_zero() {
let _ = PercentilePruner::new(0.0, Direction::Minimize);
}
#[test]
#[should_panic(expected = "percentile must be in (0.0, 100.0)")]
fn new_rejects_hundred() {
let _ = PercentilePruner::new(100.0, Direction::Minimize);
}
fn make_completed_trial(id: u64, values: &[(u64, f64)]) -> CompletedTrial {
use std::collections::HashMap;
use crate::parameter::ParamId;
CompletedTrial::with_intermediate_values(
id,
HashMap::<ParamId, crate::ParamValue>::new(),
HashMap::new(),
HashMap::new(),
0.0,
values.to_vec(),
HashMap::new(),
)
}
#[test]
fn percentile_50_matches_median_behavior() {
let pruner = PercentilePruner::new(50.0, Direction::Minimize);
let completed = vec![
make_completed_trial(0, &[(0, 1.0), (1, 2.0)]),
make_completed_trial(1, &[(0, 3.0), (1, 4.0)]),
make_completed_trial(2, &[(0, 5.0), (1, 6.0)]),
];
// Median at step 1 is 4.0
// Value 5.0 > 4.0 → prune
assert!(pruner.should_prune(3, 1, &[(0, 3.0), (1, 5.0)], &completed));
// Value 3.0 < 4.0 → keep
assert!(!pruner.should_prune(3, 1, &[(0, 3.0), (1, 3.0)], &completed));
}
#[test]
fn percentile_25_is_more_aggressive() {
let pruner_25 = PercentilePruner::new(25.0, Direction::Minimize);
let pruner_75 = PercentilePruner::new(75.0, Direction::Minimize);
let completed = vec![
make_completed_trial(0, &[(0, 1.0)]),
make_completed_trial(1, &[(0, 2.0)]),
make_completed_trial(2, &[(0, 3.0)]),
make_completed_trial(3, &[(0, 4.0)]),
];
// 25th percentile at step 0: 1.75
// 75th percentile at step 0: 3.25
// Value 2.5: above 25th (prune), below 75th (keep)
assert!(pruner_25.should_prune(4, 0, &[(0, 2.5)], &completed));
assert!(!pruner_75.should_prune(4, 0, &[(0, 2.5)], &completed));
}
#[test]
fn warmup_prevents_pruning() {
let pruner = PercentilePruner::new(50.0, Direction::Minimize).n_warmup_steps(5);
let completed = vec![make_completed_trial(0, &[(0, 1.0)])];
// Step 3 < warmup 5 → no prune even with bad value
assert!(!pruner.should_prune(1, 3, &[(3, 100.0)], &completed));
}
#[test]
fn n_min_trials_prevents_pruning() {
let pruner = PercentilePruner::new(50.0, Direction::Minimize).n_min_trials(5);
let completed = vec![
make_completed_trial(0, &[(0, 1.0)]),
make_completed_trial(1, &[(0, 2.0)]),
];
// Only 2 trials, need 5 → no prune
assert!(!pruner.should_prune(2, 0, &[(0, 100.0)], &completed));
}
#[test]
fn maximize_direction() {
let pruner = PercentilePruner::new(50.0, Direction::Maximize);
let completed = vec![
make_completed_trial(0, &[(0, 1.0)]),
make_completed_trial(1, &[(0, 3.0)]),
make_completed_trial(2, &[(0, 5.0)]),
];
// Median at step 0 is 3.0
// Value 2.0 < 3.0 → prune (maximize wants higher)
assert!(pruner.should_prune(3, 0, &[(0, 2.0)], &completed));
// Value 4.0 > 3.0 → keep
assert!(!pruner.should_prune(3, 0, &[(0, 4.0)], &completed));
}
#[test]
fn near_boundary_percentiles() {
let pruner_low = PercentilePruner::new(1.0, Direction::Minimize);
let pruner_high = PercentilePruner::new(99.0, Direction::Minimize);
let completed = vec![
make_completed_trial(0, &[(0, 1.0)]),
make_completed_trial(1, &[(0, 2.0)]),
make_completed_trial(2, &[(0, 3.0)]),
make_completed_trial(3, &[(0, 100.0)]),
];
// Percentile 1 is very aggressive (threshold near 1.0)
// Value 1.5 should be pruned
assert!(pruner_low.should_prune(4, 0, &[(0, 1.5)], &completed));
// Percentile 99 is very lenient (threshold near 100.0)
// Value 50.0 should not be pruned
assert!(!pruner_high.should_prune(4, 0, &[(0, 50.0)], &completed));
}
}
+466
View File
@@ -0,0 +1,466 @@
use super::Pruner;
use crate::sampler::CompletedTrial;
use crate::types::{Direction, TrialState};
/// Successive Halving pruner based on the SHA algorithm.
///
/// Trials are evaluated at exponentially-spaced "rungs". At each rung,
/// only the top 1/eta fraction of trials survive to the next rung.
///
/// For example, with `min_resource=1`, `max_resource=81`, `reduction_factor=3`:
/// - Rung 0: evaluate at step 1, keep top 1/3
/// - Rung 1: evaluate at step 3, keep top 1/3
/// - Rung 2: evaluate at step 9, keep top 1/3
/// - Rung 3: evaluate at step 27, keep top 1/3
/// - Rung 4: evaluate at step 81 (full budget)
///
/// # Examples
///
/// ```
/// use optimizer::Direction;
/// use optimizer::pruner::SuccessiveHalvingPruner;
///
/// let pruner = SuccessiveHalvingPruner::new()
/// .min_resource(1)
/// .max_resource(81)
/// .reduction_factor(3)
/// .direction(Direction::Minimize);
/// ```
pub struct SuccessiveHalvingPruner {
min_resource: u64,
max_resource: u64,
reduction_factor: u64,
min_early_stopping_rate: u64,
direction: Direction,
}
impl SuccessiveHalvingPruner {
/// Create a new `SuccessiveHalvingPruner` with default parameters.
///
/// Defaults: `min_resource=1`, `max_resource=81`, `reduction_factor=3`,
/// `min_early_stopping_rate=0`, `direction=Minimize`.
#[must_use]
pub fn new() -> Self {
Self {
min_resource: 1,
max_resource: 81,
reduction_factor: 3,
min_early_stopping_rate: 0,
direction: Direction::Minimize,
}
}
/// Set the minimum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn min_resource(mut self, r: u64) -> Self {
assert!(r > 0, "min_resource must be > 0, got {r}");
self.min_resource = r;
self
}
/// Set the maximum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn max_resource(mut self, r: u64) -> Self {
assert!(r > 0, "max_resource must be > 0, got {r}");
self.max_resource = r;
self
}
/// Set the reduction factor (eta). At each rung, the top 1/eta trials survive.
///
/// # Panics
///
/// Panics if `eta` is less than 2.
#[must_use]
pub fn reduction_factor(mut self, eta: u64) -> Self {
assert!(eta >= 2, "reduction_factor must be >= 2, got {eta}");
self.reduction_factor = eta;
self
}
/// Set the minimum early stopping rate. Skips the first N rungs.
#[must_use]
pub fn min_early_stopping_rate(mut self, n: u64) -> Self {
self.min_early_stopping_rate = n;
self
}
/// Set the optimization direction.
#[must_use]
pub fn direction(mut self, d: Direction) -> Self {
self.direction = d;
self
}
/// Compute the rung steps: `[min_resource * eta^(s), ...]` up to `max_resource`,
/// skipping the first `min_early_stopping_rate` rungs.
fn rung_steps(&self) -> Vec<u64> {
let eta = self.reduction_factor;
let mut steps = Vec::new();
let mut rung: u32 = 0;
while let Some(power) = eta.checked_pow(rung) {
let step = self.min_resource.saturating_mul(power);
if step > self.max_resource {
break;
}
if u64::from(rung) >= self.min_early_stopping_rate {
steps.push(step);
}
rung += 1;
}
steps
}
}
impl Default for SuccessiveHalvingPruner {
fn default() -> Self {
Self::new()
}
}
#[allow(clippy::cast_precision_loss)]
impl Pruner for SuccessiveHalvingPruner {
fn should_prune(
&self,
_trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
let rungs = self.rung_steps();
// Find the highest rung step <= current step
let Some(&rung_step) = rungs.iter().rev().find(|&&r| r <= step) else {
// No rung matches (before the first rung) → don't prune
return false;
};
// If this is the last rung (full budget), don't prune
if rung_step >= self.max_resource {
return false;
}
// Get the current trial's value at this rung step
let Some(&(_, current_value)) = intermediate_values.iter().find(|(s, _)| *s == rung_step)
else {
// Trial hasn't reported a value at this exact rung step.
// Use the latest intermediate value at or before the rung step instead.
let Some(&(_, current_value)) = intermediate_values
.iter()
.rev()
.find(|(s, _)| *s <= rung_step)
else {
return false;
};
return self.is_pruned_at_rung(current_value, rung_step, completed_trials);
};
self.is_pruned_at_rung(current_value, rung_step, completed_trials)
}
}
impl SuccessiveHalvingPruner {
/// Determine whether a trial with `current_value` should be pruned at the given rung.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn is_pruned_at_rung(
&self,
current_value: f64,
rung_step: u64,
completed_trials: &[CompletedTrial],
) -> bool {
let eta = self.reduction_factor as usize;
// Collect values at this rung step from all trials that reached it
let mut values_at_rung: Vec<f64> = completed_trials
.iter()
.filter(|t| t.state == TrialState::Complete || t.state == TrialState::Pruned)
.filter_map(|t| {
t.intermediate_values
.iter()
.find(|(s, _)| *s == rung_step)
.map(|(_, v)| *v)
})
.collect();
// Need at least eta trials to make a meaningful comparison
// (with fewer trials, we can't determine the top 1/eta fraction)
if values_at_rung.len() < eta {
return false;
}
// Include the current trial's value for ranking
values_at_rung.push(current_value);
// Sort based on direction: best values first
values_at_rung
.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(core::cmp::Ordering::Equal));
if self.direction == Direction::Maximize {
values_at_rung.reverse();
}
// Keep top 1/eta fraction
let n_keep = (values_at_rung.len() as f64 / eta as f64).ceil() as usize;
let threshold_idx = n_keep.max(1) - 1;
let threshold = values_at_rung[threshold_idx];
// Prune if current value is worse than the threshold
match self.direction {
Direction::Minimize => current_value > threshold,
Direction::Maximize => current_value < threshold,
}
}
}
#[cfg(test)]
#[allow(clippy::cast_precision_loss)]
mod tests {
use super::*;
fn make_trial(id: u64, values: &[(u64, f64)]) -> CompletedTrial {
use std::collections::HashMap;
use crate::parameter::ParamId;
CompletedTrial::with_intermediate_values(
id,
HashMap::<ParamId, crate::ParamValue>::new(),
HashMap::new(),
HashMap::new(),
0.0,
values.to_vec(),
HashMap::new(),
)
}
fn make_pruned_trial(id: u64, values: &[(u64, f64)]) -> CompletedTrial {
let mut t = make_trial(id, values);
t.state = TrialState::Pruned;
t
}
#[test]
fn rung_steps_default() {
let pruner = SuccessiveHalvingPruner::new();
let rungs = pruner.rung_steps();
// min=1, max=81, eta=3 → 1, 3, 9, 27, 81
assert_eq!(rungs, vec![1, 3, 9, 27, 81]);
}
#[test]
fn rung_steps_custom() {
let pruner = SuccessiveHalvingPruner::new()
.min_resource(2)
.max_resource(32)
.reduction_factor(2);
let rungs = pruner.rung_steps();
// 2, 4, 8, 16, 32
assert_eq!(rungs, vec![2, 4, 8, 16, 32]);
}
#[test]
fn rung_steps_with_early_stopping_rate() {
let pruner = SuccessiveHalvingPruner::new().min_early_stopping_rate(2);
let rungs = pruner.rung_steps();
// Skip rung 0 (step=1) and rung 1 (step=3), keep rung 2+ (9, 27, 81)
assert_eq!(rungs, vec![9, 27, 81]);
}
#[test]
fn no_prune_before_first_rung() {
let pruner = SuccessiveHalvingPruner::new()
.min_resource(10)
.max_resource(100)
.reduction_factor(3);
let completed = vec![
make_trial(0, &[(5, 1.0)]),
make_trial(1, &[(5, 2.0)]),
make_trial(2, &[(5, 3.0)]),
];
// Step 5 is before the first rung (10)
assert!(!pruner.should_prune(3, 5, &[(5, 100.0)], &completed));
}
#[test]
fn no_prune_with_single_trial() {
let pruner = SuccessiveHalvingPruner::new();
let completed = vec![make_trial(0, &[(1, 5.0)]), make_trial(1, &[(1, 3.0)])];
// Only 2 completed trials at rung + 1 current = 3 total, threshold = ceil(3/3) = 1
// With eta=3, we need at least 3 completed trials
assert!(!pruner.should_prune(2, 1, &[(1, 10.0)], &completed));
}
#[test]
fn prune_worst_trials_at_rung() {
let pruner = SuccessiveHalvingPruner::new().direction(Direction::Minimize);
// 9 completed trials at rung step=1, with values 1..=9
let completed: Vec<_> = (0..9)
.map(|i| make_trial(i, &[(1, (i + 1) as f64)]))
.collect();
// With eta=3, keep top 1/3. 10 total values → ceil(10/3) = 4 kept
// Best 4 values: 1, 2, 3, 4. Threshold = 4.0
// Value 3.0 → keep (in top 1/3)
assert!(!pruner.should_prune(9, 1, &[(1, 3.0)], &completed));
// Value 5.0 → prune (not in top 1/3)
assert!(pruner.should_prune(9, 1, &[(1, 5.0)], &completed));
}
#[test]
fn top_fraction_survives() {
let pruner = SuccessiveHalvingPruner::new().direction(Direction::Minimize);
// 6 completed trials at step=1
let completed: Vec<_> = (0..6)
.map(|i| make_trial(i, &[(1, (i + 1) as f64)]))
.collect();
// 7 total (6 + current). ceil(7/3) = 3 keep. Threshold = 3.0
// Value 2.0 → keep
assert!(!pruner.should_prune(6, 1, &[(1, 2.0)], &completed));
// Value 3.0 → keep (at threshold)
assert!(!pruner.should_prune(6, 1, &[(1, 3.0)], &completed));
// Value 4.0 → prune
assert!(pruner.should_prune(6, 1, &[(1, 4.0)], &completed));
}
#[test]
fn maximize_direction() {
let pruner = SuccessiveHalvingPruner::new().direction(Direction::Maximize);
let completed: Vec<_> = (0..6)
.map(|i| make_trial(i, &[(1, (i + 1) as f64)]))
.collect();
// 7 total. For maximize, best = highest. ceil(7/3)=3. Top 3: 6,5,4. Threshold=4.0
// Value 5.0 → keep
assert!(!pruner.should_prune(6, 1, &[(1, 5.0)], &completed));
// Value 4.0 → keep (at threshold)
assert!(!pruner.should_prune(6, 1, &[(1, 4.0)], &completed));
// Value 3.0 → prune
assert!(pruner.should_prune(6, 1, &[(1, 3.0)], &completed));
}
#[test]
fn reduction_factor_2() {
let pruner = SuccessiveHalvingPruner::new()
.reduction_factor(2)
.min_resource(1)
.max_resource(16)
.direction(Direction::Minimize);
// Rungs: 1, 2, 4, 8, 16
assert_eq!(pruner.rung_steps(), vec![1, 2, 4, 8, 16]);
// 4 completed trials at rung step=1
let completed: Vec<_> = (0..4)
.map(|i| make_trial(i, &[(1, (i + 1) as f64)]))
.collect();
// With eta=2, 5 total. ceil(5/2) = 3 keep. Threshold = 3.0
// Value 3.0 → keep
assert!(!pruner.should_prune(4, 1, &[(1, 3.0)], &completed));
// Value 4.0 → prune
assert!(pruner.should_prune(4, 1, &[(1, 4.0)], &completed));
}
#[test]
fn reduction_factor_4() {
let pruner = SuccessiveHalvingPruner::new()
.reduction_factor(4)
.min_resource(1)
.max_resource(64)
.direction(Direction::Minimize);
// Rungs: 1, 4, 16, 64
assert_eq!(pruner.rung_steps(), vec![1, 4, 16, 64]);
// 12 completed trials at rung step=1
let completed: Vec<_> = (0..12)
.map(|i| make_trial(i, &[(1, (i + 1) as f64)]))
.collect();
// With eta=4, 13 total. ceil(13/4) = 4 keep. Threshold = 4.0
// Value 4.0 → keep
assert!(!pruner.should_prune(12, 1, &[(1, 4.0)], &completed));
// Value 5.0 → prune
assert!(pruner.should_prune(12, 1, &[(1, 5.0)], &completed));
}
#[test]
fn non_contiguous_steps() {
let pruner = SuccessiveHalvingPruner::new().direction(Direction::Minimize);
// Trials reporting at rung step=3 (not step=1)
let completed: Vec<_> = (0..6)
.map(|i| make_trial(i, &[(3, (i + 1) as f64)]))
.collect();
// Current trial reports at step 5 (between rung 3 and rung 9)
// Highest rung <= 5 is 3. Use value at rung step 3.
// Trial has value at step 3 → use it
assert!(!pruner.should_prune(6, 5, &[(3, 2.0)], &completed));
assert!(pruner.should_prune(6, 5, &[(3, 5.0)], &completed));
}
#[test]
fn no_prune_at_max_resource() {
let pruner = SuccessiveHalvingPruner::new();
let completed: Vec<_> = (0..9)
.map(|i| make_trial(i, &[(81, (i + 1) as f64)]))
.collect();
// At the max resource rung, never prune (trial should complete)
assert!(!pruner.should_prune(9, 81, &[(81, 100.0)], &completed));
}
#[test]
fn includes_pruned_trials_in_comparison() {
let pruner = SuccessiveHalvingPruner::new().direction(Direction::Minimize);
// Mix of completed and pruned trials at rung step=1
let completed = vec![
make_trial(0, &[(1, 1.0)]),
make_trial(1, &[(1, 2.0)]),
make_pruned_trial(2, &[(1, 8.0)]),
make_pruned_trial(3, &[(1, 9.0)]),
make_pruned_trial(4, &[(1, 10.0)]),
];
// 6 total. ceil(6/3) = 2 keep. Threshold = 2.0
// Value 2.0 → keep
assert!(!pruner.should_prune(5, 1, &[(1, 2.0)], &completed));
// Value 3.0 → prune
assert!(pruner.should_prune(5, 1, &[(1, 3.0)], &completed));
}
#[test]
#[should_panic(expected = "min_resource must be > 0")]
fn rejects_zero_min_resource() {
let _ = SuccessiveHalvingPruner::new().min_resource(0);
}
#[test]
#[should_panic(expected = "max_resource must be > 0")]
fn rejects_zero_max_resource() {
let _ = SuccessiveHalvingPruner::new().max_resource(0);
}
#[test]
#[should_panic(expected = "reduction_factor must be >= 2")]
fn rejects_reduction_factor_one() {
let _ = SuccessiveHalvingPruner::new().reduction_factor(1);
}
}
+83
View File
@@ -0,0 +1,83 @@
use super::Pruner;
use crate::sampler::CompletedTrial;
/// Prune trials whose intermediate values exceed fixed thresholds.
///
/// Useful for cutting off trials that are clearly diverging or stuck
/// at bad values early in training.
///
/// # Examples
///
/// ```
/// use optimizer::pruner::ThresholdPruner;
///
/// // Prune if the intermediate value exceeds 100.0 or falls below 0.0
/// let pruner = ThresholdPruner::new().upper(100.0).lower(0.0);
/// ```
pub struct ThresholdPruner {
/// Prune if intermediate value is greater than this. `None` = no upper bound.
upper: Option<f64>,
/// Prune if intermediate value is less than this. `None` = no lower bound.
lower: Option<f64>,
}
impl ThresholdPruner {
/// Create a new `ThresholdPruner` with no thresholds set.
///
/// By default, no pruning occurs. Use [`upper`](Self::upper) and
/// [`lower`](Self::lower) to set bounds.
#[must_use]
pub fn new() -> Self {
Self {
upper: None,
lower: None,
}
}
/// Set the upper threshold. Trials with intermediate values above this
/// will be pruned.
#[must_use]
pub fn upper(mut self, threshold: f64) -> Self {
self.upper = Some(threshold);
self
}
/// Set the lower threshold. Trials with intermediate values below this
/// will be pruned.
#[must_use]
pub fn lower(mut self, threshold: f64) -> Self {
self.lower = Some(threshold);
self
}
}
impl Default for ThresholdPruner {
fn default() -> Self {
Self::new()
}
}
impl Pruner for ThresholdPruner {
fn should_prune(
&self,
_trial_id: u64,
_step: u64,
intermediate_values: &[(u64, f64)],
_completed_trials: &[CompletedTrial],
) -> bool {
let Some(&(_, latest_value)) = intermediate_values.last() else {
return false;
};
if let Some(upper) = self.upper
&& latest_value > upper
{
return true;
}
if let Some(lower) = self.lower
&& latest_value < lower
{
return true;
}
false
}
}
+531
View File
@@ -0,0 +1,531 @@
use core::cmp::Ordering;
use super::Pruner;
use crate::sampler::CompletedTrial;
use crate::types::{Direction, TrialState};
/// Prune trials using a Wilcoxon signed-rank test comparing intermediate
/// values against the best completed trial.
///
/// More principled than `MedianPruner` for noisy objectives — it accounts
/// for the paired nature of step-aligned comparisons and doesn't prune
/// on random fluctuations.
///
/// The test compares intermediate values at matching steps between the
/// current trial and the best completed trial. If the current trial is
/// statistically significantly worse (p < threshold), it is pruned.
///
/// # Examples
///
/// ```
/// use optimizer::Direction;
/// use optimizer::pruner::WilcoxonPruner;
///
/// let pruner = WilcoxonPruner::new(Direction::Minimize)
/// .p_value_threshold(0.05)
/// .n_warmup_steps(5)
/// .n_min_trials(1);
/// ```
pub struct WilcoxonPruner {
/// Significance level (default 0.05). Lower = more conservative.
p_value_threshold: f64,
/// Don't prune in the first N steps (let the trial warm up).
n_warmup_steps: u64,
/// Require at least N completed trials before pruning.
n_min_trials: usize,
/// The optimization direction.
direction: Direction,
}
impl WilcoxonPruner {
/// Create a new `WilcoxonPruner` for the given optimization direction.
///
/// By default, `p_value_threshold` is 0.05, `n_warmup_steps` is 0,
/// and `n_min_trials` is 1.
#[must_use]
pub fn new(direction: Direction) -> Self {
Self {
p_value_threshold: 0.05,
n_warmup_steps: 0,
n_min_trials: 1,
direction,
}
}
/// Set the p-value threshold for significance.
///
/// Must be in (0.0, 1.0). Lower values are more conservative (harder to prune).
///
/// # Panics
///
/// Panics if `p` is not in the open interval (0.0, 1.0).
#[must_use]
pub fn p_value_threshold(mut self, p: f64) -> Self {
assert!(
p > 0.0 && p < 1.0,
"p_value_threshold must be in (0.0, 1.0)"
);
self.p_value_threshold = p;
self
}
/// Set the number of warmup steps. No pruning occurs before this step.
#[must_use]
pub fn n_warmup_steps(mut self, n: u64) -> Self {
self.n_warmup_steps = n;
self
}
/// Set the minimum number of completed trials required before pruning.
#[must_use]
pub fn n_min_trials(mut self, n: usize) -> Self {
self.n_min_trials = n;
self
}
}
impl Pruner for WilcoxonPruner {
fn should_prune(
&self,
_trial_id: u64,
step: u64,
intermediate_values: &[(u64, f64)],
completed_trials: &[CompletedTrial],
) -> bool {
if step < self.n_warmup_steps {
return false;
}
let completed: Vec<&CompletedTrial> = completed_trials
.iter()
.filter(|t| t.state == TrialState::Complete)
.collect();
if completed.len() < self.n_min_trials {
return false;
}
// Find the best completed trial by final objective value.
let best = match self.direction {
Direction::Minimize => completed
.iter()
.min_by(|a, b| a.value.partial_cmp(&b.value).unwrap_or(Ordering::Equal)),
Direction::Maximize => completed
.iter()
.max_by(|a, b| a.value.partial_cmp(&b.value).unwrap_or(Ordering::Equal)),
};
let Some(best) = best else {
return false;
};
// Pair intermediate values at matching steps.
let pairs: Vec<(f64, f64)> = intermediate_values
.iter()
.filter_map(|&(s, current_v)| {
best.intermediate_values
.iter()
.find(|(bs, _)| *bs == s)
.map(|&(_, best_v)| (current_v, best_v))
})
.collect();
// Need at least 6 pairs for a meaningful test.
if pairs.len() < 6 {
return false;
}
// Compute signed differences: current - best.
// For minimization: positive diff means current is worse.
// For maximization: negative diff means current is worse.
let differences: Vec<f64> = pairs
.iter()
.map(|&(current, best_v)| current - best_v)
.collect();
// Run the Wilcoxon signed-rank test.
let p_value = wilcoxon_signed_rank_test(&differences, self.direction);
p_value < self.p_value_threshold
}
}
/// Perform a one-sided Wilcoxon signed-rank test.
///
/// Tests whether the values tend to be worse than zero (positive for
/// minimization, negative for maximization).
///
/// Returns a p-value. Small p-values indicate the current trial is
/// significantly worse.
fn wilcoxon_signed_rank_test(differences: &[f64], direction: Direction) -> f64 {
// 1. Remove zero differences.
let nonzero: Vec<f64> = differences.iter().copied().filter(|d| *d != 0.0).collect();
let n = nonzero.len();
if n < 6 {
return 1.0; // Not enough data
}
// 2. Rank by absolute value.
let mut abs_ranked: Vec<(usize, f64, f64)> = nonzero
.iter()
.enumerate()
.map(|(i, &d)| (i, d.abs(), d))
.collect();
abs_ranked.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal));
// 3. Assign ranks with tie correction.
let ranks = assign_ranks(&abs_ranked);
// 4. Compute W+ (sum of ranks for positive differences) and
// W- (sum of ranks for negative differences).
let mut w_plus = 0.0;
let mut w_minus = 0.0;
for (i, &(_, _, orig)) in abs_ranked.iter().enumerate() {
if orig > 0.0 {
w_plus += ranks[i];
} else {
w_minus += ranks[i];
}
}
// For one-sided test:
// - Minimization: we want to detect positive diffs (current worse).
// Large W+ means significantly worse. Test statistic = W-.
// - Maximization: we want to detect negative diffs (current worse).
// Large W- means significantly worse. Test statistic = W+.
let w = match direction {
Direction::Minimize => w_minus,
Direction::Maximize => w_plus,
};
// 5. Normal approximation for the p-value.
#[allow(clippy::cast_precision_loss)]
let n_f = n as f64;
let mean = n_f * (n_f + 1.0) / 4.0;
let variance = n_f * (n_f + 1.0) * (2.0 * n_f + 1.0) / 24.0;
// Tie correction for variance.
let tie_correction = compute_tie_correction(&ranks);
let adjusted_variance = variance - tie_correction;
if adjusted_variance <= 0.0 {
return 1.0;
}
let std_dev = adjusted_variance.sqrt();
// Continuity correction: shift W by 0.5 towards mean.
let continuity = if w < mean { 0.5 } else { -0.5 };
let z = (w + continuity - mean) / std_dev;
// One-sided p-value (lower tail): probability that the test statistic
// is this small or smaller under H0.
normal_cdf(z)
}
/// Assign average ranks, handling ties.
fn assign_ranks(sorted: &[(usize, f64, f64)]) -> Vec<f64> {
let n = sorted.len();
let mut ranks = vec![0.0; n];
let mut i = 0;
while i < n {
let mut j = i;
// Find all items tied with sorted[i].
while j < n
&& (sorted[j].1 - sorted[i].1).abs() < f64::EPSILON * sorted[i].1.max(1.0) * 100.0
{
j += 1;
}
// Average rank for the tie group. Ranks are 1-based.
#[allow(clippy::cast_precision_loss)]
let avg_rank = (i + 1 + j) as f64 / 2.0;
for rank in ranks.iter_mut().take(j).skip(i) {
*rank = avg_rank;
}
i = j;
}
ranks
}
/// Compute the tie correction term for the variance.
/// For each tie group of size t, subtract t^3 - t from the sum,
/// then divide by 48.
fn compute_tie_correction(ranks: &[f64]) -> f64 {
let mut correction = 0.0;
let mut i = 0;
while i < ranks.len() {
let mut j = i;
while j < ranks.len() && (ranks[j] - ranks[i]).abs() < f64::EPSILON {
j += 1;
}
#[allow(clippy::cast_precision_loss)]
let t = (j - i) as f64;
if t > 1.0 {
correction += t * t * t - t;
}
i = j;
}
correction / 48.0
}
/// Standard normal CDF using an approximation (Abramowitz & Stegun).
fn normal_cdf(x: f64) -> f64 {
// Use the complementary error function relationship:
// Φ(x) = 0.5 * erfc(-x / √2)
0.5 * erfc(-x / core::f64::consts::SQRT_2)
}
/// Complementary error function approximation.
/// Maximum error: 1.5 × 10⁻⁷ (Abramowitz & Stegun formula 7.1.26).
fn erfc(x: f64) -> f64 {
let t = 1.0 / (1.0 + 0.327_591_1 * x.abs());
let poly = t
* (0.254_829_592
+ t * (-0.284_496_736
+ t * (1.421_413_741 + t * (-1.453_152_027 + t * 1.061_405_429))));
let result = poly * (-x * x).exp();
if x >= 0.0 { result } else { 2.0 - result }
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
fn trial_with_values(
id: u64,
value: f64,
intermediate_values: Vec<(u64, f64)>,
) -> CompletedTrial {
CompletedTrial::with_intermediate_values(
id,
HashMap::new(),
HashMap::new(),
HashMap::new(),
value,
intermediate_values,
HashMap::new(),
)
}
#[test]
fn no_prune_during_warmup() {
let pruner = WilcoxonPruner::new(Direction::Minimize).n_warmup_steps(10);
let completed = vec![trial_with_values(
0,
0.1,
(0..20).map(|s| (s, 0.1)).collect(),
)];
let current: Vec<(u64, f64)> = (0..8).map(|s| (s, 100.0)).collect();
assert!(!pruner.should_prune(1, 7, &current, &completed));
}
#[test]
fn no_prune_with_insufficient_trials() {
let pruner = WilcoxonPruner::new(Direction::Minimize).n_min_trials(5);
let completed = vec![trial_with_values(
0,
0.1,
(0..20).map(|s| (s, 0.1)).collect(),
)];
let current: Vec<(u64, f64)> = (0..10).map(|s| (s, 100.0)).collect();
assert!(!pruner.should_prune(1, 9, &current, &completed));
}
#[test]
fn no_prune_with_fewer_than_6_pairs() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
let completed = vec![trial_with_values(
0,
0.1,
(0..5).map(|s| (s, 0.1)).collect(),
)];
// Only 5 matching steps
let current: Vec<(u64, f64)> = (0..5).map(|s| (s, 100.0)).collect();
assert!(!pruner.should_prune(1, 4, &current, &completed));
}
#[test]
fn prune_when_consistently_worse_minimize() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
// Best trial has low values.
let best_values: Vec<(u64, f64)> = (0..20).map(|s| (s, 0.1)).collect();
let completed = vec![trial_with_values(0, 0.1, best_values)];
// Current trial is consistently much worse.
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 10.0)).collect();
assert!(pruner.should_prune(1, 19, &current, &completed));
}
#[test]
fn prune_when_consistently_worse_maximize() {
let pruner = WilcoxonPruner::new(Direction::Maximize);
// Best trial has high values.
let best_values: Vec<(u64, f64)> = (0..20).map(|s| (s, 10.0)).collect();
let completed = vec![trial_with_values(0, 10.0, best_values)];
// Current trial is consistently much worse.
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 0.1)).collect();
assert!(pruner.should_prune(1, 19, &current, &completed));
}
#[test]
fn no_prune_when_statistically_similar() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
// Best trial and current trial have very similar values with noise.
let best_values: Vec<(u64, f64)> = (0..20_u64)
.map(|s| {
let noise = if s.is_multiple_of(2) { 0.01 } else { -0.01 };
(s, 1.0 + noise)
})
.collect();
let completed = vec![trial_with_values(0, 1.0, best_values)];
// Current trial is similar — alternating above/below.
let current: Vec<(u64, f64)> = (0..20_u64)
.map(|s| {
let noise = if s.is_multiple_of(2) { -0.01 } else { 0.01 };
(s, 1.0 + noise)
})
.collect();
assert!(!pruner.should_prune(1, 19, &current, &completed));
}
#[test]
fn selects_best_trial_minimize() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
// Two completed trials: trial 0 is better (lower).
let completed = vec![
trial_with_values(0, 0.1, (0..20).map(|s| (s, 0.1)).collect()),
trial_with_values(1, 5.0, (0..20).map(|s| (s, 5.0)).collect()),
];
// Current trial is worse than the best but similar to the second.
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 5.0)).collect();
assert!(pruner.should_prune(2, 19, &current, &completed));
}
#[test]
fn selects_best_trial_maximize() {
let pruner = WilcoxonPruner::new(Direction::Maximize);
// Two completed trials: trial 1 is better (higher).
let completed = vec![
trial_with_values(0, 0.1, (0..20).map(|s| (s, 0.1)).collect()),
trial_with_values(1, 10.0, (0..20).map(|s| (s, 10.0)).collect()),
];
// Current trial is worse than the best.
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 0.1)).collect();
assert!(pruner.should_prune(2, 19, &current, &completed));
}
#[test]
fn ignores_pruned_trials() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
// Only a pruned trial — no complete trials.
let mut trial = trial_with_values(0, 0.1, (0..20).map(|s| (s, 0.1)).collect());
trial.state = TrialState::Pruned;
let completed = vec![trial];
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 100.0)).collect();
assert!(!pruner.should_prune(1, 19, &current, &completed));
}
#[test]
fn lower_p_value_is_more_conservative() {
let strict = WilcoxonPruner::new(Direction::Minimize).p_value_threshold(0.001);
let lenient = WilcoxonPruner::new(Direction::Minimize).p_value_threshold(0.1);
let completed = vec![trial_with_values(
0,
0.1,
(0..20).map(|s| (s, 0.1)).collect(),
)];
// Moderately worse — should pass lenient but maybe not strict.
let current: Vec<(u64, f64)> = (0..20)
.map(|s| if s < 15 { (s, 0.2) } else { (s, 0.15) })
.collect();
let lenient_prunes = lenient.should_prune(1, 19, &current, &completed);
let strict_prunes = strict.should_prune(1, 19, &current, &completed);
// A stricter threshold should never prune when a lenient one doesn't.
if !lenient_prunes {
assert!(!strict_prunes);
}
}
#[test]
#[should_panic(expected = "p_value_threshold must be in (0.0, 1.0)")]
fn panics_on_zero_p_value() {
let _ = WilcoxonPruner::new(Direction::Minimize).p_value_threshold(0.0);
}
#[test]
#[should_panic(expected = "p_value_threshold must be in (0.0, 1.0)")]
fn panics_on_one_p_value() {
let _ = WilcoxonPruner::new(Direction::Minimize).p_value_threshold(1.0);
}
#[test]
fn correct_signed_rank_statistic() {
// Known example: differences [1, 2, 3, 4, 5, 6] (all positive).
// Ranks: 1, 2, 3, 4, 5, 6. W+ = 21, W- = 0.
// For minimization (testing if positive = worse), W- = 0.
// This should give a very small p-value.
let diffs = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let p = wilcoxon_signed_rank_test(&diffs, Direction::Minimize);
assert!(
p < 0.05,
"p-value {p} should be < 0.05 for all-positive diffs"
);
}
#[test]
fn symmetric_differences_not_significant() {
// Balanced differences: half positive, half negative.
let diffs = vec![1.0, -1.0, 2.0, -2.0, 3.0, -3.0, 4.0, -4.0];
let p = wilcoxon_signed_rank_test(&diffs, Direction::Minimize);
assert!(p > 0.05, "p-value {p} should be > 0.05 for symmetric diffs");
}
#[test]
fn normal_cdf_known_values() {
assert!((normal_cdf(0.0) - 0.5).abs() < 1e-6);
assert!(normal_cdf(-10.0) < 1e-6);
assert!((normal_cdf(10.0) - 1.0).abs() < 1e-6);
assert!((normal_cdf(-1.96) - 0.025).abs() < 0.001);
}
#[test]
fn no_intermediate_values() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
let completed = vec![trial_with_values(
0,
0.1,
(0..20).map(|s| (s, 0.1)).collect(),
)];
assert!(!pruner.should_prune(1, 0, &[], &completed));
}
#[test]
fn no_completed_trials() {
let pruner = WilcoxonPruner::new(Direction::Minimize);
let current: Vec<(u64, f64)> = (0..20).map(|s| (s, 1.0)).collect();
assert!(!pruner.should_prune(1, 19, &current, &[]));
}
}
+691
View File
@@ -0,0 +1,691 @@
//! BOHB (Bayesian Optimization + `HyperBand`) sampler.
//!
//! BOHB combines TPE's model-guided sampling with Hyperband's budget-aware
//! evaluation. Instead of building one global TPE model, BOHB conditions
//! its TPE model on trials evaluated at a specific budget level, giving
//! better-calibrated proposals for each rung of the Hyperband schedule.
//!
//! # How it works
//!
//! 1. Compute all Hyperband rung steps (budget levels) from the config.
//! 2. On each `sample()` call, scan the history's `intermediate_values`
//! to find the **largest budget level** with enough observations
//! (`>= min_points_in_model`).
//! 3. Build a filtered history where each trial's `value` is replaced
//! with its intermediate value at that budget level.
//! 4. Delegate to an internal [`TpeSampler`] for the actual sampling.
//! 5. Fall back to random sampling if no budget level has enough data.
//!
//! # Examples
//!
//! ```
//! use optimizer::sampler::bohb::BohbSampler;
//! use optimizer::{Direction, Study};
//!
//! let bohb = BohbSampler::new();
//! let pruner = bohb.matching_pruner(Direction::Minimize);
//! let study: Study<f64> = Study::with_sampler_and_pruner(Direction::Minimize, bohb, pruner);
//! ```
//!
//! Using the builder for custom configuration:
//!
//! ```
//! use optimizer::sampler::bohb::BohbSampler;
//!
//! let bohb = BohbSampler::builder()
//! .min_resource(1)
//! .max_resource(81)
//! .reduction_factor(3)
//! .min_points_in_model(10)
//! .seed(42)
//! .build()
//! .unwrap();
//! ```
use crate::distribution::Distribution;
use crate::error::Result;
use crate::param::ParamValue;
use crate::pruner::HyperbandPruner;
use crate::sampler::tpe::TpeSampler;
use crate::sampler::{CompletedTrial, Sampler};
use crate::types::Direction;
/// A BOHB sampler that combines TPE with Hyperband budget awareness.
///
/// BOHB filters trial history by budget level before delegating to TPE,
/// so the surrogate model is conditioned on trials evaluated at the same
/// resource level. This produces better-calibrated parameter proposals
/// than using a single global model across all budgets.
///
/// Use [`BohbSampler::matching_pruner`] to create a [`HyperbandPruner`]
/// with matching parameters.
pub struct BohbSampler {
min_resource: u64,
max_resource: u64,
reduction_factor: u64,
min_points_in_model: usize,
tpe: TpeSampler,
}
impl BohbSampler {
/// Creates a new BOHB sampler with default settings.
///
/// Defaults:
/// - `min_resource`: 1
/// - `max_resource`: 81
/// - `reduction_factor`: 3
/// - `min_points_in_model`: 10
/// - TPE: default settings
#[must_use]
pub fn new() -> Self {
Self {
min_resource: 1,
max_resource: 81,
reduction_factor: 3,
min_points_in_model: 10,
tpe: TpeSampler::new(),
}
}
/// Creates a builder for configuring a BOHB sampler.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::bohb::BohbSampler;
///
/// let sampler = BohbSampler::builder()
/// .min_resource(1)
/// .max_resource(27)
/// .reduction_factor(3)
/// .min_points_in_model(5)
/// .seed(42)
/// .build()
/// .unwrap();
/// ```
#[must_use]
pub fn builder() -> BohbSamplerBuilder {
BohbSamplerBuilder::new()
}
/// Creates a [`HyperbandPruner`] with matching Hyperband parameters.
///
/// This ensures the pruner's budget schedule is consistent with the
/// budget levels used by BOHB for model conditioning.
#[must_use]
pub fn matching_pruner(&self, direction: Direction) -> HyperbandPruner {
HyperbandPruner::new()
.min_resource(self.min_resource)
.max_resource(self.max_resource)
.reduction_factor(self.reduction_factor)
.direction(direction)
}
/// Compute all unique budget levels (rung steps) across all Hyperband brackets.
///
/// Returns sorted ascending.
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn all_budget_levels(&self) -> Vec<u64> {
let eta = self.reduction_factor as f64;
let ratio = self.max_resource as f64 / self.min_resource as f64;
let s_max = (ratio.ln() / eta.ln()).floor() as u64;
let mut levels = Vec::new();
for bracket in 0..=s_max {
let exponent = s_max.saturating_sub(bracket);
let min_resource_bracket =
(self.max_resource as f64 / eta.powi(exponent as i32)).ceil() as u64;
let mut rung: u32 = 0;
while let Some(power) = self.reduction_factor.checked_pow(rung) {
let step = min_resource_bracket.saturating_mul(power);
if step > self.max_resource {
break;
}
levels.push(step);
rung += 1;
}
}
levels.sort_unstable();
levels.dedup();
levels
}
/// Build a filtered history for a specific budget level.
///
/// For each trial that has an intermediate value at the given budget step,
/// creates a new `CompletedTrial` with `value` replaced by the intermediate
/// value at that step.
fn filter_history_for_budget(history: &[CompletedTrial], budget: u64) -> Vec<CompletedTrial> {
history
.iter()
.filter_map(|trial| {
trial
.intermediate_values
.iter()
.find(|(step, _)| *step == budget)
.map(|(_, iv)| CompletedTrial {
id: trial.id,
params: trial.params.clone(),
distributions: trial.distributions.clone(),
param_labels: trial.param_labels.clone(),
value: *iv,
intermediate_values: trial.intermediate_values.clone(),
state: trial.state,
user_attrs: trial.user_attrs.clone(),
constraints: trial.constraints.clone(),
})
})
.collect()
}
}
impl Default for BohbSampler {
fn default() -> Self {
Self::new()
}
}
impl Sampler for BohbSampler {
fn sample(
&self,
distribution: &Distribution,
trial_id: u64,
history: &[CompletedTrial],
) -> ParamValue {
// Find the largest budget level with enough observations
let levels = self.all_budget_levels();
for &budget in levels.iter().rev() {
let count = history
.iter()
.filter(|t| {
t.intermediate_values
.iter()
.any(|(step, _)| *step == budget)
})
.count();
if count >= self.min_points_in_model {
let filtered = Self::filter_history_for_budget(history, budget);
return self.tpe.sample(distribution, trial_id, &filtered);
}
}
// Not enough data at any budget level: delegate to TPE with empty history
// which triggers its uniform-random startup behavior.
self.tpe.sample(distribution, trial_id, &[])
}
}
/// Builder for configuring a [`BohbSampler`].
///
/// # Examples
///
/// ```
/// use optimizer::sampler::bohb::BohbSamplerBuilder;
///
/// let sampler = BohbSamplerBuilder::new()
/// .min_resource(1)
/// .max_resource(81)
/// .reduction_factor(3)
/// .gamma(0.15)
/// .seed(42)
/// .build()
/// .unwrap();
/// ```
pub struct BohbSamplerBuilder {
min_resource: u64,
max_resource: u64,
reduction_factor: u64,
min_points_in_model: usize,
tpe_builder: crate::sampler::tpe::TpeSamplerBuilder,
}
impl BohbSamplerBuilder {
/// Creates a new builder with default settings.
#[must_use]
pub fn new() -> Self {
Self {
min_resource: 1,
max_resource: 81,
reduction_factor: 3,
min_points_in_model: 10,
tpe_builder: crate::sampler::tpe::TpeSamplerBuilder::new(),
}
}
/// Sets the minimum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn min_resource(mut self, r: u64) -> Self {
assert!(r > 0, "min_resource must be > 0, got {r}");
self.min_resource = r;
self
}
/// Sets the maximum resource (budget) per trial.
///
/// # Panics
///
/// Panics if `r` is 0.
#[must_use]
pub fn max_resource(mut self, r: u64) -> Self {
assert!(r > 0, "max_resource must be > 0, got {r}");
self.max_resource = r;
self
}
/// Sets the reduction factor (eta).
///
/// # Panics
///
/// Panics if `eta` is less than 2.
#[must_use]
pub fn reduction_factor(mut self, eta: u64) -> Self {
assert!(eta >= 2, "reduction_factor must be >= 2, got {eta}");
self.reduction_factor = eta;
self
}
/// Sets the minimum number of observations at a budget level before
/// BOHB uses TPE instead of random sampling.
#[must_use]
pub fn min_points_in_model(mut self, n: usize) -> Self {
self.min_points_in_model = n;
self
}
/// Sets a fixed gamma value for the internal TPE sampler.
#[must_use]
pub fn gamma(mut self, gamma: f64) -> Self {
self.tpe_builder = self.tpe_builder.gamma(gamma);
self
}
/// Sets a custom gamma strategy for the internal TPE sampler.
#[must_use]
pub fn gamma_strategy<G: crate::sampler::tpe::GammaStrategy + 'static>(
mut self,
strategy: G,
) -> Self {
self.tpe_builder = self.tpe_builder.gamma_strategy(strategy);
self
}
/// Sets the number of EI candidates for the internal TPE sampler.
#[must_use]
pub fn n_ei_candidates(mut self, n: usize) -> Self {
self.tpe_builder = self.tpe_builder.n_ei_candidates(n);
self
}
/// Sets a fixed KDE bandwidth for the internal TPE sampler.
#[must_use]
pub fn kde_bandwidth(mut self, bandwidth: f64) -> Self {
self.tpe_builder = self.tpe_builder.kde_bandwidth(bandwidth);
self
}
/// Sets a seed for reproducible sampling.
#[must_use]
pub fn seed(mut self, seed: u64) -> Self {
self.tpe_builder = self.tpe_builder.seed(seed);
self
}
/// Builds the configured [`BohbSampler`].
///
/// # Errors
///
/// Returns an error if the TPE configuration is invalid (e.g. gamma
/// not in (0, 1) or bandwidth not positive).
pub fn build(self) -> Result<BohbSampler> {
let tpe = self.tpe_builder.build()?;
Ok(BohbSampler {
min_resource: self.min_resource,
max_resource: self.max_resource,
reduction_factor: self.reduction_factor,
min_points_in_model: self.min_points_in_model,
tpe,
})
}
}
impl Default for BohbSamplerBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[allow(clippy::cast_precision_loss)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::distribution::{FloatDistribution, IntDistribution};
use crate::parameter::ParamId;
use crate::types::TrialState;
fn make_trial_with_intermediates(
id: u64,
value: f64,
params: Vec<(ParamId, ParamValue, Distribution)>,
intermediate_values: Vec<(u64, f64)>,
) -> CompletedTrial {
let mut param_map = HashMap::new();
let mut dist_map = HashMap::new();
for (param_id, pv, dist) in params {
param_map.insert(param_id, pv);
dist_map.insert(param_id, dist);
}
CompletedTrial {
id,
params: param_map,
distributions: dist_map,
param_labels: HashMap::new(),
value,
intermediate_values,
state: TrialState::Complete,
user_attrs: HashMap::new(),
constraints: Vec::new(),
}
}
#[test]
fn budget_levels_default() {
let bohb = BohbSampler::new();
let levels = bohb.all_budget_levels();
// With min=1, max=81, eta=3:
// bracket 0: 1, 3, 9, 27, 81
// bracket 1: 3, 9, 27, 81
// bracket 2: 9, 27, 81
// bracket 3: 27, 81
// bracket 4: 81
// Unique sorted: [1, 3, 9, 27, 81]
assert_eq!(levels, vec![1, 3, 9, 27, 81]);
}
#[test]
fn budget_levels_eta2() {
let bohb = BohbSampler::builder()
.min_resource(1)
.max_resource(16)
.reduction_factor(2)
.build()
.unwrap();
let levels = bohb.all_budget_levels();
// s_max = floor(ln(16)/ln(2)) = 4
// bracket 0: 1, 2, 4, 8, 16
// bracket 1: 2, 4, 8, 16
// bracket 2: 4, 8, 16
// bracket 3: 8, 16
// bracket 4: 16
// Unique sorted: [1, 2, 4, 8, 16]
assert_eq!(levels, vec![1, 2, 4, 8, 16]);
}
#[test]
fn filter_history_selects_correct_budget() {
let x_id = ParamId::new();
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
let history = vec![
make_trial_with_intermediates(
0,
0.5,
vec![(x_id, ParamValue::Float(0.3), dist.clone())],
vec![(1, 0.9), (3, 0.7), (9, 0.5)],
),
make_trial_with_intermediates(
1,
0.4,
vec![(x_id, ParamValue::Float(0.6), dist.clone())],
vec![(1, 0.8), (3, 0.4)],
),
make_trial_with_intermediates(
2,
0.3,
vec![(x_id, ParamValue::Float(0.1), dist.clone())],
vec![(1, 0.7)],
),
];
// Budget 3: trials 0 and 1 have intermediate values at step 3
let filtered = BohbSampler::filter_history_for_budget(&history, 3);
assert_eq!(filtered.len(), 2);
assert!((filtered[0].value - 0.7).abs() < f64::EPSILON);
assert!((filtered[1].value - 0.4).abs() < f64::EPSILON);
// Budget 9: only trial 0
let filtered = BohbSampler::filter_history_for_budget(&history, 9);
assert_eq!(filtered.len(), 1);
assert!((filtered[0].value - 0.5).abs() < f64::EPSILON);
// Budget 27: nobody
let filtered = BohbSampler::filter_history_for_budget(&history, 27);
assert!(filtered.is_empty());
}
#[test]
fn matching_pruner_has_same_params() {
let bohb = BohbSampler::builder()
.min_resource(2)
.max_resource(64)
.reduction_factor(4)
.build()
.unwrap();
let pruner = bohb.matching_pruner(Direction::Minimize);
// We can't directly inspect HyperbandPruner fields, but we can
// verify it was created without panicking with the same params.
// The pruner's rung steps should match BOHB's budget levels.
// Just verify it doesn't panic.
drop(pruner);
}
#[test]
fn fallback_to_random_when_insufficient_data() {
let bohb = BohbSampler::builder()
.min_points_in_model(10)
.seed(42)
.build()
.unwrap();
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
// Only 3 trials with intermediate values (< min_points_in_model=10)
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..3)
.map(|i| {
make_trial_with_intermediates(
i,
i as f64,
vec![(x_id, ParamValue::Float(i as f64 / 3.0), dist.clone())],
vec![(1, i as f64)],
)
})
.collect();
// Should not panic, should sample within bounds
for trial_id in 0..20 {
let val = bohb.sample(&dist, trial_id, &history);
if let ParamValue::Float(v) = val {
assert!((0.0..=1.0).contains(&v));
} else {
panic!("Expected Float");
}
}
}
#[test]
fn uses_budget_level_when_enough_data() {
let bohb = BohbSampler::builder()
.min_points_in_model(5)
.seed(42)
.build()
.unwrap();
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 10.0,
log_scale: false,
step: None,
});
// Create 20 trials with intermediate values at budget 1.
// Good trials have x near 2.0, bad trials have x far from 2.0.
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
let x = i as f64 / 2.0;
let iv_at_1 = (x - 2.0).powi(2);
make_trial_with_intermediates(
i,
iv_at_1, // final value same as intermediate for simplicity
vec![(x_id, ParamValue::Float(x), dist.clone())],
vec![(1, iv_at_1)],
)
})
.collect();
// Should use TPE on filtered history at budget 1
let val = bohb.sample(&dist, 100, &history);
if let ParamValue::Float(v) = val {
assert!((0.0..=10.0).contains(&v), "Value {v} out of bounds");
} else {
panic!("Expected Float");
}
}
#[test]
fn prefers_largest_budget_level() {
let bohb = BohbSampler::builder()
.min_resource(1)
.max_resource(9)
.reduction_factor(3)
.min_points_in_model(3)
.seed(42)
.build()
.unwrap();
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 10.0,
log_scale: false,
step: None,
});
// Budget levels: [1, 3, 9]
assert_eq!(bohb.all_budget_levels(), vec![1, 3, 9]);
// Create 5 trials with intermediates at budget 1 and 3
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..5)
.map(|i| {
let x = i as f64;
make_trial_with_intermediates(
i,
x,
vec![(x_id, ParamValue::Float(x), dist.clone())],
vec![(1, x * 2.0), (3, x)],
)
})
.collect();
// Budget 3 has 5 observations (>= 3), budget 9 has 0.
// BOHB should pick budget 3 (largest with enough data).
// The filtered history at budget 3 has values [0, 1, 2, 3, 4].
let filtered_3 = BohbSampler::filter_history_for_budget(&history, 3);
assert_eq!(filtered_3.len(), 5);
let filtered_9 = BohbSampler::filter_history_for_budget(&history, 9);
assert_eq!(filtered_9.len(), 0);
// Should sample successfully
let val = bohb.sample(&dist, 100, &history);
assert!(matches!(val, ParamValue::Float(_)));
}
#[test]
fn builder_validates_tpe_params() {
// Invalid gamma
let result = BohbSampler::builder().gamma(1.5).build();
assert!(result.is_err());
// Invalid bandwidth
let result = BohbSampler::builder().kde_bandwidth(-1.0).build();
assert!(result.is_err());
}
#[test]
#[should_panic(expected = "min_resource must be > 0")]
fn builder_rejects_zero_min_resource() {
let _ = BohbSampler::builder().min_resource(0);
}
#[test]
#[should_panic(expected = "max_resource must be > 0")]
fn builder_rejects_zero_max_resource() {
let _ = BohbSampler::builder().max_resource(0);
}
#[test]
#[should_panic(expected = "reduction_factor must be >= 2")]
fn builder_rejects_small_reduction_factor() {
let _ = BohbSampler::builder().reduction_factor(1);
}
#[test]
fn int_distribution_works() {
let bohb = BohbSampler::builder()
.min_points_in_model(3)
.seed(42)
.build()
.unwrap();
let dist = Distribution::Int(IntDistribution {
low: 0,
high: 100,
log_scale: false,
step: None,
});
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..10)
.map(|i| {
make_trial_with_intermediates(
i,
i as f64,
vec![(x_id, ParamValue::Int(i.cast_signed() * 10), dist.clone())],
vec![(1, i as f64)],
)
})
.collect();
let val = bohb.sample(&dist, 100, &history);
if let ParamValue::Int(v) = val {
assert!((0..=100).contains(&v));
} else {
panic!("Expected Int");
}
}
}
File diff suppressed because it is too large Load Diff
+126 -59
View File
@@ -1,18 +1,21 @@
//! Sampler trait and implementations for parameter sampling.
pub mod bohb;
#[cfg(feature = "cma-es")]
pub mod cma_es;
pub mod grid;
pub mod multivariate_tpe;
pub mod random;
#[cfg(feature = "sobol")]
pub mod sobol;
pub mod tpe;
use std::collections::HashMap;
pub use multivariate_tpe::{
ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder,
};
use crate::distribution::Distribution;
use crate::param::ParamValue;
use crate::parameter::{ParamId, Parameter};
use crate::trial::AttrValue;
use crate::types::TrialState;
/// A completed trial with its parameters, distributions, and objective value.
///
@@ -20,32 +23,134 @@ use crate::param::ParamValue;
/// parameter values, their distributions, and the objective value returned
/// by the objective function.
#[derive(Clone, Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CompletedTrial<V = f64> {
/// The unique identifier for this trial.
pub id: u64,
/// The sampled parameter values, keyed by parameter name.
pub params: HashMap<String, ParamValue>,
/// The parameter distributions used, keyed by parameter name.
pub distributions: HashMap<String, Distribution>,
/// The sampled parameter values, keyed by parameter id.
pub params: HashMap<ParamId, ParamValue>,
/// The parameter distributions used, keyed by parameter id.
pub distributions: HashMap<ParamId, Distribution>,
/// Human-readable labels for parameters, keyed by parameter id.
pub param_labels: HashMap<ParamId, String>,
/// The objective value returned by the objective function.
pub value: V,
/// Intermediate objective values reported during the trial.
pub intermediate_values: Vec<(u64, f64)>,
/// The state of the trial (Complete, Pruned, or Failed).
pub state: TrialState,
/// User-defined attributes stored during the trial.
pub user_attrs: HashMap<String, AttrValue>,
/// Constraint values for this trial (<=0.0 means feasible).
#[cfg_attr(feature = "serde", serde(default))]
pub constraints: Vec<f64>,
}
impl<V> CompletedTrial<V> {
/// Creates a new completed trial.
pub fn new(
id: u64,
params: HashMap<String, ParamValue>,
distributions: HashMap<String, Distribution>,
params: HashMap<ParamId, ParamValue>,
distributions: HashMap<ParamId, Distribution>,
param_labels: HashMap<ParamId, String>,
value: V,
) -> Self {
Self {
id,
params,
distributions,
param_labels,
value,
intermediate_values: Vec::new(),
state: TrialState::Complete,
user_attrs: HashMap::new(),
constraints: Vec::new(),
}
}
/// Creates a new completed trial with intermediate values and user attributes.
pub fn with_intermediate_values(
id: u64,
params: HashMap<ParamId, ParamValue>,
distributions: HashMap<ParamId, Distribution>,
param_labels: HashMap<ParamId, String>,
value: V,
intermediate_values: Vec<(u64, f64)>,
user_attrs: HashMap<String, AttrValue>,
) -> Self {
Self {
id,
params,
distributions,
param_labels,
value,
intermediate_values,
state: TrialState::Complete,
user_attrs,
constraints: Vec::new(),
}
}
/// Returns the typed value for the given parameter.
///
/// Looks up the parameter by its unique id and casts the stored
/// [`ParamValue`] to the parameter's typed value.
///
/// Returns `None` if the parameter was not used in this trial.
///
/// # Panics
///
/// Panics if the stored value is incompatible with the parameter type
/// (e.g., a `Float` value stored for an `IntParam`). This indicates
/// a bug in the program, not a runtime error.
///
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x = FloatParam::new(-10.0, 10.0);
///
/// study
/// .optimize(5, |trial| {
/// let val = x.suggest(trial)?;
/// Ok::<_, optimizer::Error>(val * val)
/// })
/// .unwrap();
///
/// let best = study.best_trial().unwrap();
/// let x_val: f64 = best.get(&x).unwrap();
/// assert!((-10.0..=10.0).contains(&x_val));
/// ```
pub fn get<P: Parameter>(&self, param: &P) -> Option<P::Value> {
self.params.get(&param.id()).map(|v| {
param
.cast_param_value(v)
.expect("parameter type mismatch: stored value incompatible with parameter")
})
}
/// Returns `true` if all constraints are satisfied (values <= 0.0).
///
/// A trial with no constraints is considered feasible.
#[must_use]
pub fn is_feasible(&self) -> bool {
self.constraints.iter().all(|&c| c <= 0.0)
}
/// Gets a user attribute by key.
#[must_use]
pub fn user_attr(&self, key: &str) -> Option<&AttrValue> {
self.user_attrs.get(key)
}
/// Returns all user attributes.
#[must_use]
pub fn user_attrs(&self) -> &HashMap<String, AttrValue> {
&self.user_attrs
}
}
/// A pending (running) trial with its parameters and distributions, but no objective value yet.
@@ -53,34 +158,16 @@ impl<V> CompletedTrial<V> {
/// This struct represents a trial that has been started and has sampled parameters,
/// but is still running and hasn't returned an objective value. It is used with the
/// constant liar strategy for parallel optimization.
///
/// # Examples
///
/// ```ignore
/// use std::collections::HashMap;
/// use optimizer::sampler::PendingTrial;
/// use optimizer::param::ParamValue;
/// use optimizer::distribution::{Distribution, FloatDistribution};
///
/// let mut params = HashMap::new();
/// params.insert("x".to_string(), ParamValue::Float(0.5));
///
/// let mut distributions = HashMap::new();
/// distributions.insert("x".to_string(), Distribution::Float(FloatDistribution {
/// low: 0.0, high: 1.0, log_scale: false, step: None,
/// }));
///
/// let pending = PendingTrial::new(1, params, distributions);
/// assert_eq!(pending.id, 1);
/// ```
#[derive(Clone, Debug)]
pub struct PendingTrial {
/// The unique identifier for this trial.
pub id: u64,
/// The sampled parameter values, keyed by parameter name.
pub params: HashMap<String, ParamValue>,
/// The parameter distributions used, keyed by parameter name.
pub distributions: HashMap<String, Distribution>,
/// The sampled parameter values, keyed by parameter id.
pub params: HashMap<ParamId, ParamValue>,
/// The parameter distributions used, keyed by parameter id.
pub distributions: HashMap<ParamId, Distribution>,
/// Human-readable labels for parameters, keyed by parameter id.
pub param_labels: HashMap<ParamId, String>,
}
impl PendingTrial {
@@ -88,13 +175,15 @@ impl PendingTrial {
#[must_use]
pub fn new(
id: u64,
params: HashMap<String, ParamValue>,
distributions: HashMap<String, Distribution>,
params: HashMap<ParamId, ParamValue>,
distributions: HashMap<ParamId, Distribution>,
param_labels: HashMap<ParamId, String>,
) -> Self {
Self {
id,
params,
distributions,
param_labels,
}
}
}
@@ -104,28 +193,6 @@ impl PendingTrial {
/// Samplers are responsible for generating parameter values based on
/// the distribution and historical trial data. The trait requires
/// `Send + Sync` to support concurrent and async optimization.
///
/// # Examples
///
/// Implementing a custom sampler:
///
/// ```ignore
/// use optimizer::{Sampler, ParamValue, Distribution, CompletedTrial};
///
/// struct MySampler;
///
/// impl Sampler for MySampler {
/// fn sample(
/// &self,
/// distribution: &Distribution,
/// trial_id: u64,
/// history: &[CompletedTrial],
/// ) -> ParamValue {
/// // Custom sampling logic here
/// todo!()
/// }
/// }
/// ```
pub trait Sampler: Send + Sync {
/// Samples a parameter value from the given distribution.
///
+2 -2
View File
@@ -2,7 +2,7 @@
use parking_lot::Mutex;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
use rand::{RngExt, SeedableRng};
use crate::distribution::Distribution;
use crate::param::ParamValue;
@@ -34,7 +34,7 @@ impl RandomSampler {
#[must_use]
pub fn new() -> Self {
Self {
rng: Mutex::new(StdRng::from_os_rng()),
rng: Mutex::new(rand::make_rng()),
}
}
+420
View File
@@ -0,0 +1,420 @@
//! Quasi-random sampler using Sobol low-discrepancy sequences.
use parking_lot::Mutex;
use sobol_burley::sample;
use crate::distribution::Distribution;
use crate::param::ParamValue;
use crate::sampler::{CompletedTrial, Sampler};
/// Internal state for tracking the dimension counter within a trial.
struct SobolState {
/// The `trial_id` of the current trial (used to reset dimension counter).
current_trial: u64,
/// Next Sobol dimension to use for the current trial.
next_dimension: u32,
}
/// Quasi-random sampler using Sobol low-discrepancy sequences.
///
/// Provides better uniform coverage of the parameter space than
/// [`RandomSampler`](super::random::RandomSampler). Useful as a baseline or
/// for the startup phase of model-based samplers.
///
/// Unlike random sampling, Sobol sequences are deterministic and fill the
/// space more evenly, reducing the number of trials needed to adequately
/// cover the search space.
///
/// Each trial uses a different Sobol sequence index, and each parameter
/// within a trial maps to a different Sobol dimension. Parameters must be
/// suggested in the same order across trials for consistent dimension
/// assignment.
///
/// Sobol sequences are most effective in moderate dimensions (up to ~20).
/// For very high dimensions, the uniformity advantage diminishes.
///
/// Requires the `sobol` feature flag.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::sobol::SobolSampler;
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, SobolSampler::new());
/// ```
pub struct SobolSampler {
seed: u32,
state: Mutex<SobolState>,
}
impl SobolSampler {
/// Creates a new Sobol sampler with a default seed of 0.
#[must_use]
pub fn new() -> Self {
Self::with_seed(0)
}
/// Creates a new Sobol sampler with the given seed.
///
/// Different seeds produce statistically independent Sobol sequences.
/// Using the same seed will produce the same sequence of sampled values.
#[must_use]
#[allow(clippy::cast_possible_truncation)]
pub fn with_seed(seed: u64) -> Self {
Self {
seed: seed as u32,
state: Mutex::new(SobolState {
current_trial: u64::MAX,
next_dimension: 0,
}),
}
}
}
impl Default for SobolSampler {
fn default() -> Self {
Self::new()
}
}
impl Sampler for SobolSampler {
#[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
fn sample(
&self,
distribution: &Distribution,
trial_id: u64,
_history: &[CompletedTrial],
) -> ParamValue {
let mut state = self.state.lock();
// Reset dimension counter when a new trial starts.
if state.current_trial != trial_id {
state.current_trial = trial_id;
state.next_dimension = 0;
}
let dimension = state.next_dimension;
state.next_dimension = dimension + 1;
// Use trial_id as the Sobol sequence index.
let index = trial_id as u32;
// Generate a quasi-random point in [0, 1).
let point = f64::from(sample(index, dimension, self.seed));
map_point_to_distribution(point, distribution)
}
}
/// Maps a uniform [0, 1) point to a value within the given distribution.
#[allow(
clippy::cast_possible_truncation,
clippy::cast_precision_loss,
clippy::cast_sign_loss
)]
fn map_point_to_distribution(point: f64, distribution: &Distribution) -> ParamValue {
match distribution {
Distribution::Float(d) => {
let value = if d.log_scale {
let log_low = d.low.ln();
let log_high = d.high.ln();
(log_low + point * (log_high - log_low)).exp()
} else if let Some(step) = d.step {
let n_steps = ((d.high - d.low) / step).floor() as i64;
let k = (point * (n_steps + 1) as f64).floor() as i64;
let k = k.min(n_steps);
d.low + (k as f64) * step
} else {
d.low + point * (d.high - d.low)
};
ParamValue::Float(value)
}
Distribution::Int(d) => {
let value = if d.log_scale {
let log_low = (d.low as f64).ln();
let log_high = (d.high as f64).ln();
let raw = (log_low + point * (log_high - log_low)).exp().round() as i64;
raw.clamp(d.low, d.high)
} else if let Some(step) = d.step {
let n_steps = (d.high - d.low) / step;
let k = (point * (n_steps + 1) as f64).floor() as i64;
let k = k.min(n_steps);
d.low + k * step
} else {
let range = d.high - d.low + 1;
let k = (point * range as f64).floor() as i64;
(d.low + k).min(d.high)
};
ParamValue::Int(value)
}
Distribution::Categorical(d) => {
let index = (point * d.n_choices as f64).floor() as usize;
let index = index.min(d.n_choices - 1);
ParamValue::Categorical(index)
}
}
}
#[cfg(test)]
#[allow(
clippy::cast_possible_truncation,
clippy::cast_precision_loss,
clippy::cast_sign_loss
)]
mod tests {
use super::*;
use crate::distribution::{CategoricalDistribution, FloatDistribution, IntDistribution};
#[test]
fn float_within_bounds() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Float(FloatDistribution {
low: -5.0,
high: 5.0,
log_scale: false,
step: None,
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Float(v) = value {
assert!(
(-5.0..=5.0).contains(&v),
"value {v} out of bounds at trial {i}"
);
} else {
panic!("Expected Float value");
}
}
}
#[test]
fn float_log_scale_within_bounds() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Float(FloatDistribution {
low: 1e-5,
high: 1.0,
log_scale: true,
step: None,
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Float(v) = value {
assert!(
(1e-5..=1.0).contains(&v),
"value {v} out of bounds at trial {i}"
);
} else {
panic!("Expected Float value");
}
}
}
#[test]
fn float_step_respects_grid() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: Some(0.25),
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Float(v) = value {
assert!((0.0..=1.0).contains(&v), "value {v} out of bounds");
let k = (v / 0.25).round() as i64;
let expected = k as f64 * 0.25;
assert!((v - expected).abs() < 1e-10, "value {v} not on step grid");
} else {
panic!("Expected Float value");
}
}
}
#[test]
fn int_within_bounds() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Int(IntDistribution {
low: 0,
high: 10,
log_scale: false,
step: None,
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Int(v) = value {
assert!(
(0..=10).contains(&v),
"value {v} out of bounds at trial {i}"
);
} else {
panic!("Expected Int value");
}
}
}
#[test]
fn int_log_scale_within_bounds() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Int(IntDistribution {
low: 1,
high: 1000,
log_scale: true,
step: None,
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Int(v) = value {
assert!(
(1..=1000).contains(&v),
"value {v} out of bounds at trial {i}"
);
} else {
panic!("Expected Int value");
}
}
}
#[test]
fn int_step_respects_grid() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Int(IntDistribution {
low: 0,
high: 10,
log_scale: false,
step: Some(2),
});
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Int(v) = value {
assert!((0..=10).contains(&v), "value {v} out of bounds");
assert!(v % 2 == 0, "value {v} not on step grid");
} else {
panic!("Expected Int value");
}
}
}
#[test]
fn categorical_within_bounds() {
let sampler = SobolSampler::with_seed(42);
let dist = Distribution::Categorical(CategoricalDistribution { n_choices: 5 });
for i in 0..100 {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Categorical(idx) = value {
assert!(idx < 5, "index {idx} out of bounds at trial {i}");
} else {
panic!("Expected Categorical value");
}
}
}
#[test]
fn deterministic_with_same_seed() {
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
let sampler1 = SobolSampler::with_seed(42);
let sampler2 = SobolSampler::with_seed(42);
for i in 0..20 {
let v1 = sampler1.sample(&dist, i, &[]);
let v2 = sampler2.sample(&dist, i, &[]);
assert_eq!(v1, v2, "mismatch at trial {i}");
}
}
#[test]
fn different_seeds_produce_different_sequences() {
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
let sampler1 = SobolSampler::with_seed(0);
let sampler2 = SobolSampler::with_seed(12345);
let mut any_different = false;
for i in 0..20 {
let v1 = sampler1.sample(&dist, i, &[]);
let v2 = sampler2.sample(&dist, i, &[]);
if v1 != v2 {
any_different = true;
break;
}
}
assert!(
any_different,
"different seeds should produce different sequences"
);
}
#[test]
fn better_coverage_than_random() {
// Sobol sequence should cover [0,1] more uniformly than random.
// We measure this by checking that the Sobol samples fill all
// 10 equal-width bins with only 20 samples.
let sampler = SobolSampler::with_seed(0);
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
let n_bins = 10;
let n_samples = 20;
let mut bins = vec![0u32; n_bins];
for i in 0..n_samples {
let value = sampler.sample(&dist, i, &[]);
if let ParamValue::Float(v) = value {
let bin = ((v * n_bins as f64).floor() as usize).min(n_bins - 1);
bins[bin] += 1;
}
}
let filled_bins = bins.iter().filter(|&&c| c > 0).count();
assert!(
filled_bins >= 8,
"Expected at least 8/10 bins filled, got {filled_bins}: {bins:?}"
);
}
#[test]
fn multi_parameter_uses_different_dimensions() {
// When sampling multiple parameters per trial, each parameter
// should get a different Sobol dimension, producing different values.
let sampler = SobolSampler::with_seed(0);
let dist = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
});
// Sample two parameters for trial 0.
let v1 = sampler.sample(&dist, 0, &[]);
let v2 = sampler.sample(&dist, 0, &[]);
// They should differ (different Sobol dimensions for the same index).
assert_ne!(
v1, v2,
"multi-parameter samples should use different dimensions"
);
}
}
+4
View File
@@ -4,9 +4,13 @@
//! including support for intersection search space calculation.
mod gamma;
mod multivariate;
mod sampler;
pub mod search_space;
pub use gamma::{FixedGamma, GammaStrategy, HyperoptGamma, LinearGamma, SqrtGamma};
pub use multivariate::{
ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder,
};
pub use sampler::{TpeSampler, TpeSamplerBuilder};
pub use search_space::{GroupDecomposedSearchSpace, IntersectionSearchSpace};
File diff suppressed because it is too large Load Diff
+23 -15
View File
@@ -60,7 +60,7 @@ use std::sync::Arc;
use parking_lot::Mutex;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
use rand::{RngExt, SeedableRng};
use crate::distribution::Distribution;
use crate::error::{Error, Result};
@@ -149,7 +149,7 @@ impl TpeSampler {
n_startup_trials: 10,
n_ei_candidates: 24,
kde_bandwidth: None,
rng: Mutex::new(StdRng::from_os_rng()),
rng: Mutex::new(rand::make_rng()),
}
}
@@ -251,7 +251,7 @@ impl TpeSampler {
let rng = match seed {
Some(s) => StdRng::seed_from_u64(s),
None => StdRng::from_os_rng(),
None => rand::make_rng(),
};
Ok(Self {
@@ -280,6 +280,7 @@ impl TpeSampler {
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
#[must_use]
fn split_trials<'a>(
&self,
history: &'a [CompletedTrial],
@@ -856,7 +857,7 @@ impl TpeSamplerBuilder {
let rng = match self.seed {
Some(s) => StdRng::seed_from_u64(s),
None => StdRng::from_os_rng(),
None => rand::make_rng(),
};
Ok(TpeSampler {
@@ -1023,19 +1024,20 @@ mod tests {
use super::*;
use crate::distribution::{CategoricalDistribution, FloatDistribution, IntDistribution};
use crate::parameter::ParamId;
fn create_trial(
id: u64,
value: f64,
params: Vec<(&str, ParamValue, Distribution)>,
params: Vec<(ParamId, ParamValue, Distribution)>,
) -> CompletedTrial {
let mut param_map = HashMap::new();
let mut dist_map = HashMap::new();
for (name, pv, dist) in params {
param_map.insert(name.to_string(), pv);
dist_map.insert(name.to_string(), dist);
for (param_id, pv, dist) in params {
param_map.insert(param_id, pv);
dist_map.insert(param_id, dist);
}
CompletedTrial::new(id, param_map, dist_map, value)
CompletedTrial::new(id, param_map, dist_map, HashMap::new(), value)
}
#[test]
@@ -1103,12 +1105,13 @@ mod tests {
});
// Create 20 trials with values 0..20
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
create_trial(
i as u64,
f64::from(i),
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![(x_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
@@ -1137,6 +1140,7 @@ mod tests {
});
// Create history where low values (near 0.2) are "good"
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
let x = f64::from(i) / 20.0;
@@ -1145,7 +1149,7 @@ mod tests {
create_trial(
i as u64,
value,
vec![("x", ParamValue::Float(x), dist.clone())],
vec![(x_id, ParamValue::Float(x), dist.clone())],
)
})
.collect();
@@ -1174,6 +1178,7 @@ mod tests {
let dist = Distribution::Categorical(CategoricalDistribution { n_choices: 4 });
// Create history where category 1 is consistently good
let cat_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
let category = i % 4;
@@ -1183,7 +1188,7 @@ mod tests {
i as u64,
value,
vec![(
"cat",
cat_id,
ParamValue::Categorical(category as usize),
dist.clone(),
)],
@@ -1219,6 +1224,7 @@ mod tests {
});
// Create history where values near 30 are good
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
let x = i * 5; // 0, 5, 10, ..., 95
@@ -1226,7 +1232,7 @@ mod tests {
create_trial(
i as u64,
value,
vec![("x", ParamValue::Int(x), dist.clone())],
vec![(x_id, ParamValue::Int(x), dist.clone())],
)
})
.collect();
@@ -1251,12 +1257,13 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20)
.map(|i| {
create_trial(
i as u64,
f64::from(i),
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![(x_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
@@ -1331,12 +1338,13 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let history: Vec<CompletedTrial> = (0..20u32)
.map(|i| {
create_trial(
u64::from(i),
f64::from(i),
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![(x_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
+232 -191
View File
@@ -11,6 +11,7 @@
use std::collections::{HashMap, HashSet, VecDeque};
use crate::distribution::Distribution;
use crate::parameter::ParamId;
use crate::sampler::CompletedTrial;
/// Computes the intersection of search spaces across completed trials.
@@ -87,34 +88,30 @@ impl IntersectionSearchSpace {
/// - For dynamic search spaces where not all parameters are sampled in every trial,
/// this helps identify the "stable" set of parameters that can be modeled jointly.
#[must_use]
pub fn calculate(trials: &[CompletedTrial]) -> HashMap<String, Distribution> {
pub fn calculate(trials: &[CompletedTrial]) -> HashMap<ParamId, Distribution> {
if trials.is_empty() {
return HashMap::new();
}
// Get parameter names from the first trial as the initial candidate set
// Get parameter ids from the first trial as the initial candidate set
let first_trial = &trials[0];
let mut candidate_params: HashSet<&str> = first_trial
.distributions
.keys()
.map(String::as_str)
.collect();
let mut candidate_params: HashSet<ParamId> =
first_trial.distributions.keys().copied().collect();
// Intersect with parameter sets from all other trials
for trial in trials.iter().skip(1) {
let trial_params: HashSet<&str> =
trial.distributions.keys().map(String::as_str).collect();
let trial_params: HashSet<ParamId> = trial.distributions.keys().copied().collect();
candidate_params.retain(|param| trial_params.contains(param));
}
// Build the result map using distributions from the first trial
// that contains each parameter
let mut result = HashMap::new();
for param_name in candidate_params {
for param_id in candidate_params {
// Find the first trial that has this parameter and use its distribution
for trial in trials {
if let Some(dist) = trial.distributions.get(param_name) {
result.insert(param_name.to_string(), dist.clone());
if let Some(dist) = trial.distributions.get(&param_id) {
result.insert(param_id, dist.clone());
break;
}
}
@@ -193,7 +190,7 @@ impl GroupDecomposedSearchSpace {
///
/// # Returns
///
/// A `Vec<HashSet<String>>` where each `HashSet` represents an independent group
/// A `Vec<HashSet<ParamId>>` where each `HashSet` represents an independent group
/// of parameters that co-occur. Parameters within the same group have appeared
/// together in at least one trial (directly or transitively). Parameters in
/// different groups have never appeared in the same trial.
@@ -212,16 +209,16 @@ impl GroupDecomposedSearchSpace {
/// - This is useful for `MultivariateTpeSampler` when `group=true` to sample
/// independent parameter groups separately.
#[must_use]
pub fn calculate(trials: &[CompletedTrial]) -> Vec<HashSet<String>> {
pub fn calculate(trials: &[CompletedTrial]) -> Vec<HashSet<ParamId>> {
if trials.is_empty() {
return Vec::new();
}
// Collect all unique parameter names
let mut all_params: HashSet<String> = HashSet::new();
// Collect all unique parameter ids
let mut all_params: HashSet<ParamId> = HashSet::new();
for trial in trials {
for param_name in trial.distributions.keys() {
all_params.insert(param_name.clone());
for &param_id in trial.distributions.keys() {
all_params.insert(param_id);
}
}
@@ -231,51 +228,51 @@ impl GroupDecomposedSearchSpace {
// Build adjacency list for co-occurrence graph
// Two parameters are adjacent if they appear in the same trial
let mut adjacency: HashMap<String, HashSet<String>> = HashMap::new();
for param in &all_params {
adjacency.insert(param.clone(), HashSet::new());
let mut adjacency: HashMap<ParamId, HashSet<ParamId>> = HashMap::new();
for &param in &all_params {
adjacency.insert(param, HashSet::new());
}
for trial in trials {
let trial_params: Vec<&String> = trial.distributions.keys().collect();
let trial_params: Vec<ParamId> = trial.distributions.keys().copied().collect();
// Connect all pairs of parameters in this trial
for (i, param1) in trial_params.iter().enumerate() {
for param2 in trial_params.iter().skip(i + 1) {
for (i, &param1) in trial_params.iter().enumerate() {
for &param2 in trial_params.iter().skip(i + 1) {
adjacency
.get_mut(*param1)
.get_mut(&param1)
.expect("param should exist in adjacency map")
.insert((*param2).clone());
.insert(param2);
adjacency
.get_mut(*param2)
.get_mut(&param2)
.expect("param should exist in adjacency map")
.insert((*param1).clone());
.insert(param1);
}
}
}
// Find connected components using BFS
let mut visited: HashSet<String> = HashSet::new();
let mut groups: Vec<HashSet<String>> = Vec::new();
let mut visited: HashSet<ParamId> = HashSet::new();
let mut groups: Vec<HashSet<ParamId>> = Vec::new();
for param in &all_params {
if visited.contains(param) {
for &param in &all_params {
if visited.contains(&param) {
continue;
}
// BFS to find all parameters in this component
let mut component: HashSet<String> = HashSet::new();
let mut queue: VecDeque<String> = VecDeque::new();
queue.push_back(param.clone());
visited.insert(param.clone());
let mut component: HashSet<ParamId> = HashSet::new();
let mut queue: VecDeque<ParamId> = VecDeque::new();
queue.push_back(param);
visited.insert(param);
while let Some(current) = queue.pop_front() {
component.insert(current.clone());
component.insert(current);
if let Some(neighbors) = adjacency.get(&current) {
for neighbor in neighbors {
if !visited.contains(neighbor) {
visited.insert(neighbor.clone());
queue.push_back(neighbor.clone());
for &neighbor in neighbors {
if !visited.contains(&neighbor) {
visited.insert(neighbor);
queue.push_back(neighbor);
}
}
}
@@ -293,19 +290,20 @@ mod tests {
use super::*;
use crate::distribution::{CategoricalDistribution, FloatDistribution, IntDistribution};
use crate::param::ParamValue;
use crate::parameter::ParamId;
fn create_trial(
id: u64,
params: Vec<(&str, ParamValue, Distribution)>,
params: Vec<(ParamId, ParamValue, Distribution)>,
value: f64,
) -> CompletedTrial {
let mut param_map = HashMap::new();
let mut dist_map = HashMap::new();
for (name, pv, dist) in params {
param_map.insert(name.to_string(), pv);
dist_map.insert(name.to_string(), dist);
for (param_id, pv, dist) in params {
param_map.insert(param_id, pv);
dist_map.insert(param_id, dist);
}
CompletedTrial::new(id, param_map, dist_map, value)
CompletedTrial::new(id, param_map, dist_map, HashMap::new(), value)
}
#[test]
@@ -317,6 +315,8 @@ mod tests {
#[test]
fn test_single_trial() {
let x_id = ParamId::new();
let y_id = ParamId::new();
let dist_x = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -333,22 +333,24 @@ mod tests {
let trial = create_trial(
0,
vec![
("x", ParamValue::Float(0.5), dist_x.clone()),
("y", ParamValue::Int(5), dist_y.clone()),
(x_id, ParamValue::Float(0.5), dist_x.clone()),
(y_id, ParamValue::Int(5), dist_y.clone()),
],
1.0,
);
let result = IntersectionSearchSpace::calculate(&[trial]);
assert_eq!(result.len(), 2);
assert!(result.contains_key("x"));
assert!(result.contains_key("y"));
assert_eq!(result.get("x"), Some(&dist_x));
assert_eq!(result.get("y"), Some(&dist_y));
assert!(result.contains_key(&x_id));
assert!(result.contains_key(&y_id));
assert_eq!(result.get(&x_id), Some(&dist_x));
assert_eq!(result.get(&y_id), Some(&dist_y));
}
#[test]
fn test_all_trials_same_params() {
let x_id = ParamId::new();
let y_id = ParamId::new();
let dist_x = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -369,8 +371,8 @@ mod tests {
create_trial(
i,
vec![
("x", ParamValue::Float(val), dist_x.clone()),
("y", ParamValue::Float(val - 0.5), dist_y.clone()),
(x_id, ParamValue::Float(val), dist_x.clone()),
(y_id, ParamValue::Float(val - 0.5), dist_y.clone()),
],
val * val,
)
@@ -379,12 +381,15 @@ mod tests {
let result = IntersectionSearchSpace::calculate(&trials);
assert_eq!(result.len(), 2);
assert!(result.contains_key("x"));
assert!(result.contains_key("y"));
assert!(result.contains_key(&x_id));
assert!(result.contains_key(&y_id));
}
#[test]
fn test_partial_overlap() {
let x_id = ParamId::new();
let y_id = ParamId::new();
let z_id = ParamId::new();
let dist_x = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -408,8 +413,8 @@ mod tests {
let trial1 = create_trial(
0,
vec![
("x", ParamValue::Float(0.5), dist_x.clone()),
("y", ParamValue::Float(0.3), dist_y.clone()),
(x_id, ParamValue::Float(0.5), dist_x.clone()),
(y_id, ParamValue::Float(0.3), dist_y.clone()),
],
1.0,
);
@@ -418,8 +423,8 @@ mod tests {
let trial2 = create_trial(
1,
vec![
("x", ParamValue::Float(0.7), dist_x.clone()),
("z", ParamValue::Float(0.2), dist_z.clone()),
(x_id, ParamValue::Float(0.7), dist_x.clone()),
(z_id, ParamValue::Float(0.2), dist_z.clone()),
],
0.5,
);
@@ -428,24 +433,26 @@ mod tests {
let trial3 = create_trial(
2,
vec![
("x", ParamValue::Float(0.6), dist_x.clone()),
("y", ParamValue::Float(0.4), dist_y.clone()),
("z", ParamValue::Float(0.1), dist_z.clone()),
(x_id, ParamValue::Float(0.6), dist_x.clone()),
(y_id, ParamValue::Float(0.4), dist_y.clone()),
(z_id, ParamValue::Float(0.1), dist_z.clone()),
],
0.8,
);
let result = IntersectionSearchSpace::calculate(&[trial1, trial2, trial3]);
// Only "x" appears in all three trials
// Only x appears in all three trials
assert_eq!(result.len(), 1);
assert!(result.contains_key("x"));
assert!(!result.contains_key("y"));
assert!(!result.contains_key("z"));
assert!(result.contains_key(&x_id));
assert!(!result.contains_key(&y_id));
assert!(!result.contains_key(&z_id));
}
#[test]
fn test_no_common_params() {
let x_id = ParamId::new();
let y_id = ParamId::new();
let dist_x = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -460,10 +467,10 @@ mod tests {
});
// Trial 1: only x
let trial1 = create_trial(0, vec![("x", ParamValue::Float(0.5), dist_x.clone())], 1.0);
let trial1 = create_trial(0, vec![(x_id, ParamValue::Float(0.5), dist_x.clone())], 1.0);
// Trial 2: only y
let trial2 = create_trial(1, vec![("y", ParamValue::Float(0.3), dist_y.clone())], 0.5);
let trial2 = create_trial(1, vec![(y_id, ParamValue::Float(0.3), dist_y.clone())], 0.5);
let result = IntersectionSearchSpace::calculate(&[trial1, trial2]);
@@ -473,6 +480,9 @@ mod tests {
#[test]
fn test_mixed_distribution_types() {
let lr_id = ParamId::new();
let n_layers_id = ParamId::new();
let optimizer_id = ParamId::new();
let dist_float = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -490,9 +500,9 @@ mod tests {
let trial1 = create_trial(
0,
vec![
("learning_rate", ParamValue::Float(0.01), dist_float.clone()),
("n_layers", ParamValue::Int(3), dist_int.clone()),
("optimizer", ParamValue::Categorical(0), dist_cat.clone()),
(lr_id, ParamValue::Float(0.01), dist_float.clone()),
(n_layers_id, ParamValue::Int(3), dist_int.clone()),
(optimizer_id, ParamValue::Categorical(0), dist_cat.clone()),
],
1.0,
);
@@ -500,13 +510,9 @@ mod tests {
let trial2 = create_trial(
1,
vec![
(
"learning_rate",
ParamValue::Float(0.001),
dist_float.clone(),
),
("n_layers", ParamValue::Int(5), dist_int.clone()),
("optimizer", ParamValue::Categorical(1), dist_cat.clone()),
(lr_id, ParamValue::Float(0.001), dist_float.clone()),
(n_layers_id, ParamValue::Int(5), dist_int.clone()),
(optimizer_id, ParamValue::Categorical(1), dist_cat.clone()),
],
0.8,
);
@@ -514,19 +520,20 @@ mod tests {
let result = IntersectionSearchSpace::calculate(&[trial1, trial2]);
assert_eq!(result.len(), 3);
assert!(matches!(result.get(&lr_id), Some(Distribution::Float(_))));
assert!(matches!(
result.get("learning_rate"),
Some(Distribution::Float(_))
result.get(&n_layers_id),
Some(Distribution::Int(_))
));
assert!(matches!(result.get("n_layers"), Some(Distribution::Int(_))));
assert!(matches!(
result.get("optimizer"),
result.get(&optimizer_id),
Some(Distribution::Categorical(_))
));
}
#[test]
fn test_distribution_from_first_trial() {
let x_id = ParamId::new();
// Test that when distributions differ, the first trial's distribution is used
let dist_x_v1 = Distribution::Float(FloatDistribution {
low: 0.0,
@@ -543,13 +550,13 @@ mod tests {
let trial1 = create_trial(
0,
vec![("x", ParamValue::Float(0.5), dist_x_v1.clone())],
vec![(x_id, ParamValue::Float(0.5), dist_x_v1.clone())],
1.0,
);
let trial2 = create_trial(
1,
vec![("x", ParamValue::Float(5.0), dist_x_v2.clone())],
vec![(x_id, ParamValue::Float(5.0), dist_x_v2.clone())],
0.5,
);
@@ -557,13 +564,15 @@ mod tests {
assert_eq!(result.len(), 1);
// Should use the distribution from the first trial
assert_eq!(result.get("x"), Some(&dist_x_v1));
assert_eq!(result.get(&x_id), Some(&dist_x_v1));
}
#[test]
fn test_many_trials_with_conditional_params() {
let lr_id = ParamId::new();
let use_dropout_id = ParamId::new();
let dropout_rate_id = ParamId::new();
// Simulate a scenario with conditional parameters
// e.g., "use_dropout" is a bool, and "dropout_rate" only exists when use_dropout=true
let dist_lr = Distribution::Float(FloatDistribution {
low: 1e-5,
high: 1e-1,
@@ -582,14 +591,14 @@ mod tests {
let trial1 = create_trial(
0,
vec![
("lr", ParamValue::Float(0.01), dist_lr.clone()),
(lr_id, ParamValue::Float(0.01), dist_lr.clone()),
(
"use_dropout",
use_dropout_id,
ParamValue::Categorical(1),
dist_dropout.clone(),
),
(
"dropout_rate",
dropout_rate_id,
ParamValue::Float(0.2),
dist_dropout_rate.clone(),
),
@@ -601,9 +610,9 @@ mod tests {
let trial2 = create_trial(
1,
vec![
("lr", ParamValue::Float(0.001), dist_lr.clone()),
(lr_id, ParamValue::Float(0.001), dist_lr.clone()),
(
"use_dropout",
use_dropout_id,
ParamValue::Categorical(0),
dist_dropout.clone(),
),
@@ -615,14 +624,14 @@ mod tests {
let trial3 = create_trial(
2,
vec![
("lr", ParamValue::Float(0.005), dist_lr.clone()),
(lr_id, ParamValue::Float(0.005), dist_lr.clone()),
(
"use_dropout",
use_dropout_id,
ParamValue::Categorical(1),
dist_dropout.clone(),
),
(
"dropout_rate",
dropout_rate_id,
ParamValue::Float(0.3),
dist_dropout_rate.clone(),
),
@@ -632,11 +641,11 @@ mod tests {
let result = IntersectionSearchSpace::calculate(&[trial1, trial2, trial3]);
// Only "lr" and "use_dropout" appear in all trials
// Only lr and use_dropout appear in all trials
assert_eq!(result.len(), 2);
assert!(result.contains_key("lr"));
assert!(result.contains_key("use_dropout"));
assert!(!result.contains_key("dropout_rate")); // Not in trial2
assert!(result.contains_key(&lr_id));
assert!(result.contains_key(&use_dropout_id));
assert!(!result.contains_key(&dropout_rate_id)); // Not in trial2
}
// ==================== GroupDecomposedSearchSpace Tests ====================
@@ -657,12 +666,13 @@ mod tests {
step: None,
});
let trial = create_trial(0, vec![("x", ParamValue::Float(0.5), dist)], 1.0);
let x_id = ParamId::new();
let trial = create_trial(0, vec![(x_id, ParamValue::Float(0.5), dist)], 1.0);
let groups = GroupDecomposedSearchSpace::calculate(&[trial]);
assert_eq!(groups.len(), 1);
assert!(groups[0].contains("x"));
assert!(groups[0].contains(&x_id));
assert_eq!(groups[0].len(), 1);
}
@@ -675,12 +685,15 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let y_id = ParamId::new();
let z_id = ParamId::new();
let trial = create_trial(
0,
vec![
("x", ParamValue::Float(0.5), dist.clone()),
("y", ParamValue::Float(0.3), dist.clone()),
("z", ParamValue::Float(0.7), dist),
(x_id, ParamValue::Float(0.5), dist.clone()),
(y_id, ParamValue::Float(0.3), dist.clone()),
(z_id, ParamValue::Float(0.7), dist),
],
1.0,
);
@@ -690,9 +703,9 @@ mod tests {
// All params appear together, so one group
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].len(), 3);
assert!(groups[0].contains("x"));
assert!(groups[0].contains("y"));
assert!(groups[0].contains("z"));
assert!(groups[0].contains(&x_id));
assert!(groups[0].contains(&y_id));
assert!(groups[0].contains(&z_id));
}
#[test]
@@ -704,12 +717,17 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let y_id = ParamId::new();
let a_id = ParamId::new();
let b_id = ParamId::new();
// Trial 1: x, y (group 1)
let trial1 = create_trial(
0,
vec![
("x", ParamValue::Float(0.1), dist.clone()),
("y", ParamValue::Float(0.2), dist.clone()),
(x_id, ParamValue::Float(0.1), dist.clone()),
(y_id, ParamValue::Float(0.2), dist.clone()),
],
1.0,
);
@@ -718,8 +736,8 @@ mod tests {
let trial2 = create_trial(
1,
vec![
("a", ParamValue::Float(0.3), dist.clone()),
("b", ParamValue::Float(0.4), dist),
(a_id, ParamValue::Float(0.3), dist.clone()),
(b_id, ParamValue::Float(0.4), dist),
],
0.5,
);
@@ -730,8 +748,8 @@ mod tests {
assert_eq!(groups.len(), 2);
// Find which group has x/y and which has a/b
let group_xy = groups.iter().find(|g| g.contains("x"));
let group_ab = groups.iter().find(|g| g.contains("a"));
let group_xy = groups.iter().find(|g| g.contains(&x_id));
let group_ab = groups.iter().find(|g| g.contains(&a_id));
assert!(group_xy.is_some());
assert!(group_ab.is_some());
@@ -740,12 +758,12 @@ mod tests {
let group_ab = group_ab.expect("group with a should exist");
assert_eq!(group_xy.len(), 2);
assert!(group_xy.contains("x"));
assert!(group_xy.contains("y"));
assert!(group_xy.contains(&x_id));
assert!(group_xy.contains(&y_id));
assert_eq!(group_ab.len(), 2);
assert!(group_ab.contains("a"));
assert!(group_ab.contains("b"));
assert!(group_ab.contains(&a_id));
assert!(group_ab.contains(&b_id));
}
#[test]
@@ -757,12 +775,16 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let y_id = ParamId::new();
let z_id = ParamId::new();
// Trial 1: x, y
let trial1 = create_trial(
0,
vec![
("x", ParamValue::Float(0.1), dist.clone()),
("y", ParamValue::Float(0.2), dist.clone()),
(x_id, ParamValue::Float(0.1), dist.clone()),
(y_id, ParamValue::Float(0.2), dist.clone()),
],
1.0,
);
@@ -771,8 +793,8 @@ mod tests {
let trial2 = create_trial(
1,
vec![
("y", ParamValue::Float(0.3), dist.clone()),
("z", ParamValue::Float(0.4), dist),
(y_id, ParamValue::Float(0.3), dist.clone()),
(z_id, ParamValue::Float(0.4), dist),
],
0.5,
);
@@ -782,9 +804,9 @@ mod tests {
// All should be in one group due to transitive connection via y
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].len(), 3);
assert!(groups[0].contains("x"));
assert!(groups[0].contains("y"));
assert!(groups[0].contains("z"));
assert!(groups[0].contains(&x_id));
assert!(groups[0].contains(&y_id));
assert!(groups[0].contains(&z_id));
}
#[test]
@@ -796,33 +818,35 @@ mod tests {
step: None,
});
let a_id = ParamId::new();
let b_id = ParamId::new();
let c_id = ParamId::new();
let d_id = ParamId::new();
// Create a chain: a-b, b-c, c-d
// Trial 1: a, b
let trial1 = create_trial(
0,
vec![
("a", ParamValue::Float(0.1), dist.clone()),
("b", ParamValue::Float(0.2), dist.clone()),
(a_id, ParamValue::Float(0.1), dist.clone()),
(b_id, ParamValue::Float(0.2), dist.clone()),
],
1.0,
);
// Trial 2: b, c
let trial2 = create_trial(
1,
vec![
("b", ParamValue::Float(0.3), dist.clone()),
("c", ParamValue::Float(0.4), dist.clone()),
(b_id, ParamValue::Float(0.3), dist.clone()),
(c_id, ParamValue::Float(0.4), dist.clone()),
],
0.5,
);
// Trial 3: c, d
let trial3 = create_trial(
2,
vec![
("c", ParamValue::Float(0.5), dist.clone()),
("d", ParamValue::Float(0.6), dist),
(c_id, ParamValue::Float(0.5), dist.clone()),
(d_id, ParamValue::Float(0.6), dist),
],
0.3,
);
@@ -832,10 +856,10 @@ mod tests {
// All should be connected in one group
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].len(), 4);
assert!(groups[0].contains("a"));
assert!(groups[0].contains("b"));
assert!(groups[0].contains("c"));
assert!(groups[0].contains("d"));
assert!(groups[0].contains(&a_id));
assert!(groups[0].contains(&b_id));
assert!(groups[0].contains(&c_id));
assert!(groups[0].contains(&d_id));
}
#[test]
@@ -847,10 +871,14 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let y_id = ParamId::new();
let z_id = ParamId::new();
// Each trial has exactly one parameter (all isolated)
let trial1 = create_trial(0, vec![("x", ParamValue::Float(0.1), dist.clone())], 1.0);
let trial2 = create_trial(1, vec![("y", ParamValue::Float(0.2), dist.clone())], 0.5);
let trial3 = create_trial(2, vec![("z", ParamValue::Float(0.3), dist)], 0.3);
let trial1 = create_trial(0, vec![(x_id, ParamValue::Float(0.1), dist.clone())], 1.0);
let trial2 = create_trial(1, vec![(y_id, ParamValue::Float(0.2), dist.clone())], 0.5);
let trial3 = create_trial(2, vec![(z_id, ParamValue::Float(0.3), dist)], 0.3);
let groups = GroupDecomposedSearchSpace::calculate(&[trial1, trial2, trial3]);
@@ -861,10 +889,10 @@ mod tests {
}
// Verify all params are covered
let all_params: HashSet<String> = groups.iter().flatten().cloned().collect();
assert!(all_params.contains("x"));
assert!(all_params.contains("y"));
assert!(all_params.contains("z"));
let all_params: HashSet<ParamId> = groups.iter().flatten().copied().collect();
assert!(all_params.contains(&x_id));
assert!(all_params.contains(&y_id));
assert!(all_params.contains(&z_id));
}
#[test]
@@ -876,51 +904,54 @@ mod tests {
step: None,
});
let x_id = ParamId::new();
let y_id = ParamId::new();
let z_id = ParamId::new();
let a_id = ParamId::new();
let b_id = ParamId::new();
let w_id = ParamId::new();
// Group 1: x, y, z connected
// Group 2: a, b connected
// w is isolated
// Trial 1: x, y
let trial1 = create_trial(
0,
vec![
("x", ParamValue::Float(0.1), dist.clone()),
("y", ParamValue::Float(0.2), dist.clone()),
(x_id, ParamValue::Float(0.1), dist.clone()),
(y_id, ParamValue::Float(0.2), dist.clone()),
],
1.0,
);
// Trial 2: y, z
let trial2 = create_trial(
1,
vec![
("y", ParamValue::Float(0.3), dist.clone()),
("z", ParamValue::Float(0.4), dist.clone()),
(y_id, ParamValue::Float(0.3), dist.clone()),
(z_id, ParamValue::Float(0.4), dist.clone()),
],
0.5,
);
// Trial 3: a, b
let trial3 = create_trial(
2,
vec![
("a", ParamValue::Float(0.5), dist.clone()),
("b", ParamValue::Float(0.6), dist.clone()),
(a_id, ParamValue::Float(0.5), dist.clone()),
(b_id, ParamValue::Float(0.6), dist.clone()),
],
0.3,
);
// Trial 4: w (isolated)
let trial4 = create_trial(3, vec![("w", ParamValue::Float(0.7), dist)], 0.2);
let trial4 = create_trial(3, vec![(w_id, ParamValue::Float(0.7), dist)], 0.2);
let groups = GroupDecomposedSearchSpace::calculate(&[trial1, trial2, trial3, trial4]);
// Should have 3 groups: {x,y,z}, {a,b}, {w}
assert_eq!(groups.len(), 3);
let group_xyz = groups.iter().find(|g| g.contains("x"));
let group_ab = groups.iter().find(|g| g.contains("a"));
let group_w = groups.iter().find(|g| g.contains("w"));
let group_xyz = groups.iter().find(|g| g.contains(&x_id));
let group_ab = groups.iter().find(|g| g.contains(&a_id));
let group_w = groups.iter().find(|g| g.contains(&w_id));
assert!(group_xyz.is_some());
assert!(group_ab.is_some());
@@ -928,18 +959,18 @@ mod tests {
let group_xyz = group_xyz.expect("group with x should exist");
assert_eq!(group_xyz.len(), 3);
assert!(group_xyz.contains("x"));
assert!(group_xyz.contains("y"));
assert!(group_xyz.contains("z"));
assert!(group_xyz.contains(&x_id));
assert!(group_xyz.contains(&y_id));
assert!(group_xyz.contains(&z_id));
let group_ab = group_ab.expect("group with a should exist");
assert_eq!(group_ab.len(), 2);
assert!(group_ab.contains("a"));
assert!(group_ab.contains("b"));
assert!(group_ab.contains(&a_id));
assert!(group_ab.contains(&b_id));
let group_w = group_w.expect("group with w should exist");
assert_eq!(group_w.len(), 1);
assert!(group_w.contains("w"));
assert!(group_w.contains(&w_id));
}
#[test]
@@ -951,14 +982,19 @@ mod tests {
step: None,
});
let a_id = ParamId::new();
let b_id = ParamId::new();
let c_id = ParamId::new();
let d_id = ParamId::new();
// All params in single trial
let trial = create_trial(
0,
vec![
("a", ParamValue::Float(0.1), dist.clone()),
("b", ParamValue::Float(0.2), dist.clone()),
("c", ParamValue::Float(0.3), dist.clone()),
("d", ParamValue::Float(0.4), dist),
(a_id, ParamValue::Float(0.1), dist.clone()),
(b_id, ParamValue::Float(0.2), dist.clone()),
(c_id, ParamValue::Float(0.3), dist.clone()),
(d_id, ParamValue::Float(0.4), dist),
],
1.0,
);
@@ -972,6 +1008,9 @@ mod tests {
#[test]
fn test_group_with_mixed_distribution_types() {
let lr_id = ParamId::new();
let n_layers_id = ParamId::new();
let optimizer_id = ParamId::new();
let dist_float = Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
@@ -991,23 +1030,23 @@ mod tests {
let trial1 = create_trial(
0,
vec![
("learning_rate", ParamValue::Float(0.01), dist_float.clone()),
("n_layers", ParamValue::Int(3), dist_int.clone()),
(lr_id, ParamValue::Float(0.01), dist_float.clone()),
(n_layers_id, ParamValue::Int(3), dist_int.clone()),
],
1.0,
);
let trial2 = create_trial(
1,
vec![("optimizer", ParamValue::Categorical(1), dist_cat)],
vec![(optimizer_id, ParamValue::Categorical(1), dist_cat)],
0.5,
);
let trial3 = create_trial(
2,
vec![
("learning_rate", ParamValue::Float(0.001), dist_float),
("n_layers", ParamValue::Int(5), dist_int),
(lr_id, ParamValue::Float(0.001), dist_float),
(n_layers_id, ParamValue::Int(5), dist_int),
],
0.8,
);
@@ -1017,20 +1056,20 @@ mod tests {
// Should have 2 groups: {learning_rate, n_layers} and {optimizer}
assert_eq!(groups.len(), 2);
let group_lr = groups.iter().find(|g| g.contains("learning_rate"));
let group_opt = groups.iter().find(|g| g.contains("optimizer"));
let group_lr = groups.iter().find(|g| g.contains(&lr_id));
let group_opt = groups.iter().find(|g| g.contains(&optimizer_id));
assert!(group_lr.is_some());
assert!(group_opt.is_some());
let group_lr = group_lr.expect("group with learning_rate should exist");
assert_eq!(group_lr.len(), 2);
assert!(group_lr.contains("learning_rate"));
assert!(group_lr.contains("n_layers"));
assert!(group_lr.contains(&lr_id));
assert!(group_lr.contains(&n_layers_id));
let group_opt = group_opt.expect("group with optimizer should exist");
assert_eq!(group_opt.len(), 1);
assert!(group_opt.contains("optimizer"));
assert!(group_opt.contains(&optimizer_id));
}
#[test]
@@ -1042,15 +1081,17 @@ mod tests {
step: None,
});
let center_id = ParamId::new();
let a_id = ParamId::new();
let b_id = ParamId::new();
let c_id = ParamId::new();
// Star topology: center connects to all others
// Trial 1: center, a
// Trial 2: center, b
// Trial 3: center, c
let trial1 = create_trial(
0,
vec![
("center", ParamValue::Float(0.1), dist.clone()),
("a", ParamValue::Float(0.2), dist.clone()),
(center_id, ParamValue::Float(0.1), dist.clone()),
(a_id, ParamValue::Float(0.2), dist.clone()),
],
1.0,
);
@@ -1058,8 +1099,8 @@ mod tests {
let trial2 = create_trial(
1,
vec![
("center", ParamValue::Float(0.3), dist.clone()),
("b", ParamValue::Float(0.4), dist.clone()),
(center_id, ParamValue::Float(0.3), dist.clone()),
(b_id, ParamValue::Float(0.4), dist.clone()),
],
0.5,
);
@@ -1067,8 +1108,8 @@ mod tests {
let trial3 = create_trial(
2,
vec![
("center", ParamValue::Float(0.5), dist.clone()),
("c", ParamValue::Float(0.6), dist),
(center_id, ParamValue::Float(0.5), dist.clone()),
(c_id, ParamValue::Float(0.6), dist),
],
0.3,
);
@@ -1078,9 +1119,9 @@ mod tests {
// All connected via center
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].len(), 4);
assert!(groups[0].contains("center"));
assert!(groups[0].contains("a"));
assert!(groups[0].contains("b"));
assert!(groups[0].contains("c"));
assert!(groups[0].contains(&center_id));
assert!(groups[0].contains(&a_id));
assert!(groups[0].contains(&b_id));
assert!(groups[0].contains(&c_id));
}
}
+1337 -289
View File
File diff suppressed because it is too large Load Diff
+206 -750
View File
File diff suppressed because it is too large Load Diff
+4
View File
@@ -2,6 +2,7 @@
/// The direction of optimization.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum Direction {
/// Minimize the objective value.
Minimize,
@@ -11,6 +12,7 @@ pub enum Direction {
/// The state of a trial in its lifecycle.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum TrialState {
/// The trial is currently running.
Running,
@@ -18,4 +20,6 @@ pub enum TrialState {
Complete,
/// The trial failed with an error.
Failed,
/// The trial was pruned (stopped early).
Pruned,
}
+61 -23
View File
@@ -4,6 +4,7 @@
#![cfg(feature = "async")]
use optimizer::parameter::{FloatParam, Parameter};
use optimizer::sampler::random::RandomSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, Error, Study};
@@ -13,10 +14,15 @@ async fn test_optimize_async_basic() {
let sampler = RandomSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-10.0, 10.0);
study
.optimize_async(10, |mut trial| async move {
let x = trial.suggest_float("x", -10.0, 10.0)?;
Ok::<_, Error>((trial, x * x))
.optimize_async(10, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x))
}
})
.await
.expect("async optimization should succeed");
@@ -27,7 +33,7 @@ async fn test_optimize_async_basic() {
}
#[tokio::test]
async fn test_optimize_async_with_sampler() {
async fn test_optimize_async_with_tpe() {
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(5)
@@ -36,10 +42,15 @@ async fn test_optimize_async_with_sampler() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
study
.optimize_async_with_sampler(15, |mut trial| async move {
let x = trial.suggest_float("x", -5.0, 5.0)?;
Ok::<_, Error>((trial, x * x))
.optimize_async(15, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x))
}
})
.await
.expect("async optimization with sampler should succeed");
@@ -54,10 +65,15 @@ async fn test_optimize_parallel() {
let sampler = RandomSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-10.0, 10.0);
study
.optimize_parallel(20, 4, |mut trial| async move {
let x = trial.suggest_float("x", -10.0, 10.0)?;
Ok::<_, Error>((trial, x * x))
.optimize_parallel(20, 4, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x))
}
})
.await
.expect("parallel optimization should succeed");
@@ -66,7 +82,7 @@ async fn test_optimize_parallel() {
}
#[tokio::test]
async fn test_optimize_parallel_with_sampler() {
async fn test_optimize_parallel_with_tpe() {
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(5)
@@ -75,11 +91,18 @@ async fn test_optimize_parallel_with_sampler() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_parallel_with_sampler(15, 3, |mut trial| async move {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
Ok::<_, Error>((trial, x * x + y * y))
.optimize_parallel(15, 3, move |mut trial| {
let x_param = x_param.clone();
let y_param = y_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
let y = y_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x + y * y))
}
})
.await
.expect("parallel optimization with sampler should succeed");
@@ -105,6 +128,7 @@ async fn test_optimize_async_all_failures() {
}
#[tokio::test]
#[allow(deprecated)]
async fn test_optimize_async_with_sampler_all_failures() {
let study: Study<f64> = Study::new(Direction::Minimize);
@@ -139,6 +163,7 @@ async fn test_optimize_parallel_all_failures() {
}
#[tokio::test]
#[allow(deprecated)]
async fn test_optimize_parallel_with_sampler_all_failures() {
let study: Study<f64> = Study::new(Direction::Minimize);
@@ -162,12 +187,15 @@ async fn test_optimize_async_partial_failures() {
let counter = std::sync::atomic::AtomicUsize::new(0);
let x_param = FloatParam::new(0.0, 10.0);
study
.optimize_async(10, |mut trial| {
.optimize_async(10, move |mut trial| {
let count = counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let x_param = x_param.clone();
async move {
if count.is_multiple_of(2) {
let x = trial.suggest_float("x", 0.0, 10.0)?;
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x))
} else {
Err(Error::NoCompletedTrials) // Use as error type
@@ -186,11 +214,16 @@ async fn test_optimize_parallel_high_concurrency() {
let sampler = RandomSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(0.0, 10.0);
// Run with concurrency higher than n_trials
study
.optimize_parallel(5, 10, |mut trial| async move {
let x = trial.suggest_float("x", 0.0, 10.0)?;
Ok::<_, Error>((trial, x))
.optimize_parallel(5, 10, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x))
}
})
.await
.expect("should handle high concurrency");
@@ -203,11 +236,16 @@ async fn test_optimize_parallel_single_concurrency() {
let sampler = RandomSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(0.0, 10.0);
// Run with concurrency of 1 (sequential)
study
.optimize_parallel(10, 1, |mut trial| async move {
let x = trial.suggest_float("x", 0.0, 10.0)?;
Ok::<_, Error>((trial, x))
.optimize_parallel(10, 1, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x))
}
})
.await
.expect("should work with single concurrency");
+133
View File
@@ -0,0 +1,133 @@
//! Integration tests for the BOHB sampler.
#![allow(
clippy::cast_sign_loss,
clippy::cast_precision_loss,
clippy::cast_possible_truncation
)]
use optimizer::parameter::{FloatParam, Parameter};
use optimizer::sampler::bohb::BohbSampler;
use optimizer::{Direction, Error, Study, TrialPruned};
#[test]
fn bohb_converges_on_quadratic() {
let bohb = BohbSampler::builder()
.min_resource(1)
.max_resource(9)
.reduction_factor(3)
.min_points_in_model(5)
.seed(42)
.build()
.unwrap();
let pruner = bohb.matching_pruner(Direction::Minimize);
let study: Study<f64> = Study::with_sampler_and_pruner(Direction::Minimize, bohb, pruner);
let x_param = FloatParam::new(-10.0, 10.0);
study
.optimize(60, |trial| {
let x = x_param.suggest(trial)?;
// Report intermediate values at budget steps 1, 3, 9
let obj = (x - 3.0).powi(2);
// Simulate budget-based evaluation with noise decreasing at higher budgets
trial.report(1, obj + 5.0);
trial.report(3, obj + 1.0);
trial.report(9, obj);
Ok::<_, Error>(obj)
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have trials");
assert!(
best.value < 10.0,
"BOHB should find a reasonable solution, got {}",
best.value
);
}
#[test]
fn bohb_with_pruning() {
let bohb = BohbSampler::builder()
.min_resource(1)
.max_resource(27)
.reduction_factor(3)
.min_points_in_model(3)
.seed(123)
.build()
.unwrap();
let pruner = bohb.matching_pruner(Direction::Minimize);
let study: Study<f64> = Study::with_sampler_and_pruner(Direction::Minimize, bohb, pruner);
let x_param = FloatParam::new(-5.0, 5.0);
study
.optimize(40, |trial| {
let x = x_param.suggest(trial)?;
let obj = x * x;
// Report at each rung step and check for pruning
for &step in &[1u64, 3, 9, 27] {
let noisy_obj = obj + 10.0 / step as f64;
trial.report(step, noisy_obj);
if trial.should_prune() {
return Err(TrialPruned.into());
}
}
Ok::<_, Error>(obj)
})
.expect("optimization should succeed");
// Verify we have completed trials
let best = study.best_trial().expect("should have at least one trial");
assert!(
best.value < 25.0,
"best value {} should be reasonable",
best.value
);
}
#[test]
fn bohb_uses_budget_conditioned_history() {
// Verify that BOHB conditions on budget level by testing that samples
// are influenced by intermediate values, not just final values.
let bohb = BohbSampler::builder()
.min_resource(1)
.max_resource(9)
.reduction_factor(3)
.min_points_in_model(3)
.seed(42)
.build()
.unwrap();
let pruner = bohb.matching_pruner(Direction::Minimize);
let study: Study<f64> = Study::with_sampler_and_pruner(Direction::Minimize, bohb, pruner);
let x_param = FloatParam::new(0.0, 10.0);
study
.optimize(30, |trial| {
let x = x_param.suggest(trial)?;
// Intermediate values that guide optimization toward x=2
trial.report(1, (x - 2.0).powi(2) + 1.0);
trial.report(3, (x - 2.0).powi(2) + 0.5);
trial.report(9, (x - 2.0).powi(2));
Ok::<_, Error>((x - 2.0).powi(2))
})
.expect("optimization should succeed");
let best = study.best_trial().unwrap();
let best_x: f64 = best.get(&x_param).unwrap();
// Should find x reasonably close to 2.0
assert!(
(best_x - 2.0).abs() < 5.0,
"BOHB should explore near x=2, got x={best_x}"
);
}
+259
View File
@@ -0,0 +1,259 @@
#![cfg(feature = "cma-es")]
use optimizer::prelude::*;
use optimizer::sampler::cma_es::CmaEsSampler;
#[test]
fn sphere_function() {
let sampler = CmaEsSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-5.0, 5.0).name("x");
let y = FloatParam::new(-5.0, 5.0).name("y");
study
.optimize(200, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
Ok::<_, Error>(xv * xv + yv * yv)
})
.unwrap();
let best = study.best_trial().unwrap();
assert!(
best.value < 1.0,
"sphere best value should be < 1.0, got {}",
best.value
);
}
#[test]
fn rosenbrock_function() {
let sampler = CmaEsSampler::builder().population_size(20).seed(42).build();
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-5.0, 5.0).name("x");
let y = FloatParam::new(-5.0, 5.0).name("y");
study
.optimize(300, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
let val = (1.0 - xv).powi(2) + 100.0 * (yv - xv * xv).powi(2);
Ok::<_, Error>(val)
})
.unwrap();
let best = study.best_trial().unwrap();
// Rosenbrock minimum is 0 at (1, 1); we just check reasonable convergence
assert!(
best.value < 50.0,
"rosenbrock best value should be < 50.0, got {}",
best.value
);
}
#[test]
fn bounds_respected() {
let sampler = CmaEsSampler::with_seed(123);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-2.0, 3.0).name("x");
let y = FloatParam::new(0.0, 10.0).name("y");
study
.optimize(100, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
Ok::<_, Error>(xv + yv)
})
.unwrap();
for trial in study.trials() {
let xv: f64 = trial.get(&x).unwrap();
let yv: f64 = trial.get(&y).unwrap();
assert!((-2.0..=3.0).contains(&xv), "x = {xv} out of bounds [-2, 3]");
assert!((0.0..=10.0).contains(&yv), "y = {yv} out of bounds [0, 10]");
}
}
#[test]
fn mixed_params_float_and_categorical() {
let sampler = CmaEsSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-5.0, 5.0).name("x");
let cat = CategoricalParam::new(vec!["a", "b", "c"]).name("cat");
study
.optimize(50, |trial| {
let xv = x.suggest(trial)?;
let cv = cat.suggest(trial)?;
let penalty = match cv {
"a" => 0.0,
"b" => 1.0,
_ => 2.0,
};
Ok::<_, Error>(xv * xv + penalty)
})
.unwrap();
let best = study.best_trial().unwrap();
// Should find a reasonable value
assert!(
best.value < 10.0,
"best value should be < 10.0, got {}",
best.value
);
}
#[test]
fn seeded_reproducibility() {
let x = FloatParam::new(-5.0, 5.0).name("x");
let y = FloatParam::new(-5.0, 5.0).name("y");
let run = |seed: u64| {
let sampler = CmaEsSampler::with_seed(seed);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
study
.optimize(50, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
Ok::<_, Error>(xv * xv + yv * yv)
})
.unwrap();
study.trials().iter().map(|t| t.value).collect::<Vec<_>>()
};
let results1 = run(42);
let results2 = run(42);
assert_eq!(results1, results2, "same seed should produce same results");
}
#[test]
fn different_seeds_different_results() {
let x = FloatParam::new(-5.0, 5.0).name("x");
let y = FloatParam::new(-5.0, 5.0).name("y");
let run = |seed: u64| {
let sampler = CmaEsSampler::with_seed(seed);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
study
.optimize(20, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
Ok::<_, Error>(xv * xv + yv * yv)
})
.unwrap();
study.trials().iter().map(|t| t.value).collect::<Vec<_>>()
};
let results1 = run(42);
let results2 = run(99);
assert_ne!(
results1, results2,
"different seeds should produce different results"
);
}
#[test]
fn single_dimension() {
let sampler = CmaEsSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-10.0, 10.0).name("x");
study
.optimize(100, |trial| {
let xv = x.suggest(trial)?;
Ok::<_, Error>((xv - 3.0).powi(2))
})
.unwrap();
let best = study.best_trial().unwrap();
assert!(
best.value < 1.0,
"1-D optimization should converge, got {}",
best.value
);
}
#[test]
fn integer_params() {
let sampler = CmaEsSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let n = IntParam::new(1, 20).name("n");
study
.optimize(100, |trial| {
let nv = n.suggest(trial)?;
// Minimum at n = 10
Ok::<_, Error>(((nv - 10) * (nv - 10)) as f64)
})
.unwrap();
let best = study.best_trial().unwrap();
let best_n: i64 = best.get(&n).unwrap();
assert!(
(1..=20).contains(&best_n),
"integer value {best_n} out of bounds"
);
assert!(
best.value < 10.0,
"integer optimization should converge, got {}",
best.value
);
}
#[test]
fn log_scale_params() {
let sampler = CmaEsSampler::with_seed(42);
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let lr = FloatParam::new(1e-5, 1.0).log_scale().name("lr");
study
.optimize(100, |trial| {
let lrv = lr.suggest(trial)?;
// Minimum at lr = 0.01
Ok::<_, Error>((lrv.ln() - 0.01_f64.ln()).powi(2))
})
.unwrap();
for trial in study.trials() {
let lrv: f64 = trial.get(&lr).unwrap();
assert!(
(1e-5..=1.0).contains(&lrv),
"log-scale value {lrv} out of bounds"
);
}
}
#[test]
fn custom_population_size_and_sigma() {
let sampler = CmaEsSampler::builder()
.sigma0(1.0)
.population_size(10)
.seed(42)
.build();
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x = FloatParam::new(-5.0, 5.0).name("x");
let y = FloatParam::new(-5.0, 5.0).name("y");
study
.optimize(100, |trial| {
let xv = x.suggest(trial)?;
let yv = y.suggest(trial)?;
Ok::<_, Error>(xv * xv + yv * yv)
})
.unwrap();
let best = study.best_trial().unwrap();
assert!(
best.value < 5.0,
"custom config optimization should work, got {}",
best.value
);
}
+61
View File
@@ -0,0 +1,61 @@
use optimizer::Trial;
use optimizer::parameter::{EnumParam, Parameter};
use optimizer_derive::Categorical;
#[derive(Clone, Debug, PartialEq, Categorical)]
enum Color {
Red,
Green,
Blue,
}
#[derive(Clone, Debug, PartialEq, Categorical)]
enum SingleVariant {
Only,
}
#[test]
fn derive_categorical_n_choices() {
use optimizer::parameter::Categorical;
assert_eq!(Color::N_CHOICES, 3);
assert_eq!(SingleVariant::N_CHOICES, 1);
}
#[test]
fn derive_categorical_roundtrip() {
use optimizer::parameter::Categorical;
for i in 0..Color::N_CHOICES {
let val = Color::from_index(i);
assert_eq!(val.to_index(), i);
}
}
#[test]
fn derive_categorical_values() {
use optimizer::parameter::Categorical;
assert_eq!(Color::from_index(0), Color::Red);
assert_eq!(Color::from_index(1), Color::Green);
assert_eq!(Color::from_index(2), Color::Blue);
assert_eq!(Color::Red.to_index(), 0);
assert_eq!(Color::Green.to_index(), 1);
assert_eq!(Color::Blue.to_index(), 2);
}
#[test]
fn derive_categorical_with_enum_param() {
let mut trial = Trial::new(0);
let param = EnumParam::<Color>::new();
let color = param.suggest(&mut trial).unwrap();
assert!([Color::Red, Color::Green, Color::Blue].contains(&color));
// Cached (same param id)
let color2 = param.suggest(&mut trial).unwrap();
assert_eq!(color, color2);
}
#[test]
fn derive_categorical_suggest_via_trial() {
let mut trial = Trial::new(0);
let color = trial.suggest_param(&EnumParam::<Color>::new()).unwrap();
assert!([Color::Red, Color::Green, Color::Blue].contains(&color));
}
+166
View File
@@ -0,0 +1,166 @@
use optimizer::parameter::{CategoricalParam, FloatParam, IntParam, Parameter};
use optimizer::{Direction, Study};
#[test]
fn known_perfect_correlation() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 100.0).name("x");
// Objective = x, so perfect correlation.
for _ in 0..30 {
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 1);
assert_eq!(importance[0].0, "x");
assert!(
(importance[0].1 - 1.0).abs() < 1e-10,
"single param should be 1.0"
);
}
#[test]
fn no_effect_parameter() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 100.0).name("x");
let noise = FloatParam::new(0.0, 100.0).name("noise");
// Objective depends only on x; noise is unused in objective.
for _ in 0..50 {
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
let _nv = noise.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 2);
// x should have much higher importance than noise.
let x_score = importance.iter().find(|(l, _)| l == "x").unwrap().1;
let noise_score = importance.iter().find(|(l, _)| l == "noise").unwrap().1;
assert!(
x_score > noise_score,
"x ({x_score}) should outrank noise ({noise_score})"
);
// x should dominate
assert!(x_score > 0.7, "x importance {x_score} should be dominant");
}
#[test]
fn multiple_parameters_varying_importance() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 10.0).name("x");
let y = FloatParam::new(0.0, 10.0).name("y");
// Objective = 10*x + 0.01*y → x should be far more important.
for _ in 0..50 {
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
let yv = y.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(10.0 * xv + 0.01 * yv));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 2);
assert_eq!(importance[0].0, "x", "x should rank first");
assert!(importance[0].1 > importance[1].1);
}
#[test]
fn fewer_than_two_trials_returns_empty() {
let study: Study<f64> = Study::new(Direction::Minimize);
// 0 trials
assert!(study.param_importance().is_empty());
// 1 trial
let x = FloatParam::new(0.0, 1.0).name("x");
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv));
assert!(study.param_importance().is_empty());
}
#[test]
fn int_parameter() {
let study: Study<f64> = Study::new(Direction::Minimize);
let n = IntParam::new(1, 100).name("n");
for _ in 0..30 {
let mut trial = study.ask();
let nv = n.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(nv as f64));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 1);
assert_eq!(importance[0].0, "n");
assert!(importance[0].1 > 0.9);
}
#[test]
fn categorical_parameter() {
let study: Study<f64> = Study::new(Direction::Minimize);
let cat = CategoricalParam::new(vec!["a", "b", "c"]).name("cat");
let x = FloatParam::new(0.0, 100.0).name("x");
// Objective depends only on x; categorical is random noise.
for _ in 0..50 {
let mut trial = study.ask();
let _c = cat.suggest(&mut trial).unwrap();
let xv = x.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 2);
let x_score = importance.iter().find(|(l, _)| l == "x").unwrap().1;
assert!(x_score > 0.5, "x should dominate over categorical noise");
}
#[test]
fn normalization_sums_to_one() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 10.0).name("x");
let y = FloatParam::new(0.0, 10.0).name("y");
let z = FloatParam::new(0.0, 10.0).name("z");
for _ in 0..50 {
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
let yv = y.suggest(&mut trial).unwrap();
let zv = z.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv + 0.5 * yv + 0.1 * zv));
}
let importance = study.param_importance();
let sum: f64 = importance.iter().map(|(_, s)| *s).sum();
assert!(
(sum - 1.0).abs() < 1e-10,
"scores should sum to 1.0, got {sum}"
);
}
#[test]
fn label_when_unnamed_uses_debug() {
let study: Study<f64> = Study::new(Direction::Minimize);
// No .name() call → label defaults to Debug representation.
let x = FloatParam::new(0.0, 10.0);
for _ in 0..10 {
let mut trial = study.ask();
let xv = x.suggest(&mut trial).unwrap();
study.tell(trial, Ok::<_, &str>(xv));
}
let importance = study.param_importance();
assert_eq!(importance.len(), 1);
assert!(
importance[0].0.starts_with("FloatParam"),
"expected Debug label, got {:?}",
importance[0].0
);
}
+1182 -343
View File
File diff suppressed because it is too large Load Diff
+228
View File
@@ -0,0 +1,228 @@
use std::collections::HashMap;
use optimizer::Direction;
use optimizer::pruner::{MedianPruner, Pruner};
use optimizer::sampler::CompletedTrial;
/// Helper to build a completed trial with given intermediate values.
fn trial_with_values(id: u64, intermediate_values: Vec<(u64, f64)>) -> CompletedTrial {
CompletedTrial::with_intermediate_values(
id,
HashMap::new(),
HashMap::new(),
HashMap::new(),
0.0,
intermediate_values,
HashMap::new(),
)
}
// --- Minimize direction ---
#[test]
fn prune_when_worse_than_median_minimize() {
let pruner = MedianPruner::new(Direction::Minimize);
// 3 completed trials with values at step 2: [1.0, 2.0, 3.0] => median = 2.0
let completed = vec![
trial_with_values(0, vec![(0, 0.5), (1, 0.8), (2, 1.0)]),
trial_with_values(1, vec![(0, 0.6), (1, 1.5), (2, 2.0)]),
trial_with_values(2, vec![(0, 0.7), (1, 2.0), (2, 3.0)]),
];
// Current trial value at step 2 is 2.5 > median 2.0 => prune
let current = vec![(0, 0.5), (1, 1.0), (2, 2.5)];
assert!(pruner.should_prune(3, 2, &current, &completed));
}
#[test]
fn no_prune_when_better_than_median_minimize() {
let pruner = MedianPruner::new(Direction::Minimize);
let completed = vec![
trial_with_values(0, vec![(0, 0.5), (1, 0.8), (2, 1.0)]),
trial_with_values(1, vec![(0, 0.6), (1, 1.5), (2, 2.0)]),
trial_with_values(2, vec![(0, 0.7), (1, 2.0), (2, 3.0)]),
];
// Current trial value at step 2 is 1.5 < median 2.0 => don't prune
let current = vec![(0, 0.5), (1, 1.0), (2, 1.5)];
assert!(!pruner.should_prune(3, 2, &current, &completed));
}
// --- Maximize direction ---
#[test]
fn prune_when_worse_than_median_maximize() {
let pruner = MedianPruner::new(Direction::Maximize);
// Values at step 1: [5.0, 7.0, 9.0] => median = 7.0
let completed = vec![
trial_with_values(0, vec![(0, 3.0), (1, 5.0)]),
trial_with_values(1, vec![(0, 4.0), (1, 7.0)]),
trial_with_values(2, vec![(0, 5.0), (1, 9.0)]),
];
// Current value 6.0 < median 7.0 => prune (worse for maximize)
let current = vec![(0, 4.0), (1, 6.0)];
assert!(pruner.should_prune(3, 1, &current, &completed));
}
#[test]
fn no_prune_when_better_than_median_maximize() {
let pruner = MedianPruner::new(Direction::Maximize);
let completed = vec![
trial_with_values(0, vec![(0, 3.0), (1, 5.0)]),
trial_with_values(1, vec![(0, 4.0), (1, 7.0)]),
trial_with_values(2, vec![(0, 5.0), (1, 9.0)]),
];
// Current value 8.0 > median 7.0 => don't prune
let current = vec![(0, 4.0), (1, 8.0)];
assert!(!pruner.should_prune(3, 1, &current, &completed));
}
// --- Warmup steps ---
#[test]
fn no_prune_during_warmup() {
let pruner = MedianPruner::new(Direction::Minimize).n_warmup_steps(5);
let completed = vec![trial_with_values(0, vec![(0, 1.0), (1, 1.0), (2, 1.0)])];
// Step 2 < warmup 5 => never prune, even if value is terrible
let current = vec![(0, 100.0), (1, 100.0), (2, 100.0)];
assert!(!pruner.should_prune(1, 2, &current, &completed));
}
#[test]
fn prune_after_warmup() {
let pruner = MedianPruner::new(Direction::Minimize).n_warmup_steps(2);
let completed = vec![trial_with_values(0, vec![(0, 1.0), (1, 1.0), (2, 1.0)])];
// Step 2 >= warmup 2 => pruning allowed; current 100.0 > median 1.0
let current = vec![(0, 100.0), (1, 100.0), (2, 100.0)];
assert!(pruner.should_prune(1, 2, &current, &completed));
}
// --- n_min_trials ---
#[test]
fn no_prune_when_fewer_than_n_min_trials() {
let pruner = MedianPruner::new(Direction::Minimize).n_min_trials(3);
// Only 2 completed trials — below threshold of 3
let completed = vec![
trial_with_values(0, vec![(0, 1.0)]),
trial_with_values(1, vec![(0, 2.0)]),
];
let current = vec![(0, 100.0)];
assert!(!pruner.should_prune(2, 0, &current, &completed));
}
#[test]
fn prune_when_at_least_n_min_trials() {
let pruner = MedianPruner::new(Direction::Minimize).n_min_trials(3);
// 3 completed trials with step 0: [1.0, 2.0, 3.0] => median 2.0
let completed = vec![
trial_with_values(0, vec![(0, 1.0)]),
trial_with_values(1, vec![(0, 2.0)]),
trial_with_values(2, vec![(0, 3.0)]),
];
// 5.0 > median 2.0 => prune
let current = vec![(0, 5.0)];
assert!(pruner.should_prune(3, 0, &current, &completed));
}
// --- No completed trials with values at step ---
#[test]
fn no_prune_when_no_completed_trials_at_step() {
let pruner = MedianPruner::new(Direction::Minimize);
// Completed trials only have values at step 0, not step 5
let completed = vec![
trial_with_values(0, vec![(0, 1.0)]),
trial_with_values(1, vec![(0, 2.0)]),
];
let current = vec![(0, 0.5), (5, 100.0)];
assert!(!pruner.should_prune(2, 5, &current, &completed));
}
// --- Median calculation edge cases ---
#[test]
fn correct_median_with_even_number_of_trials() {
let pruner = MedianPruner::new(Direction::Minimize);
// 4 trials at step 0: [1.0, 2.0, 3.0, 4.0] => median = 2.5
let completed = vec![
trial_with_values(0, vec![(0, 1.0)]),
trial_with_values(1, vec![(0, 2.0)]),
trial_with_values(2, vec![(0, 3.0)]),
trial_with_values(3, vec![(0, 4.0)]),
];
// 2.6 > 2.5 => prune
let current = vec![(0, 2.6)];
assert!(pruner.should_prune(4, 0, &current, &completed));
// 2.4 < 2.5 => don't prune
let current = vec![(0, 2.4)];
assert!(!pruner.should_prune(4, 0, &current, &completed));
}
#[test]
fn correct_median_with_odd_number_of_trials() {
let pruner = MedianPruner::new(Direction::Minimize);
// 5 trials at step 0: [1.0, 2.0, 3.0, 4.0, 5.0] => median = 3.0
let completed = vec![
trial_with_values(0, vec![(0, 1.0)]),
trial_with_values(1, vec![(0, 2.0)]),
trial_with_values(2, vec![(0, 3.0)]),
trial_with_values(3, vec![(0, 4.0)]),
trial_with_values(4, vec![(0, 5.0)]),
];
// 3.5 > 3.0 => prune
let current = vec![(0, 3.5)];
assert!(pruner.should_prune(5, 0, &current, &completed));
// 2.5 < 3.0 => don't prune
let current = vec![(0, 2.5)];
assert!(!pruner.should_prune(5, 0, &current, &completed));
}
// --- Non-contiguous step numbers ---
#[test]
fn works_with_non_contiguous_steps() {
let pruner = MedianPruner::new(Direction::Minimize);
// Steps are 0, 10, 100 — non-contiguous
let completed = vec![
trial_with_values(0, vec![(0, 1.0), (10, 2.0), (100, 3.0)]),
trial_with_values(1, vec![(0, 1.5), (10, 2.5), (100, 4.0)]),
trial_with_values(2, vec![(0, 2.0), (10, 3.0), (100, 5.0)]),
];
// At step 100: [3.0, 4.0, 5.0] => median = 4.0
let current = vec![(0, 1.0), (10, 2.0), (100, 4.5)];
assert!(pruner.should_prune(3, 100, &current, &completed));
let current = vec![(0, 1.0), (10, 2.0), (100, 3.5)];
assert!(!pruner.should_prune(3, 100, &current, &completed));
}
// --- No intermediate values for current trial ---
#[test]
fn no_prune_when_no_intermediate_values() {
let pruner = MedianPruner::new(Direction::Minimize);
let completed = vec![trial_with_values(0, vec![(0, 1.0)])];
assert!(!pruner.should_prune(1, 0, &[], &completed));
}
// --- Pruned trials are excluded from median calculation ---
#[test]
fn pruned_trials_excluded_from_median() {
use optimizer::TrialState;
let pruner = MedianPruner::new(Direction::Minimize);
let mut pruned = trial_with_values(0, vec![(0, 0.1)]);
pruned.state = TrialState::Pruned;
// Only the completed trial (value 5.0) counts. Pruned trial (0.1) is excluded.
let completed = vec![pruned, trial_with_values(1, vec![(0, 5.0)])];
// 3.0 < 5.0 => don't prune (only 1 completed trial with median 5.0)
let current = vec![(0, 3.0)];
assert!(!pruner.should_prune(2, 0, &current, &completed));
// 6.0 > 5.0 => prune
let current = vec![(0, 6.0)];
assert!(pruner.should_prune(2, 0, &current, &completed));
}
+64 -33
View File
@@ -9,8 +9,8 @@
clippy::cast_possible_truncation
)]
use optimizer::sampler::MultivariateTpeSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::parameter::{CategoricalParam, FloatParam, IntParam, Parameter};
use optimizer::sampler::tpe::{MultivariateTpeSampler, TpeSampler};
use optimizer::{Direction, Error, Study};
// =============================================================================
@@ -51,10 +51,13 @@ fn test_multivariate_tpe_rosenbrock_finds_good_solution() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-2.0, 2.0);
let y_param = FloatParam::new(-2.0, 4.0);
study
.optimize_with_sampler(100, |trial| {
let x = trial.suggest_float("x", -2.0, 2.0)?;
let y = trial.suggest_float("y", -2.0, 4.0)?;
.optimize(100, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.expect("optimization should succeed");
@@ -83,10 +86,13 @@ fn test_independent_tpe_rosenbrock() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-2.0, 2.0);
let y_param = FloatParam::new(-2.0, 4.0);
study
.optimize_with_sampler(100, |trial| {
let x = trial.suggest_float("x", -2.0, 2.0)?;
let y = trial.suggest_float("y", -2.0, 4.0)?;
.optimize(100, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.expect("optimization should succeed");
@@ -123,10 +129,13 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, multivariate_sampler);
let x_param = FloatParam::new(-2.0, 2.0);
let y_param = FloatParam::new(-2.0, 4.0);
study
.optimize_with_sampler(n_trials, |trial| {
let x = trial.suggest_float("x", -2.0, 2.0)?;
let y = trial.suggest_float("y", -2.0, 4.0)?;
.optimize(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.unwrap();
@@ -143,10 +152,13 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, independent_sampler);
let x_param = FloatParam::new(-2.0, 2.0);
let y_param = FloatParam::new(-2.0, 4.0);
study
.optimize_with_sampler(n_trials, |trial| {
let x = trial.suggest_float("x", -2.0, 2.0)?;
let y = trial.suggest_float("y", -2.0, 4.0)?;
.optimize(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.unwrap();
@@ -202,10 +214,13 @@ fn test_multivariate_tpe_independent_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_with_sampler(50, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
.optimize(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
@@ -231,10 +246,13 @@ fn test_independent_tpe_independent_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_with_sampler(50, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
.optimize(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
@@ -268,10 +286,13 @@ fn test_both_samplers_work_on_independent_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_with_sampler(n_trials, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
.optimize(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.unwrap();
@@ -287,10 +308,13 @@ fn test_both_samplers_work_on_independent_problem() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_with_sampler(n_trials, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
.optimize(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.unwrap();
@@ -332,10 +356,13 @@ fn test_multivariate_tpe_with_group_decomposition() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let y_param = FloatParam::new(-5.0, 5.0);
study
.optimize_with_sampler(50, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let y = trial.suggest_float("y", -5.0, 5.0)?;
.optimize(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
@@ -364,11 +391,15 @@ fn test_multivariate_tpe_mixed_parameter_types() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
let n_param = IntParam::new(1, 10);
let mode_param = CategoricalParam::new(vec!["a", "b", "c"]);
study
.optimize_with_sampler(50, |trial| {
let x = trial.suggest_float("x", -5.0, 5.0)?;
let n = trial.suggest_int("n", 1, 10)?;
let mode = trial.suggest_categorical("mode", &["a", "b", "c"])?;
.optimize(50, |trial| {
let x = x_param.suggest(trial)?;
let n = n_param.suggest(trial)?;
let mode = mode_param.suggest(trial)?;
// Objective depends on all parameters
let mode_factor = match mode {
+240
View File
@@ -0,0 +1,240 @@
use optimizer::parameter::{
BoolParam, Categorical, CategoricalParam, EnumParam, FloatParam, IntParam, Parameter,
};
use optimizer::{Direction, Study, Trial};
#[test]
fn suggest_float_param_via_trial() {
let param = FloatParam::new(0.0, 1.0);
let mut trial = Trial::new(0);
let x = trial.suggest_param(&param).unwrap();
assert!((0.0..=1.0).contains(&x));
// Cached
let x2 = trial.suggest_param(&param).unwrap();
assert_eq!(x, x2);
}
#[test]
fn suggest_float_log_param_via_trial() {
let param = FloatParam::new(1e-5, 1e-1).log_scale();
let mut trial = Trial::new(0);
let lr = trial.suggest_param(&param).unwrap();
assert!((1e-5..=1e-1).contains(&lr));
}
#[test]
fn suggest_float_step_param_via_trial() {
let param = FloatParam::new(0.0, 1.0).step(0.25);
let mut trial = Trial::new(0);
let x = trial.suggest_param(&param).unwrap();
assert!((0.0..=1.0).contains(&x));
}
#[test]
fn suggest_int_param_via_trial() {
let param = IntParam::new(1, 10);
let mut trial = Trial::new(0);
let n = trial.suggest_param(&param).unwrap();
assert!((1..=10).contains(&n));
// Cached
let n2 = trial.suggest_param(&param).unwrap();
assert_eq!(n, n2);
}
#[test]
fn suggest_int_log_param_via_trial() {
let param = IntParam::new(1, 1024).log_scale();
let mut trial = Trial::new(0);
let batch = trial.suggest_param(&param).unwrap();
assert!((1..=1024).contains(&batch));
}
#[test]
fn suggest_int_step_param_via_trial() {
let param = IntParam::new(32, 512).step(32);
let mut trial = Trial::new(0);
let units = trial.suggest_param(&param).unwrap();
assert!((32..=512).contains(&units));
assert_eq!((units - 32) % 32, 0);
}
#[test]
fn suggest_categorical_param_via_trial() {
let choices = vec!["sgd", "adam", "rmsprop"];
let param = CategoricalParam::new(choices.clone());
let mut trial = Trial::new(0);
let opt = trial.suggest_param(&param).unwrap();
assert!(choices.contains(&opt));
// Cached
let opt2 = trial.suggest_param(&param).unwrap();
assert_eq!(opt, opt2);
}
#[test]
fn suggest_bool_param_via_trial() {
let param = BoolParam::new();
let mut trial = Trial::new(0);
let val = trial.suggest_param(&param).unwrap();
let _ = val;
// Cached
let val2 = trial.suggest_param(&param).unwrap();
assert_eq!(val, val2);
}
#[derive(Clone, Debug, PartialEq)]
enum Activation {
Relu,
Sigmoid,
Tanh,
}
impl Categorical for Activation {
const N_CHOICES: usize = 3;
fn from_index(index: usize) -> Self {
match index {
0 => Activation::Relu,
1 => Activation::Sigmoid,
2 => Activation::Tanh,
_ => panic!("invalid index"),
}
}
fn to_index(&self) -> usize {
match self {
Activation::Relu => 0,
Activation::Sigmoid => 1,
Activation::Tanh => 2,
}
}
}
#[test]
fn suggest_enum_param_via_trial() {
let param = EnumParam::<Activation>::new();
let mut trial = Trial::new(0);
let act = trial.suggest_param(&param).unwrap();
assert!([Activation::Relu, Activation::Sigmoid, Activation::Tanh].contains(&act));
// Cached
let act2 = trial.suggest_param(&param).unwrap();
assert_eq!(act, act2);
}
#[test]
fn parameter_conflict_detection() {
let float_param = FloatParam::new(0.0, 1.0);
let int_param = IntParam::new(0, 10);
let mut trial = Trial::new(0);
let _ = trial.suggest_param(&float_param).unwrap();
// Different param type with different id - no conflict
let result = trial.suggest_param(&int_param);
assert!(result.is_ok());
// Different bounds for same param type but different id - no conflict
let float_param2 = FloatParam::new(0.0, 2.0);
let result = trial.suggest_param(&float_param2);
assert!(result.is_ok());
}
#[test]
fn validation_prevents_suggest() {
let mut trial = Trial::new(0);
assert!(trial.suggest_param(&FloatParam::new(1.0, 0.0)).is_err());
assert!(
trial
.suggest_param(&FloatParam::new(-1.0, 1.0).log_scale())
.is_err()
);
assert!(
trial
.suggest_param(&FloatParam::new(0.0, 1.0).step(-0.1))
.is_err()
);
assert!(trial.suggest_param(&IntParam::new(10, 1)).is_err());
assert!(
trial
.suggest_param(&IntParam::new(0, 10).log_scale())
.is_err()
);
assert!(trial.suggest_param(&IntParam::new(0, 10).step(-1)).is_err());
assert!(
trial
.suggest_param(&CategoricalParam::<&str>::new(vec![]))
.is_err()
);
}
#[test]
fn parameter_api_with_study() {
let x_param = FloatParam::new(-5.0, 5.0);
let n_param = IntParam::new(1, 10);
let dropout_param = BoolParam::new();
let opt_param = CategoricalParam::new(vec!["sgd", "adam"]);
let study: Study<f64> = Study::new(Direction::Minimize);
study
.optimize(5, |trial| {
let x = x_param.suggest(trial)?;
let n = n_param.suggest(trial)?;
let dropout = dropout_param.suggest(trial)?;
let opt = opt_param.suggest(trial)?;
let _ = (n, dropout, opt);
Ok::<_, optimizer::Error>(x * x)
})
.unwrap();
let best = study.best_trial().unwrap();
assert!(best.value >= 0.0);
}
#[test]
fn parameter_suggest_method() {
let param = FloatParam::new(0.0, 1.0);
let mut trial = Trial::new(0);
let x = param.suggest(&mut trial).unwrap();
assert!((0.0..=1.0).contains(&x));
}
#[test]
fn existing_suggest_methods_still_work() {
let mut trial = Trial::new(0);
let x_param = FloatParam::new(0.0, 1.0);
let x = x_param.suggest(&mut trial).unwrap();
assert!((0.0..=1.0).contains(&x));
let lr_param = FloatParam::new(1e-5, 1e-1).log_scale();
let lr = lr_param.suggest(&mut trial).unwrap();
assert!((1e-5..=1e-1).contains(&lr));
let step_param = FloatParam::new(0.0, 1.0).step(0.25);
let step = step_param.suggest(&mut trial).unwrap();
assert!((0.0..=1.0).contains(&step));
let n_param = IntParam::new(1, 10);
let n = n_param.suggest(&mut trial).unwrap();
assert!((1..=10).contains(&n));
let batch_param = IntParam::new(1, 1024).log_scale();
let batch = batch_param.suggest(&mut trial).unwrap();
assert!((1..=1024).contains(&batch));
let units_param = IntParam::new(32, 512).step(32);
let units = units_param.suggest(&mut trial).unwrap();
assert!((32..=512).contains(&units));
let opt_param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]);
let opt = opt_param.suggest(&mut trial).unwrap();
assert!(["sgd", "adam", "rmsprop"].contains(&opt));
let flag_param = BoolParam::new();
let flag = flag_param.suggest(&mut trial).unwrap();
let _ = flag;
}
+295
View File
@@ -0,0 +1,295 @@
#![cfg(feature = "serde")]
use std::collections::HashMap;
use optimizer::parameter::{FloatParam, IntParam, Parameter};
use optimizer::sampler::CompletedTrial;
use optimizer::{Direction, ParamValue, Study, StudySnapshot, TrialState};
#[test]
fn round_trip_save_load() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(-10.0, 10.0).name("x");
let n = IntParam::new(1, 100).name("n");
study
.optimize(5, |trial| {
let x_val = x.suggest(trial)?;
let n_val = n.suggest(trial)?;
Ok::<_, optimizer::Error>(x_val * x_val + n_val as f64)
})
.unwrap();
let dir = tempdir();
let path = dir.join("study.json");
study.save(&path).unwrap();
let loaded: Study<f64> = Study::load(&path).unwrap();
assert_eq!(loaded.direction(), study.direction());
assert_eq!(loaded.n_trials(), study.n_trials());
let orig_trials = study.trials();
let loaded_trials = loaded.trials();
for (orig, loaded) in orig_trials.iter().zip(loaded_trials.iter()) {
assert_eq!(orig.id, loaded.id);
assert!((orig.value - loaded.value).abs() < 1e-10);
assert_eq!(orig.state, loaded.state);
assert_eq!(orig.params.len(), loaded.params.len());
assert_eq!(orig.distributions, loaded.distributions);
assert_eq!(orig.param_labels, loaded.param_labels);
}
// Clean up
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn json_output_is_human_readable() {
let study: Study<f64> = Study::new(Direction::Maximize);
let x = FloatParam::new(0.0, 1.0).name("x");
study
.optimize(2, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v)
})
.unwrap();
let dir = tempdir();
let path = dir.join("study.json");
study.save(&path).unwrap();
let contents = std::fs::read_to_string(&path).unwrap();
// Verify it's pretty-printed JSON with recognizable fields
assert!(contents.contains("\"version\""));
assert!(contents.contains("\"direction\""));
assert!(contents.contains("\"trials\""));
assert!(contents.contains("\"next_trial_id\""));
assert!(contents.contains("\"Maximize\""));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn round_trip_empty_study() {
let study: Study<f64> = Study::new(Direction::Minimize);
let dir = tempdir();
let path = dir.join("empty.json");
study.save(&path).unwrap();
let loaded: Study<f64> = Study::load(&path).unwrap();
assert_eq!(loaded.direction(), Direction::Minimize);
assert_eq!(loaded.n_trials(), 0);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn snapshot_version_field_is_present() {
let study: Study<f64> = Study::new(Direction::Minimize);
let dir = tempdir();
let path = dir.join("version.json");
study.save(&path).unwrap();
let contents = std::fs::read_to_string(&path).unwrap();
let snapshot: StudySnapshot<f64> = serde_json::from_str(&contents).unwrap();
assert_eq!(snapshot.version, 1);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn completed_trial_serde_round_trip() {
let trial = CompletedTrial::new(42, HashMap::new(), HashMap::new(), HashMap::new(), 2.78);
let json = serde_json::to_string(&trial).unwrap();
let deserialized: CompletedTrial<f64> = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.id, 42);
assert_eq!(deserialized.value, 2.78);
assert_eq!(deserialized.state, TrialState::Complete);
}
#[test]
fn param_value_serde_round_trip() {
let values = vec![
ParamValue::Float(1.23),
ParamValue::Int(42),
ParamValue::Categorical(2),
];
for val in &values {
let json = serde_json::to_string(val).unwrap();
let deserialized: ParamValue = serde_json::from_str(&json).unwrap();
assert_eq!(&deserialized, val);
}
}
#[test]
fn direction_serde_round_trip() {
let min_json = serde_json::to_string(&Direction::Minimize).unwrap();
let max_json = serde_json::to_string(&Direction::Maximize).unwrap();
assert_eq!(
serde_json::from_str::<Direction>(&min_json).unwrap(),
Direction::Minimize
);
assert_eq!(
serde_json::from_str::<Direction>(&max_json).unwrap(),
Direction::Maximize
);
}
#[test]
fn round_trip_preserves_trial_id_counter() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(10, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v)
})
.unwrap();
let dir = tempdir();
let path = dir.join("counter.json");
study.save(&path).unwrap();
let loaded: Study<f64> = Study::load(&path).unwrap();
// Creating a new trial should use an ID >= 10
let trial = loaded.create_trial();
assert!(trial.id() >= 10);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn checkpoint_file_created_at_interval() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(-10.0, 10.0).name("x");
let dir = tempdir();
let checkpoint = dir.join("checkpoint.json");
study
.optimize_with_checkpoint(10, 3, &checkpoint, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v * v)
})
.unwrap();
// Checkpoint should exist (written at trials 3, 6, 9)
assert!(checkpoint.exists(), "checkpoint file was not created");
// Load it and verify it's valid
let loaded: Study<f64> = Study::load(&checkpoint).unwrap();
// Last checkpoint was at trial 9, so it should have 9 trials
assert_eq!(loaded.n_trials(), 9);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn checkpoint_overwrites_previous() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
let dir = tempdir();
let checkpoint = dir.join("checkpoint.json");
study
.optimize_with_checkpoint(6, 3, &checkpoint, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v)
})
.unwrap();
// The checkpoint at trial 6 should overwrite the one from trial 3
let loaded: Study<f64> = Study::load(&checkpoint).unwrap();
assert_eq!(loaded.n_trials(), 6);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resume_from_checkpoint_continues_trial_ids() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(-5.0, 5.0).name("x");
let dir = tempdir();
let checkpoint = dir.join("resume.json");
// Run 10 trials with checkpointing
study
.optimize_with_checkpoint(10, 5, &checkpoint, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v * v)
})
.unwrap();
// Load and continue
let loaded: Study<f64> = Study::load(&checkpoint).unwrap();
assert_eq!(loaded.n_trials(), 10);
let remaining = 15 - loaded.n_trials();
loaded
.optimize(remaining, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v * v)
})
.unwrap();
assert_eq!(loaded.n_trials(), 15);
// Verify no duplicate trial IDs
let trials = loaded.trials();
let mut ids: Vec<u64> = trials.iter().map(|t| t.id).collect();
ids.sort_unstable();
ids.dedup();
assert_eq!(ids.len(), 15, "duplicate trial IDs found");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn atomic_write_no_temp_file_left_behind() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
let dir = tempdir();
let checkpoint = dir.join("atomic.json");
study
.optimize_with_checkpoint(3, 3, &checkpoint, |trial| {
let v = x.suggest(trial)?;
Ok::<_, optimizer::Error>(v)
})
.unwrap();
// The temp file should have been renamed, not left behind
let tmp_path = dir.join(".atomic.json.tmp");
assert!(!tmp_path.exists(), "temp file was not cleaned up");
assert!(checkpoint.exists(), "checkpoint file was not created");
std::fs::remove_dir_all(&dir).ok();
}
/// Helper to create a unique temporary directory.
fn tempdir() -> std::path::PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let dir =
std::env::temp_dir().join(format!("optimizer_serde_test_{}_{id}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
dir
}
+45
View File
@@ -0,0 +1,45 @@
#[path = "../benches/test_functions.rs"]
mod test_functions;
use test_functions::*;
const TOL: f64 = 1e-10;
#[test]
fn sphere_at_optimum() {
assert!(sphere(&[0.0, 0.0]).abs() < TOL);
assert!(sphere(&[0.0; 10]).abs() < TOL);
}
#[test]
fn rosenbrock_at_optimum() {
assert!(rosenbrock(&[1.0, 1.0]).abs() < TOL);
assert!(rosenbrock(&[1.0; 5]).abs() < TOL);
}
#[test]
fn rastrigin_at_optimum() {
assert!(rastrigin(&[0.0, 0.0]).abs() < TOL);
assert!(rastrigin(&[0.0; 10]).abs() < TOL);
}
#[test]
fn ackley_at_optimum() {
assert!(ackley(&[0.0, 0.0]).abs() < 1e-8);
assert!(ackley(&[0.0; 10]).abs() < 1e-8);
}
#[test]
fn branin_at_optimum() {
let target = 0.397_887_357_729_738_1;
let val = branin(&[std::f64::consts::PI, 2.275]);
assert!((val - target).abs() < 1e-3);
}
#[test]
fn hartmann6_at_optimum() {
let target = -3.3224;
let x_opt = [0.20169, 0.150011, 0.476874, 0.275332, 0.311652, 0.6573];
let val = hartmann6(&x_opt);
assert!((val - target).abs() < 0.01);
}
+64
View File
@@ -0,0 +1,64 @@
use optimizer::pruner::{Pruner, ThresholdPruner};
#[test]
fn prune_when_value_exceeds_upper_threshold() {
let pruner = ThresholdPruner::new().upper(10.0);
let values = vec![(0, 5.0), (1, 8.0), (2, 11.0)];
assert!(pruner.should_prune(0, 2, &values, &[]));
}
#[test]
fn prune_when_value_falls_below_lower_threshold() {
let pruner = ThresholdPruner::new().lower(0.0);
let values = vec![(0, 5.0), (1, 2.0), (2, -1.0)];
assert!(pruner.should_prune(0, 2, &values, &[]));
}
#[test]
fn no_prune_when_value_within_bounds() {
let pruner = ThresholdPruner::new().upper(10.0).lower(0.0);
let values = vec![(0, 3.0), (1, 5.0), (2, 7.0)];
assert!(!pruner.should_prune(0, 2, &values, &[]));
}
#[test]
fn no_prune_when_no_intermediate_values() {
let pruner = ThresholdPruner::new().upper(10.0).lower(0.0);
assert!(!pruner.should_prune(0, 0, &[], &[]));
}
#[test]
fn works_with_only_upper_set() {
let pruner = ThresholdPruner::new().upper(5.0);
let below = vec![(0, 3.0)];
let above = vec![(0, 6.0)];
assert!(!pruner.should_prune(0, 0, &below, &[]));
assert!(pruner.should_prune(0, 0, &above, &[]));
}
#[test]
fn works_with_only_lower_set() {
let pruner = ThresholdPruner::new().lower(2.0);
let above = vec![(0, 5.0)];
let below = vec![(0, 1.0)];
assert!(!pruner.should_prune(0, 0, &above, &[]));
assert!(pruner.should_prune(0, 0, &below, &[]));
}
#[test]
fn works_with_both_thresholds_set() {
let pruner = ThresholdPruner::new().upper(10.0).lower(0.0);
// Within bounds
assert!(!pruner.should_prune(0, 0, &[(0, 5.0)], &[]));
// Exceeds upper
assert!(pruner.should_prune(0, 0, &[(0, 15.0)], &[]));
// Below lower
assert!(pruner.should_prune(0, 0, &[(0, -3.0)], &[]));
// At exact boundary (not pruned — strictly greater/less)
assert!(!pruner.should_prune(0, 0, &[(0, 10.0)], &[]));
assert!(!pruner.should_prune(0, 0, &[(0, 0.0)], &[]));
}
+148
View File
@@ -0,0 +1,148 @@
use optimizer::parameter::{FloatParam, Parameter};
use optimizer::{AttrValue, Direction, Study};
#[test]
fn set_and_get_float_attr() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("score", 42.5);
assert_eq!(trial.user_attr("score"), Some(&AttrValue::Float(42.5)));
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
}
#[test]
fn set_and_get_int_attr() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("epoch", 42_i64);
assert_eq!(trial.user_attr("epoch"), Some(&AttrValue::Int(42)));
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
}
#[test]
fn set_and_get_string_attr() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("model", "resnet50");
assert_eq!(
trial.user_attr("model"),
Some(&AttrValue::String("resnet50".to_owned()))
);
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
}
#[test]
fn set_and_get_bool_attr() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("converged", true);
assert_eq!(trial.user_attr("converged"), Some(&AttrValue::Bool(true)));
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
}
#[test]
fn attrs_propagate_to_completed_trial() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("time_secs", 1.5);
trial.set_user_attr("tag", "baseline");
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
let best = study.best_trial().unwrap();
assert_eq!(best.user_attr("time_secs"), Some(&AttrValue::Float(1.5)));
assert_eq!(
best.user_attr("tag"),
Some(&AttrValue::String("baseline".to_owned()))
);
}
#[test]
fn overwrite_attr_replaces_value() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("key", "old");
trial.set_user_attr("key", "new");
assert_eq!(
trial.user_attr("key"),
Some(&AttrValue::String("new".to_owned()))
);
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
let best = study.best_trial().unwrap();
assert_eq!(
best.user_attr("key"),
Some(&AttrValue::String("new".to_owned()))
);
}
#[test]
fn missing_attr_returns_none() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
assert_eq!(trial.user_attr("nonexistent"), None);
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
let best = study.best_trial().unwrap();
assert_eq!(best.user_attr("nonexistent"), None);
}
#[test]
fn user_attrs_map_returns_all() {
let study: Study<f64> = Study::new(Direction::Minimize);
let x = FloatParam::new(0.0, 1.0);
study
.optimize(1, |trial| {
let _ = x.suggest(trial)?;
trial.set_user_attr("a", 1.0);
trial.set_user_attr("b", true);
assert_eq!(trial.user_attrs().len(), 2);
Ok::<_, optimizer::Error>(1.0)
})
.unwrap();
let best = study.best_trial().unwrap();
assert_eq!(best.user_attrs().len(), 2);
}
+1
View File
@@ -1,2 +1,3 @@
[default.extend-words]
Tpe = "Tpe"
consts = "consts"