4 Commits

Author SHA1 Message Date
Manuel Raimann 762f1b3cde chore: release v0.3.1 2026-01-31 11:31:03 +01:00
Manuel Raimann b3534b3a3b feat: add tests for suggest_bool and suggest_range methods 2026-01-31 11:30:37 +01:00
Manuel Raimann 9944e69920 feat: add SuggestableRange trait and suggest_range method for parameter suggestion from ranges 2026-01-31 11:30:33 +01:00
Manuel Raimann e8e446fb0e feat: add suggest_bool method for boolean parameter suggestion 2026-01-31 11:19:55 +01:00
29 changed files with 1608 additions and 12985 deletions
-16
View File
@@ -61,21 +61,6 @@ jobs:
cargo test --verbose --features "${{ matrix.features }}"
fi
examples:
name: Examples
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- name: Install Rust
run: |
rustup override set stable
rustup update stable
- uses: Swatinem/rust-cache@v2
- name: Run sync example (ml_hyperparameter_tuning)
run: cargo run --example ml_hyperparameter_tuning
- name: Run async example (async_api_optimization)
run: cargo run --example async_api_optimization --features async
docs:
name: Docs
runs-on: ubuntu-latest
@@ -258,7 +243,6 @@ jobs:
- fmt
- clippy
- test
- examples
- docs
- feature-check
- coverage
+2 -22
View File
@@ -1,9 +1,6 @@
[workspace]
members = ["optimizer-derive"]
[package]
name = "optimizer"
version = "0.5.0"
version = "0.3.1"
edition = "2024"
rust-version = "1.88"
license = "MIT"
@@ -20,27 +17,10 @@ rand = "0.9"
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 }
[features]
default = []
async = ["dep:tokio"]
derive = ["dep:optimizer-derive"]
[dev-dependencies]
tokio = { version = "1", features = ["rt-multi-thread", "macros", "time"] }
optimizer-derive = { version = "0.1.0", path = "optimizer-derive" }
[[example]]
name = "async_api_optimization"
path = "examples/async_api_optimization.rs"
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"]
tokio = { version = "1", features = ["rt-multi-thread", "macros"] }
+4 -66
View File
@@ -13,39 +13,28 @@ 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, categorical, boolean, and enum parameter types
- Float, integer, and categorical 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::parameter::{FloatParam, Parameter};
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, Study};
use optimizer::sampler::tpe::TpeSampler;
// 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 = x_param.suggest(trial)?;
let x = trial.suggest_float("x", -10.0, 10.0)?;
Ok::<_, optimizer::Error>(x * x)
})
.unwrap();
// Get the best result
let best = study.best_trial().unwrap();
println!("Best value: {}", best.value);
for (id, label) in &best.param_labels {
println!(" {}: {:?}", label, best.params[id]);
}
println!("Best value: {} at x={:?}", best.value, best.params);
```
## Samplers
@@ -79,56 +68,6 @@ let sampler = TpeSampler::builder()
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
```
#### Gamma Strategies
The gamma parameter controls what fraction of trials are considered "good" when building the TPE model. Instead of a fixed value, you can use adaptive strategies:
| Strategy | Description | Formula |
|----------|-------------|---------|
| `FixedGamma` | Constant value (default: 0.25) | `γ = constant` |
| `LinearGamma` | Linear interpolation over trials | `γ = γ_min + (γ_max - γ_min) * min(n/n_max, 1)` |
| `SqrtGamma` | Optuna-style inverse sqrt scaling | `γ = min(γ_max, factor/√n / n)` |
| `HyperoptGamma` | Hyperopt-style adaptive | `γ = min(γ_max, (base + 1) / n)` |
```rust
use optimizer::sampler::tpe::{TpeSampler, SqrtGamma, LinearGamma};
// Optuna-style gamma that decreases with more trials
let sampler = TpeSampler::builder()
.gamma_strategy(SqrtGamma::default())
.build()
.unwrap();
// Linear interpolation from 0.1 to 0.3 over 100 trials
let sampler = TpeSampler::builder()
.gamma_strategy(LinearGamma::new(0.1, 0.3, 100).unwrap())
.build()
.unwrap();
```
You can also implement custom strategies:
```rust
use optimizer::sampler::tpe::{TpeSampler, GammaStrategy};
#[derive(Debug, Clone)]
struct MyGamma { base: f64 }
impl GammaStrategy for MyGamma {
fn gamma(&self, n_trials: usize) -> f64 {
(self.base + 0.01 * n_trials as f64).min(0.5)
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(self.clone())
}
}
let sampler = TpeSampler::builder()
.gamma_strategy(MyGamma { base: 0.1 })
.build()
.unwrap();
```
### Grid Search
```rust
@@ -145,7 +84,6 @@ 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
-368
View File
@@ -1,368 +0,0 @@
//! Async API Parameter Optimization Example
//!
//! This example shows how to use async/parallel optimization to tune
//! configuration parameters for a web service. Each evaluation simulates
//! an async operation (like deploying and load-testing a service).
//!
//! # Key Concepts Demonstrated
//!
//! - Async optimization with `optimize_parallel`
//! - Running multiple trials concurrently for faster optimization
//! - Boolean and categorical parameter types
//! - Measuring speedup from parallelism
//!
//! # When to Use Async Optimization
//!
//! Use async/parallel optimization when your objective function involves:
//! - Network requests (API calls, database queries)
//! - File I/O operations
//! - External service calls
//! - Any operation where you're waiting for I/O rather than computing
//!
//! With parallelism, you can evaluate multiple configurations simultaneously,
//! significantly reducing total optimization time.
//!
//! Run with: `cargo run --example async_api_optimization --features async`
use std::time::{Duration, Instant};
use optimizer::prelude::*;
// ============================================================================
// Configuration: Service parameters we want to tune
// ============================================================================
/// Configuration for a web service.
///
/// In a real application, these parameters would control:
/// - Memory allocation (cache sizes)
/// - Connection management (pool sizes, timeouts)
/// - Request handling (batching, compression)
/// - Protocol options (HTTP version, load balancing)
struct ServiceConfig {
cache_size_mb: i64,
connection_pool_size: i64,
request_timeout_ms: i64,
retry_count: i64,
batch_size: i64,
compression_level: i64,
use_http2: bool,
load_balancing: String,
}
// ============================================================================
// Objective Function: Evaluate a service configuration
// ============================================================================
/// Simulates deploying and load-testing a service configuration.
///
/// In a real scenario, this function might:
/// 1. Deploy the configuration to a staging environment
/// 2. Run load tests against the service
/// 3. Collect metrics (latency, throughput, error rate)
/// 4. Return a composite score
///
/// The async sleep simulates the I/O time of these operations.
/// This is where parallel execution helps - while one trial is waiting
/// for I/O, other trials can run.
#[allow(clippy::too_many_arguments)]
async fn evaluate_service(config: &ServiceConfig) -> f64 {
// Simulate async I/O (deployment, load testing, metric collection)
tokio::time::sleep(Duration::from_millis(50)).await;
// Calculate a score based on how close we are to optimal values
// Lower score = better configuration
let mut score = 0.0;
// Cache size: too small = cache misses, too large = wasted memory
// Optimal around 512MB
let cache_optimal = 512.0;
score += ((config.cache_size_mb as f64 - cache_optimal) / 256.0).powi(2);
// Connection pool: too small = contention, too large = resource waste
// Optimal around 100
let pool_optimal = 100.0;
score += ((config.connection_pool_size as f64 - pool_optimal) / 50.0).powi(2);
// Timeout: too short = false failures, too long = slow recovery
// Optimal around 5000ms
let timeout_optimal = 5000.0;
score += ((config.request_timeout_ms as f64 - timeout_optimal) / 2000.0).powi(2);
// Retries: too few = fragile, too many = amplifies failures
// Optimal around 3
let retry_optimal = 3.0;
score += ((config.retry_count as f64 - retry_optimal) / 2.0).powi(2);
// Batch size: trade-off between latency and throughput
// Optimal around 64
let batch_optimal = 64.0;
score += ((config.batch_size as f64 - batch_optimal) / 32.0).powi(2);
// Compression level: trade-off between CPU and bandwidth
// Optimal around 6
let compression_optimal = 6.0;
score += ((config.compression_level as f64 - compression_optimal) / 3.0).powi(2);
// HTTP/2 is generally better for our use case
if !config.use_http2 {
score += 0.5;
}
// Load balancing strategy affects performance
score += match config.load_balancing.as_str() {
"round_robin" => 0.0, // Best for our use case
"least_connections" => 0.1, // Good alternative
"ip_hash" => 0.2, // OK for session affinity
"random" => 0.3, // Not ideal
_ => 1.0,
};
// Add noise to simulate real-world variability
let noise = (config.cache_size_mb as f64 * 0.1).sin() * 0.05;
score + noise
}
/// The async objective function for each trial.
///
/// For async optimization, the objective function must:
/// 1. Take ownership of the Trial (not a mutable reference)
/// 2. Return a Future
/// 3. Return both the Trial and the result value as a tuple
///
/// This ownership pattern allows the trial to be used across await points.
#[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 {
cache_size_mb,
connection_pool_size,
request_timeout_ms,
retry_count,
batch_size,
compression_level,
use_http2,
load_balancing: load_balancing.to_string(),
};
// Evaluate (this is the async part)
let score = evaluate_service(&config).await;
// Return both the trial and the score
Ok((trial, score))
}
// ============================================================================
// Helper Functions
// ============================================================================
/// Prints the results of the optimization.
fn print_results(study: &Study<f64>, elapsed: Duration, n_trials: usize) {
println!("\n{}", "=".repeat(60));
println!("\nOptimization completed!");
println!("Total trials: {}", study.n_trials());
println!("Time elapsed: {elapsed:.2?}");
// Calculate speedup from parallelism
// Each trial takes ~50ms, so sequential would take n_trials * 50ms
let sequential_time = n_trials as f64 * 0.050;
let actual_time = elapsed.as_secs_f64();
println!(
"Effective parallelism: {:.1}x speedup",
sequential_time / actual_time
);
}
/// Prints the best configuration found.
#[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:");
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(())
}
/// Prints the top N trials.
fn print_top_trials(study: &Study<f64>, n: usize) {
println!("\nTop {n} trials:");
let mut trials = study.trials();
trials.sort_by(|a, b| a.value.partial_cmp(&b.value).unwrap());
for (i, trial) in trials.iter().take(n).enumerate() {
println!(
" {}. Trial #{}: score = {:.6}",
i + 1,
trial.id,
trial.value
);
}
}
// ============================================================================
// Main: Set up and run the async optimization
// ============================================================================
#[tokio::main]
async fn main() -> optimizer::Result<()> {
println!("=== Async API Parameter Optimization Example ===\n");
// Step 1: Create a TPE sampler
let sampler = TpeSampler::builder()
.n_startup_trials(8)
.gamma(0.2)
.seed(123)
.build()
.expect("Failed to build TPE sampler");
// Step 2: Create a study to minimize the score
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
// 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
println!("Starting parallel optimization with {concurrency} concurrent evaluations...\n");
let start = Instant::now();
// Step 5: Run parallel async optimization
//
// 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 sampler gets access to trial history for informed sampling.
study
.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,
&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(())
}
-275
View File
@@ -1,275 +0,0 @@
//! Machine Learning Hyperparameter Tuning Example
//!
//! This example shows how to use the optimizer library to find the best
//! hyperparameters for a machine learning model. We simulate a gradient
//! boosting model (like XGBoost or LightGBM) and search for optimal settings.
//!
//! # Key Concepts Demonstrated
//!
//! - Creating a Study with a TPE (Tree-Parzen Estimator) sampler
//! - Defining an objective function that the optimizer will minimize
//! - Using different parameter types: floats, integers, log-scale, stepped
//! - Using callbacks to monitor progress and implement early stopping
//!
//! # How It Works
//!
//! 1. Create a `Study` - this manages the optimization process
//! 2. Define an objective function that takes a `Trial` and returns a score
//! 3. Inside the objective, use `trial.suggest_*()` to sample parameters
//! 4. The optimizer runs many trials, learning which parameter regions work best
//! 5. After optimization, retrieve the best parameters found
//!
//! Run with: `cargo run --example ml_hyperparameter_tuning`
use std::ops::ControlFlow;
use optimizer::prelude::*;
// ============================================================================
// Configuration: Hyperparameters we want to tune
// ============================================================================
/// Holds all the hyperparameters for our model.
///
/// In a real application, you would pass these to your ML framework
/// (e.g., XGBoost, LightGBM, scikit-learn).
struct ModelConfig {
learning_rate: f64,
max_depth: i64,
n_estimators: i64,
subsample: f64,
colsample_bytree: f64,
min_child_weight: i64,
reg_alpha: f64,
reg_lambda: f64,
}
// ============================================================================
// Objective Function: What we want to optimize
// ============================================================================
/// Simulates training a model and returns the validation loss.
///
/// In a real scenario, this function would:
/// 1. Create a model with the given hyperparameters
/// 2. Train it on your training data
/// 3. Evaluate it on validation data
/// 4. Return the validation metric (e.g., RMSE, log loss, accuracy)
///
/// The optimizer will try to MINIMIZE this value (we set Direction::Minimize).
#[allow(clippy::too_many_arguments)]
fn evaluate_model(config: &ModelConfig) -> f64 {
// Simulated optimal hyperparameters:
// learning_rate ~ 0.05, max_depth ~ 6, n_estimators ~ 200
// subsample ~ 0.8, colsample_bytree ~ 0.8, min_child_weight ~ 3
// reg_alpha ~ 0.1, reg_lambda ~ 1.0
let mut loss = 0.15; // Base loss
// Each term penalizes deviation from the optimal value
loss += (config.learning_rate - 0.05).powi(2) * 100.0;
loss += ((config.max_depth - 6) as f64).powi(2) * 0.01;
loss += ((config.n_estimators - 200) as f64).powi(2) * 0.00001;
loss += (config.subsample - 0.8).powi(2) * 10.0;
loss += (config.colsample_bytree - 0.8).powi(2) * 10.0;
loss += ((config.min_child_weight - 3) as f64).powi(2) * 0.05;
loss += (config.reg_alpha - 0.1).powi(2) * 5.0;
loss += (config.reg_lambda - 1.0).powi(2) * 2.0;
// Add some noise to simulate real-world variability
let noise = (config.learning_rate * 1000.0).sin() * 0.01;
loss + noise
}
/// The objective function that the optimizer calls for each trial.
///
/// This function:
/// 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.
#[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 {
learning_rate,
max_depth,
n_estimators,
subsample,
colsample_bytree,
min_child_weight,
reg_alpha,
reg_lambda,
};
let loss = evaluate_model(&config);
Ok(loss)
}
// ============================================================================
// Callback Function: Monitor progress and implement early stopping
// ============================================================================
/// Called after each successful trial completes.
///
/// Use callbacks to:
/// - Log progress to console or file
/// - Save checkpoints
/// - Implement early stopping when a good solution is found
/// - Track metrics over time
///
/// Return `ControlFlow::Continue(())` to keep optimizing.
/// Return `ControlFlow::Break(())` to stop early.
fn on_trial_complete(study: &Study<f64>, trial: &CompletedTrial<f64>) -> ControlFlow<()> {
// 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 {
println!("\nEarly stopping: found excellent solution!");
return ControlFlow::Break(());
}
ControlFlow::Continue(())
}
// ============================================================================
// Main: Set up and run the optimization
// ============================================================================
fn main() -> optimizer::Result<()> {
println!("=== ML Hyperparameter Tuning Example ===\n");
// Step 1: Create a sampler
//
// TPE (Tree-Parzen Estimator) is a Bayesian optimization algorithm.
// It learns from previous trials to suggest better parameters.
// - n_startup_trials: Number of random trials before TPE kicks in
// - gamma: What fraction of trials are considered "good" (lower = more selective)
// - seed: For reproducibility
let sampler = TpeSampler::builder()
.n_startup_trials(10)
.gamma(0.25)
.seed(42)
.build()
.expect("Failed to build TPE sampler");
// Step 2: Create a study
//
// The study manages the optimization process. We want to MINIMIZE
// the loss (lower is better). Use Direction::Maximize for metrics
// where higher is better (like accuracy).
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
// Print header
println!("Starting hyperparameter optimization...\n");
println!(
"{:>5} {:>12} (parameters...) {:>12}",
"Trial", "Params", "Loss"
);
println!("{}", "-".repeat(60));
// 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 runs the objective function for up to
// n_trials iterations. After each trial, it calls the callback.
// The sampler gets access to trial history for informed sampling.
let n_trials = 50;
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));
println!("\nOptimization completed!");
println!("Total trials: {}", study.n_trials());
let best = study.best_trial()?;
println!("\nBest trial:");
println!(" Loss: {:.6}", best.value);
println!(" Parameters:");
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)
//
// Now you would take best.params and use them to train your final model
// on the full dataset.
Ok(())
}
-56
View File
@@ -1,56 +0,0 @@
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());
}
-16
View File
@@ -1,16 +0,0 @@
[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"
proc-macro2 = "1"
-71
View File
@@ -1,71 +0,0 @@
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()
}
-26
View File
@@ -46,32 +46,6 @@ pub enum Error {
#[error("KDE requires at least one sample")]
EmptySamples,
/// Returned when multivariate KDE samples have zero dimensions.
#[error("multivariate KDE samples must have at least one dimension")]
ZeroDimensions,
/// Returned when multivariate KDE samples have inconsistent dimensions.
#[error(
"dimension mismatch: expected {expected} dimensions but sample {sample_index} has {got}"
)]
DimensionMismatch {
/// The expected number of dimensions.
expected: usize,
/// The actual number of dimensions in the sample.
got: usize,
/// The index of the sample with mismatched dimensions.
sample_index: usize,
},
/// Returned when bandwidth vector length doesn't match the number of dimensions.
#[error("bandwidth dimension mismatch: expected {expected} bandwidths but got {got}")]
BandwidthDimensionMismatch {
/// The expected number of bandwidths.
expected: usize,
/// The actual number of bandwidths provided.
got: usize,
},
/// Returned when an internal invariant is violated.
#[error("internal error: {0}")]
Internal(&'static str),
-13
View File
@@ -1,13 +0,0 @@
//! Kernel Density Estimation for parameter distributions.
//!
//! This module provides kernel density estimators used by TPE samplers
//! to model probability distributions over good and bad trial regions.
//!
//! - [`univariate`] - Univariate (single-parameter) KDE
//! - [`multivariate`] - Multivariate (joint-parameter) KDE for capturing parameter dependencies
mod multivariate;
mod univariate;
pub(crate) use multivariate::MultivariateKDE;
pub(crate) use univariate::KernelDensityEstimator;
-921
View File
@@ -1,921 +0,0 @@
//! Multivariate Kernel Density Estimation for joint parameter distributions.
//!
//! This module provides a multivariate kernel density estimator that captures
//! dependencies between parameters. Unlike the univariate KDE which models each
//! parameter independently, the multivariate KDE models the joint distribution
//! to better capture correlations between parameters.
use rand::Rng;
use crate::error::{Error, Result};
/// A multivariate Gaussian kernel density estimator for joint distributions.
///
/// `MultivariateKDE` estimates a joint probability density function from a set
/// of multi-dimensional samples. This is used in multivariate TPE to model
/// the joint distributions l(x) and g(x) across multiple parameters simultaneously,
/// capturing their dependencies.
///
/// # Examples
///
/// ```ignore
/// use crate::kde::MultivariateKDE;
///
/// // Create samples: 3 samples with 2 dimensions each
/// let samples = vec![
/// vec![1.0, 2.0],
/// vec![1.5, 2.5],
/// vec![2.0, 3.0],
/// ];
/// let kde = MultivariateKDE::new(samples).unwrap();
///
/// // Get dimensionality
/// assert_eq!(kde.n_dims(), 2);
/// ```
#[allow(dead_code)] // Fields and methods will be used in subsequent stories (US-003, US-004)
#[derive(Clone, Debug)]
pub(crate) struct MultivariateKDE {
/// The sample points used to construct the KDE.
/// Each inner Vec is one sample with `n_dims` values.
samples: Vec<Vec<f64>>,
/// The bandwidth (standard deviation) for each dimension.
/// Uses a diagonal bandwidth matrix (independent bandwidths per dimension).
bandwidths: Vec<f64>,
/// The number of dimensions.
n_dims: usize,
}
#[allow(dead_code)] // Methods will be used in subsequent stories (US-003, US-004)
impl MultivariateKDE {
/// Creates a new multivariate KDE with automatic bandwidth selection using Scott's rule.
///
/// Scott's rule for multivariate KDE sets bandwidth per dimension as:
/// `h_j = n^(-1/(d+4)) * sigma_j`
///
/// where n is the number of samples, d is the dimensionality, and `sigma_j` is the
/// standard deviation of the j-th dimension.
///
/// # Errors
///
/// Returns `Error::EmptySamples` if `samples` is empty.
/// Returns `Error::DimensionMismatch` if samples have inconsistent dimensions.
/// Returns `Error::ZeroDimensions` if samples have zero dimensions.
#[allow(clippy::cast_precision_loss)]
pub(crate) fn new(samples: Vec<Vec<f64>>) -> Result<Self> {
if samples.is_empty() {
return Err(Error::EmptySamples);
}
let n_dims = samples[0].len();
if n_dims == 0 {
return Err(Error::ZeroDimensions);
}
// Check all samples have the same dimensionality
for (i, sample) in samples.iter().enumerate() {
if sample.len() != n_dims {
return Err(Error::DimensionMismatch {
expected: n_dims,
got: sample.len(),
sample_index: i,
});
}
}
let bandwidths = Self::scotts_rule_multivariate(&samples, n_dims);
Ok(Self {
samples,
bandwidths,
n_dims,
})
}
/// Creates a new multivariate KDE with specified bandwidths.
///
/// Use this when you want explicit control over the smoothing parameters.
///
/// # Errors
///
/// Returns `Error::EmptySamples` if `samples` is empty.
/// Returns `Error::DimensionMismatch` if samples have inconsistent dimensions.
/// Returns `Error::ZeroDimensions` if samples have zero dimensions.
/// Returns `Error::BandwidthDimensionMismatch` if bandwidths length doesn't match dimensions.
/// Returns `Error::InvalidBandwidth` if any bandwidth is not positive.
pub(crate) fn with_bandwidths(samples: Vec<Vec<f64>>, bandwidths: Vec<f64>) -> Result<Self> {
if samples.is_empty() {
return Err(Error::EmptySamples);
}
let n_dims = samples[0].len();
if n_dims == 0 {
return Err(Error::ZeroDimensions);
}
// Check all samples have the same dimensionality
for (i, sample) in samples.iter().enumerate() {
if sample.len() != n_dims {
return Err(Error::DimensionMismatch {
expected: n_dims,
got: sample.len(),
sample_index: i,
});
}
}
// Check bandwidths length matches dimensions
if bandwidths.len() != n_dims {
return Err(Error::BandwidthDimensionMismatch {
expected: n_dims,
got: bandwidths.len(),
});
}
// Check all bandwidths are positive
for &bw in &bandwidths {
if bw <= 0.0 {
return Err(Error::InvalidBandwidth(bw));
}
}
Ok(Self {
samples,
bandwidths,
n_dims,
})
}
/// Computes bandwidths using Scott's rule for multivariate KDE.
///
/// Scott's rule for d dimensions: `h_j = n^(-1/(d+4)) * sigma_j`
#[allow(clippy::cast_precision_loss)]
fn scotts_rule_multivariate(samples: &[Vec<f64>], n_dims: usize) -> Vec<f64> {
let n = samples.len() as f64;
let d = n_dims as f64;
// Scott's rule exponent for multivariate: -1/(d+4)
let exponent = -1.0 / (d + 4.0);
let scale_factor = n.powf(exponent);
(0..n_dims)
.map(|dim| {
let std_dev = Self::dimension_std_dev(samples, dim);
// For degenerate case where all samples are identical in this dimension,
// use a small positive bandwidth
if std_dev < f64::EPSILON {
1.0
} else {
scale_factor * std_dev
}
})
.collect()
}
/// Computes the sample standard deviation for a single dimension.
#[allow(clippy::cast_precision_loss)]
fn dimension_std_dev(samples: &[Vec<f64>], dim: usize) -> f64 {
let n = samples.len() as f64;
let values: Vec<f64> = samples.iter().map(|s| s[dim]).collect();
let mean = values.iter().sum::<f64>() / n;
let variance = values.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / n;
variance.sqrt()
}
/// Returns the number of dimensions.
pub(crate) fn n_dims(&self) -> usize {
self.n_dims
}
/// Returns the number of samples.
#[allow(dead_code)] // Will be used in subsequent stories
pub(crate) fn n_samples(&self) -> usize {
self.samples.len()
}
/// Returns the bandwidths for each dimension.
#[cfg(test)]
pub(crate) fn bandwidths(&self) -> &[f64] {
&self.bandwidths
}
/// Returns a reference to the samples.
#[allow(dead_code)] // Will be used in subsequent stories
pub(crate) fn samples(&self) -> &[Vec<f64>] {
&self.samples
}
/// Returns the log probability density at point `x`.
///
/// This computes the log-density for numerical stability, using the formula:
///
/// `log f(x) = log((1/n) * Σ_i Π_j K_hj((x_j - x_ij) / h_j))`
///
/// The computation uses the log-sum-exp trick for numerical stability:
///
/// `log(Σ exp(a_i)) = max(a) + log(Σ exp(a_i - max(a)))`
///
/// # Panics
///
/// Panics if `x.len() != self.n_dims`.
#[allow(clippy::cast_precision_loss)]
pub(crate) fn log_pdf(&self, x: &[f64]) -> f64 {
assert_eq!(
x.len(),
self.n_dims,
"Point dimension {} doesn't match KDE dimension {}",
x.len(),
self.n_dims
);
let n = self.samples.len() as f64;
// Precompute log normalization constant for each dimension
// For a Gaussian kernel: K_h(z) = (1/(h*sqrt(2*pi))) * exp(-0.5*z^2)
// log(K_h(z)) = -log(h) - 0.5*log(2*pi) - 0.5*z^2
let log_2pi = (2.0 * core::f64::consts::PI).ln();
let log_norm_per_dim: Vec<f64> = self
.bandwidths
.iter()
.map(|&h| -h.ln() - 0.5 * log_2pi)
.collect();
// Compute log of kernel contribution for each sample
// log(prod_j K_hj(z_j)) = sum_j log(K_hj(z_j))
let log_kernels: Vec<f64> = self
.samples
.iter()
.map(|sample| {
let mut log_kernel_sum = 0.0;
for j in 0..self.n_dims {
let z = (x[j] - sample[j]) / self.bandwidths[j];
log_kernel_sum += log_norm_per_dim[j] - 0.5 * z * z;
}
log_kernel_sum
})
.collect();
// Use log-sum-exp trick for numerical stability
// log(sum(exp(log_kernels))) = max + log(sum(exp(log_kernels - max)))
let max_log_kernel = log_kernels
.iter()
.copied()
.fold(f64::NEG_INFINITY, f64::max);
// Handle case where all log_kernels are -inf (extremely unlikely point)
if max_log_kernel.is_infinite() && max_log_kernel < 0.0 {
return f64::NEG_INFINITY;
}
let sum_exp: f64 = log_kernels
.iter()
.map(|&lk| (lk - max_log_kernel).exp())
.sum();
// log((1/n) * sum) = -log(n) + log(sum)
-n.ln() + max_log_kernel + sum_exp.ln()
}
/// Returns the probability density at point `x`.
///
/// The density is computed as the average of multivariate Gaussian kernels
/// centered at each sample point:
///
/// `f(x) = (1/n) * Σ_i Π_j K_hj((x_j - x_ij) / h_j)`
///
/// where `K_hj` is the univariate Gaussian kernel with bandwidth `h_j`.
///
/// This method computes in log-space for numerical stability and then
/// exponentiates the result.
///
/// # Panics
///
/// Panics if `x.len() != self.n_dims`.
pub(crate) fn pdf(&self, x: &[f64]) -> f64 {
self.log_pdf(x).exp()
}
/// Samples a point from the estimated joint density distribution.
///
/// Sampling works by:
/// 1. Uniformly selecting one of the kernel centers (samples)
/// 2. Adding independent Gaussian noise to each dimension with that
/// dimension's bandwidth as the standard deviation
///
/// This is equivalent to sampling from a mixture of multivariate Gaussians
/// with diagonal covariance matrices, where each mixture component is
/// centered at a sample point.
///
/// # Returns
///
/// A `Vec<f64>` of length `n_dims` representing a sample from the KDE.
pub(crate) fn sample<R: Rng>(&self, rng: &mut R) -> Vec<f64> {
// Select a random sample to center the kernel on
let idx = rng.random_range(0..self.samples.len());
let center = &self.samples[idx];
// Add independent Gaussian noise to each dimension
// Using Box-Muller transform for Gaussian sampling
center
.iter()
.zip(self.bandwidths.iter())
.map(|(&center_j, &bandwidth_j)| {
let u1: f64 = rng.random();
let u2: f64 = rng.random();
// Box-Muller transform: generates standard normal variate
let z = (-2.0 * u1.ln()).sqrt() * (2.0 * core::f64::consts::PI * u2).cos();
// Scale by bandwidth and shift by center
center_j + z * bandwidth_j
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_multivariate_kde_new_basic() {
let samples = vec![vec![1.0, 2.0], vec![1.5, 2.5], vec![2.0, 3.0]];
let kde = MultivariateKDE::new(samples).unwrap();
assert_eq!(kde.n_dims(), 2);
assert_eq!(kde.n_samples(), 3);
assert_eq!(kde.bandwidths().len(), 2);
}
#[test]
fn test_multivariate_kde_new_single_sample() {
let samples = vec![vec![1.0, 2.0, 3.0]];
let kde = MultivariateKDE::new(samples).unwrap();
assert_eq!(kde.n_dims(), 3);
assert_eq!(kde.n_samples(), 1);
// With single sample, bandwidths should default to 1.0 (degenerate case)
for &bw in kde.bandwidths() {
assert!((bw - 1.0).abs() < f64::EPSILON);
}
}
#[test]
fn test_multivariate_kde_new_single_dimension() {
let samples = vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0], vec![5.0]];
let kde = MultivariateKDE::new(samples).unwrap();
assert_eq!(kde.n_dims(), 1);
assert_eq!(kde.n_samples(), 5);
assert_eq!(kde.bandwidths().len(), 1);
// Bandwidth should be positive
assert!(kde.bandwidths()[0] > 0.0);
}
#[test]
fn test_multivariate_kde_scotts_rule() {
// Create samples with known statistics
// 10 samples, 2 dimensions
let samples: Vec<Vec<f64>> = (0..10)
.map(|i| {
let x = f64::from(i);
vec![x, x * 2.0] // Second dimension has 2x variance
})
.collect();
let kde = MultivariateKDE::new(samples).unwrap();
// Scott's rule: h = n^(-1/(d+4)) * sigma
// n=10, d=2: exponent = -1/6 ≈ -0.167
// 10^(-1/6) ≈ 0.681
// First dim std_dev ≈ 2.87, second ≈ 5.74
// Expected bandwidths: ~1.95 and ~3.91
let bw = kde.bandwidths();
assert!(
bw[0] > 1.0 && bw[0] < 3.0,
"First bandwidth {} unexpected",
bw[0]
);
assert!(
bw[1] > 2.0 && bw[1] < 6.0,
"Second bandwidth {} unexpected",
bw[1]
);
// Second bandwidth should be approximately 2x the first
assert!(
(bw[1] / bw[0] - 2.0).abs() < 0.1,
"Ratio {} not close to 2",
bw[1] / bw[0]
);
}
#[test]
fn test_multivariate_kde_with_bandwidths() {
let samples = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let bandwidths = vec![0.5, 1.0];
let kde = MultivariateKDE::with_bandwidths(samples, bandwidths).unwrap();
assert_eq!(kde.n_dims(), 2);
assert!((kde.bandwidths()[0] - 0.5).abs() < f64::EPSILON);
assert!((kde.bandwidths()[1] - 1.0).abs() < f64::EPSILON);
}
#[test]
fn test_multivariate_kde_empty_samples() {
let samples: Vec<Vec<f64>> = vec![];
let result = MultivariateKDE::new(samples);
assert!(matches!(result, Err(Error::EmptySamples)));
}
#[test]
fn test_multivariate_kde_zero_dimensions() {
let samples = vec![vec![], vec![]];
let result = MultivariateKDE::new(samples);
assert!(matches!(result, Err(Error::ZeroDimensions)));
}
#[test]
fn test_multivariate_kde_dimension_mismatch() {
let samples = vec![vec![1.0, 2.0], vec![3.0]]; // Second sample has wrong dimensions
let result = MultivariateKDE::new(samples);
assert!(matches!(
result,
Err(Error::DimensionMismatch {
expected: 2,
got: 1,
sample_index: 1
})
));
}
#[test]
fn test_multivariate_kde_with_bandwidths_wrong_length() {
let samples = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let bandwidths = vec![0.5]; // Only 1 bandwidth for 2 dimensions
let result = MultivariateKDE::with_bandwidths(samples, bandwidths);
assert!(matches!(
result,
Err(Error::BandwidthDimensionMismatch {
expected: 2,
got: 1
})
));
}
#[test]
fn test_multivariate_kde_with_bandwidths_zero() {
let samples = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let bandwidths = vec![0.5, 0.0]; // Second bandwidth is zero
let result = MultivariateKDE::with_bandwidths(samples, bandwidths);
assert!(matches!(result, Err(Error::InvalidBandwidth(bw)) if bw == 0.0));
}
#[test]
fn test_multivariate_kde_with_bandwidths_negative() {
let samples = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let bandwidths = vec![0.5, -1.0]; // Negative bandwidth
let result = MultivariateKDE::with_bandwidths(samples, bandwidths);
assert!(
matches!(result, Err(Error::InvalidBandwidth(bw)) if (bw - (-1.0)).abs() < f64::EPSILON)
);
}
#[test]
fn test_multivariate_kde_identical_samples() {
// All samples identical - should handle degenerate case
let samples = vec![vec![5.0, 10.0], vec![5.0, 10.0], vec![5.0, 10.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Bandwidths should default to 1.0 for degenerate dimensions
for &bw in kde.bandwidths() {
assert!((bw - 1.0).abs() < f64::EPSILON);
}
}
#[test]
fn test_multivariate_kde_high_dimensional() {
// Test with higher dimensions
let samples: Vec<Vec<f64>> = (0..20)
.map(|i| {
let x = f64::from(i);
vec![x, x * 0.5, x * 2.0, x * 0.1, x * 10.0]
})
.collect();
let kde = MultivariateKDE::new(samples).unwrap();
assert_eq!(kde.n_dims(), 5);
assert_eq!(kde.n_samples(), 20);
assert_eq!(kde.bandwidths().len(), 5);
// All bandwidths should be positive
for &bw in kde.bandwidths() {
assert!(bw > 0.0);
}
}
#[test]
fn test_multivariate_kde_samples_accessor() {
let samples = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let kde = MultivariateKDE::new(samples.clone()).unwrap();
assert_eq!(kde.samples(), &samples);
}
// ==================== PDF Tests ====================
#[test]
fn test_multivariate_kde_pdf_basic() {
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0], vec![2.0, 2.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Density should be positive everywhere
assert!(kde.pdf(&[0.0, 0.0]) > 0.0);
assert!(kde.pdf(&[1.0, 1.0]) > 0.0);
assert!(kde.pdf(&[2.0, 2.0]) > 0.0);
assert!(kde.pdf(&[0.5, 0.5]) > 0.0);
// Density should be higher near sample points
let near_density = kde.pdf(&[1.0, 1.0]);
let far_density = kde.pdf(&[10.0, 10.0]);
assert!(near_density > far_density);
}
#[test]
fn test_multivariate_kde_pdf_with_custom_bandwidths() {
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0]];
let bandwidths = vec![0.5, 0.5];
let kde = MultivariateKDE::with_bandwidths(samples, bandwidths).unwrap();
// Density should be positive
assert!(kde.pdf(&[0.5, 0.5]) > 0.0);
assert!(kde.pdf(&[0.0, 0.0]) > 0.0);
}
#[test]
fn test_multivariate_kde_pdf_single_sample() {
let samples = vec![vec![5.0, 10.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Should have positive density near the sample
assert!(kde.pdf(&[5.0, 10.0]) > 0.0);
assert!(kde.pdf(&[4.5, 9.5]) > 0.0);
}
#[test]
fn test_multivariate_kde_log_pdf_consistency() {
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0], vec![2.0, 2.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// log_pdf and pdf should be consistent: exp(log_pdf(x)) == pdf(x)
let test_points = vec![
vec![0.0, 0.0],
vec![1.0, 1.0],
vec![0.5, 0.5],
vec![3.0, 3.0],
];
for point in test_points {
let log_p = kde.log_pdf(&point);
let p = kde.pdf(&point);
let p_from_log = log_p.exp();
assert!(
(p - p_from_log).abs() < 1e-10,
"pdf={p}, exp(log_pdf)={p_from_log}"
);
}
}
#[test]
fn test_multivariate_kde_pdf_integrates_to_one_1d() {
// Test with 1D case first (should match univariate KDE behavior)
let samples = vec![vec![0.0], vec![1.0], vec![2.0], vec![3.0], vec![4.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Numerical integration over a wide range
let n_points = 1000;
let low = -10.0;
let high = 15.0;
let dx = (high - low) / f64::from(n_points);
let integral: f64 = (0..n_points)
.map(|i| {
let x = low + (f64::from(i) + 0.5) * dx;
kde.pdf(&[x]) * dx
})
.sum();
// Should be approximately 1.0 (within numerical error)
assert!(
(integral - 1.0).abs() < 0.02,
"1D integral = {integral}, expected ~1.0"
);
}
#[test]
fn test_multivariate_kde_pdf_integrates_to_one_2d() {
// Test with 2D case
let samples = vec![
vec![0.0, 0.0],
vec![1.0, 0.0],
vec![0.0, 1.0],
vec![1.0, 1.0],
vec![0.5, 0.5],
];
let kde = MultivariateKDE::new(samples).unwrap();
// Numerical integration over a 2D grid
let n_points = 100; // 100x100 = 10000 points
let low = -5.0;
let high = 6.0;
let dx = (high - low) / f64::from(n_points);
let mut integral = 0.0;
for i in 0..n_points {
for j in 0..n_points {
let x = low + (f64::from(i) + 0.5) * dx;
let y = low + (f64::from(j) + 0.5) * dx;
integral += kde.pdf(&[x, y]) * dx * dx;
}
}
// Should be approximately 1.0 (within numerical error)
// 2D integration has more error, so use larger tolerance
assert!(
(integral - 1.0).abs() < 0.05,
"2D integral = {integral}, expected ~1.0"
);
}
#[test]
fn test_multivariate_kde_pdf_symmetry() {
// With symmetric samples, PDF should be symmetric
let samples = vec![
vec![1.0, 0.0],
vec![-1.0, 0.0],
vec![0.0, 1.0],
vec![0.0, -1.0],
];
let kde = MultivariateKDE::new(samples).unwrap();
// Density at symmetric points should be equal
let d1 = kde.pdf(&[0.5, 0.0]);
let d2 = kde.pdf(&[-0.5, 0.0]);
assert!(
(d1 - d2).abs() < 1e-10,
"Symmetric points have different densities: {d1} vs {d2}"
);
let d3 = kde.pdf(&[0.0, 0.5]);
let d4 = kde.pdf(&[0.0, -0.5]);
assert!(
(d3 - d4).abs() < 1e-10,
"Symmetric points have different densities: {d3} vs {d4}"
);
}
#[test]
fn test_multivariate_kde_pdf_high_dimensional() {
// Test with higher dimensions
let samples: Vec<Vec<f64>> = (0..10)
.map(|i| {
let x = f64::from(i) * 0.1;
vec![x, x, x, x, x] // 5D
})
.collect();
let kde = MultivariateKDE::new(samples).unwrap();
// Density should be positive
assert!(kde.pdf(&[0.5, 0.5, 0.5, 0.5, 0.5]) > 0.0);
assert!(kde.pdf(&[0.0, 0.0, 0.0, 0.0, 0.0]) > 0.0);
// Log PDF should be finite
let log_p = kde.log_pdf(&[0.5, 0.5, 0.5, 0.5, 0.5]);
assert!(log_p.is_finite());
}
#[test]
fn test_multivariate_kde_pdf_numerical_stability() {
// Test numerical stability with points far from samples
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Very far point should have very small but non-negative density
let far_pdf = kde.pdf(&[100.0, 100.0]);
assert!(far_pdf >= 0.0);
assert!(far_pdf.is_finite() || far_pdf == 0.0);
// Log PDF should be finite (or -inf for zero density)
let far_log_pdf = kde.log_pdf(&[100.0, 100.0]);
assert!(far_log_pdf.is_finite() || far_log_pdf.is_infinite());
}
#[test]
#[should_panic(expected = "Point dimension")]
fn test_multivariate_kde_pdf_wrong_dimension() {
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0]];
let kde = MultivariateKDE::new(samples).unwrap();
// Should panic with wrong dimension
let _ = kde.pdf(&[0.0]); // Only 1 value for 2D KDE
}
// ==================== Sampling Tests ====================
#[test]
fn test_multivariate_kde_sample_basic() {
let samples = vec![vec![0.0, 0.0], vec![1.0, 1.0], vec![2.0, 2.0]];
let kde = MultivariateKDE::new(samples).unwrap();
let mut rng = rand::rng();
// Sample should have correct dimensionality
let sample = kde.sample(&mut rng);
assert_eq!(sample.len(), 2);
}
#[test]
fn test_multivariate_kde_sample_in_reasonable_range() {
let samples = vec![
vec![0.0, 0.0],
vec![1.0, 1.0],
vec![2.0, 2.0],
vec![3.0, 3.0],
vec![4.0, 4.0],
];
let kde = MultivariateKDE::new(samples).unwrap();
let mut rng = rand::rng();
// Samples should generally be in a reasonable range around the data
for _ in 0..100 {
let s = kde.sample(&mut rng);
// With high probability, samples should be within a few bandwidths
// of the data range. Use a generous range to avoid flaky tests.
assert!(
s[0] > -10.0 && s[0] < 15.0,
"Sample dimension 0: {} outside expected range",
s[0]
);
assert!(
s[1] > -10.0 && s[1] < 15.0,
"Sample dimension 1: {} outside expected range",
s[1]
);
}
}
#[test]
fn test_multivariate_kde_sample_single_sample() {
// When KDE has only one sample, all samples should be centered around it
let samples = vec![vec![5.0, 10.0]];
let kde = MultivariateKDE::new(samples).unwrap();
let mut rng = rand::rng();
// Generate many samples and check they cluster around (5.0, 10.0)
let n_samples = 100;
let mut sum_x = 0.0;
let mut sum_y = 0.0;
for _ in 0..n_samples {
let s = kde.sample(&mut rng);
sum_x += s[0];
sum_y += s[1];
}
let mean_x = sum_x / f64::from(n_samples);
let mean_y = sum_y / f64::from(n_samples);
// Mean should be close to the single sample point
// With 100 samples and bandwidth=1.0, we expect mean within ~0.3 of center
assert!(
(mean_x - 5.0).abs() < 1.0,
"Mean x={mean_x}, expected close to 5.0"
);
assert!(
(mean_y - 10.0).abs() < 1.0,
"Mean y={mean_y}, expected close to 10.0"
);
}
#[test]
fn test_multivariate_kde_sample_high_dimensional() {
// Test sampling in higher dimensions
let samples: Vec<Vec<f64>> = (0..10)
.map(|i| {
let x = f64::from(i) * 0.5;
vec![x, x * 2.0, x * 0.5, x + 1.0, x - 1.0] // 5D
})
.collect();
let kde = MultivariateKDE::new(samples).unwrap();
let mut rng = rand::rng();
// Sample should have correct dimensionality
for _ in 0..50 {
let sample = kde.sample(&mut rng);
assert_eq!(sample.len(), 5);
// All values should be finite
for &val in &sample {
assert!(val.is_finite(), "Sample value is not finite: {val}");
}
}
}
#[test]
#[allow(clippy::cast_precision_loss)]
fn test_multivariate_kde_sample_respects_bandwidth() {
// Create samples all at origin with large custom bandwidth
let data = vec![vec![0.0, 0.0], vec![0.0, 0.0], vec![0.0, 0.0]];
let bandwidths = vec![0.1, 10.0]; // Small bandwidth in x, large in y
let kde = MultivariateKDE::with_bandwidths(data, bandwidths).unwrap();
let mut rng = rand::rng();
// Generate samples and check variance in each dimension
let n_samples = 1000;
let mut values_x: Vec<f64> = Vec::with_capacity(n_samples);
let mut values_y: Vec<f64> = Vec::with_capacity(n_samples);
for _ in 0..n_samples {
let s = kde.sample(&mut rng);
values_x.push(s[0]);
values_y.push(s[1]);
}
// Compute sample variances
let n = n_samples as f64;
let mean_x: f64 = values_x.iter().sum::<f64>() / n;
let mean_y: f64 = values_y.iter().sum::<f64>() / n;
let var_x: f64 = values_x.iter().map(|x| (x - mean_x).powi(2)).sum::<f64>() / n;
let var_y: f64 = values_y.iter().map(|y| (y - mean_y).powi(2)).sum::<f64>() / n;
// Variance should be approximately bandwidth^2
// x: bandwidth=0.1, expected variance ~0.01
// y: bandwidth=10.0, expected variance ~100.0
assert!(
var_x < 0.05,
"X variance {var_x} too large for bandwidth 0.1"
);
assert!(
var_y > 50.0 && var_y < 200.0,
"Y variance {var_y} unexpected for bandwidth 10.0"
);
}
#[test]
fn test_multivariate_kde_sample_distribution_shape() {
// Create samples along a diagonal line
let data = vec![
vec![0.0, 0.0],
vec![1.0, 1.0],
vec![2.0, 2.0],
vec![3.0, 3.0],
vec![4.0, 4.0],
];
let kde = MultivariateKDE::new(data).unwrap();
let mut rng = rand::rng();
// Sample many points and verify the mean is near the center
let n_samples = 500;
let mut sum = [0.0, 0.0];
for _ in 0..n_samples {
let s = kde.sample(&mut rng);
sum[0] += s[0];
sum[1] += s[1];
}
let mean_x = sum[0] / f64::from(n_samples);
let mean_y = sum[1] / f64::from(n_samples);
// Mean should be close to (2.0, 2.0) - the center of the samples
assert!(
(mean_x - 2.0).abs() < 0.5,
"Mean x={mean_x}, expected close to 2.0"
);
assert!(
(mean_y - 2.0).abs() < 0.5,
"Mean y={mean_y}, expected close to 2.0"
);
}
#[test]
fn test_multivariate_kde_sample_deterministic_with_seeded_rng() {
use rand::SeedableRng;
let data = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let kde = MultivariateKDE::new(data).unwrap();
// Use a seeded RNG for reproducibility
let mut rng1 = rand::rngs::StdRng::seed_from_u64(42);
let mut rng2 = rand::rngs::StdRng::seed_from_u64(42);
// Same seed should produce same samples
let result1 = kde.sample(&mut rng1);
let result2 = kde.sample(&mut rng2);
assert!(
(result1[0] - result2[0]).abs() < f64::EPSILON,
"Samples with same seed differ in dimension 0"
);
assert!(
(result1[1] - result2[1]).abs() < f64::EPSILON,
"Samples with same seed differ in dimension 1"
);
}
}
+27 -78
View File
@@ -28,26 +28,24 @@
//! # Quick Start
//!
//! ```
//! use optimizer::prelude::*;
//! 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 = FloatParam::new(-10.0, 10.0).name("x");
//!
//! // Optimize x^2 for 20 trials
//! study
//! .optimize(20, |trial| {
//! let x_val = x.suggest(trial)?;
//! Ok::<_, Error>(x_val * x_val)
//! .optimize_with_sampler(20, |trial| {
//! let x = trial.suggest_float("x", -10.0, 10.0)?;
//! Ok::<_, optimizer::Error>(x * x)
//! })
//! .unwrap();
//!
//! // Get the best result
//! let best = study.best_trial().unwrap();
//! println!("x = {}", best.get(&x).unwrap());
//! println!("Best value: {} at x={:?}", best.value, best.params);
//! ```
//!
//! # Creating a Study
@@ -71,35 +69,29 @@
//!
//! # Suggesting Parameters
//!
//! Within the objective function, use parameter types to suggest values:
//! Within the objective function, use [`Trial`] to suggest parameter 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| {
//! 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)?;
//! // 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)?;
//!
//! // 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();
@@ -155,78 +147,35 @@
//!
//! ```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| {
//! let x_param = x_param.clone();
//! async move {
//! let x = x_param.suggest(&mut trial)?;
//! Ok((trial, x * x))
//! }
//! study.optimize_async(10, |mut trial| async move {
//! let x = trial.suggest_float("x", 0.0, 1.0)?;
//! Ok((trial, x * x))
//! }).await?;
//!
//! // Parallel with bounded concurrency
//! 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))
//! }
//! study.optimize_parallel(10, 4, |mut trial| async move {
//! let x = trial.suggest_float("x", 0.0, 1.0)?;
//! Ok((trial, x * x))
//! }).await?;
//! ```
//!
//! # Feature Flags
//!
//! - `async`: Enable async optimization methods (requires tokio)
//! - `derive`: Enable `#[derive(Categorical)]` for enum parameters
mod distribution;
mod error;
mod kde;
mod param;
pub mod parameter;
pub mod sampler;
mod study;
mod trial;
mod types;
pub use error::{Error, Result};
#[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 sampler::CompletedTrial;
pub use sampler::grid::GridSearchSampler;
pub use sampler::random::RandomSampler;
pub use sampler::tpe::TpeSampler;
pub use study::Study;
pub use trial::Trial;
pub use trial::{SuggestableRange, 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};
pub use crate::param::ParamValue;
pub use crate::parameter::{
BoolParam, Categorical, CategoricalParam, EnumParam, FloatParam, IntParam, Parameter,
};
pub use crate::sampler::CompletedTrial;
pub use crate::sampler::grid::GridSearchSampler;
pub use crate::sampler::random::RandomSampler;
pub use crate::sampler::tpe::TpeSampler;
pub use crate::study::Study;
pub use crate::trial::Trial;
pub use crate::types::Direction;
}
-10
View File
@@ -14,13 +14,3 @@ 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})"),
}
}
}
-995
View File
@@ -1,995 +0,0 @@
//! Central parameter trait and built-in parameter types.
//!
//! The [`Parameter`] trait provides a unified way to define parameter types
//! and suggest values from a [`Trial`]. Built-in implementations
//! cover floats, integers, categoricals, booleans, and enum types.
//!
//! # Example
//!
//! ```
//! use optimizer::Trial;
//! use optimizer::parameter::{BoolParam, FloatParam, IntParam, Parameter};
//!
//! let mut trial = Trial::new(0);
//!
//! let lr = FloatParam::new(1e-5, 1e-1)
//! .log_scale()
//! .suggest(&mut trial)
//! .unwrap();
//! let layers = IntParam::new(1, 10).suggest(&mut trial).unwrap();
//! let dropout = BoolParam::new().suggest(&mut trial).unwrap();
//! ```
use core::fmt::Debug;
use core::sync::atomic::{AtomicU64, Ordering};
use crate::distribution::{
CategoricalDistribution, Distribution, FloatDistribution, IntDistribution,
};
use crate::error::{Error, Result};
use crate::param::ParamValue;
use crate::trial::Trial;
static NEXT_PARAM_ID: AtomicU64 = AtomicU64::new(0);
/// A unique identifier for a parameter instance.
///
/// Each parameter is assigned a unique `ParamId` at creation time. Cloning a parameter
/// copies its `ParamId`, so clones refer to the same logical parameter.
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct ParamId(u64);
impl ParamId {
/// Creates a new unique `ParamId`.
pub fn new() -> Self {
Self(NEXT_PARAM_ID.fetch_add(1, Ordering::Relaxed))
}
}
impl Default for ParamId {
fn default() -> Self {
Self::new()
}
}
impl core::fmt::Display for ParamId {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "param_{}", self.0)
}
}
/// A trait for defining parameter types that can be suggested by a [`Trial`].
///
/// Implementors specify the distribution to sample from and how to convert
/// the raw [`ParamValue`] back into a typed value.
pub trait Parameter: Debug {
/// The typed value returned after sampling.
type Value;
/// Returns the unique identifier for this parameter.
fn id(&self) -> ParamId;
/// Returns the distribution that this parameter samples from.
fn distribution(&self) -> Distribution;
/// Converts a raw [`ParamValue`] into the typed value.
///
/// # Errors
///
/// Returns an error if the `ParamValue` variant doesn't match what this parameter expects.
fn cast_param_value(&self, param_value: &ParamValue) -> Result<Self::Value>;
/// Validates the parameter configuration.
///
/// Called before sampling. The default implementation accepts all configurations.
///
/// # Errors
///
/// Returns an error if the parameter configuration is invalid.
fn validate(&self) -> Result<()> {
Ok(())
}
/// Returns a human-readable label for this parameter.
///
/// Defaults to the `Debug` output of the parameter.
fn label(&self) -> String {
format!("{self:?}")
}
/// Suggests a value for this parameter from the given trial.
///
/// This is a convenience method that delegates to [`Trial::suggest_param`].
///
/// # Errors
///
/// Returns an error if validation fails, the parameter conflicts with
/// a previously suggested parameter of the same id, or sampling fails.
fn suggest(&self, trial: &mut Trial) -> Result<Self::Value>
where
Self: Sized,
{
trial.suggest_param(self)
}
}
/// A floating-point parameter with optional log-scale and step size.
///
/// # Example
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{FloatParam, Parameter};
///
/// let mut trial = Trial::new(0);
///
/// // Simple range
/// let x = FloatParam::new(0.0, 1.0).suggest(&mut trial).unwrap();
///
/// // Log-scale
/// let lr = FloatParam::new(1e-5, 1e-1)
/// .log_scale()
/// .suggest(&mut trial)
/// .unwrap();
///
/// // Stepped
/// let step = FloatParam::new(0.0, 1.0)
/// .step(0.25)
/// .suggest(&mut trial)
/// .unwrap();
/// ```
#[derive(Clone, Debug)]
pub struct FloatParam {
id: ParamId,
low: f64,
high: f64,
log_scale: bool,
step: Option<f64>,
name: Option<String>,
}
impl FloatParam {
/// Creates a new float parameter with the given bounds.
#[must_use]
pub fn new(low: f64, high: f64) -> Self {
Self {
id: ParamId::new(),
low,
high,
log_scale: false,
step: None,
name: None,
}
}
/// Enables log-scale sampling.
#[must_use]
pub fn log_scale(mut self) -> Self {
self.log_scale = true;
self
}
/// Sets a step size for discretized sampling.
#[must_use]
pub fn step(mut self, step: f64) -> Self {
self.step = Some(step);
self
}
/// Sets a human-readable name for this parameter.
///
/// When set, this name is used as the parameter's label instead of
/// the default `Debug` output.
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
}
impl Parameter for FloatParam {
type Value = f64;
fn id(&self) -> ParamId {
self.id
}
fn distribution(&self) -> Distribution {
Distribution::Float(FloatDistribution {
low: self.low,
high: self.high,
log_scale: self.log_scale,
step: self.step,
})
}
fn cast_param_value(&self, param_value: &ParamValue) -> Result<f64> {
match param_value {
ParamValue::Float(v) => Ok(*v),
_ => Err(Error::Internal(
"Float distribution should return Float value",
)),
}
}
fn validate(&self) -> Result<()> {
if self.low > self.high {
return Err(Error::InvalidBounds {
low: self.low,
high: self.high,
});
}
if self.log_scale && self.low <= 0.0 {
return Err(Error::InvalidLogBounds);
}
if let Some(step) = self.step
&& step <= 0.0
{
return Err(Error::InvalidStep);
}
Ok(())
}
fn label(&self) -> String {
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
}
}
/// An integer parameter with optional log-scale and step size.
///
/// # Example
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{IntParam, Parameter};
///
/// let mut trial = Trial::new(0);
///
/// // Simple range
/// let n = IntParam::new(1, 10).suggest(&mut trial).unwrap();
///
/// // Log-scale
/// let batch = IntParam::new(1, 1024)
/// .log_scale()
/// .suggest(&mut trial)
/// .unwrap();
///
/// // Stepped
/// let units = IntParam::new(32, 512).step(32).suggest(&mut trial).unwrap();
/// ```
#[derive(Clone, Debug)]
pub struct IntParam {
id: ParamId,
low: i64,
high: i64,
log_scale: bool,
step: Option<i64>,
name: Option<String>,
}
impl IntParam {
/// Creates a new integer parameter with the given bounds.
#[must_use]
pub fn new(low: i64, high: i64) -> Self {
Self {
id: ParamId::new(),
low,
high,
log_scale: false,
step: None,
name: None,
}
}
/// Enables log-scale sampling.
#[must_use]
pub fn log_scale(mut self) -> Self {
self.log_scale = true;
self
}
/// Sets a step size for discretized sampling.
#[must_use]
pub fn step(mut self, step: i64) -> Self {
self.step = Some(step);
self
}
/// Sets a human-readable name for this parameter.
///
/// When set, this name is used as the parameter's label instead of
/// the default `Debug` output.
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
}
impl Parameter for IntParam {
type Value = i64;
fn id(&self) -> ParamId {
self.id
}
fn distribution(&self) -> Distribution {
Distribution::Int(IntDistribution {
low: self.low,
high: self.high,
log_scale: self.log_scale,
step: self.step,
})
}
#[allow(clippy::cast_precision_loss)]
fn cast_param_value(&self, param_value: &ParamValue) -> Result<i64> {
match param_value {
ParamValue::Int(v) => Ok(*v),
_ => Err(Error::Internal("Int distribution should return Int value")),
}
}
#[allow(clippy::cast_precision_loss)]
fn validate(&self) -> Result<()> {
if self.low > self.high {
return Err(Error::InvalidBounds {
low: self.low as f64,
high: self.high as f64,
});
}
if self.log_scale && self.low < 1 {
return Err(Error::InvalidLogBounds);
}
if let Some(step) = self.step
&& step <= 0
{
return Err(Error::InvalidStep);
}
Ok(())
}
fn label(&self) -> String {
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
}
}
/// A categorical parameter that selects from a list of choices.
///
/// # Example
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{CategoricalParam, Parameter};
///
/// let mut trial = Trial::new(0);
/// let opt = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"])
/// .suggest(&mut trial)
/// .unwrap();
/// ```
#[derive(Clone, Debug)]
pub struct CategoricalParam<T: Clone> {
id: ParamId,
choices: Vec<T>,
name: Option<String>,
}
impl<T: Clone> CategoricalParam<T> {
/// Creates a new categorical parameter with the given choices.
#[must_use]
pub fn new(choices: Vec<T>) -> Self {
Self {
id: ParamId::new(),
choices,
name: None,
}
}
/// Sets a human-readable name for this parameter.
///
/// When set, this name is used as the parameter's label instead of
/// the default `Debug` output.
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
}
impl<T: Clone + Debug> Parameter for CategoricalParam<T> {
type Value = T;
fn id(&self) -> ParamId {
self.id
}
fn distribution(&self) -> Distribution {
Distribution::Categorical(CategoricalDistribution {
n_choices: self.choices.len(),
})
}
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
match param_value {
ParamValue::Categorical(index) => Ok(self.choices[*index].clone()),
_ => Err(Error::Internal(
"Categorical distribution should return Categorical value",
)),
}
}
fn validate(&self) -> Result<()> {
if self.choices.is_empty() {
return Err(Error::EmptyChoices);
}
Ok(())
}
fn label(&self) -> String {
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
}
}
/// A boolean parameter (equivalent to a categorical with `[false, true]`).
///
/// # Example
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{BoolParam, Parameter};
///
/// let mut trial = Trial::new(0);
/// let dropout = BoolParam::new().suggest(&mut trial).unwrap();
/// ```
#[derive(Clone, Debug)]
pub struct BoolParam {
id: ParamId,
name: Option<String>,
}
impl BoolParam {
/// Creates a new boolean parameter.
#[must_use]
pub fn new() -> Self {
Self {
id: ParamId::new(),
name: None,
}
}
/// Sets a human-readable name for this parameter.
///
/// When set, this name is used as the parameter's label instead of
/// the default `Debug` output.
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
}
impl Default for BoolParam {
fn default() -> Self {
Self::new()
}
}
impl Parameter for BoolParam {
type Value = bool;
fn id(&self) -> ParamId {
self.id
}
fn distribution(&self) -> Distribution {
Distribution::Categorical(CategoricalDistribution { n_choices: 2 })
}
fn cast_param_value(&self, param_value: &ParamValue) -> Result<bool> {
match param_value {
ParamValue::Categorical(index) => Ok(*index != 0),
_ => Err(Error::Internal(
"Categorical distribution should return Categorical value",
)),
}
}
fn label(&self) -> String {
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
}
}
/// A trait for enum types that can be used as categorical parameters.
///
/// This trait maps enum variants to sequential indices and back. It can be
/// derived automatically for fieldless enums using `#[derive(Categorical)]`
/// when the `derive` feature is enabled.
///
/// # Example
///
/// Manual implementation:
///
/// ```
/// use optimizer::parameter::Categorical;
///
/// #[derive(Clone)]
/// 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,
/// }
/// }
/// }
/// ```
pub trait Categorical: Sized + Clone {
/// The number of variants in the enum.
const N_CHOICES: usize;
/// Creates an instance from a variant index.
///
/// # Panics
///
/// Panics if `index >= N_CHOICES`.
fn from_index(index: usize) -> Self;
/// Returns the index of this variant.
fn to_index(&self) -> usize;
}
/// A parameter that selects from the variants of an enum implementing [`Categorical`].
///
/// # Example
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{Categorical, EnumParam, Parameter};
///
/// #[derive(Clone, Debug)]
/// enum Optimizer {
/// Sgd,
/// Adam,
/// Rmsprop,
/// }
///
/// impl Categorical for Optimizer {
/// const N_CHOICES: usize = 3;
/// fn from_index(index: usize) -> Self {
/// match index {
/// 0 => Optimizer::Sgd,
/// 1 => Optimizer::Adam,
/// 2 => Optimizer::Rmsprop,
/// _ => panic!("invalid index"),
/// }
/// }
/// fn to_index(&self) -> usize {
/// match self {
/// Optimizer::Sgd => 0,
/// Optimizer::Adam => 1,
/// Optimizer::Rmsprop => 2,
/// }
/// }
/// }
///
/// let mut trial = Trial::new(0);
/// let opt = EnumParam::<Optimizer>::new().suggest(&mut trial).unwrap();
/// ```
#[derive(Clone, Debug)]
pub struct EnumParam<T: Categorical> {
id: ParamId,
name: Option<String>,
_marker: core::marker::PhantomData<T>,
}
impl<T: Categorical> EnumParam<T> {
/// Creates a new enum parameter.
#[must_use]
pub fn new() -> Self {
Self {
id: ParamId::new(),
name: None,
_marker: core::marker::PhantomData,
}
}
/// Sets a human-readable name for this parameter.
///
/// When set, this name is used as the parameter's label instead of
/// the default `Debug` output.
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
}
impl<T: Categorical> Default for EnumParam<T> {
fn default() -> Self {
Self::new()
}
}
impl<T: Categorical + Debug> Parameter for EnumParam<T> {
type Value = T;
fn id(&self) -> ParamId {
self.id
}
fn distribution(&self) -> Distribution {
Distribution::Categorical(CategoricalDistribution {
n_choices: T::N_CHOICES,
})
}
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
match param_value {
ParamValue::Categorical(index) => Ok(T::from_index(*index)),
_ => Err(Error::Internal(
"Categorical distribution should return Categorical value",
)),
}
}
fn label(&self) -> String {
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn float_param_distribution() {
let param = FloatParam::new(0.0, 1.0);
assert_eq!(
param.distribution(),
Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: None,
})
);
}
#[test]
fn float_param_log_scale() {
let param = FloatParam::new(1e-5, 1e-1).log_scale();
assert_eq!(
param.distribution(),
Distribution::Float(FloatDistribution {
low: 1e-5,
high: 1e-1,
log_scale: true,
step: None,
})
);
}
#[test]
fn float_param_step() {
let param = FloatParam::new(0.0, 1.0).step(0.25);
assert_eq!(
param.distribution(),
Distribution::Float(FloatDistribution {
low: 0.0,
high: 1.0,
log_scale: false,
step: Some(0.25),
})
);
}
#[test]
fn float_param_validate_invalid_bounds() {
let param = FloatParam::new(1.0, 0.0);
assert!(param.validate().is_err());
}
#[test]
fn float_param_validate_invalid_log() {
let param = FloatParam::new(-1.0, 1.0).log_scale();
assert!(param.validate().is_err());
}
#[test]
fn float_param_validate_invalid_step() {
let param = FloatParam::new(0.0, 1.0).step(-0.1);
assert!(param.validate().is_err());
}
#[test]
#[allow(clippy::float_cmp)]
fn float_param_cast_param_value() {
let param = FloatParam::new(0.0, 1.0);
assert_eq!(
param.cast_param_value(&ParamValue::Float(0.5)).unwrap(),
0.5
);
assert!(param.cast_param_value(&ParamValue::Int(1)).is_err());
}
#[test]
fn int_param_distribution() {
let param = IntParam::new(1, 100);
assert_eq!(
param.distribution(),
Distribution::Int(IntDistribution {
low: 1,
high: 100,
log_scale: false,
step: None,
})
);
}
#[test]
fn int_param_log_scale() {
let param = IntParam::new(1, 1024).log_scale();
assert_eq!(
param.distribution(),
Distribution::Int(IntDistribution {
low: 1,
high: 1024,
log_scale: true,
step: None,
})
);
}
#[test]
fn int_param_step() {
let param = IntParam::new(100, 500).step(50);
assert_eq!(
param.distribution(),
Distribution::Int(IntDistribution {
low: 100,
high: 500,
log_scale: false,
step: Some(50),
})
);
}
#[test]
fn int_param_validate_invalid_bounds() {
let param = IntParam::new(10, 1);
assert!(param.validate().is_err());
}
#[test]
fn int_param_validate_invalid_log() {
let param = IntParam::new(0, 10).log_scale();
assert!(param.validate().is_err());
}
#[test]
fn int_param_validate_invalid_step() {
let param = IntParam::new(0, 10).step(-1);
assert!(param.validate().is_err());
}
#[test]
fn int_param_cast_param_value() {
let param = IntParam::new(1, 10);
assert_eq!(param.cast_param_value(&ParamValue::Int(5)).unwrap(), 5);
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
}
#[test]
fn categorical_param_distribution() {
let param = CategoricalParam::new(vec!["a", "b", "c"]);
assert_eq!(
param.distribution(),
Distribution::Categorical(CategoricalDistribution { n_choices: 3 })
);
}
#[test]
fn categorical_param_validate_empty() {
let param = CategoricalParam::<&str>::new(vec![]);
assert!(param.validate().is_err());
}
#[test]
fn categorical_param_cast_param_value() {
let param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]);
assert_eq!(
param.cast_param_value(&ParamValue::Categorical(1)).unwrap(),
"adam"
);
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
}
#[test]
fn bool_param_distribution() {
let param = BoolParam::new();
assert_eq!(
param.distribution(),
Distribution::Categorical(CategoricalDistribution { n_choices: 2 })
);
}
#[test]
fn bool_param_cast_param_value() {
let param = BoolParam::new();
assert!(!param.cast_param_value(&ParamValue::Categorical(0)).unwrap());
assert!(param.cast_param_value(&ParamValue::Categorical(1)).unwrap());
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
}
#[derive(Clone, Debug, PartialEq)]
enum TestEnum {
A,
B,
C,
}
impl Categorical for TestEnum {
const N_CHOICES: usize = 3;
fn from_index(index: usize) -> Self {
match index {
0 => TestEnum::A,
1 => TestEnum::B,
2 => TestEnum::C,
_ => panic!("invalid index"),
}
}
fn to_index(&self) -> usize {
match self {
TestEnum::A => 0,
TestEnum::B => 1,
TestEnum::C => 2,
}
}
}
#[test]
fn enum_param_distribution() {
let param = EnumParam::<TestEnum>::new();
assert_eq!(
param.distribution(),
Distribution::Categorical(CategoricalDistribution { n_choices: 3 })
);
}
#[test]
fn enum_param_cast_param_value() {
let param = EnumParam::<TestEnum>::new();
assert_eq!(
param.cast_param_value(&ParamValue::Categorical(0)).unwrap(),
TestEnum::A
);
assert_eq!(
param.cast_param_value(&ParamValue::Categorical(2)).unwrap(),
TestEnum::C
);
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
}
#[test]
fn float_param_suggest_via_trial() {
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));
// Cached value (same param id) - exact equality expected for cached values
let x2 = param.suggest(&mut trial).unwrap();
assert!((x - x2).abs() < f64::EPSILON);
}
#[test]
fn int_param_suggest_via_trial() {
let param = IntParam::new(1, 10);
let mut trial = Trial::new(0);
let n = param.suggest(&mut trial).unwrap();
assert!((1..=10).contains(&n));
// Cached value
let n2 = param.suggest(&mut trial).unwrap();
assert_eq!(n, n2);
}
#[test]
fn categorical_param_suggest_via_trial() {
let choices = vec!["sgd", "adam", "rmsprop"];
let param = CategoricalParam::new(choices.clone());
let mut trial = Trial::new(0);
let opt = param.suggest(&mut trial).unwrap();
assert!(choices.contains(&opt));
// Cached value
let opt2 = param.suggest(&mut trial).unwrap();
assert_eq!(opt, opt2);
}
#[test]
fn bool_param_suggest_via_trial() {
let param = BoolParam::new();
let mut trial = Trial::new(0);
let val = param.suggest(&mut trial).unwrap();
let _ = val; // just check it doesn't error
// Cached value
let val2 = param.suggest(&mut trial).unwrap();
assert_eq!(val, val2);
}
#[test]
fn enum_param_suggest_via_trial() {
let param = EnumParam::<TestEnum>::new();
let mut trial = Trial::new(0);
let val = param.suggest(&mut trial).unwrap();
assert!([TestEnum::A, TestEnum::B, TestEnum::C].contains(&val));
// Cached value
let val2 = param.suggest(&mut trial).unwrap();
assert_eq!(val, val2);
}
#[test]
fn parameter_conflict_detection() {
let param_float = FloatParam::new(0.0, 1.0);
let param_int = IntParam::new(0, 10);
let mut trial = Trial::new(0);
let _ = param_float.suggest(&mut trial).unwrap();
// Different param with different distribution but same id won't happen
// since each param gets a unique id. But different param object = different id = no conflict.
let result = param_int.suggest(&mut trial);
assert!(result.is_ok()); // Different id, no conflict
}
#[test]
fn float_param_validation_prevents_suggest() {
let param = FloatParam::new(1.0, 0.0);
let mut trial = Trial::new(0);
let result = param.suggest(&mut trial);
assert!(result.is_err());
}
#[test]
fn categorical_trait_roundtrip() {
for i in 0..TestEnum::N_CHOICES {
let val = TestEnum::from_index(i);
assert_eq!(val.to_index(), i);
}
}
#[test]
fn param_id_uniqueness() {
let id1 = ParamId::new();
let id2 = ParamId::new();
assert_ne!(id1, id2);
}
#[test]
fn param_clone_preserves_id() {
let param = FloatParam::new(0.0, 1.0);
let cloned = param.clone();
assert_eq!(param.id(), cloned.id());
}
}
+28 -87
View File
@@ -8,7 +8,6 @@ use std::collections::HashMap;
use crate::distribution::Distribution;
use crate::param::ParamValue;
use crate::parameter::{ParamId, Parameter};
/// A completed trial with its parameters, distributions, and objective value.
///
@@ -19,12 +18,10 @@ use crate::parameter::{ParamId, Parameter};
pub struct CompletedTrial<V = f64> {
/// The unique identifier for this trial.
pub id: u64,
/// 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 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 objective value returned by the objective function.
pub value: V,
}
@@ -33,95 +30,17 @@ impl<V> CompletedTrial<V> {
/// Creates a new completed trial.
pub fn new(
id: u64,
params: HashMap<ParamId, ParamValue>,
distributions: HashMap<ParamId, Distribution>,
param_labels: HashMap<ParamId, String>,
params: HashMap<String, ParamValue>,
distributions: HashMap<String, Distribution>,
value: V,
) -> Self {
Self {
id,
params,
distributions,
param_labels,
value,
}
}
/// 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")
})
}
}
/// A pending (running) trial with its parameters and distributions, but no objective value yet.
///
/// 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.
#[derive(Clone, Debug)]
pub struct PendingTrial {
/// The unique identifier for this trial.
pub id: u64,
/// 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 {
/// Creates a new pending trial.
#[must_use]
pub fn new(
id: u64,
params: HashMap<ParamId, ParamValue>,
distributions: HashMap<ParamId, Distribution>,
param_labels: HashMap<ParamId, String>,
) -> Self {
Self {
id,
params,
distributions,
param_labels,
}
}
}
/// Trait for pluggable parameter sampling strategies.
@@ -129,6 +48,28 @@ 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.
///
+41 -291
View File
@@ -3,60 +3,6 @@
//! TPE is a Bayesian optimization algorithm that models the objective function
//! using two probability distributions: one for promising (good) parameter values
//! and one for unpromising (bad) parameter values.
//!
//! # Gamma Strategies
//!
//! The gamma parameter controls what fraction of trials are considered "good".
//! This module provides several built-in strategies via the [`GammaStrategy`] trait:
//!
//! - [`FixedGamma`]: Constant gamma value (default: 0.25)
//! - [`LinearGamma`]: Linear interpolation between min and max based on trial count
//! - [`SqrtGamma`]: Gamma decreases as 1/√n (similar to Optuna)
//! - [`HyperoptGamma`]: Hyperopt-style adaptive gamma
//!
//! You can also implement your own strategy by implementing the [`GammaStrategy`] trait.
//!
//! # Examples
//!
//! Using a built-in gamma strategy:
//!
//! ```
//! use optimizer::sampler::tpe::{SqrtGamma, TpeSampler};
//!
//! let sampler = TpeSampler::builder()
//! .gamma_strategy(SqrtGamma::default())
//! .build()
//! .unwrap();
//! ```
//!
//! Implementing a custom gamma strategy:
//!
//! ```
//! use optimizer::sampler::tpe::{GammaStrategy, TpeSampler};
//!
//! #[derive(Debug, Clone)]
//! struct MyGamma {
//! base: f64,
//! }
//!
//! impl GammaStrategy for MyGamma {
//! fn gamma(&self, n_trials: usize) -> f64 {
//! (self.base + 0.01 * n_trials as f64).min(0.5)
//! }
//!
//! fn clone_box(&self) -> Box<dyn GammaStrategy> {
//! Box::new(self.clone())
//! }
//! }
//!
//! let sampler = TpeSampler::builder()
//! .gamma_strategy(MyGamma { base: 0.1 })
//! .build()
//! .unwrap();
//! ```
use core::fmt::Debug;
use std::sync::Arc;
use parking_lot::Mutex;
use rand::rngs::StdRng;
@@ -66,17 +12,8 @@ use crate::distribution::Distribution;
use crate::error::{Error, Result};
use crate::kde::KernelDensityEstimator;
use crate::param::ParamValue;
use crate::sampler::tpe::gamma::{FixedGamma, GammaStrategy};
use crate::sampler::{CompletedTrial, Sampler};
// ============================================================================
// Gamma Strategy Trait and Implementations
// ============================================================================
// ============================================================================
// TPE Sampler
// ============================================================================
/// A Tree-Parzen Estimator (TPE) sampler for Bayesian optimization.
///
/// TPE works by splitting completed trials into two groups based on their
@@ -88,42 +25,26 @@ use crate::sampler::{CompletedTrial, Sampler};
/// During the startup phase (when fewer than `n_startup_trials` are completed),
/// TPE falls back to random sampling to gather initial data.
///
/// # Gamma Strategies
///
/// The gamma quantile can be configured using different strategies via the
/// [`GammaStrategy`] trait.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::TpeSampler;
///
/// // Create with default settings (FixedGamma at 0.25)
/// // Create with default settings
/// let sampler = TpeSampler::new();
///
/// // Create with custom settings using the builder
/// let sampler = TpeSampler::builder()
/// .gamma(0.15) // Shorthand for Fixednew(0.15)
/// .gamma(0.15)
/// .n_startup_trials(20)
/// .n_ei_candidates(32)
/// .seed(42)
/// .build()
/// .unwrap();
/// ```
///
/// Using a different gamma strategy:
///
/// ```
/// use optimizer::sampler::tpe::{SqrtGamma, TpeSampler};
///
/// let sampler = TpeSampler::builder()
/// .gamma_strategy(SqrtGamma::default())
/// .build()
/// .unwrap();
/// ```
pub struct TpeSampler {
/// Strategy for computing the gamma quantile.
gamma_strategy: Arc<dyn GammaStrategy>,
/// Fraction of trials to consider as "good" (gamma quantile).
gamma: f64,
/// Number of trials before TPE kicks in (uses random sampling before this).
n_startup_trials: usize,
/// Number of candidate samples to evaluate when selecting the next point.
@@ -138,14 +59,14 @@ impl TpeSampler {
/// Creates a new TPE sampler with default settings.
///
/// Default settings:
/// - gamma strategy: [`FixedGamma`] with gamma = 0.25
/// - gamma: 0.25 (top 25% of trials are considered "good")
/// - `n_startup_trials`: 10 (random sampling for first 10 trials)
/// - `n_ei_candidates`: 24 (evaluate 24 candidates per sample)
/// - `kde_bandwidth`: None (uses Scott's rule for automatic bandwidth)
#[must_use]
pub fn new() -> Self {
Self {
gamma_strategy: Arc::new(FixedGamma::default()),
gamma: 0.25,
n_startup_trials: 10,
n_ei_candidates: 24,
kde_bandwidth: None,
@@ -175,10 +96,6 @@ impl TpeSampler {
/// Creates a new TPE sampler with custom configuration.
///
/// This method uses a fixed gamma value. For more advanced gamma strategies,
/// use [`TpeSampler::with_strategy`] or the builder pattern with
/// [`TpeSamplerBuilder::gamma_strategy`].
///
/// # Arguments
///
/// * `gamma` - Fraction of trials to consider "good" (0.0 to 1.0).
@@ -198,51 +115,9 @@ impl TpeSampler {
kde_bandwidth: Option<f64>,
seed: Option<u64>,
) -> Result<Self> {
let gamma_strategy = FixedGamma::new(gamma)?;
Self::with_strategy(
gamma_strategy,
n_startup_trials,
n_ei_candidates,
kde_bandwidth,
seed,
)
}
/// Creates a new TPE sampler with a custom gamma strategy.
///
/// # Arguments
///
/// * `gamma_strategy` - The strategy for computing the gamma quantile.
/// * `n_startup_trials` - Number of random trials before TPE sampling.
/// * `n_ei_candidates` - Number of candidates to evaluate per sample.
/// * `kde_bandwidth` - Optional fixed bandwidth for KDE. If None, uses Scott's rule.
/// * `seed` - Optional seed for reproducibility.
///
/// # Errors
///
/// Returns `Error::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::{SqrtGamma, TpeSampler};
///
/// let sampler = TpeSampler::with_strategy(
/// SqrtGamma::default(),
/// 10, // n_startup_trials
/// 24, // n_ei_candidates
/// None, // kde_bandwidth
/// Some(42), // seed
/// )
/// .unwrap();
/// ```
pub fn with_strategy<G: GammaStrategy + 'static>(
gamma_strategy: G,
n_startup_trials: usize,
n_ei_candidates: usize,
kde_bandwidth: Option<f64>,
seed: Option<u64>,
) -> Result<Self> {
if gamma <= 0.0 || gamma >= 1.0 {
return Err(Error::InvalidGamma(gamma));
}
if let Some(bw) = kde_bandwidth
&& bw <= 0.0
{
@@ -255,7 +130,7 @@ impl TpeSampler {
};
Ok(Self {
gamma_strategy: Arc::new(gamma_strategy),
gamma,
n_startup_trials,
n_ei_candidates,
kde_bandwidth,
@@ -263,16 +138,8 @@ impl TpeSampler {
})
}
/// Returns the gamma strategy used by this sampler.
#[must_use]
pub fn gamma_strategy(&self) -> &dyn GammaStrategy {
self.gamma_strategy.as_ref()
}
/// Splits trials into good and bad groups based on the gamma quantile.
///
/// The gamma value is computed dynamically using the configured [`GammaStrategy`].
///
/// Returns (`good_trials`, `bad_trials`) where `good_trials` contains trials
/// with values below the gamma quantile (for minimization).
#[allow(
@@ -280,7 +147,6 @@ impl TpeSampler {
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
#[must_use]
fn split_trials<'a>(
&self,
history: &'a [CompletedTrial],
@@ -298,15 +164,9 @@ impl TpeSampler {
.unwrap_or(core::cmp::Ordering::Equal)
});
// Compute gamma using the strategy and clamp to valid range
let gamma = self
.gamma_strategy
.gamma(history.len())
.clamp(f64::EPSILON, 1.0 - f64::EPSILON);
// Calculate the split point (gamma quantile)
// Ensure at least 1 trial in each group if possible
let n_good = ((history.len() as f64 * gamma).ceil() as usize)
let n_good = ((history.len() as f64 * self.gamma).ceil() as usize)
.max(1)
.min(history.len() - 1);
@@ -561,8 +421,6 @@ impl Default for TpeSampler {
///
/// # Examples
///
/// Using a fixed gamma value:
///
/// ```
/// use optimizer::sampler::tpe::TpeSamplerBuilder;
///
@@ -574,23 +432,9 @@ impl Default for TpeSampler {
/// .build()
/// .unwrap();
/// ```
///
/// Using a custom gamma strategy:
///
/// ```
/// use optimizer::sampler::tpe::{SqrtGamma, TpeSamplerBuilder};
///
/// let sampler = TpeSamplerBuilder::new()
/// .gamma_strategy(SqrtGamma::default())
/// .n_startup_trials(20)
/// .build()
/// .unwrap();
/// ```
#[derive(Debug, Clone)]
pub struct TpeSamplerBuilder {
gamma_strategy: Box<dyn GammaStrategy>,
/// Raw gamma value for deferred validation (Some if `gamma()` was called)
raw_gamma: Option<f64>,
gamma: f64,
n_startup_trials: usize,
n_ei_candidates: usize,
kde_bandwidth: Option<f64>,
@@ -601,7 +445,7 @@ impl TpeSamplerBuilder {
/// Creates a new builder with default settings.
///
/// Default settings:
/// - gamma strategy: [`FixedGamma`] with gamma = 0.25
/// - gamma: 0.25 (top 25% of trials are considered "good")
/// - `n_startup_trials`: 10 (random sampling for first 10 trials)
/// - `n_ei_candidates`: 24 (evaluate 24 candidates per sample)
/// - `kde_bandwidth`: None (uses Scott's rule for automatic bandwidth)
@@ -609,8 +453,7 @@ impl TpeSamplerBuilder {
#[must_use]
pub fn new() -> Self {
Self {
gamma_strategy: Box::new(FixedGamma::default()),
raw_gamma: None,
gamma: 0.25,
n_startup_trials: 10,
n_ei_candidates: 24,
kde_bandwidth: None,
@@ -618,10 +461,7 @@ impl TpeSamplerBuilder {
}
}
/// Sets a fixed gamma value for splitting trials into good/bad groups.
///
/// This is a convenience method that creates a [`FixedGamma`] strategy.
/// For more advanced gamma strategies, use [`gamma_strategy`](Self::gamma_strategy).
/// Sets the gamma quantile for splitting trials into good/bad groups.
///
/// A gamma of 0.25 means the top 25% of trials (by objective value) are
/// considered "good" and used to build the l(x) distribution.
@@ -647,68 +487,7 @@ impl TpeSamplerBuilder {
/// `build()` will return `Err(Error::InvalidGamma)`.
#[must_use]
pub fn gamma(mut self, gamma: f64) -> Self {
// We defer validation to build() time for consistency with the existing API
// Store the raw value for validation later
self.raw_gamma = Some(gamma);
self
}
/// Sets a custom gamma strategy for splitting trials into good/bad groups.
///
/// The gamma strategy determines what fraction of trials are considered
/// "good" based on the number of completed trials. This allows the gamma
/// value to adapt dynamically during optimization.
///
/// # Arguments
///
/// * `strategy` - A type implementing [`GammaStrategy`].
///
/// # Examples
///
/// Using built-in strategies:
///
/// ```
/// use optimizer::sampler::tpe::{LinearGamma, SqrtGamma, TpeSamplerBuilder};
///
/// // Square root strategy (Optuna-style)
/// let sampler = TpeSamplerBuilder::new()
/// .gamma_strategy(SqrtGamma::default())
/// .build()
/// .unwrap();
///
/// // Linear interpolation strategy
/// let sampler = TpeSamplerBuilder::new()
/// .gamma_strategy(LinearGamma::new(0.1, 0.3, 50).unwrap())
/// .build()
/// .unwrap();
/// ```
///
/// Using a custom strategy:
///
/// ```
/// use optimizer::sampler::tpe::{GammaStrategy, TpeSamplerBuilder};
///
/// #[derive(Debug, Clone)]
/// struct MyGamma;
///
/// impl GammaStrategy for MyGamma {
/// fn gamma(&self, n_trials: usize) -> f64 {
/// 0.25 // Always return 0.25
/// }
/// fn clone_box(&self) -> Box<dyn GammaStrategy> {
/// Box::new(self.clone())
/// }
/// }
///
/// let sampler = TpeSamplerBuilder::new()
/// .gamma_strategy(MyGamma)
/// .build()
/// .unwrap();
/// ```
#[must_use]
pub fn gamma_strategy<G: GammaStrategy + 'static>(mut self, strategy: G) -> Self {
self.gamma_strategy = Box::new(strategy);
self.raw_gamma = None; // Clear any raw gamma set by gamma()
self.gamma = gamma;
self
}
@@ -823,7 +602,7 @@ impl TpeSamplerBuilder {
///
/// # Errors
///
/// Returns `Error::InvalidGamma` if a fixed gamma value was set and is not in (0.0, 1.0).
/// Returns `Error::InvalidGamma` if gamma is not in (0.0, 1.0).
/// Returns `Error::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
///
/// # Examples
@@ -840,33 +619,13 @@ impl TpeSamplerBuilder {
/// .unwrap();
/// ```
pub fn build(self) -> Result<TpeSampler> {
// Determine the gamma strategy to use
let gamma_strategy: Arc<dyn GammaStrategy> = if let Some(raw) = self.raw_gamma {
// Validate and create FixedGamma from raw value
Arc::new(FixedGamma::new(raw)?)
} else {
Arc::from(self.gamma_strategy)
};
// Validate bandwidth
if let Some(bw) = self.kde_bandwidth
&& bw <= 0.0
{
return Err(Error::InvalidBandwidth(bw));
}
let rng = match self.seed {
Some(s) => StdRng::seed_from_u64(s),
None => StdRng::from_os_rng(),
};
Ok(TpeSampler {
gamma_strategy,
n_startup_trials: self.n_startup_trials,
n_ei_candidates: self.n_ei_candidates,
kde_bandwidth: self.kde_bandwidth,
rng: Mutex::new(rng),
})
TpeSampler::with_config(
self.gamma,
self.n_startup_trials,
self.n_ei_candidates,
self.kde_bandwidth,
self.seed,
)
}
}
@@ -1024,27 +783,25 @@ mod tests {
use super::*;
use crate::distribution::{CategoricalDistribution, FloatDistribution, IntDistribution};
use crate::parameter::ParamId;
fn create_trial(
id: u64,
value: f64,
params: Vec<(ParamId, ParamValue, Distribution)>,
params: Vec<(&str, ParamValue, Distribution)>,
) -> 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);
for (name, pv, dist) in params {
param_map.insert(name.to_string(), pv);
dist_map.insert(name.to_string(), dist);
}
CompletedTrial::new(id, param_map, dist_map, HashMap::new(), value)
CompletedTrial::new(id, param_map, dist_map, value)
}
#[test]
fn test_tpe_sampler_new() {
let sampler = TpeSampler::new();
// Default uses FixedGamma with 0.25
assert!((sampler.gamma_strategy().gamma(0) - 0.25).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.25).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 10);
assert_eq!(sampler.n_ei_candidates, 24);
}
@@ -1052,8 +809,7 @@ mod tests {
#[test]
fn test_tpe_sampler_with_config() {
let sampler = TpeSampler::with_config(0.15, 20, 32, None, Some(42)).unwrap();
// with_config uses FixedGamma
assert!((sampler.gamma_strategy().gamma(0) - 0.15).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.15).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 20);
assert_eq!(sampler.n_ei_candidates, 32);
}
@@ -1105,13 +861,12 @@ 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_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
@@ -1140,7 +895,6 @@ 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;
@@ -1149,7 +903,7 @@ mod tests {
create_trial(
i as u64,
value,
vec![(x_id, ParamValue::Float(x), dist.clone())],
vec![("x", ParamValue::Float(x), dist.clone())],
)
})
.collect();
@@ -1178,7 +932,6 @@ 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;
@@ -1188,7 +941,7 @@ mod tests {
i as u64,
value,
vec![(
cat_id,
"cat",
ParamValue::Categorical(category as usize),
dist.clone(),
)],
@@ -1224,7 +977,6 @@ 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
@@ -1232,7 +984,7 @@ mod tests {
create_trial(
i as u64,
value,
vec![(x_id, ParamValue::Int(x), dist.clone())],
vec![("x", ParamValue::Int(x), dist.clone())],
)
})
.collect();
@@ -1257,13 +1009,12 @@ 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_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
@@ -1282,7 +1033,7 @@ mod tests {
fn test_tpe_sampler_builder_default() {
let builder = TpeSamplerBuilder::new();
let sampler = builder.build().unwrap();
assert!((sampler.gamma_strategy().gamma(0) - 0.25).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.25).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 10);
assert_eq!(sampler.n_ei_candidates, 24);
}
@@ -1296,7 +1047,7 @@ mod tests {
.seed(42)
.build()
.unwrap();
assert!((sampler.gamma_strategy().gamma(0) - 0.15).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.15).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 20);
assert_eq!(sampler.n_ei_candidates, 32);
}
@@ -1309,7 +1060,7 @@ mod tests {
.n_ei_candidates(48)
.build()
.unwrap();
assert!((sampler.gamma_strategy().gamma(0) - 0.10).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.10).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 15);
assert_eq!(sampler.n_ei_candidates, 48);
}
@@ -1318,7 +1069,7 @@ mod tests {
fn test_tpe_sampler_builder_partial() {
// Test setting only some options
let sampler = TpeSamplerBuilder::new().gamma(0.20).build().unwrap();
assert!((sampler.gamma_strategy().gamma(0) - 0.20).abs() < f64::EPSILON);
assert!((sampler.gamma - 0.20).abs() < f64::EPSILON);
assert_eq!(sampler.n_startup_trials, 10); // default
assert_eq!(sampler.n_ei_candidates, 24); // default
}
@@ -1338,13 +1089,12 @@ 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_id, ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
vec![("x", ParamValue::Float(f64::from(i) / 20.0), dist.clone())],
)
})
.collect();
-670
View File
@@ -1,670 +0,0 @@
use core::fmt::Debug;
use crate::Error;
/// A strategy for computing the gamma quantile in TPE.
///
/// The gamma value determines what fraction of trials are considered "good"
/// when splitting the trial history. Different strategies can adapt this
/// fraction based on the number of completed trials.
///
/// # Implementation Notes
///
/// - The returned gamma must be in the range (0.0, 1.0)
/// - Implementations should be deterministic for reproducibility
/// - The `clone_box` method enables trait object cloning
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::GammaStrategy;
///
/// #[derive(Debug, Clone)]
/// struct ConstantGamma(f64);
///
/// impl GammaStrategy for ConstantGamma {
/// fn gamma(&self, _n_trials: usize) -> f64 {
/// self.0
/// }
///
/// fn clone_box(&self) -> Box<dyn GammaStrategy> {
/// Box::new(self.clone())
/// }
/// }
/// ```
pub trait GammaStrategy: Send + Sync + Debug {
/// Computes the gamma quantile based on the number of completed trials.
///
/// # Arguments
///
/// * `n_trials` - The number of completed trials in the history.
///
/// # Returns
///
/// A gamma value in the range (0.0, 1.0). Values outside this range
/// will be clamped by the sampler.
fn gamma(&self, n_trials: usize) -> f64;
/// Creates a boxed clone of this strategy.
///
/// This method enables cloning of trait objects, which is necessary
/// for the builder pattern and sampler configuration.
fn clone_box(&self) -> Box<dyn GammaStrategy>;
}
impl Clone for Box<dyn GammaStrategy> {
fn clone(&self) -> Self {
self.clone_box()
}
}
/// A fixed gamma strategy that returns a constant value.
///
/// This is the simplest strategy and the default behavior of TPE.
/// The gamma value remains constant regardless of the number of trials.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::{FixedGamma, TpeSampler};
///
/// // Use 15% of trials as "good"
/// let sampler = TpeSampler::builder()
/// .gamma_strategy(FixedGamma::new(0.15).unwrap())
/// .build()
/// .unwrap();
/// ```
#[derive(Debug, Clone, Copy)]
pub struct FixedGamma {
gamma: f64,
}
impl FixedGamma {
/// Creates a new fixed gamma strategy.
///
/// # Arguments
///
/// * `gamma` - The constant gamma value to use.
///
/// # Errors
///
/// Returns `Error::InvalidGamma` if gamma is not in (0.0, 1.0).
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::FixedGamma;
///
/// let strategy = FixedGamma::new(0.25).unwrap();
/// assert!((strategy.value() - 0.25).abs() < f64::EPSILON);
/// ```
pub fn new(gamma: f64) -> crate::Result<Self> {
if gamma <= 0.0 || gamma >= 1.0 {
return Err(Error::InvalidGamma(gamma));
}
Ok(Self { gamma })
}
/// Returns the fixed gamma value.
#[must_use]
pub fn value(&self) -> f64 {
self.gamma
}
}
impl Default for FixedGamma {
/// Creates a fixed gamma strategy with the default value of 0.25.
fn default() -> Self {
Self { gamma: 0.25 }
}
}
impl GammaStrategy for FixedGamma {
fn gamma(&self, _n_trials: usize) -> f64 {
self.gamma
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(*self)
}
}
/// A linear gamma strategy that interpolates between min and max values.
///
/// The gamma value increases linearly from `gamma_min` to `gamma_max` as the
/// number of trials grows from 0 to `n_trials_max`. Beyond `n_trials_max`,
/// gamma remains at `gamma_max`.
///
/// This strategy is useful when you want to be more explorative early on
/// (smaller gamma = fewer "good" trials) and more exploitative later
/// (larger gamma = more "good" trials).
///
/// # Formula
///
/// ```text
/// gamma = gamma_min + (gamma_max - gamma_min) * min(n_trials / n_trials_max, 1.0)
/// ```
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::{GammaStrategy, LinearGamma, TpeSampler};
///
/// let strategy = LinearGamma::new(0.1, 0.4, 100).unwrap();
///
/// // At 0 trials: gamma = 0.1
/// assert!((strategy.gamma(0) - 0.1).abs() < f64::EPSILON);
///
/// // At 50 trials: gamma = 0.25 (midpoint)
/// assert!((strategy.gamma(50) - 0.25).abs() < f64::EPSILON);
///
/// // At 100+ trials: gamma = 0.4
/// assert!((strategy.gamma(100) - 0.4).abs() < f64::EPSILON);
/// assert!((strategy.gamma(200) - 0.4).abs() < f64::EPSILON);
/// ```
#[derive(Debug, Clone, Copy)]
pub struct LinearGamma {
gamma_min: f64,
gamma_max: f64,
n_trials_max: usize,
}
impl LinearGamma {
/// Creates a new linear gamma strategy.
///
/// # Arguments
///
/// * `gamma_min` - The minimum gamma value (at 0 trials).
/// * `gamma_max` - The maximum gamma value (at `n_trials_max` trials).
/// * `n_trials_max` - The number of trials at which gamma reaches its maximum.
///
/// # Errors
///
/// Returns `Error::InvalidGamma` if:
/// - `gamma_min` is not in (0.0, 1.0)
/// - `gamma_max` is not in (0.0, 1.0)
/// - `gamma_min > gamma_max`
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::LinearGamma;
///
/// // Gamma goes from 0.1 to 0.3 over 50 trials
/// let strategy = LinearGamma::new(0.1, 0.3, 50).unwrap();
/// ```
pub fn new(gamma_min: f64, gamma_max: f64, n_trials_max: usize) -> crate::Result<Self> {
if gamma_min <= 0.0 || gamma_min >= 1.0 {
return Err(Error::InvalidGamma(gamma_min));
}
if gamma_max <= 0.0 || gamma_max >= 1.0 {
return Err(Error::InvalidGamma(gamma_max));
}
if gamma_min > gamma_max {
return Err(Error::InvalidGamma(gamma_min));
}
Ok(Self {
gamma_min,
gamma_max,
n_trials_max,
})
}
/// Returns the minimum gamma value.
#[must_use]
pub fn gamma_min(&self) -> f64 {
self.gamma_min
}
/// Returns the maximum gamma value.
#[must_use]
pub fn gamma_max(&self) -> f64 {
self.gamma_max
}
/// Returns the number of trials at which gamma reaches its maximum.
#[must_use]
pub fn n_trials_max(&self) -> usize {
self.n_trials_max
}
}
impl Default for LinearGamma {
/// Creates a linear gamma strategy with default values:
/// - `gamma_min`: 0.10
/// - `gamma_max`: 0.25
/// - `n_trials_max`: 100
fn default() -> Self {
Self {
gamma_min: 0.10,
gamma_max: 0.25,
n_trials_max: 100,
}
}
}
impl GammaStrategy for LinearGamma {
#[allow(clippy::cast_precision_loss)]
fn gamma(&self, n_trials: usize) -> f64 {
if self.n_trials_max == 0 {
return self.gamma_max;
}
let t = (n_trials as f64 / self.n_trials_max as f64).min(1.0);
self.gamma_min + (self.gamma_max - self.gamma_min) * t
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(*self)
}
}
/// A square root gamma strategy inspired by Optuna's default behavior.
///
/// The gamma value is computed based on the inverse square root of the number
/// of trials, providing a balance between exploration and exploitation that
/// naturally adapts as more data becomes available.
///
/// # Formula
///
/// ```text
/// n_good = max(1, floor(gamma_factor / sqrt(n_trials)))
/// gamma = min(gamma_max, n_good / n_trials)
/// ```
///
/// When `n_trials` is 0, returns `gamma_max`.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::{GammaStrategy, SqrtGamma, TpeSampler};
///
/// let strategy = SqrtGamma::default();
///
/// // Gamma decreases as trials increase
/// let g10 = strategy.gamma(10);
/// let g100 = strategy.gamma(100);
/// assert!(g10 > g100, "Gamma should decrease with more trials");
/// ```
#[derive(Debug, Clone, Copy)]
pub struct SqrtGamma {
gamma_factor: f64,
gamma_max: f64,
}
impl SqrtGamma {
/// Creates a new square root gamma strategy.
///
/// # Arguments
///
/// * `gamma_factor` - The factor controlling how quickly gamma decreases.
/// Higher values mean more "good" trials at any given point.
/// * `gamma_max` - The maximum gamma value (used when `n_trials` is small).
///
/// # Errors
///
/// Returns `Error::InvalidGamma` if:
/// - `gamma_factor` is not positive
/// - `gamma_max` is not in (0.0, 1.0)
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::SqrtGamma;
///
/// let strategy = SqrtGamma::new(1.0, 0.25).unwrap();
/// ```
pub fn new(gamma_factor: f64, gamma_max: f64) -> crate::Result<Self> {
if gamma_factor <= 0.0 {
return Err(Error::InvalidGamma(gamma_factor));
}
if gamma_max <= 0.0 || gamma_max >= 1.0 {
return Err(Error::InvalidGamma(gamma_max));
}
Ok(Self {
gamma_factor,
gamma_max,
})
}
/// Returns the gamma factor.
#[must_use]
pub fn gamma_factor(&self) -> f64 {
self.gamma_factor
}
/// Returns the maximum gamma value.
#[must_use]
pub fn gamma_max(&self) -> f64 {
self.gamma_max
}
}
impl Default for SqrtGamma {
/// Creates a square root gamma strategy with default values:
/// - `gamma_factor`: 1.0
/// - `gamma_max`: 0.25
fn default() -> Self {
Self {
gamma_factor: 1.0,
gamma_max: 0.25,
}
}
}
impl GammaStrategy for SqrtGamma {
#[allow(clippy::cast_precision_loss)]
fn gamma(&self, n_trials: usize) -> f64 {
if n_trials == 0 {
return self.gamma_max;
}
let n_good = (self.gamma_factor / (n_trials as f64).sqrt()).max(1.0);
(n_good / n_trials as f64).min(self.gamma_max)
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(*self)
}
}
/// A Hyperopt-style gamma strategy.
///
/// This strategy computes gamma as `min(gamma_max, (gamma_base + 1) / n_trials)`,
/// which is inspired by the original Hyperopt TPE implementation.
///
/// # Formula
///
/// ```text
/// gamma = min(gamma_max, (gamma_base + 1) / n_trials)
/// ```
///
/// When `n_trials` is 0, returns `gamma_max`.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::{GammaStrategy, HyperoptGamma};
///
/// // With gamma_base=24 and gamma_max=0.5:
/// // - At n=25: gamma = min(0.5, 25/25) = 0.5 (capped)
/// // - At n=100: gamma = min(0.5, 25/100) = 0.25
/// let strategy = HyperoptGamma::new(24.0, 0.5).unwrap();
///
/// // Early trials have higher gamma
/// let g50 = strategy.gamma(50);
/// let g200 = strategy.gamma(200);
/// assert!(g50 > g200, "Gamma should decrease with more trials");
/// ```
#[derive(Debug, Clone, Copy)]
pub struct HyperoptGamma {
gamma_base: f64,
gamma_max: f64,
}
impl HyperoptGamma {
/// Creates a new Hyperopt-style gamma strategy.
///
/// # Arguments
///
/// * `gamma_base` - The base value added to 1 in the numerator.
/// * `gamma_max` - The maximum gamma value.
///
/// # Errors
///
/// Returns `Error::InvalidGamma` if:
/// - `gamma_base` is negative
/// - `gamma_max` is not in (0.0, 1.0)
///
/// # Examples
///
/// ```
/// use optimizer::sampler::tpe::HyperoptGamma;
///
/// let strategy = HyperoptGamma::new(24.0, 0.25).unwrap();
/// ```
pub fn new(gamma_base: f64, gamma_max: f64) -> crate::Result<Self> {
if gamma_base < 0.0 {
return Err(Error::InvalidGamma(gamma_base));
}
if gamma_max <= 0.0 || gamma_max >= 1.0 {
return Err(Error::InvalidGamma(gamma_max));
}
Ok(Self {
gamma_base,
gamma_max,
})
}
/// Returns the gamma base value.
#[must_use]
pub fn gamma_base(&self) -> f64 {
self.gamma_base
}
/// Returns the maximum gamma value.
#[must_use]
pub fn gamma_max(&self) -> f64 {
self.gamma_max
}
}
impl Default for HyperoptGamma {
/// Creates a Hyperopt-style gamma strategy with default values:
/// - `gamma_base`: 24.0
/// - `gamma_max`: 0.25
fn default() -> Self {
Self {
gamma_base: 24.0,
gamma_max: 0.25,
}
}
}
impl GammaStrategy for HyperoptGamma {
#[allow(clippy::cast_precision_loss)]
fn gamma(&self, n_trials: usize) -> f64 {
if n_trials == 0 {
return self.gamma_max;
}
((self.gamma_base + 1.0) / n_trials as f64).min(self.gamma_max)
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(*self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sampler::tpe::TpeSampler;
#[test]
fn test_fixed_gamma_default() {
let strategy = FixedGamma::default();
assert!((strategy.gamma(0) - 0.25).abs() < f64::EPSILON);
assert!((strategy.gamma(100) - 0.25).abs() < f64::EPSILON);
assert!((strategy.value() - 0.25).abs() < f64::EPSILON);
}
#[test]
fn test_fixed_gamma_custom() {
let strategy = FixedGamma::new(0.15).unwrap();
assert!((strategy.gamma(0) - 0.15).abs() < f64::EPSILON);
assert!((strategy.gamma(50) - 0.15).abs() < f64::EPSILON);
assert!((strategy.gamma(1000) - 0.15).abs() < f64::EPSILON);
}
#[test]
fn test_fixed_gamma_invalid() {
assert!(FixedGamma::new(0.0).is_err());
assert!(FixedGamma::new(1.0).is_err());
assert!(FixedGamma::new(-0.1).is_err());
assert!(FixedGamma::new(1.5).is_err());
}
#[test]
fn test_linear_gamma_default() {
let strategy = LinearGamma::default();
assert!((strategy.gamma(0) - 0.10).abs() < f64::EPSILON);
assert!((strategy.gamma(50) - 0.175).abs() < f64::EPSILON); // midpoint
assert!((strategy.gamma(100) - 0.25).abs() < f64::EPSILON);
assert!((strategy.gamma(200) - 0.25).abs() < f64::EPSILON); // capped
}
#[test]
fn test_linear_gamma_custom() {
let strategy = LinearGamma::new(0.1, 0.4, 100).unwrap();
assert!((strategy.gamma(0) - 0.1).abs() < f64::EPSILON);
assert!((strategy.gamma(50) - 0.25).abs() < f64::EPSILON);
assert!((strategy.gamma(100) - 0.4).abs() < f64::EPSILON);
assert!((strategy.gamma(200) - 0.4).abs() < f64::EPSILON);
}
#[test]
fn test_linear_gamma_invalid() {
assert!(LinearGamma::new(0.0, 0.5, 100).is_err());
assert!(LinearGamma::new(0.1, 1.0, 100).is_err());
assert!(LinearGamma::new(0.5, 0.2, 100).is_err()); // min > max
}
#[test]
fn test_sqrt_gamma_default() {
let strategy = SqrtGamma::default();
// At n=0, returns gamma_max
assert!((strategy.gamma(0) - 0.25).abs() < f64::EPSILON);
// gamma decreases with more trials
let g10 = strategy.gamma(10);
let g100 = strategy.gamma(100);
assert!(g10 > g100);
}
#[test]
fn test_sqrt_gamma_custom() {
let strategy = SqrtGamma::new(2.0, 0.5).unwrap();
assert!((strategy.gamma(0) - 0.5).abs() < f64::EPSILON);
// At n=4: n_good = max(1, 2/2) = 1, gamma = 1/4 = 0.25
let g4 = strategy.gamma(4);
assert!((g4 - 0.25).abs() < f64::EPSILON);
}
#[test]
fn test_sqrt_gamma_invalid() {
assert!(SqrtGamma::new(0.0, 0.25).is_err()); // factor must be positive
assert!(SqrtGamma::new(-1.0, 0.25).is_err());
assert!(SqrtGamma::new(1.0, 0.0).is_err());
assert!(SqrtGamma::new(1.0, 1.0).is_err());
}
#[test]
fn test_hyperopt_gamma_default() {
let strategy = HyperoptGamma::default();
// At n=0, returns gamma_max
assert!((strategy.gamma(0) - 0.25).abs() < f64::EPSILON);
// At n=100: (24+1)/100 = 0.25, so capped to 0.25
assert!((strategy.gamma(100) - 0.25).abs() < f64::EPSILON);
// At n=200: (24+1)/200 = 0.125
assert!((strategy.gamma(200) - 0.125).abs() < f64::EPSILON);
}
#[test]
fn test_hyperopt_gamma_custom() {
let strategy = HyperoptGamma::new(9.0, 0.5).unwrap();
// At n=20: (9+1)/20 = 0.5, capped to 0.5
assert!((strategy.gamma(20) - 0.5).abs() < f64::EPSILON);
// At n=100: (9+1)/100 = 0.1
assert!((strategy.gamma(100) - 0.1).abs() < f64::EPSILON);
}
#[test]
fn test_hyperopt_gamma_invalid() {
assert!(HyperoptGamma::new(-1.0, 0.25).is_err());
assert!(HyperoptGamma::new(24.0, 0.0).is_err());
assert!(HyperoptGamma::new(24.0, 1.0).is_err());
}
#[test]
fn test_gamma_strategy_clone_box() {
let fixed: Box<dyn GammaStrategy> = Box::new(FixedGamma::new(0.3).unwrap());
let cloned = fixed.clone();
assert!((cloned.gamma(0) - 0.3).abs() < f64::EPSILON);
let linear: Box<dyn GammaStrategy> = Box::new(LinearGamma::default());
let cloned = linear.clone();
assert!((cloned.gamma(0) - 0.10).abs() < f64::EPSILON);
}
#[test]
fn test_tpe_with_linear_gamma_strategy() {
let sampler = TpeSampler::builder()
.gamma_strategy(LinearGamma::new(0.1, 0.3, 50).unwrap())
.n_startup_trials(5)
.seed(42)
.build()
.unwrap();
// Verify the strategy is applied
let g = sampler.gamma_strategy().gamma(25);
assert!((g - 0.2).abs() < f64::EPSILON); // midpoint of 0.1 to 0.3
}
#[test]
fn test_gamma_overrides_gamma_strategy() {
// When gamma() is called after gamma_strategy(), it should take precedence
let sampler = TpeSampler::builder()
.gamma_strategy(SqrtGamma::default())
.gamma(0.15) // This should override
.build()
.unwrap();
// Should use fixed gamma of 0.15
assert!((sampler.gamma_strategy().gamma(0) - 0.15).abs() < f64::EPSILON);
assert!((sampler.gamma_strategy().gamma(100) - 0.15).abs() < f64::EPSILON);
}
#[test]
fn test_gamma_strategy_overrides_gamma() {
// When gamma_strategy() is called after gamma(), it should take precedence
let sampler = TpeSampler::builder()
.gamma(0.15)
.gamma_strategy(SqrtGamma::default()) // This should override
.build()
.unwrap();
// Should use SqrtGamma - gamma decreases with trials
let g10 = sampler.gamma_strategy().gamma(10);
let g100 = sampler.gamma_strategy().gamma(100);
assert!(g10 > g100, "SqrtGamma should decrease with more trials");
}
#[test]
fn test_custom_gamma_strategy() {
#[derive(Debug, Clone)]
struct DoubleGamma;
impl GammaStrategy for DoubleGamma {
fn gamma(&self, n_trials: usize) -> f64 {
// Double the trial count-based calculation, capped at 0.5
#[allow(clippy::cast_precision_loss)]
(0.01 * n_trials as f64).min(0.5)
}
fn clone_box(&self) -> Box<dyn GammaStrategy> {
Box::new(self.clone())
}
}
let sampler = TpeSampler::builder()
.gamma_strategy(DoubleGamma)
.build()
.unwrap();
assert!((sampler.gamma_strategy().gamma(10) - 0.1).abs() < f64::EPSILON);
assert!((sampler.gamma_strategy().gamma(50) - 0.5).abs() < f64::EPSILON);
assert!((sampler.gamma_strategy().gamma(100) - 0.5).abs() < f64::EPSILON);
}
}
-16
View File
@@ -1,16 +0,0 @@
//! Tree-Parzen Estimator (TPE) sampler implementation and utilities.
//!
//! This module provides TPE-based sampling for Bayesian optimization,
//! 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
File diff suppressed because it is too large Load Diff
+366 -159
View File
@@ -1,6 +1,5 @@
//! Study implementation for managing optimization trials.
use core::any::Any;
#[cfg(feature = "async")]
use core::future::Future;
use core::ops::ControlFlow;
@@ -44,10 +43,6 @@ where
completed_trials: Arc<RwLock<Vec<CompletedTrial<V>>>>,
/// Counter for generating unique trial IDs.
next_trial_id: AtomicU64,
/// Optional factory for creating sampler-aware trials.
/// Set automatically for `Study<f64>` so that `create_trial()` and all
/// optimization methods use the sampler without requiring `_with_sampler` suffixes.
trial_factory: Option<Arc<dyn Fn(u64) -> Trial + Send + Sync>>,
}
impl<V> Study<V>
@@ -71,10 +66,7 @@ where
/// assert_eq!(study.direction(), Direction::Minimize);
/// ```
#[must_use]
pub fn new(direction: Direction) -> Self
where
V: 'static,
{
pub fn new(direction: Direction) -> Self {
Self::with_sampler(direction, RandomSampler::new())
}
@@ -95,49 +87,15 @@ where
/// let study: Study<f64> = Study::with_sampler(Direction::Maximize, sampler);
/// assert_eq!(study.direction(), Direction::Maximize);
/// ```
pub fn with_sampler(direction: Direction, sampler: impl Sampler + 'static) -> Self
where
V: 'static,
{
let sampler: Arc<dyn Sampler> = Arc::new(sampler);
let completed_trials = Arc::new(RwLock::new(Vec::new()));
// For Study<f64>, set up a trial factory that provides sampler integration.
// This uses Any downcasting to check at runtime whether V = f64.
let trial_factory = Self::make_trial_factory(&sampler, &completed_trials);
pub fn with_sampler(direction: Direction, sampler: impl Sampler + 'static) -> Self {
Self {
direction,
sampler,
completed_trials,
sampler: Arc::new(sampler),
completed_trials: Arc::new(RwLock::new(Vec::new())),
next_trial_id: AtomicU64::new(0),
trial_factory,
}
}
/// Builds a trial factory for sampler integration when `V = f64`.
fn make_trial_factory(
sampler: &Arc<dyn Sampler>,
completed_trials: &Arc<RwLock<Vec<CompletedTrial<V>>>>,
) -> Option<Arc<dyn Fn(u64) -> Trial + Send + Sync>>
where
V: 'static,
{
// Try to downcast the completed_trials Arc to the f64 specialization.
// This succeeds only when V = f64, enabling automatic sampler integration.
let any_ref: &dyn Any = completed_trials;
let f64_trials: Option<&Arc<RwLock<Vec<CompletedTrial<f64>>>>> = any_ref.downcast_ref();
f64_trials.map(|trials| {
let sampler = Arc::clone(sampler);
let trials = Arc::clone(trials);
let factory: Arc<dyn Fn(u64) -> Trial + Send + Sync> = Arc::new(move |id| {
Trial::with_sampler(id, Arc::clone(&sampler), Arc::clone(&trials))
});
factory
})
}
/// Returns the optimization direction.
pub fn direction(&self) -> Direction {
self.direction
@@ -158,12 +116,8 @@ where
/// let mut study: Study<f64> = Study::new(Direction::Minimize);
/// study.set_sampler(TpeSampler::new());
/// ```
pub fn set_sampler(&mut self, sampler: impl Sampler + 'static)
where
V: 'static,
{
pub fn set_sampler(&mut self, sampler: impl Sampler + 'static) {
self.sampler = Arc::new(sampler);
self.trial_factory = Self::make_trial_factory(&self.sampler, &self.completed_trials);
}
/// Generates the next unique trial ID.
@@ -177,9 +131,9 @@ where
/// parameter values. After the objective function is evaluated, call
/// `complete_trial` or `fail_trial` to record the result.
///
/// For `Study<f64>`, this method automatically integrates with the study's
/// sampler and trial history, so there is no need to call a separate
/// `create_trial_with_sampler()` method.
/// Note: For `Study<f64>`, this method creates a trial without sampler
/// integration. Use `create_trial_with_sampler()` to create trials that
/// use the study's sampler and have access to trial history.
///
/// # Examples
///
@@ -195,11 +149,7 @@ where
/// ```
pub fn create_trial(&self) -> Trial {
let id = self.next_trial_id();
if let Some(factory) = &self.trial_factory {
factory(id)
} else {
Trial::new(id)
}
Trial::new(id)
}
/// Records a completed trial with its objective value.
@@ -216,13 +166,11 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x_param = FloatParam::new(0.0, 1.0);
/// let mut trial = study.create_trial();
/// let x = x_param.suggest(&mut trial).unwrap();
/// let x = trial.suggest_float("x", 0.0, 1.0).unwrap();
/// let objective_value = x * x;
/// study.complete_trial(trial, objective_value);
///
@@ -234,7 +182,6 @@ where
trial.id(),
trial.params().clone(),
trial.distributions().clone(),
trial.param_labels().clone(),
value,
);
self.completed_trials.write().push(completed);
@@ -281,13 +228,11 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// let x_param = FloatParam::new(0.0, 1.0);
/// let mut trial = study.create_trial();
/// let _ = x_param.suggest(&mut trial);
/// let _ = trial.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial, 0.5);
///
/// for completed in study.trials() {
@@ -308,15 +253,13 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
/// assert_eq!(study.n_trials(), 0);
///
/// let x_param = FloatParam::new(0.0, 1.0);
/// let mut trial = study.create_trial();
/// let _ = x_param.suggest(&mut trial);
/// let _ = trial.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial, 0.5);
/// assert_eq!(study.n_trials(), 1);
/// ```
@@ -337,7 +280,6 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Minimize);
@@ -345,14 +287,12 @@ where
/// // Error when no trials completed
/// assert!(study.best_trial().is_err());
///
/// let x_param = FloatParam::new(0.0, 1.0);
///
/// let mut trial1 = study.create_trial();
/// let _ = x_param.suggest(&mut trial1);
/// let _ = trial1.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial1, 0.8);
///
/// let mut trial2 = study.create_trial();
/// let _ = x_param.suggest(&mut trial2);
/// let _ = trial2.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial2, 0.3);
///
/// let best = study.best_trial().unwrap();
@@ -403,7 +343,6 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::{Direction, Study};
///
/// let study: Study<f64> = Study::new(Direction::Maximize);
@@ -411,14 +350,12 @@ where
/// // Error when no trials completed
/// assert!(study.best_value().is_err());
///
/// let x_param = FloatParam::new(0.0, 1.0);
///
/// let mut trial1 = study.create_trial();
/// let _ = x_param.suggest(&mut trial1);
/// let _ = trial1.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial1, 0.3);
///
/// let mut trial2 = study.create_trial();
/// let _ = x_param.suggest(&mut trial2);
/// let _ = trial2.suggest_float("x", 0.0, 1.0);
/// study.complete_trial(trial2, 0.8);
///
/// let best = study.best_value().unwrap();
@@ -455,7 +392,6 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
@@ -463,11 +399,9 @@ where
/// 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(10, |trial| {
/// let x = x_param.suggest(trial)?;
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// Ok::<_, optimizer::Error>(x * x)
/// })
/// .unwrap();
@@ -526,7 +460,6 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
@@ -536,17 +469,12 @@ where
/// 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| {
/// let x_param = x_param.clone();
/// async move {
/// let x = x_param.suggest(&mut trial)?;
/// // Simulate async work (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// }
/// .optimize_async(10, |mut trial| async move {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// // Simulate async work (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// })
/// .await?;
///
@@ -614,7 +542,6 @@ where
/// # Examples
///
/// ```
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
@@ -624,17 +551,12 @@ where
/// 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(10, 4, move |mut trial| {
/// let x_param = x_param.clone();
/// async move {
/// let x = x_param.suggest(&mut trial)?;
/// // Async objective function (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// }
/// .optimize_parallel(10, 4, |mut trial| async move {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// // Async objective function (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// })
/// .await?;
///
@@ -730,7 +652,6 @@ where
/// ```
/// use std::ops::ControlFlow;
///
/// use optimizer::parameter::{FloatParam, Parameter};
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
@@ -738,13 +659,11 @@ where
/// 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_with_callback(
/// 100,
/// |trial| {
/// let x = x_param.suggest(trial)?;
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// Ok::<_, optimizer::Error>(x * x)
/// },
/// |_study, completed_trial| {
@@ -813,72 +732,256 @@ where
}
}
// Specialized implementation for Study<f64> that provides deprecated `_with_sampler` aliases.
//
// For Study<f64>, the generic methods from `impl<V> Study<V>` (like `optimize()`,
// `create_trial()`) now automatically use the sampler via the `trial_factory`.
// The `_with_sampler` method names are deprecated in favor of the generic names.
#[allow(clippy::missing_errors_doc)]
// Specialized implementation for Study<f64> that provides full sampler integration.
impl Study<f64> {
/// Deprecated: use `create_trial()` instead.
/// Creates a new trial with sampler integration.
///
/// The generic `create_trial()` now automatically integrates with the sampler
/// for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `create_trial()` instead — it now uses the sampler automatically for Study<f64>"
)]
/// This method creates a trial that uses the study's sampler and has access
/// to the history of completed trials for informed parameter suggestions.
/// This is the recommended way to create trials when using `Study<f64>`.
///
/// The trial's `suggest_*` methods will delegate to the sampler (e.g., TPE)
/// which can use historical trial data to make informed sampling decisions.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// // With a seeded sampler for reproducibility
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
/// let mut trial = study.create_trial_with_sampler();
///
/// // Parameter suggestions now use the study's sampler and history
/// let x = trial.suggest_float("x", 0.0, 1.0).unwrap();
/// ```
pub fn create_trial_with_sampler(&self) -> Trial {
self.create_trial()
let id = self.next_trial_id();
Trial::with_sampler(
id,
Arc::clone(&self.sampler),
Arc::clone(&self.completed_trials),
)
}
/// Deprecated: use `optimize()` instead.
/// Runs optimization with full sampler integration.
///
/// The generic `optimize()` now automatically integrates with the sampler
/// for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `optimize()` instead — it now uses the sampler automatically for Study<f64>"
)]
pub fn optimize_with_sampler<F, E>(&self, n_trials: usize, objective: F) -> crate::Result<()>
/// This method is similar to the generic `optimizer` method but creates trials
/// using `create_trial_with_sampler()`, giving the sampler access to the history
/// of completed trials for informed parameter suggestions.
///
/// This is the recommended way to run optimization when using `Study<f64>`
/// with advanced samplers like TPE.
///
/// # Arguments
///
/// * `n_trials` - The number of trials to run.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
///
/// # Examples
///
/// ```
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// // Minimize x^2 with sampler integration
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// study
/// .optimize_with_sampler(10, |trial| {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// Ok::<_, optimizer::Error>(x * x)
/// })
/// .unwrap();
///
/// // At least one trial should have completed
/// assert!(study.n_trials() > 0);
/// ```
pub fn optimize_with_sampler<F, E>(
&self,
n_trials: usize,
mut objective: F,
) -> crate::Result<()>
where
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
E: ToString,
{
self.optimize(n_trials, objective)
for _ in 0..n_trials {
let mut trial = self.create_trial_with_sampler();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
}
Err(e) => {
self.fail_trial(trial, e.to_string());
}
}
}
// Return error if no trials succeeded
if self.n_trials() == 0 {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Deprecated: use `optimize_with_callback()` instead.
/// Runs optimization with a callback and full sampler integration.
///
/// The generic `optimize_with_callback()` now automatically integrates with the
/// sampler for `Study<f64>`.
#[deprecated(
since = "0.2.0",
note = "use `optimize_with_callback()` instead — it now uses the sampler automatically for Study<f64>"
)]
/// This method combines the benefits of `optimize_with_sampler` (sampler access
/// to trial history) with `optimize_with_callback` (progress monitoring and
/// early stopping).
///
/// # Arguments
///
/// * `n_trials` - The maximum number of trials to run.
/// * `objective` - A closure that takes a mutable reference to a `Trial` and
/// returns the objective value or an error.
/// * `callback` - A closure called after each successful trial. Returns
/// `ControlFlow::Continue(())` to proceed or `ControlFlow::Break(())` to stop.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if no trials completed successfully.
/// Returns `Error::Internal` if a completed trial is not found after adding (internal invariant violation).
///
/// # Examples
///
/// ```
/// use std::ops::ControlFlow;
///
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// // Optimize with sampler integration and early stopping
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// study
/// .optimize_with_callback_sampler(
/// 100,
/// |trial| {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// Ok::<_, optimizer::Error>(x * x)
/// },
/// |study, _completed_trial| {
/// // Stop after finding 5 good trials
/// if study.n_trials() >= 5 {
/// ControlFlow::Break(())
/// } else {
/// ControlFlow::Continue(())
/// }
/// },
/// )
/// .unwrap();
///
/// assert!(study.n_trials() >= 5);
/// ```
pub fn optimize_with_callback_sampler<F, C, E>(
&self,
n_trials: usize,
objective: F,
callback: C,
mut objective: F,
mut callback: C,
) -> crate::Result<()>
where
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
C: FnMut(&Study<f64>, &CompletedTrial<f64>) -> ControlFlow<()>,
E: ToString,
{
self.optimize_with_callback(n_trials, objective, callback)
for _ in 0..n_trials {
let mut trial = self.create_trial_with_sampler();
match objective(&mut trial) {
Ok(value) => {
self.complete_trial(trial, value);
// Get the just-completed trial for the callback
let trials = self.completed_trials.read();
let Some(completed) = trials.last() else {
return Err(crate::Error::Internal(
"completed trial not found after adding",
));
};
// Call the callback and check if we should stop
// Note: We need to drop the read lock before calling callback
// to avoid potential deadlock if callback accesses the study
let completed_clone = completed.clone();
drop(trials);
if let ControlFlow::Break(()) = callback(self, &completed_clone) {
break;
}
}
Err(e) => {
self.fail_trial(trial, e.to_string());
}
}
}
// Return error if no trials succeeded
if self.n_trials() == 0 {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Deprecated: use `optimize_async()` instead.
/// Runs optimization asynchronously with full sampler integration.
///
/// The generic `optimize_async()` now automatically integrates with the sampler
/// for `Study<f64>`.
/// This method combines async execution with the TPE sampler's ability to use
/// historical trial data for informed parameter suggestions.
///
/// The objective function takes ownership of the `Trial` and must return it
/// along with the result. This allows async operations to use the trial
/// across await points.
///
/// # Arguments
///
/// * `n_trials` - The number of trials to run.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, Result<f64, E>)`.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
///
/// # Examples
///
/// ```
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// # #[cfg(feature = "async")]
/// # async fn example() -> optimizer::Result<()> {
/// // Minimize x^2 with async objective and sampler integration
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// study
/// .optimize_async_with_sampler(10, |mut trial| async move {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// // Simulate async work (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// })
/// .await?;
///
/// // At least one trial should have completed
/// assert!(study.n_trials() > 0);
/// # Ok(())
/// # }
/// ```
#[cfg(feature = "async")]
#[deprecated(
since = "0.2.0",
note = "use `optimize_async()` instead — it now uses the sampler automatically for Study<f64>"
)]
pub async fn optimize_async_with_sampler<F, Fut, E>(
&self,
n_trials: usize,
@@ -889,18 +992,78 @@ impl Study<f64> {
Fut: Future<Output = core::result::Result<(Trial, f64), E>>,
E: ToString,
{
self.optimize_async(n_trials, objective).await
for _ in 0..n_trials {
let trial = self.create_trial_with_sampler();
match objective(trial).await {
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
// For async, we don't have the trial back on error
// We'll just count this as a failed trial without recording it
let _ = e.to_string();
}
}
}
// Return error if no trials succeeded
if self.n_trials() == 0 {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
/// Deprecated: use `optimize_parallel()` instead.
/// Runs optimization with bounded parallelism and full sampler integration.
///
/// The generic `optimize_parallel()` now automatically integrates with the
/// sampler for `Study<f64>`.
/// This method combines parallel async execution with the TPE sampler's ability
/// to use historical trial data for informed parameter suggestions. Up to
/// `concurrency` trials run simultaneously.
///
/// The objective function takes ownership of the `Trial` and must return it
/// along with the result. This allows async operations to use the trial
/// across await points.
///
/// # Arguments
///
/// * `n_trials` - The total number of trials to run.
/// * `concurrency` - The maximum number of trials to run simultaneously.
/// * `objective` - A function that takes a `Trial` and returns a `Future`
/// that resolves to a tuple of `(Trial, f64)` or an error.
///
/// # Errors
///
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
/// Returns `Error::TaskError` if the semaphore is closed or a spawned task panics.
///
/// # Examples
///
/// ```
/// use optimizer::sampler::random::RandomSampler;
/// use optimizer::{Direction, Study};
///
/// # #[cfg(feature = "async")]
/// # async fn example() -> optimizer::Result<()> {
/// // Minimize x^2 with parallel async evaluation and sampler integration
/// let sampler = RandomSampler::with_seed(42);
/// let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
///
/// study
/// .optimize_parallel_with_sampler(10, 4, |mut trial| async move {
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
/// // Async objective function (e.g., network request)
/// let value = x * x;
/// Ok::<_, optimizer::Error>((trial, value))
/// })
/// .await?;
///
/// // All trials should have completed
/// assert_eq!(study.n_trials(), 10);
/// # Ok(())
/// # }
/// ```
#[cfg(feature = "async")]
#[deprecated(
since = "0.2.0",
note = "use `optimize_parallel()` instead — it now uses the sampler automatically for Study<f64>"
)]
pub async fn optimize_parallel_with_sampler<F, Fut, E>(
&self,
n_trials: usize,
@@ -912,7 +1075,51 @@ impl Study<f64> {
Fut: Future<Output = core::result::Result<(Trial, f64), E>> + Send,
E: ToString + Send + 'static,
{
self.optimize_parallel(n_trials, concurrency, objective)
.await
use tokio::sync::Semaphore;
let semaphore = Arc::new(Semaphore::new(concurrency));
let objective = Arc::new(objective);
let mut handles = Vec::with_capacity(n_trials);
for _ in 0..n_trials {
let permit = semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?;
let trial = self.create_trial_with_sampler();
let objective = Arc::clone(&objective);
let handle = tokio::spawn(async move {
let result = objective(trial).await;
drop(permit); // Release semaphore permit when done
result
});
handles.push(handle);
}
// Wait for all tasks and record results
for handle in handles {
match handle
.await
.map_err(|e| crate::Error::TaskError(e.to_string()))?
{
Ok((trial, value)) => {
self.complete_trial(trial, value);
}
Err(e) => {
let _ = e.to_string();
}
}
}
// Return error if no trials succeeded
if self.n_trials() == 0 {
return Err(crate::Error::NoCompletedTrials);
}
Ok(())
}
}
+763 -53
View File
@@ -1,17 +1,77 @@
//! Trial implementation for tracking sampled parameters and trial state.
use core::ops::{Range, RangeInclusive};
use std::collections::HashMap;
use std::sync::Arc;
use parking_lot::RwLock;
use crate::distribution::Distribution;
use crate::distribution::{
CategoricalDistribution, Distribution, FloatDistribution, IntDistribution,
};
use crate::error::{Error, Result};
use crate::param::ParamValue;
use crate::parameter::{ParamId, Parameter};
use crate::sampler::{CompletedTrial, Sampler};
use crate::types::TrialState;
/// A trait for types that can be used with [`Trial::suggest_range`].
///
/// This trait is implemented for [`Range`] and [`RangeInclusive`] over `f64` and `i64`.
/// It allows using Rust's range syntax directly with the optimizer.
///
/// # Supported Range Types
///
/// | Range Type | Example | Description |
/// |------------|---------|-------------|
/// | `Range<f64>` | `0.0..1.0` | Float range (end-exclusive, treated as inclusive for continuous sampling) |
/// | `RangeInclusive<f64>` | `0.0..=1.0` | Float range (end-inclusive) |
/// | `Range<i64>` | `1..10` | Integer range from 1 to 9 (end-exclusive) |
/// | `RangeInclusive<i64>` | `1..=10` | Integer range from 1 to 10 (end-inclusive) |
pub trait SuggestableRange {
/// The output type when suggesting from this range.
type Output;
/// Suggests a value from this range using the given trial.
///
/// # Errors
///
/// Returns an error if the range is invalid (e.g., empty or low > high).
fn suggest(self, trial: &mut Trial, name: String) -> Result<Self::Output>;
}
impl SuggestableRange for Range<f64> {
type Output = f64;
fn suggest(self, trial: &mut Trial, name: String) -> Result<f64> {
trial.suggest_float(name, self.start, self.end)
}
}
impl SuggestableRange for RangeInclusive<f64> {
type Output = f64;
fn suggest(self, trial: &mut Trial, name: String) -> Result<f64> {
trial.suggest_float(name, *self.start(), *self.end())
}
}
impl SuggestableRange for Range<i64> {
type Output = i64;
fn suggest(self, trial: &mut Trial, name: String) -> Result<i64> {
// Range is exclusive on the end, so subtract 1
trial.suggest_int(name, self.start, self.end.saturating_sub(1))
}
}
impl SuggestableRange for RangeInclusive<i64> {
type Output = i64;
fn suggest(self, trial: &mut Trial, name: String) -> Result<i64> {
trial.suggest_int(name, *self.start(), *self.end())
}
}
/// A trial represents a single evaluation of the objective function.
///
/// Each trial has a unique ID and stores the sampled parameters along with
@@ -26,12 +86,10 @@ pub struct Trial {
id: u64,
/// Current state of the trial.
state: TrialState,
/// Sampled parameter values, keyed by parameter id.
params: HashMap<ParamId, ParamValue>,
/// Parameter distributions, keyed by parameter id.
distributions: HashMap<ParamId, Distribution>,
/// Human-readable labels for parameters, keyed by parameter id.
param_labels: HashMap<ParamId, String>,
/// Sampled parameter values, keyed by parameter name.
params: HashMap<String, ParamValue>,
/// Parameter distributions, keyed by parameter name.
distributions: HashMap<String, Distribution>,
/// The sampler to use for generating parameter values.
sampler: Option<Arc<dyn Sampler>>,
/// Access to the history of completed trials (shared with Study).
@@ -45,7 +103,6 @@ impl core::fmt::Debug for Trial {
.field("state", &self.state)
.field("params", &self.params)
.field("distributions", &self.distributions)
.field("param_labels", &self.param_labels)
.field("has_sampler", &self.sampler.is_some())
.field("has_history", &self.history.is_some())
.finish()
@@ -57,7 +114,7 @@ impl Trial {
///
/// The trial starts in the `Running` state with no parameters sampled.
/// This constructor creates a trial without a sampler, which will use
/// local random sampling for suggest methods.
/// local random sampling for suggest_* methods.
///
/// For trials that use the study's sampler, use `Trial::with_sampler` instead.
///
@@ -80,7 +137,6 @@ impl Trial {
state: TrialState::Running,
params: HashMap::new(),
distributions: HashMap::new(),
param_labels: HashMap::new(),
sampler: None,
history: None,
}
@@ -106,7 +162,6 @@ impl Trial {
state: TrialState::Running,
params: HashMap::new(),
distributions: HashMap::new(),
param_labels: HashMap::new(),
sampler: Some(sampler),
history: Some(history),
}
@@ -143,22 +198,16 @@ impl Trial {
/// Returns a reference to the sampled parameters.
#[must_use]
pub fn params(&self) -> &HashMap<ParamId, ParamValue> {
pub fn params(&self) -> &HashMap<String, ParamValue> {
&self.params
}
/// Returns a reference to the parameter distributions.
#[must_use]
pub fn distributions(&self) -> &HashMap<ParamId, Distribution> {
pub fn distributions(&self) -> &HashMap<String, Distribution> {
&self.distributions
}
/// Returns a reference to the parameter labels.
#[must_use]
pub fn param_labels(&self) -> &HashMap<ParamId, String> {
&self.param_labels
}
/// Sets the trial state to Complete.
pub(crate) fn set_complete(&mut self) {
self.state = TrialState::Complete;
@@ -169,69 +218,730 @@ impl Trial {
self.state = TrialState::Failed;
}
/// Suggests a parameter value using a [`Parameter`] definition.
/// Suggests a float parameter with the given bounds.
///
/// This is the primary entry point for sampling parameters. It handles
/// validation, caching, conflict detection, sampling, and conversion.
/// If the parameter has already been sampled with the same bounds, the cached value is returned.
/// If the parameter was sampled with different bounds, a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `param` - The parameter definition.
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive).
/// * `high` - The upper bound (inclusive).
///
/// # Errors
///
/// Returns an error if:
/// - The parameter fails validation
/// - The parameter conflicts with a previously suggested parameter of the same id
/// - Sampling or conversion fails
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different bounds.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
/// use optimizer::parameter::{BoolParam, FloatParam, IntParam, Parameter};
///
/// let x_param = FloatParam::new(0.0, 1.0);
/// let n_param = IntParam::new(1, 10);
/// let flag_param = BoolParam::new();
///
/// let mut trial = Trial::new(0);
/// let x = trial.suggest_float("x", 0.0, 1.0).unwrap();
/// assert!(x >= 0.0 && x <= 1.0);
///
/// let x = trial.suggest_param(&x_param).unwrap();
/// let n = trial.suggest_param(&n_param).unwrap();
/// let flag = trial.suggest_param(&flag_param).unwrap();
/// // Calling again with same bounds returns cached value
/// let x2 = trial.suggest_float("x", 0.0, 1.0).unwrap();
/// assert_eq!(x, x2);
/// ```
pub fn suggest_param<P: Parameter>(&mut self, param: &P) -> Result<P::Value> {
param.validate()?;
pub fn suggest_float(&mut self, name: impl Into<String>, low: f64, high: f64) -> Result<f64> {
if low > high {
return Err(Error::InvalidBounds { low, high });
}
let param_id = param.id();
let distribution = param.distribution();
let name = name.into();
let distribution = FloatDistribution {
low,
high,
log_scale: false,
step: None,
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&param_id) {
if *existing_dist == distribution {
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Float(existing) = existing_dist
&& (existing.low - low).abs() < f64::EPSILON
&& (existing.high - high).abs() < f64::EPSILON
&& !existing.log_scale
&& existing.step.is_none()
{
// Same distribution, return cached value
if let Some(value) = self.params.get(&param_id) {
return param.cast_param_value(value);
if let Some(ParamValue::Float(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name: param.label(),
reason: "parameter was previously sampled with different configuration or type"
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler
let value = self.sample_value(&distribution);
let result = param.cast_param_value(&value)?;
let dist = Distribution::Float(distribution);
let ParamValue::Float(value) = self.sample_value(&dist) else {
return Err(Error::Internal(
"Float distribution should return Float value",
));
};
// Store distribution, value, and label
self.distributions.insert(param_id, distribution);
self.params.insert(param_id, value);
self.param_labels.insert(param_id, param.label());
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Float(value));
Ok(result)
Ok(value)
}
/// Suggests a float parameter sampled on a logarithmic scale.
///
/// The value is sampled uniformly in log space, which is useful for parameters
/// that span multiple orders of magnitude (e.g., learning rates).
///
/// If the parameter has already been sampled with the same bounds and `log_scale=true`,
/// the cached value is returned. If the parameter was sampled with different configuration,
/// a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive, must be positive).
/// * `high` - The upper bound (inclusive).
///
/// # Errors
///
/// Returns `InvalidLogBounds` if `low <= 0`.
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different configuration.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let lr = trial
/// .suggest_float_log("learning_rate", 1e-5, 1e-1)
/// .unwrap();
/// assert!(lr >= 1e-5 && lr <= 1e-1);
///
/// // Calling again with same bounds returns cached value
/// let lr2 = trial
/// .suggest_float_log("learning_rate", 1e-5, 1e-1)
/// .unwrap();
/// assert_eq!(lr, lr2);
/// ```
pub fn suggest_float_log(
&mut self,
name: impl Into<String>,
low: f64,
high: f64,
) -> Result<f64> {
if low <= 0.0 {
return Err(Error::InvalidLogBounds);
}
if low > high {
return Err(Error::InvalidBounds { low, high });
}
let name = name.into();
let distribution = FloatDistribution {
low,
high,
log_scale: true,
step: None,
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Float(existing) = existing_dist
&& (existing.low - low).abs() < f64::EPSILON
&& (existing.high - high).abs() < f64::EPSILON
&& existing.log_scale
&& existing.step.is_none()
{
// Same distribution, return cached value
if let Some(ParamValue::Float(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler (sampler handles log-scale transformation)
let dist = Distribution::Float(distribution);
let ParamValue::Float(value) = self.sample_value(&dist) else {
return Err(Error::Internal(
"Float distribution should return Float value",
));
};
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Float(value));
Ok(value)
}
/// Suggests a float parameter that snaps to a step grid.
///
/// The value is sampled from the discrete set {low, low + step, low + 2*step, ...}
/// where each value is <= high.
///
/// If the parameter has already been sampled with the same configuration,
/// the cached value is returned. If the parameter was sampled with different configuration,
/// a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive).
/// * `high` - The upper bound (inclusive).
/// * `step` - The step size (must be positive).
///
/// # Errors
///
/// Returns `InvalidStep` if `step <= 0`.
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different configuration.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let x = trial.suggest_float_step("x", 0.0, 1.0, 0.25).unwrap();
/// // x will be one of: 0.0, 0.25, 0.5, 0.75, 1.0
/// assert!(x >= 0.0 && x <= 1.0);
/// assert!((x / 0.25).fract().abs() < 1e-10 || (x / 0.25).fract().abs() > 1.0 - 1e-10);
///
/// // Calling again with same bounds returns cached value
/// let x2 = trial.suggest_float_step("x", 0.0, 1.0, 0.25).unwrap();
/// assert_eq!(x, x2);
/// ```
pub fn suggest_float_step(
&mut self,
name: impl Into<String>,
low: f64,
high: f64,
step: f64,
) -> Result<f64> {
if step <= 0.0 {
return Err(Error::InvalidStep);
}
if low > high {
return Err(Error::InvalidBounds { low, high });
}
let name = name.into();
let distribution = FloatDistribution {
low,
high,
log_scale: false,
step: Some(step),
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Float(existing) = existing_dist
&& (existing.low - low).abs() < f64::EPSILON
&& (existing.high - high).abs() < f64::EPSILON
&& !existing.log_scale
&& existing.step == Some(step)
{
// Same distribution, return cached value
if let Some(ParamValue::Float(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler (sampler handles step-grid)
let dist = Distribution::Float(distribution);
let ParamValue::Float(value) = self.sample_value(&dist) else {
return Err(Error::Internal(
"Float distribution should return Float value",
));
};
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Float(value));
Ok(value)
}
/// Suggests an integer parameter with the given bounds.
///
/// The value is sampled uniformly from the range [low, high] inclusive.
///
/// If the parameter has already been sampled with the same bounds, the cached value is returned.
/// If the parameter was sampled with different bounds, a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive).
/// * `high` - The upper bound (inclusive).
///
/// # Errors
///
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different bounds.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let n = trial.suggest_int("n_layers", 1, 10).unwrap();
/// assert!(n >= 1 && n <= 10);
///
/// // Calling again with same bounds returns cached value
/// let n2 = trial.suggest_int("n_layers", 1, 10).unwrap();
/// assert_eq!(n, n2);
/// ```
#[allow(clippy::cast_precision_loss)]
pub fn suggest_int(&mut self, name: impl Into<String>, low: i64, high: i64) -> Result<i64> {
if low > high {
return Err(Error::InvalidBounds {
low: low as f64,
high: high as f64,
});
}
let name = name.into();
let distribution = IntDistribution {
low,
high,
log_scale: false,
step: None,
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Int(existing) = existing_dist
&& existing.low == low
&& existing.high == high
&& !existing.log_scale
&& existing.step.is_none()
{
// Same distribution, return cached value
if let Some(ParamValue::Int(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler
let dist = Distribution::Int(distribution);
let ParamValue::Int(value) = self.sample_value(&dist) else {
return Err(Error::Internal("Int distribution should return Int value"));
};
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Int(value));
Ok(value)
}
/// Suggests an integer parameter sampled on a logarithmic scale.
///
/// The value is sampled uniformly in log space, which is useful for parameters
/// that span multiple orders of magnitude (e.g., batch sizes).
///
/// If the parameter has already been sampled with the same bounds and `log_scale=true`,
/// the cached value is returned. If the parameter was sampled with different configuration,
/// a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive, must be >= 1).
/// * `high` - The upper bound (inclusive).
///
/// # Errors
///
/// Returns `InvalidLogBounds` if `low < 1`.
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different configuration.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let batch_size = trial.suggest_int_log("batch_size", 1, 1024).unwrap();
/// assert!(batch_size >= 1 && batch_size <= 1024);
///
/// // Calling again with same bounds returns cached value
/// let batch_size2 = trial.suggest_int_log("batch_size", 1, 1024).unwrap();
/// assert_eq!(batch_size, batch_size2);
/// ```
#[allow(clippy::cast_precision_loss)]
pub fn suggest_int_log(&mut self, name: impl Into<String>, low: i64, high: i64) -> Result<i64> {
if low < 1 {
return Err(Error::InvalidLogBounds);
}
if low > high {
return Err(Error::InvalidBounds {
low: low as f64,
high: high as f64,
});
}
let name = name.into();
let distribution = IntDistribution {
low,
high,
log_scale: true,
step: None,
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Int(existing) = existing_dist
&& existing.low == low
&& existing.high == high
&& existing.log_scale
&& existing.step.is_none()
{
// Same distribution, return cached value
if let Some(ParamValue::Int(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler (sampler handles log-scale transformation)
let dist = Distribution::Int(distribution);
let ParamValue::Int(value) = self.sample_value(&dist) else {
return Err(Error::Internal("Int distribution should return Int value"));
};
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Int(value));
Ok(value)
}
/// Suggests an integer parameter that snaps to a step grid.
///
/// The value is sampled from the discrete set {low, low + step, low + 2*step, ...}
/// where each value is <= high.
///
/// If the parameter has already been sampled with the same configuration,
/// the cached value is returned. If the parameter was sampled with different configuration,
/// a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `low` - The lower bound (inclusive).
/// * `high` - The upper bound (inclusive).
/// * `step` - The step size (must be positive).
///
/// # Errors
///
/// Returns `InvalidStep` if `step <= 0`.
/// Returns `InvalidBounds` if `low > high`.
/// Returns `ParameterConflict` if the parameter was previously sampled with different configuration.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let n = trial
/// .suggest_int_step("n_estimators", 100, 500, 50)
/// .unwrap();
/// // n will be one of: 100, 150, 200, 250, 300, 350, 400, 450, 500
/// assert!(n >= 100 && n <= 500);
/// assert!((n - 100) % 50 == 0);
///
/// // Calling again with same bounds returns cached value
/// let n2 = trial
/// .suggest_int_step("n_estimators", 100, 500, 50)
/// .unwrap();
/// assert_eq!(n, n2);
/// ```
#[allow(clippy::cast_precision_loss)]
pub fn suggest_int_step(
&mut self,
name: impl Into<String>,
low: i64,
high: i64,
step: i64,
) -> Result<i64> {
if step <= 0 {
return Err(Error::InvalidStep);
}
if low > high {
return Err(Error::InvalidBounds {
low: low as f64,
high: high as f64,
});
}
let name = name.into();
let distribution = IntDistribution {
low,
high,
log_scale: false,
step: Some(step),
};
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Int(existing) = existing_dist
&& existing.low == low
&& existing.high == high
&& !existing.log_scale
&& existing.step == Some(step)
{
// Same distribution, return cached value
if let Some(ParamValue::Int(value)) = self.params.get(&name) {
return Ok(*value);
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different bounds or type"
.to_string(),
});
}
// Sample using the sampler (sampler handles step-grid)
let dist = Distribution::Int(distribution);
let ParamValue::Int(value) = self.sample_value(&dist) else {
return Err(Error::Internal("Int distribution should return Int value"));
};
// Store distribution and value
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Int(value));
Ok(value)
}
/// Suggests a categorical parameter from the given choices.
///
/// The value is selected uniformly at random from the provided choices.
/// Internally, the index of the selected choice is stored.
///
/// If the parameter has already been sampled with the same number of choices,
/// the cached value (same index) is returned. If the parameter was sampled with
/// a different number of choices, a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `choices` - A slice of choices to select from.
///
/// # Type Parameters
///
/// * `T` - The type of the choices. Only requires `Clone`.
///
/// # Errors
///
/// Returns `EmptyChoices` if `choices` is empty.
/// Returns `ParameterConflict` if the parameter was previously sampled with a different number of choices.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let optimizer = trial
/// .suggest_categorical("optimizer", &["sgd", "adam", "rmsprop"])
/// .unwrap();
/// assert!(["sgd", "adam", "rmsprop"].contains(&optimizer));
///
/// // Calling again with same choices returns cached value
/// let optimizer2 = trial
/// .suggest_categorical("optimizer", &["sgd", "adam", "rmsprop"])
/// .unwrap();
/// assert_eq!(optimizer, optimizer2);
/// ```
pub fn suggest_categorical<T: Clone>(
&mut self,
name: impl Into<String>,
choices: &[T],
) -> Result<T> {
if choices.is_empty() {
return Err(Error::EmptyChoices);
}
let name = name.into();
let n_choices = choices.len();
let distribution = CategoricalDistribution { n_choices };
// Check if parameter already exists
if let Some(existing_dist) = self.distributions.get(&name) {
// Verify the distribution matches
if let Distribution::Categorical(existing) = existing_dist
&& existing.n_choices == n_choices
{
// Same distribution, return cached value
if let Some(ParamValue::Categorical(index)) = self.params.get(&name) {
return Ok(choices[*index].clone());
}
}
// Distribution exists but doesn't match
return Err(Error::ParameterConflict {
name,
reason: "parameter was previously sampled with different number of choices or type"
.to_string(),
});
}
// Sample using the sampler
let dist = Distribution::Categorical(distribution);
let ParamValue::Categorical(index) = self.sample_value(&dist) else {
return Err(Error::Internal(
"Categorical distribution should return Categorical value",
));
};
// Store distribution and value (store the index)
self.distributions.insert(name.clone(), dist);
self.params.insert(name, ParamValue::Categorical(index));
Ok(choices[index].clone())
}
/// Suggests a boolean parameter.
///
/// The value is selected uniformly at random from `{false, true}`.
/// This is equivalent to calling `suggest_categorical(name, &[false, true])`.
///
/// If the parameter has already been sampled, the cached value is returned.
/// If the parameter was sampled with a different type, a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
///
/// # Errors
///
/// Returns `ParameterConflict` if the parameter was previously sampled with a different type.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
/// let use_dropout = trial.suggest_bool("use_dropout").unwrap();
/// assert!(use_dropout == true || use_dropout == false);
///
/// // Calling again returns cached value
/// let use_dropout2 = trial.suggest_bool("use_dropout").unwrap();
/// assert_eq!(use_dropout, use_dropout2);
/// ```
pub fn suggest_bool(&mut self, name: impl Into<String>) -> Result<bool> {
self.suggest_categorical(name, &[false, true])
}
/// Suggests a parameter value from a range.
///
/// This method accepts both [`Range`] (`..`) and [`RangeInclusive`] (`..=`)
/// for both `f64` and `i64` types, allowing natural Rust range syntax.
///
/// For integer ranges, note that `Range` (`..`) is end-exclusive while
/// `RangeInclusive` (`..=`) is end-inclusive, matching Rust's semantics.
///
/// If the parameter has already been sampled with the same bounds, the cached value is returned.
/// If the parameter was sampled with different bounds, a `ParameterConflict` error is returned.
///
/// # Arguments
///
/// * `name` - The name of the parameter.
/// * `range` - The range to sample from.
///
/// # Type Parameters
///
/// * `R` - A range type implementing [`SuggestableRange`].
///
/// # Errors
///
/// Returns `InvalidBounds` if the range is invalid (e.g., low > high or empty integer range).
/// Returns `ParameterConflict` if the parameter was previously sampled with different bounds.
///
/// # Examples
///
/// ```
/// use optimizer::Trial;
///
/// let mut trial = Trial::new(0);
///
/// // Float ranges
/// let x = trial.suggest_range("x", 0.0..1.0).unwrap();
/// assert!(x >= 0.0 && x <= 1.0);
///
/// let y = trial.suggest_range("y", 0.0..=1.0).unwrap();
/// assert!(y >= 0.0 && y <= 1.0);
///
/// // Integer ranges
/// let n = trial.suggest_range("n", 1_i64..10).unwrap(); // 1 to 9 inclusive
/// assert!(n >= 1 && n <= 9);
///
/// let m = trial.suggest_range("m", 1_i64..=10).unwrap(); // 1 to 10 inclusive
/// assert!(m >= 1 && m <= 10);
///
/// // Calling again with same range returns cached value
/// let x2 = trial.suggest_range("x", 0.0..1.0).unwrap();
/// assert_eq!(x, x2);
/// ```
pub fn suggest_range<R: SuggestableRange>(
&mut self,
name: impl Into<String>,
range: R,
) -> Result<R::Output> {
range.suggest(self, name.into())
}
}
+23 -61
View File
@@ -4,7 +4,6 @@
#![cfg(feature = "async")]
use optimizer::parameter::{FloatParam, Parameter};
use optimizer::sampler::random::RandomSampler;
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, Error, Study};
@@ -14,15 +13,10 @@ 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, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x))
}
.optimize_async(10, |mut trial| async move {
let x = trial.suggest_float("x", -10.0, 10.0)?;
Ok::<_, Error>((trial, x * x))
})
.await
.expect("async optimization should succeed");
@@ -33,7 +27,7 @@ async fn test_optimize_async_basic() {
}
#[tokio::test]
async fn test_optimize_async_with_tpe() {
async fn test_optimize_async_with_sampler() {
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(5)
@@ -42,15 +36,10 @@ async fn test_optimize_async_with_tpe() {
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
let x_param = FloatParam::new(-5.0, 5.0);
study
.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))
}
.optimize_async_with_sampler(15, |mut trial| async move {
let x = trial.suggest_float("x", -5.0, 5.0)?;
Ok::<_, Error>((trial, x * x))
})
.await
.expect("async optimization with sampler should succeed");
@@ -65,15 +54,10 @@ 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, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x * x))
}
.optimize_parallel(20, 4, |mut trial| async move {
let x = trial.suggest_float("x", -10.0, 10.0)?;
Ok::<_, Error>((trial, x * x))
})
.await
.expect("parallel optimization should succeed");
@@ -82,7 +66,7 @@ async fn test_optimize_parallel() {
}
#[tokio::test]
async fn test_optimize_parallel_with_tpe() {
async fn test_optimize_parallel_with_sampler() {
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(5)
@@ -91,18 +75,11 @@ async fn test_optimize_parallel_with_tpe() {
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(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))
}
.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))
})
.await
.expect("parallel optimization with sampler should succeed");
@@ -128,7 +105,6 @@ 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);
@@ -163,7 +139,6 @@ 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);
@@ -187,15 +162,12 @@ 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, move |mut trial| {
.optimize_async(10, |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 = x_param.suggest(&mut trial)?;
let x = trial.suggest_float("x", 0.0, 10.0)?;
Ok::<_, Error>((trial, x))
} else {
Err(Error::NoCompletedTrials) // Use as error type
@@ -214,16 +186,11 @@ 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, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x))
}
.optimize_parallel(5, 10, |mut trial| async move {
let x = trial.suggest_float("x", 0.0, 10.0)?;
Ok::<_, Error>((trial, x))
})
.await
.expect("should handle high concurrency");
@@ -236,16 +203,11 @@ 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, move |mut trial| {
let x_param = x_param.clone();
async move {
let x = x_param.suggest(&mut trial)?;
Ok::<_, Error>((trial, x))
}
.optimize_parallel(10, 1, |mut trial| async move {
let x = trial.suggest_float("x", 0.0, 10.0)?;
Ok::<_, Error>((trial, x))
})
.await
.expect("should work with single concurrency");
-61
View File
@@ -1,61 +0,0 @@
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));
}
+354 -306
View File
File diff suppressed because it is too large Load Diff
-424
View File
@@ -1,424 +0,0 @@
//! Integration tests for the Multivariate TPE sampler.
//!
//! These tests compare the performance of `MultivariateTpeSampler` against
//! the standard `TpeSampler` on problems with and without parameter correlations.
#![allow(
clippy::cast_sign_loss,
clippy::cast_precision_loss,
clippy::cast_possible_truncation
)]
use optimizer::parameter::{CategoricalParam, FloatParam, IntParam, Parameter};
use optimizer::sampler::tpe::{MultivariateTpeSampler, TpeSampler};
use optimizer::{Direction, Error, Study};
// =============================================================================
// Rosenbrock function: f(x,y) = (a-x)^2 + b*(y-x^2)^2
// This is a classic benchmark with strong parameter correlation.
// Optimal: (x, y) = (a, a^2) with f(x, y) = 0
// Standard parameters: a = 1, b = 100
// The "banana-shaped" valley makes this hard for independent samplers.
// =============================================================================
/// Computes the Rosenbrock function value.
///
/// f(x, y) = (a - x)^2 + b * (y - x^2)^2
///
/// With standard parameters a = 1, b = 100:
/// - Optimal point: (1, 1)
/// - Optimal value: 0
fn rosenbrock(x: f64, y: f64) -> f64 {
let a = 1.0;
let b = 100.0;
(a - x).powi(2) + b * (y - x * x).powi(2)
}
// =============================================================================
// Test: Multivariate TPE on Rosenbrock function (correlated parameters)
// =============================================================================
#[test]
fn test_multivariate_tpe_rosenbrock_finds_good_solution() {
// Multivariate TPE should find a good solution on Rosenbrock
// because it can model the correlation between x and y
let sampler = MultivariateTpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(100, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
// Multivariate TPE should find a reasonably good solution
// The global minimum is 0, but getting close is challenging
assert!(
best.value < 10.0,
"Multivariate TPE should find good Rosenbrock solution: best value {} should be < 10.0",
best.value
);
}
#[test]
fn test_independent_tpe_rosenbrock() {
// Independent TPE (standard TpeSampler) on Rosenbrock
// This establishes a baseline for comparison
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(100, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
// Independent TPE should still find a decent solution,
// but may not be as good as multivariate TPE on this correlated problem
assert!(
best.value < 50.0,
"Independent TPE should find reasonable Rosenbrock solution: best value {} should be < 50.0",
best.value
);
}
#[test]
fn test_multivariate_tpe_outperforms_on_correlated_problem() {
// Run multiple seeds and compare average performance
// Multivariate TPE should generally find better solutions on Rosenbrock
let n_runs = 5;
let n_trials = 80;
let mut multivariate_best_values = Vec::new();
let mut independent_best_values = Vec::new();
for seed in 0..n_runs {
// Multivariate TPE
let multivariate_sampler = MultivariateTpeSampler::builder()
.seed(seed as u64)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.unwrap();
multivariate_best_values.push(study.best_trial().unwrap().value);
// Independent TPE
let independent_sampler = TpeSampler::builder()
.seed(seed as u64)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(rosenbrock(x, y))
})
.unwrap();
independent_best_values.push(study.best_trial().unwrap().value);
}
let multivariate_mean: f64 = multivariate_best_values.iter().sum::<f64>() / n_runs as f64;
let independent_mean: f64 = independent_best_values.iter().sum::<f64>() / n_runs as f64;
// Log results for debugging (these won't show in normal test runs,
// but are useful when running with --nocapture)
eprintln!("Multivariate TPE mean best: {multivariate_mean:.4}");
eprintln!("Independent TPE mean best: {independent_mean:.4}");
eprintln!("Multivariate best values: {multivariate_best_values:?}");
eprintln!("Independent best values: {independent_best_values:?}");
// Both methods should find reasonable solutions
assert!(
multivariate_mean < 20.0,
"Multivariate TPE mean {multivariate_mean:.4} should be < 20.0"
);
assert!(
independent_mean < 100.0,
"Independent TPE mean {independent_mean:.4} should be < 100.0"
);
}
// =============================================================================
// Independent parameter problem: f(x,y) = x^2 + y^2
// No correlation between parameters - both methods should work equally well.
// =============================================================================
/// Simple sphere function with independent parameters.
///
/// f(x, y) = x^2 + y^2
///
/// Optimal point: (0, 0)
/// Optimal value: 0
fn sphere(x: f64, y: f64) -> f64 {
x * x + y * y
}
#[test]
fn test_multivariate_tpe_independent_problem() {
// On an independent problem, multivariate TPE should still work well
let sampler = MultivariateTpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
assert!(
best.value < 5.0,
"Multivariate TPE should find good solution on sphere: best value {} should be < 5.0",
best.value
);
}
#[test]
fn test_independent_tpe_independent_problem() {
// Baseline: independent TPE on sphere function
let sampler = TpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.n_ei_candidates(24)
.build()
.unwrap();
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(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
assert!(
best.value < 5.0,
"Independent TPE should find good solution on sphere: best value {} should be < 5.0",
best.value
);
}
#[test]
fn test_both_samplers_work_on_independent_problem() {
// Run both samplers on the independent sphere function
// and verify they both achieve similar performance
let n_runs = 5;
let n_trials = 50;
let mut multivariate_results = Vec::new();
let mut independent_results = Vec::new();
for seed in 0..n_runs {
// Multivariate TPE
let sampler = MultivariateTpeSampler::builder()
.seed(seed as u64)
.n_startup_trials(10)
.build()
.unwrap();
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(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.unwrap();
multivariate_results.push(study.best_trial().unwrap().value);
// Independent TPE
let sampler = TpeSampler::builder()
.seed(seed as u64)
.n_startup_trials(10)
.build()
.unwrap();
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(n_trials, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.unwrap();
independent_results.push(study.best_trial().unwrap().value);
}
let multivariate_mean: f64 = multivariate_results.iter().sum::<f64>() / n_runs as f64;
let independent_mean: f64 = independent_results.iter().sum::<f64>() / n_runs as f64;
eprintln!("Sphere function results:");
eprintln!(" Multivariate TPE mean: {multivariate_mean:.4}");
eprintln!(" Independent TPE mean: {independent_mean:.4}");
// Both should find good solutions on this simple problem
assert!(
multivariate_mean < 5.0,
"Multivariate TPE mean {multivariate_mean:.4} should be < 5.0 on sphere"
);
assert!(
independent_mean < 5.0,
"Independent TPE mean {independent_mean:.4} should be < 5.0 on sphere"
);
}
// =============================================================================
// Test: Multivariate TPE with group decomposition
// =============================================================================
#[test]
fn test_multivariate_tpe_with_group_decomposition() {
// Test that group decomposition works correctly
let sampler = MultivariateTpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.group(true) // Enable group decomposition
.build()
.unwrap();
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(50, |trial| {
let x = x_param.suggest(trial)?;
let y = y_param.suggest(trial)?;
Ok::<_, Error>(sphere(x, y))
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
assert!(
best.value < 10.0,
"Multivariate TPE with groups should find good solution: best value {} should be < 10.0",
best.value
);
}
// =============================================================================
// Test: Multivariate TPE with mixed parameter types
// =============================================================================
#[test]
fn test_multivariate_tpe_mixed_parameter_types() {
// Test with float, int, and categorical parameters
let sampler = MultivariateTpeSampler::builder()
.seed(42)
.n_startup_trials(10)
.build()
.unwrap();
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(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 {
"a" => 1.0,
"b" => 0.5,
"c" => 2.0,
_ => unreachable!(),
};
Ok::<_, Error>(x * x + (n as f64 - 5.0).powi(2) * mode_factor)
})
.expect("optimization should succeed");
let best = study.best_trial().expect("should have at least one trial");
// Should find a reasonable solution
assert!(
best.value < 25.0,
"Multivariate TPE should handle mixed types: best value {} should be < 25.0",
best.value
);
}
-240
View File
@@ -1,240 +0,0 @@
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;
}
-1
View File
@@ -1,3 +1,2 @@
[default.extend-words]
Tpe = "Tpe"
consts = "consts"