Implement Parameters API
- Add `.name()` builder method on all 5 parameter types for custom labels - Add `CompletedTrial::get(¶m)` for typed parameter access - Add `Display` impl on `ParamValue` - Add prelude module at `optimizer::prelude::*` - Shadow `_with_sampler` methods on `Study<f64>` so `optimize()` auto-uses the configured sampler; deprecate `_with_sampler` variants - Use runtime `Any` downcasting with `trial_factory` to avoid E0592 - Update all examples and tests to use the new API Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -6,7 +6,7 @@
|
||||
//!
|
||||
//! # Key Concepts Demonstrated
|
||||
//!
|
||||
//! - Async optimization with `optimize_parallel_with_sampler`
|
||||
//! - Async optimization with `optimize_parallel`
|
||||
//! - Running multiple trials concurrently for faster optimization
|
||||
//! - Boolean and categorical parameter types
|
||||
//! - Measuring speedup from parallelism
|
||||
@@ -26,9 +26,7 @@
|
||||
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use optimizer::parameter::{BoolParam, CategoricalParam, IntParam, Parameter};
|
||||
use optimizer::sampler::tpe::TpeSampler;
|
||||
use optimizer::{Direction, ParamValue, Study, Trial};
|
||||
use optimizer::prelude::*;
|
||||
|
||||
// ============================================================================
|
||||
// Configuration: Service parameters we want to tune
|
||||
@@ -179,15 +177,6 @@ async fn objective(
|
||||
// Helper Functions
|
||||
// ============================================================================
|
||||
|
||||
/// Formats a parameter value for display.
|
||||
fn format_param(value: &ParamValue) -> String {
|
||||
match value {
|
||||
ParamValue::Float(v) => format!("{v:.4}"),
|
||||
ParamValue::Int(v) => format!("{v}"),
|
||||
ParamValue::Categorical(idx) => format!("category_{idx}"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Prints the results of the optimization.
|
||||
fn print_results(study: &Study<f64>, elapsed: Duration, n_trials: usize) {
|
||||
println!("\n{}", "=".repeat(60));
|
||||
@@ -206,21 +195,46 @@ fn print_results(study: &Study<f64>, elapsed: Duration, n_trials: usize) {
|
||||
}
|
||||
|
||||
/// Prints the best configuration found.
|
||||
fn print_best_config(study: &Study<f64>) -> optimizer::Result<()> {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn print_best_config(
|
||||
study: &Study<f64>,
|
||||
cache_size_mb_param: &IntParam,
|
||||
connection_pool_size_param: &IntParam,
|
||||
request_timeout_ms_param: &IntParam,
|
||||
retry_count_param: &IntParam,
|
||||
batch_size_param: &IntParam,
|
||||
compression_level_param: &IntParam,
|
||||
use_http2_param: &BoolParam,
|
||||
load_balancing_param: &CategoricalParam<&str>,
|
||||
) -> optimizer::Result<()> {
|
||||
let best = study.best_trial()?;
|
||||
|
||||
println!("\nBest configuration found:");
|
||||
println!(" Score: {:.6}", best.value);
|
||||
println!("\n Parameters:");
|
||||
|
||||
for (id, value) in &best.params {
|
||||
let label = best
|
||||
.param_labels
|
||||
.get(id)
|
||||
.map_or_else(|| format!("{id}"), |l| l.clone());
|
||||
let display = format_param(value);
|
||||
println!(" {label}: {display}");
|
||||
}
|
||||
println!(
|
||||
" cache_size_mb: {}",
|
||||
best.get(cache_size_mb_param).unwrap()
|
||||
);
|
||||
println!(
|
||||
" connection_pool_size: {}",
|
||||
best.get(connection_pool_size_param).unwrap()
|
||||
);
|
||||
println!(
|
||||
" request_timeout_ms: {}",
|
||||
best.get(request_timeout_ms_param).unwrap()
|
||||
);
|
||||
println!(" retry_count: {}", best.get(retry_count_param).unwrap());
|
||||
println!(" batch_size: {}", best.get(batch_size_param).unwrap());
|
||||
println!(
|
||||
" compression_level: {}",
|
||||
best.get(compression_level_param).unwrap()
|
||||
);
|
||||
println!(" use_http2: {}", best.get(use_http2_param).unwrap());
|
||||
println!(
|
||||
" load_balancing: {}",
|
||||
best.get(load_balancing_param).unwrap()
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -262,19 +276,32 @@ async fn main() -> optimizer::Result<()> {
|
||||
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).step(64);
|
||||
let connection_pool_size_param = IntParam::new(10, 200).step(10);
|
||||
let request_timeout_ms_param = IntParam::new(1000, 10000).step(500);
|
||||
let retry_count_param = IntParam::new(0, 5);
|
||||
let batch_size_param = IntParam::new(1, 256).log_scale();
|
||||
let compression_level_param = IntParam::new(0, 9);
|
||||
let use_http2_param = BoolParam::new();
|
||||
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;
|
||||
@@ -286,16 +313,15 @@ async fn main() -> optimizer::Result<()> {
|
||||
|
||||
// Step 5: Run parallel async optimization
|
||||
//
|
||||
// optimize_parallel_with_sampler:
|
||||
// optimize_parallel:
|
||||
// - Runs up to `concurrency` trials simultaneously
|
||||
// - Each trial calls the objective function
|
||||
// - Uses a semaphore to limit concurrent evaluations
|
||||
// - Collects results as trials complete
|
||||
//
|
||||
// The "_with_sampler" suffix means the TPE sampler gets access to
|
||||
// trial history for informed sampling.
|
||||
// The sampler gets access to trial history for informed sampling.
|
||||
study
|
||||
.optimize_parallel_with_sampler(n_trials, concurrency, move |trial| {
|
||||
.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();
|
||||
@@ -325,7 +351,17 @@ async fn main() -> optimizer::Result<()> {
|
||||
|
||||
// Step 5: Print results
|
||||
print_results(&study, elapsed, n_trials);
|
||||
print_best_config(&study)?;
|
||||
print_best_config(
|
||||
&study,
|
||||
&cache_size_mb_p,
|
||||
&connection_pool_size_p,
|
||||
&request_timeout_ms_p,
|
||||
&retry_count_p,
|
||||
&batch_size_p,
|
||||
&compression_level_p,
|
||||
&use_http2_p,
|
||||
&load_balancing_p,
|
||||
)?;
|
||||
print_top_trials(&study, 5);
|
||||
|
||||
Ok(())
|
||||
|
||||
@@ -23,10 +23,7 @@
|
||||
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
use optimizer::parameter::{FloatParam, IntParam, Parameter};
|
||||
use optimizer::sampler::CompletedTrial;
|
||||
use optimizer::sampler::tpe::TpeSampler;
|
||||
use optimizer::{Direction, ParamValue, Study, Trial};
|
||||
use optimizer::prelude::*;
|
||||
|
||||
// ============================================================================
|
||||
// Configuration: Hyperparameters we want to tune
|
||||
@@ -150,11 +147,7 @@ fn on_trial_complete(study: &Study<f64>, trial: &CompletedTrial<f64>) -> Control
|
||||
// Print trial number and objective value
|
||||
print!("{:>5} ", study.n_trials());
|
||||
for value in trial.params.values() {
|
||||
match value {
|
||||
ParamValue::Float(v) => print!("{v:>12.5} "),
|
||||
ParamValue::Int(v) => print!("{v:>12} "),
|
||||
ParamValue::Categorical(v) => print!("{v:>12} "),
|
||||
}
|
||||
print!("{value:>12} ");
|
||||
}
|
||||
println!("{:>12.6}", trial.value);
|
||||
|
||||
@@ -204,24 +197,25 @@ fn main() -> optimizer::Result<()> {
|
||||
println!("{}", "-".repeat(60));
|
||||
|
||||
// Step 3: Define parameter search spaces
|
||||
let learning_rate_param = FloatParam::new(0.001, 0.3).log_scale();
|
||||
let max_depth_param = IntParam::new(3, 12);
|
||||
let n_estimators_param = IntParam::new(50, 500).step(50);
|
||||
let subsample_param = FloatParam::new(0.5, 1.0);
|
||||
let colsample_bytree_param = FloatParam::new(0.5, 1.0);
|
||||
let min_child_weight_param = IntParam::new(1, 10);
|
||||
let reg_alpha_param = FloatParam::new(1e-3, 10.0).log_scale();
|
||||
let reg_lambda_param = FloatParam::new(1e-3, 10.0).log_scale();
|
||||
let learning_rate_param = FloatParam::new(0.001, 0.3)
|
||||
.name("learning_rate")
|
||||
.log_scale();
|
||||
let max_depth_param = IntParam::new(3, 12).name("max_depth");
|
||||
let n_estimators_param = IntParam::new(50, 500).name("n_estimators").step(50);
|
||||
let subsample_param = FloatParam::new(0.5, 1.0).name("subsample");
|
||||
let colsample_bytree_param = FloatParam::new(0.5, 1.0).name("colsample_bytree");
|
||||
let min_child_weight_param = IntParam::new(1, 10).name("min_child_weight");
|
||||
let reg_alpha_param = FloatParam::new(1e-3, 10.0).name("reg_alpha").log_scale();
|
||||
let reg_lambda_param = FloatParam::new(1e-3, 10.0).name("reg_lambda").log_scale();
|
||||
|
||||
// Step 4: Run optimization
|
||||
//
|
||||
// optimize_with_callback_sampler runs the objective function for up to
|
||||
// optimize_with_callback runs the objective function for up to
|
||||
// n_trials iterations. After each trial, it calls the callback.
|
||||
// The "_sampler" suffix means the TPE sampler gets access to trial
|
||||
// history for informed sampling.
|
||||
// The sampler gets access to trial history for informed sampling.
|
||||
let n_trials = 50;
|
||||
|
||||
study.optimize_with_callback_sampler(
|
||||
study.optimize_with_callback(
|
||||
n_trials,
|
||||
|trial| {
|
||||
objective(
|
||||
@@ -248,18 +242,29 @@ fn main() -> optimizer::Result<()> {
|
||||
println!("\nBest trial:");
|
||||
println!(" Loss: {:.6}", best.value);
|
||||
println!(" Parameters:");
|
||||
|
||||
for (id, value) in &best.params {
|
||||
let label = best
|
||||
.param_labels
|
||||
.get(id)
|
||||
.map_or_else(|| format!("{id}"), |l| l.clone());
|
||||
match value {
|
||||
ParamValue::Float(v) => println!(" {label}: {v:.6}"),
|
||||
ParamValue::Int(v) => println!(" {label}: {v}"),
|
||||
ParamValue::Categorical(v) => println!(" {label}: category {v}"),
|
||||
}
|
||||
}
|
||||
println!(
|
||||
" learning_rate: {:.6}",
|
||||
best.get(&learning_rate_param).unwrap()
|
||||
);
|
||||
println!(" max_depth: {}", best.get(&max_depth_param).unwrap());
|
||||
println!(
|
||||
" n_estimators: {}",
|
||||
best.get(&n_estimators_param).unwrap()
|
||||
);
|
||||
println!(" subsample: {:.6}", best.get(&subsample_param).unwrap());
|
||||
println!(
|
||||
" colsample_bytree: {:.6}",
|
||||
best.get(&colsample_bytree_param).unwrap()
|
||||
);
|
||||
println!(
|
||||
" min_child_weight: {}",
|
||||
best.get(&min_child_weight_param).unwrap()
|
||||
);
|
||||
println!(" reg_alpha: {:.6}", best.get(®_alpha_param).unwrap());
|
||||
println!(
|
||||
" reg_lambda: {:.6}",
|
||||
best.get(®_lambda_param).unwrap()
|
||||
);
|
||||
|
||||
// Step 5: Use the best parameters (in a real app)
|
||||
//
|
||||
|
||||
+16
-15
@@ -1,7 +1,4 @@
|
||||
use optimizer::parameter::{
|
||||
BoolParam, CategoricalParam, EnumParam, FloatParam, IntParam, Parameter,
|
||||
};
|
||||
use optimizer::{Direction, Study};
|
||||
use optimizer::prelude::*;
|
||||
use optimizer_derive::Categorical;
|
||||
|
||||
#[derive(Clone, Debug, Categorical)]
|
||||
@@ -15,13 +12,13 @@ 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).log_scale();
|
||||
let n_layers_param = IntParam::new(1, 5);
|
||||
let units_param = IntParam::new(32, 512).step(32);
|
||||
let optimizer_param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]);
|
||||
let activation_param = EnumParam::<Activation>::new();
|
||||
let batch_size_param = IntParam::new(16, 256).log_scale();
|
||||
let use_dropout_param = BoolParam::new();
|
||||
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| {
|
||||
@@ -43,13 +40,17 @@ fn main() {
|
||||
trial.id()
|
||||
);
|
||||
|
||||
Ok::<_, optimizer::Error>(loss)
|
||||
Ok::<_, Error>(loss)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
println!("\nBest trial: value={:.4}", best.value);
|
||||
for (id, label) in &best.param_labels {
|
||||
println!(" {}: {:?}", label, best.params[id]);
|
||||
}
|
||||
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());
|
||||
}
|
||||
|
||||
+33
-11
@@ -28,31 +28,26 @@
|
||||
//! # Quick Start
|
||||
//!
|
||||
//! ```
|
||||
//! use optimizer::parameter::{FloatParam, Parameter};
|
||||
//! use optimizer::sampler::tpe::TpeSampler;
|
||||
//! use optimizer::{Direction, Study};
|
||||
//! use optimizer::prelude::*;
|
||||
//!
|
||||
//! // Create a study with TPE sampler
|
||||
//! let sampler = TpeSampler::builder().seed(42).build().unwrap();
|
||||
//! let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
||||
//!
|
||||
//! // Define parameter search space
|
||||
//! let x_param = FloatParam::new(-10.0, 10.0);
|
||||
//! let x = FloatParam::new(-10.0, 10.0).name("x");
|
||||
//!
|
||||
//! // Optimize x^2 for 20 trials
|
||||
//! study
|
||||
//! .optimize_with_sampler(20, |trial| {
|
||||
//! let x = x_param.suggest(trial)?;
|
||||
//! Ok::<_, optimizer::Error>(x * x)
|
||||
//! .optimize(20, |trial| {
|
||||
//! let x_val = x.suggest(trial)?;
|
||||
//! Ok::<_, Error>(x_val * x_val)
|
||||
//! })
|
||||
//! .unwrap();
|
||||
//!
|
||||
//! // Get the best result
|
||||
//! let best = study.best_trial().unwrap();
|
||||
//! println!("Best value: {}", best.value);
|
||||
//! for (id, label) in &best.param_labels {
|
||||
//! println!(" {}: {:?}", label, best.params[id]);
|
||||
//! }
|
||||
//! println!("x = {}", best.get(&x).unwrap());
|
||||
//! ```
|
||||
//!
|
||||
//! # Creating a Study
|
||||
@@ -205,6 +200,33 @@ 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 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;
|
||||
}
|
||||
|
||||
@@ -14,3 +14,13 @@ pub enum ParamValue {
|
||||
/// A categorical parameter value, stored as an index into the choices array.
|
||||
Categorical(usize),
|
||||
}
|
||||
|
||||
impl core::fmt::Display for ParamValue {
|
||||
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
|
||||
match self {
|
||||
Self::Float(v) => write!(f, "{v}"),
|
||||
Self::Int(v) => write!(f, "{v}"),
|
||||
Self::Categorical(v) => write!(f, "category({v})"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+83
-1
@@ -145,6 +145,7 @@ pub struct FloatParam {
|
||||
high: f64,
|
||||
log_scale: bool,
|
||||
step: Option<f64>,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
impl FloatParam {
|
||||
@@ -157,6 +158,7 @@ impl FloatParam {
|
||||
high,
|
||||
log_scale: false,
|
||||
step: None,
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -173,6 +175,16 @@ impl FloatParam {
|
||||
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 {
|
||||
@@ -217,6 +229,10 @@ impl Parameter for FloatParam {
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn label(&self) -> String {
|
||||
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
/// An integer parameter with optional log-scale and step size.
|
||||
@@ -248,6 +264,7 @@ pub struct IntParam {
|
||||
high: i64,
|
||||
log_scale: bool,
|
||||
step: Option<i64>,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
impl IntParam {
|
||||
@@ -260,6 +277,7 @@ impl IntParam {
|
||||
high,
|
||||
log_scale: false,
|
||||
step: None,
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -276,6 +294,16 @@ impl IntParam {
|
||||
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 {
|
||||
@@ -320,6 +348,10 @@ impl Parameter for IntParam {
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn label(&self) -> String {
|
||||
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
/// A categorical parameter that selects from a list of choices.
|
||||
@@ -339,6 +371,7 @@ impl Parameter for IntParam {
|
||||
pub struct CategoricalParam<T: Clone> {
|
||||
id: ParamId,
|
||||
choices: Vec<T>,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
impl<T: Clone> CategoricalParam<T> {
|
||||
@@ -348,8 +381,19 @@ impl<T: Clone> CategoricalParam<T> {
|
||||
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> {
|
||||
@@ -380,6 +424,10 @@ impl<T: Clone + Debug> Parameter for CategoricalParam<T> {
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn label(&self) -> String {
|
||||
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
/// A boolean parameter (equivalent to a categorical with `[false, true]`).
|
||||
@@ -396,13 +444,27 @@ impl<T: Clone + Debug> Parameter for CategoricalParam<T> {
|
||||
#[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() }
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -431,6 +493,10 @@ impl Parameter for BoolParam {
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn label(&self) -> String {
|
||||
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
/// A trait for enum types that can be used as categorical parameters.
|
||||
@@ -529,6 +595,7 @@ pub trait Categorical: Sized + Clone {
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct EnumParam<T: Categorical> {
|
||||
id: ParamId,
|
||||
name: Option<String>,
|
||||
_marker: core::marker::PhantomData<T>,
|
||||
}
|
||||
|
||||
@@ -538,9 +605,20 @@ impl<T: Categorical> EnumParam<T> {
|
||||
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> {
|
||||
@@ -570,6 +648,10 @@ impl<T: Categorical + Debug> Parameter for EnumParam<T> {
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn label(&self) -> String {
|
||||
self.name.clone().unwrap_or_else(|| format!("{self:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
+42
-1
@@ -8,7 +8,7 @@ use std::collections::HashMap;
|
||||
|
||||
use crate::distribution::Distribution;
|
||||
use crate::param::ParamValue;
|
||||
use crate::parameter::ParamId;
|
||||
use crate::parameter::{ParamId, Parameter};
|
||||
|
||||
/// A completed trial with its parameters, distributions, and objective value.
|
||||
///
|
||||
@@ -46,6 +46,47 @@ impl<V> CompletedTrial<V> {
|
||||
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(¶m.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.
|
||||
|
||||
+109
-367
@@ -1,5 +1,6 @@
|
||||
//! Study implementation for managing optimization trials.
|
||||
|
||||
use core::any::Any;
|
||||
#[cfg(feature = "async")]
|
||||
use core::future::Future;
|
||||
use core::ops::ControlFlow;
|
||||
@@ -43,6 +44,10 @@ 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>
|
||||
@@ -66,7 +71,10 @@ where
|
||||
/// assert_eq!(study.direction(), Direction::Minimize);
|
||||
/// ```
|
||||
#[must_use]
|
||||
pub fn new(direction: Direction) -> Self {
|
||||
pub fn new(direction: Direction) -> Self
|
||||
where
|
||||
V: 'static,
|
||||
{
|
||||
Self::with_sampler(direction, RandomSampler::new())
|
||||
}
|
||||
|
||||
@@ -87,15 +95,49 @@ 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 {
|
||||
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);
|
||||
|
||||
Self {
|
||||
direction,
|
||||
sampler: Arc::new(sampler),
|
||||
completed_trials: Arc::new(RwLock::new(Vec::new())),
|
||||
sampler,
|
||||
completed_trials,
|
||||
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
|
||||
@@ -116,8 +158,12 @@ 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) {
|
||||
pub fn set_sampler(&mut self, sampler: impl Sampler + 'static)
|
||||
where
|
||||
V: 'static,
|
||||
{
|
||||
self.sampler = Arc::new(sampler);
|
||||
self.trial_factory = Self::make_trial_factory(&self.sampler, &self.completed_trials);
|
||||
}
|
||||
|
||||
/// Generates the next unique trial ID.
|
||||
@@ -131,9 +177,9 @@ where
|
||||
/// parameter values. After the objective function is evaluated, call
|
||||
/// `complete_trial` or `fail_trial` to record the result.
|
||||
///
|
||||
/// 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.
|
||||
/// 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.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -149,7 +195,11 @@ where
|
||||
/// ```
|
||||
pub fn create_trial(&self) -> Trial {
|
||||
let id = self.next_trial_id();
|
||||
Trial::new(id)
|
||||
if let Some(factory) = &self.trial_factory {
|
||||
factory(id)
|
||||
} else {
|
||||
Trial::new(id)
|
||||
}
|
||||
}
|
||||
|
||||
/// Records a completed trial with its objective value.
|
||||
@@ -763,270 +813,72 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
// Specialized implementation for Study<f64> that provides full sampler integration.
|
||||
// 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)]
|
||||
impl Study<f64> {
|
||||
/// Creates a new trial with sampler integration.
|
||||
/// Deprecated: use `create_trial()` instead.
|
||||
///
|
||||
/// 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::parameter::{FloatParam, Parameter};
|
||||
/// 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_param = FloatParam::new(0.0, 1.0);
|
||||
/// let x = x_param.suggest(&mut trial).unwrap();
|
||||
/// ```
|
||||
/// 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>"
|
||||
)]
|
||||
pub fn create_trial_with_sampler(&self) -> Trial {
|
||||
let id = self.next_trial_id();
|
||||
Trial::with_sampler(
|
||||
id,
|
||||
Arc::clone(&self.sampler),
|
||||
Arc::clone(&self.completed_trials),
|
||||
)
|
||||
self.create_trial()
|
||||
}
|
||||
|
||||
/// Runs optimization with full sampler integration.
|
||||
/// Deprecated: use `optimize()` instead.
|
||||
///
|
||||
/// 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::parameter::{FloatParam, Parameter};
|
||||
/// 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);
|
||||
///
|
||||
/// let x_param = FloatParam::new(-10.0, 10.0);
|
||||
///
|
||||
/// study
|
||||
/// .optimize_with_sampler(10, |trial| {
|
||||
/// let x = x_param.suggest(trial)?;
|
||||
/// 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<()>
|
||||
/// 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<()>
|
||||
where
|
||||
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
|
||||
E: ToString,
|
||||
{
|
||||
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(())
|
||||
self.optimize(n_trials, objective)
|
||||
}
|
||||
|
||||
/// Runs optimization with a callback and full sampler integration.
|
||||
/// Deprecated: use `optimize_with_callback()` instead.
|
||||
///
|
||||
/// 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::parameter::{FloatParam, Parameter};
|
||||
/// 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);
|
||||
///
|
||||
/// let x_param = FloatParam::new(-10.0, 10.0);
|
||||
///
|
||||
/// study
|
||||
/// .optimize_with_callback_sampler(
|
||||
/// 100,
|
||||
/// |trial| {
|
||||
/// let x = x_param.suggest(trial)?;
|
||||
/// 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);
|
||||
/// ```
|
||||
/// 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>"
|
||||
)]
|
||||
pub fn optimize_with_callback_sampler<F, C, E>(
|
||||
&self,
|
||||
n_trials: usize,
|
||||
mut objective: F,
|
||||
mut callback: C,
|
||||
objective: F,
|
||||
callback: C,
|
||||
) -> crate::Result<()>
|
||||
where
|
||||
F: FnMut(&mut Trial) -> core::result::Result<f64, E>,
|
||||
C: FnMut(&Study<f64>, &CompletedTrial<f64>) -> ControlFlow<()>,
|
||||
E: ToString,
|
||||
{
|
||||
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(())
|
||||
self.optimize_with_callback(n_trials, objective, callback)
|
||||
}
|
||||
|
||||
/// Runs optimization asynchronously with full sampler integration.
|
||||
/// Deprecated: use `optimize_async()` instead.
|
||||
///
|
||||
/// 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::parameter::{FloatParam, Parameter};
|
||||
/// 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);
|
||||
///
|
||||
/// let x_param = FloatParam::new(-10.0, 10.0);
|
||||
///
|
||||
/// study
|
||||
/// .optimize_async_with_sampler(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))
|
||||
/// }
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
/// // At least one trial should have completed
|
||||
/// assert!(study.n_trials() > 0);
|
||||
/// # Ok(())
|
||||
/// # }
|
||||
/// ```
|
||||
/// The generic `optimize_async()` now automatically integrates with the sampler
|
||||
/// for `Study<f64>`.
|
||||
#[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,
|
||||
@@ -1037,84 +889,18 @@ impl Study<f64> {
|
||||
Fut: Future<Output = core::result::Result<(Trial, f64), E>>,
|
||||
E: ToString,
|
||||
{
|
||||
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(())
|
||||
self.optimize_async(n_trials, objective).await
|
||||
}
|
||||
|
||||
/// Runs optimization with bounded parallelism and full sampler integration.
|
||||
/// Deprecated: use `optimize_parallel()` instead.
|
||||
///
|
||||
/// 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::parameter::{FloatParam, Parameter};
|
||||
/// 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);
|
||||
///
|
||||
/// let x_param = FloatParam::new(-10.0, 10.0);
|
||||
///
|
||||
/// study
|
||||
/// .optimize_parallel_with_sampler(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))
|
||||
/// }
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
/// // All trials should have completed
|
||||
/// assert_eq!(study.n_trials(), 10);
|
||||
/// # Ok(())
|
||||
/// # }
|
||||
/// ```
|
||||
/// The generic `optimize_parallel()` now automatically integrates with the
|
||||
/// sampler for `Study<f64>`.
|
||||
#[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,
|
||||
@@ -1126,51 +912,7 @@ impl Study<f64> {
|
||||
Fut: Future<Output = core::result::Result<(Trial, f64), E>> + Send,
|
||||
E: ToString + Send + 'static,
|
||||
{
|
||||
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(())
|
||||
self.optimize_parallel(n_trials, concurrency, objective)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ async fn test_optimize_async_basic() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_optimize_async_with_sampler() {
|
||||
async fn test_optimize_async_with_tpe() {
|
||||
let sampler = TpeSampler::builder()
|
||||
.seed(42)
|
||||
.n_startup_trials(5)
|
||||
@@ -45,7 +45,7 @@ async fn test_optimize_async_with_sampler() {
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_async_with_sampler(15, move |mut trial| {
|
||||
.optimize_async(15, move |mut trial| {
|
||||
let x_param = x_param.clone();
|
||||
async move {
|
||||
let x = x_param.suggest(&mut trial)?;
|
||||
@@ -82,7 +82,7 @@ async fn test_optimize_parallel() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_optimize_parallel_with_sampler() {
|
||||
async fn test_optimize_parallel_with_tpe() {
|
||||
let sampler = TpeSampler::builder()
|
||||
.seed(42)
|
||||
.n_startup_trials(5)
|
||||
@@ -95,7 +95,7 @@ async fn test_optimize_parallel_with_sampler() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_parallel_with_sampler(15, 3, move |mut trial| {
|
||||
.optimize_parallel(15, 3, move |mut trial| {
|
||||
let x_param = x_param.clone();
|
||||
let y_param = y_param.clone();
|
||||
async move {
|
||||
@@ -128,6 +128,7 @@ async fn test_optimize_async_all_failures() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[allow(deprecated)]
|
||||
async fn test_optimize_async_with_sampler_all_failures() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
@@ -162,6 +163,7 @@ async fn test_optimize_parallel_all_failures() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[allow(deprecated)]
|
||||
async fn test_optimize_parallel_with_sampler_all_failures() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
|
||||
+60
-21
@@ -31,7 +31,7 @@ fn test_tpe_optimizes_quadratic_function() {
|
||||
let x_param = FloatParam::new(-10.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>((x - 3.0).powi(2))
|
||||
})
|
||||
@@ -64,7 +64,7 @@ fn test_tpe_optimizes_multivariate_function() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(100, |trial| {
|
||||
.optimize(100, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x + y * y)
|
||||
@@ -96,7 +96,7 @@ fn test_tpe_maximization() {
|
||||
let x_param = FloatParam::new(-10.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(-(x - 2.0).powi(2) + 10.0)
|
||||
})
|
||||
@@ -231,7 +231,7 @@ fn test_random_sampler_reproducibility() {
|
||||
let x_param2 = FloatParam::new(0.0, 100.0);
|
||||
|
||||
study1
|
||||
.optimize_with_sampler(100, |trial| {
|
||||
.optimize(100, |trial| {
|
||||
let x = x_param1.suggest(trial)?;
|
||||
values1.push(x);
|
||||
Ok::<_, Error>(x)
|
||||
@@ -239,7 +239,7 @@ fn test_random_sampler_reproducibility() {
|
||||
.unwrap();
|
||||
|
||||
study2
|
||||
.optimize_with_sampler(100, |trial| {
|
||||
.optimize(100, |trial| {
|
||||
let x = x_param2.suggest(trial)?;
|
||||
values2.push(x);
|
||||
Ok::<_, Error>(x)
|
||||
@@ -516,7 +516,7 @@ fn test_tpe_with_categorical_parameter() {
|
||||
|
||||
// Optimization where the best choice depends on the categorical
|
||||
study
|
||||
.optimize_with_sampler(30, |trial| {
|
||||
.optimize(30, |trial| {
|
||||
let choice = model_param.suggest(trial)?;
|
||||
let x = x_param.suggest(trial)?;
|
||||
|
||||
@@ -553,7 +553,7 @@ fn test_tpe_with_integer_parameters() {
|
||||
|
||||
// Minimize (n - 7)^2 where n in [1, 10]
|
||||
study
|
||||
.optimize_with_sampler(30, |trial| {
|
||||
.optimize(30, |trial| {
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>(((n - 7) as f64).powi(2))
|
||||
})
|
||||
@@ -730,7 +730,7 @@ fn test_study_set_sampler() {
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(10, |trial| {
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x)
|
||||
})
|
||||
@@ -787,6 +787,7 @@ fn test_optimize_with_callback_all_trials_fail() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[allow(deprecated)]
|
||||
fn test_optimize_with_sampler_all_trials_fail() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
|
||||
@@ -799,6 +800,7 @@ fn test_optimize_with_sampler_all_trials_fail() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[allow(deprecated)]
|
||||
fn test_optimize_with_callback_sampler_all_trials_fail() {
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
@@ -840,7 +842,7 @@ fn test_tpe_sampler_builder_default_trait() {
|
||||
let x_param = FloatParam::new(0.0, 1.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(5, |trial| {
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
@@ -857,7 +859,7 @@ fn test_tpe_sampler_default_trait() {
|
||||
let x_param = FloatParam::new(0.0, 1.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(5, |trial| {
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
@@ -879,7 +881,7 @@ fn test_tpe_with_fixed_kde_bandwidth() {
|
||||
let x_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(20, |trial| {
|
||||
.optimize(20, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x)
|
||||
})
|
||||
@@ -907,7 +909,7 @@ fn test_tpe_split_trials_with_two_trials() {
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(5, |trial| {
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
@@ -928,7 +930,7 @@ fn test_tpe_with_log_scale_int() {
|
||||
let batch_param = IntParam::new(1, 1024).log_scale();
|
||||
|
||||
study
|
||||
.optimize_with_sampler(20, |trial| {
|
||||
.optimize(20, |trial| {
|
||||
let batch_size = batch_param.suggest(trial)?;
|
||||
Ok::<_, Error>(((batch_size as f64).log2() - 5.0).powi(2))
|
||||
})
|
||||
@@ -951,7 +953,7 @@ fn test_tpe_with_step_distributions() {
|
||||
let n_param = IntParam::new(0, 100).step(10);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(20, |trial| {
|
||||
.optimize(20, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>((x - 5.0).powi(2) + ((n - 50) as f64).powi(2))
|
||||
@@ -963,15 +965,16 @@ fn test_tpe_with_step_distributions() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[allow(deprecated)]
|
||||
fn test_create_trial_vs_create_trial_with_sampler() {
|
||||
let sampler = RandomSampler::with_seed(42);
|
||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
||||
|
||||
// create_trial() creates trial without sampler integration
|
||||
// create_trial() creates trial with sampler integration for Study<f64>
|
||||
let trial1 = study.create_trial();
|
||||
assert_eq!(trial1.id(), 0);
|
||||
|
||||
// create_trial_with_sampler() creates trial with sampler
|
||||
// create_trial_with_sampler() is deprecated but still works
|
||||
let trial2 = study.create_trial_with_sampler();
|
||||
assert_eq!(trial2.id(), 1);
|
||||
|
||||
@@ -1034,7 +1037,7 @@ fn test_tpe_empty_good_or_bad_values_fallback() {
|
||||
|
||||
// First optimize with one parameter
|
||||
study
|
||||
.optimize_with_sampler(10, |trial| {
|
||||
.optimize(10, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x)
|
||||
})
|
||||
@@ -1042,7 +1045,7 @@ fn test_tpe_empty_good_or_bad_values_fallback() {
|
||||
|
||||
// Now try with a different parameter - TPE won't have history for "y"
|
||||
study
|
||||
.optimize_with_sampler(5, |trial| {
|
||||
.optimize(5, |trial| {
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(y)
|
||||
})
|
||||
@@ -1084,7 +1087,7 @@ fn test_callback_sampler_early_stopping() {
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize_with_callback_sampler(
|
||||
.optimize_with_callback(
|
||||
100,
|
||||
|trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
@@ -1197,7 +1200,7 @@ fn test_suggest_bool_with_tpe() {
|
||||
let x_param = FloatParam::new(0.0, 10.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(20, |trial| {
|
||||
.optimize(20, |trial| {
|
||||
let use_large = use_large_param.suggest(trial)?;
|
||||
let x = x_param.suggest(trial)?;
|
||||
// The value depends on use_large flag
|
||||
@@ -1293,7 +1296,7 @@ fn test_params_with_tpe() {
|
||||
let n_param = IntParam::new(1, 10);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(30, |trial| {
|
||||
.optimize(30, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x + (n as f64 - 5.0).powi(2))
|
||||
@@ -1324,3 +1327,39 @@ fn test_single_value_float_range() {
|
||||
"single-value range should return that value"
|
||||
);
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Tests for new API features
|
||||
// =============================================================================
|
||||
|
||||
#[test]
|
||||
fn test_param_name() {
|
||||
let param = FloatParam::new(0.0, 1.0).name("learning_rate");
|
||||
let mut trial = Trial::new(0);
|
||||
param.suggest(&mut trial).unwrap();
|
||||
|
||||
let labels = trial.param_labels();
|
||||
let label = labels.values().next().unwrap();
|
||||
assert_eq!(label, "learning_rate");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_completed_trial_get() {
|
||||
let study: Study<f64> = Study::new(Direction::Minimize);
|
||||
let x_param = FloatParam::new(-10.0, 10.0).name("x");
|
||||
let n_param = IntParam::new(1, 10).name("n");
|
||||
|
||||
study
|
||||
.optimize(5, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
Ok::<_, Error>(x * x + n as f64)
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
let best = study.best_trial().unwrap();
|
||||
let x_val: f64 = best.get(&x_param).unwrap();
|
||||
let n_val: i64 = best.get(&n_param).unwrap();
|
||||
assert!((-10.0..=10.0).contains(&x_val));
|
||||
assert!((1..=10).contains(&n_val));
|
||||
}
|
||||
|
||||
@@ -55,7 +55,7 @@ fn test_multivariate_tpe_rosenbrock_finds_good_solution() {
|
||||
let y_param = FloatParam::new(-2.0, 4.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(100, |trial| {
|
||||
.optimize(100, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(rosenbrock(x, y))
|
||||
@@ -90,7 +90,7 @@ fn test_independent_tpe_rosenbrock() {
|
||||
let y_param = FloatParam::new(-2.0, 4.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(100, |trial| {
|
||||
.optimize(100, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(rosenbrock(x, y))
|
||||
@@ -133,7 +133,7 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() {
|
||||
let y_param = FloatParam::new(-2.0, 4.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(n_trials, |trial| {
|
||||
.optimize(n_trials, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(rosenbrock(x, y))
|
||||
@@ -156,7 +156,7 @@ fn test_multivariate_tpe_outperforms_on_correlated_problem() {
|
||||
let y_param = FloatParam::new(-2.0, 4.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(n_trials, |trial| {
|
||||
.optimize(n_trials, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(rosenbrock(x, y))
|
||||
@@ -218,7 +218,7 @@ fn test_multivariate_tpe_independent_problem() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(sphere(x, y))
|
||||
@@ -250,7 +250,7 @@ fn test_independent_tpe_independent_problem() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(sphere(x, y))
|
||||
@@ -290,7 +290,7 @@ fn test_both_samplers_work_on_independent_problem() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(n_trials, |trial| {
|
||||
.optimize(n_trials, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(sphere(x, y))
|
||||
@@ -312,7 +312,7 @@ fn test_both_samplers_work_on_independent_problem() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(n_trials, |trial| {
|
||||
.optimize(n_trials, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(sphere(x, y))
|
||||
@@ -360,7 +360,7 @@ fn test_multivariate_tpe_with_group_decomposition() {
|
||||
let y_param = FloatParam::new(-5.0, 5.0);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let y = y_param.suggest(trial)?;
|
||||
Ok::<_, Error>(sphere(x, y))
|
||||
@@ -396,7 +396,7 @@ fn test_multivariate_tpe_mixed_parameter_types() {
|
||||
let mode_param = CategoricalParam::new(vec!["a", "b", "c"]);
|
||||
|
||||
study
|
||||
.optimize_with_sampler(50, |trial| {
|
||||
.optimize(50, |trial| {
|
||||
let x = x_param.suggest(trial)?;
|
||||
let n = n_param.suggest(trial)?;
|
||||
let mode = mode_param.suggest(trial)?;
|
||||
|
||||
Reference in New Issue
Block a user