Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 12 additions & 16 deletions rustuna_pyo3/src/sampler/nsgaii.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use pyo3::Py;

use rustuna_core::sampler::Sampler;
use rustuna_core::trial::TrialStateValues;
use rustuna_sampler::nsgaii::NSGAIISampler;
use rustuna_sampler::nsgaii::{NSGAIISampler, NsgaiiBuilder};

use crate::distribution::PyDistribution;
use crate::sampler::{extract_storage, PySamplerContext};
Expand All @@ -30,21 +30,17 @@ impl PyNSGAIISampler {
crossover_prob: f64,
swapping_prob: f64,
) -> PyResult<Self> {
let rs_sampler = match seed {
Some(seed) => NSGAIISampler::seed_from_u64(
seed,
population_size,
mutation_prob,
crossover_prob,
swapping_prob,
),
None => NSGAIISampler::new(
population_size,
mutation_prob,
crossover_prob,
swapping_prob,
),
};
let mut builder = NsgaiiBuilder::new()
.population_size(population_size)
.crossover_prob(crossover_prob)
.swapping_prob(swapping_prob);
if let Some(seed) = seed {
builder = builder.seed(seed);
}
if let Some(mutation_prob) = mutation_prob {
builder = builder.mutation_prob(mutation_prob);
}
let rs_sampler = builder.build();
Ok(PyNSGAIISampler {
sampler: Arc::new(rs_sampler),
})
Expand Down
15 changes: 9 additions & 6 deletions rustuna_pyo3/src/sampler/tpe.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use pyo3::Py;

use rustuna_core::sampler::Sampler;
use rustuna_core::trial::TrialStateValues;
use rustuna_sampler::tpe::{TpeConfig, TpeSampler};
use rustuna_sampler::tpe::{TpeBuilder, TpeSampler};

use crate::distribution::PyDistribution;
use crate::sampler::{extract_storage, PySamplerContext};
Expand All @@ -28,11 +28,14 @@ impl PyTpeSampler {
n_startup_trials: usize,
multivariate: Option<bool>,
) -> PyResult<Self> {
let rs_sampler = TpeSampler::from_config(TpeConfig {
seed,
n_startup_trials,
multivariate,
});
let mut builder = TpeBuilder::new().n_startup_trials(n_startup_trials);
if let Some(seed) = seed {
builder = builder.seed(seed);
}
if let Some(multivariate) = multivariate {
builder = builder.multivariate(multivariate);
}
let rs_sampler = builder.build();
Ok(PyTpeSampler {
sampler: Arc::new(rs_sampler),
})
Expand Down
121 changes: 107 additions & 14 deletions rustuna_sampler/src/nsgaii.rs
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,98 @@ impl Default for NSGAIISampler {
}
}

/// Builder for [`NSGAIISampler`], following the API style of [`std::thread::Builder`].
///
/// # Examples
///
/// ```
/// use rustuna_sampler::nsgaii::NsgaiiBuilder;
///
/// let sampler = NsgaiiBuilder::new()
/// .population_size(100)
/// .mutation_prob(0.1)
/// .crossover_prob(0.9)
/// .swapping_prob(0.5)
/// .seed(42)
/// .build();
/// ```
pub struct NsgaiiBuilder {
population_size: usize,
mutation_prob: Option<f64>,
crossover_prob: f64,
swapping_prob: f64,
seed: Option<u64>,
}
impl Default for NsgaiiBuilder {
fn default() -> Self {
Self::new()
}
}
impl NsgaiiBuilder {
/// Creates a builder with the default configuration.
pub fn new() -> Self {
Self {
population_size: 50,
mutation_prob: None,
crossover_prob: 0.9,
swapping_prob: 0.5,
seed: None,
}
}

pub fn population_size(self, population_size: usize) -> Self {
Self {
population_size,
..self
}
}

/// Sets the per-parameter mutation probability. When unset, the automatic default
/// `1 / n_params` is used.
pub fn mutation_prob(self, mutation_prob: f64) -> Self {
Self {
mutation_prob: Some(mutation_prob),
..self
}
}

pub fn crossover_prob(self, crossover_prob: f64) -> Self {
Self {
crossover_prob,
..self
}
}

pub fn swapping_prob(self, swapping_prob: f64) -> Self {
Self {
swapping_prob,
..self
}
}

pub fn seed(self, seed: u64) -> Self {
Self {
seed: Some(seed),
..self
}
}

pub fn build(self) -> NSGAIISampler {
let rng = match self.seed {
Some(seed) => StdRng::seed_from_u64(seed),
None => StdRng::from_seed(Default::default()),
};
NSGAIISampler {
rng: Mutex::new(rng),
population_size: self.population_size,
mutation_prob: self.mutation_prob,
crossover_prob: self.crossover_prob,
swapping_prob: self.swapping_prob,
generation_to_numbers: RwLock::new(HashMap::new()),
}
}
}

impl NSGAIISampler {
/// Creates an NSGA-II sampler.
///
Expand All @@ -87,14 +179,14 @@ impl NSGAIISampler {
crossover_prob: f64,
swapping_prob: f64,
) -> NSGAIISampler {
NSGAIISampler {
rng: Mutex::new(StdRng::from_seed(Default::default())),
population_size,
mutation_prob,
crossover_prob,
swapping_prob,
generation_to_numbers: RwLock::new(HashMap::new()),
let mut builder = NsgaiiBuilder::new()
.population_size(population_size)
.crossover_prob(crossover_prob)
.swapping_prob(swapping_prob);
if let Some(mutation_prob) = mutation_prob {
builder = builder.mutation_prob(mutation_prob);
}
builder.build()
}
/// Creates a reproducibly seeded NSGA-II sampler.
///
Expand All @@ -107,14 +199,15 @@ impl NSGAIISampler {
crossover_prob: f64,
swapping_prob: f64,
) -> NSGAIISampler {
NSGAIISampler {
rng: Mutex::new(StdRng::seed_from_u64(seed)),
population_size,
mutation_prob,
crossover_prob,
swapping_prob,
generation_to_numbers: RwLock::new(HashMap::new()),
let mut builder = NsgaiiBuilder::new()
.population_size(population_size)
.crossover_prob(crossover_prob)
.swapping_prob(swapping_prob)
.seed(seed);
if let Some(mutation_prob) = mutation_prob {
builder = builder.mutation_prob(mutation_prob);
}
builder.build()
}
fn get_rng_lock(&self) -> Result<MutexGuard<'_, StdRng>> {
self.rng.lock().map_err(|e| {
Expand Down
2 changes: 1 addition & 1 deletion rustuna_sampler/src/tpe/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,4 +5,4 @@

mod sampler;

pub use sampler::{TpeConfig, TpeSampler};
pub use sampler::{TpeBuilder, TpeSampler};
121 changes: 70 additions & 51 deletions rustuna_sampler/src/tpe/sampler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,27 +14,79 @@ use rustuna_core::trial::TrialStateValues;
use rustuna_core::Result;
use rustuna_core::{Error, ErrorKind};

/// Configuration for [`TpeSampler`].
pub struct TpeConfig {
/// Whether to use multivariate TPE for joint suggestions over the inferred search space.
///
/// `Some(true)` forces multivariate (joint) sampling and `Some(false)` forces independent
/// (univariate) sampling. `None` selects automatically, matching Optuna: multivariate for
/// single-objective studies and independent for multi-objective studies.
pub multivariate: Option<bool>,
/// Number of completed trials to collect before switching from random sampling to TPE.
pub n_startup_trials: usize,
/// Optional RNG seed.
pub seed: Option<u64>,
/// Builder for [`TpeSampler`], following the API style of [`std::thread::Builder`].
///
/// # Examples
///
/// ```
/// use rustuna_sampler::tpe::TpeBuilder;
///
/// let sampler = TpeBuilder::new()
/// .n_startup_trials(20)
/// .multivariate(true)
/// .seed(42)
/// .build();
/// ```
pub struct TpeBuilder {
multivariate: Option<bool>,
n_startup_trials: usize,
seed: Option<u64>,
}
impl Default for TpeConfig {
impl Default for TpeBuilder {
fn default() -> Self {
Self::new()
}
}
impl TpeBuilder {
/// Creates a builder with the default configuration.
pub fn new() -> Self {
Self {
multivariate: None,
n_startup_trials: 10,
seed: None,
}
}

/// Sets whether to force multivariate (joint) sampling. When unset, it is selected
/// automatically (multivariate for single-objective, independent for multi-objective,
/// matching Optuna).
pub fn multivariate(self, multivariate: bool) -> Self {
Self {
multivariate: Some(multivariate),
..self
}
}

/// Sets the number of completed trials before switching from random sampling to TPE.
pub fn n_startup_trials(self, n_startup_trials: usize) -> Self {
Self {
n_startup_trials,
..self
}
}

pub fn seed(self, seed: u64) -> Self {
Self {
seed: Some(seed),
..self
}
}

pub fn build(self) -> TpeSampler {
let mut rng = match self.seed {
Some(seed) => StdRng::seed_from_u64(seed),
None => StdRng::from_seed(Default::default()),
};
let seed_for_random_sampler = rng.gen();
TpeSampler {
rng: Mutex::new(rng),
multivariate: self.multivariate,
n_startup_trials: self.n_startup_trials,
random_sampler: RandomSampler::seed_from_u64(seed_for_random_sampler),
split_cache: RwLock::new(HashMap::new()),
observations_cache: RwLock::new(HashMap::new()),
}
}
}

type SplitKey = (Vec<u32>, usize);
Expand Down Expand Up @@ -178,42 +230,21 @@ impl Default for TpeSampler {
}
}
impl TpeSampler {
/// Creates a sampler from an explicit configuration.
pub fn from_config(cfg: TpeConfig) -> TpeSampler {
let mut rng = match cfg.seed {
Some(s) => StdRng::seed_from_u64(s),
None => StdRng::from_seed(Default::default()),
};
let seed_for_random_sampler = rng.gen();
Self {
rng: Mutex::new(rng),
multivariate: cfg.multivariate,
n_startup_trials: cfg.n_startup_trials,
random_sampler: RandomSampler::seed_from_u64(seed_for_random_sampler),
split_cache: RwLock::new(HashMap::new()),
observations_cache: RwLock::new(HashMap::new()),
}
}

/// Creates a sampler with the default configuration.
///
/// The default configuration selects multivariate TPE automatically (multivariate for
/// single-objective, independent for multi-objective, matching Optuna) and uses random
/// sampling for the first 10 completed trials.
pub fn new() -> TpeSampler {
Self::from_config(TpeConfig::default())
TpeBuilder::new().build()
}

/// Creates a reproducibly seeded sampler.
///
/// This is equivalent to [`TpeSampler::new`] but initializes the internal random number
/// generator from the provided seed.
pub fn seed_from_u64(seed: u64) -> TpeSampler {
Self::from_config(TpeConfig {
multivariate: None,
seed: Some(seed),
n_startup_trials: 10,
})
TpeBuilder::new().seed(seed).build()
}

fn sample(
Expand Down Expand Up @@ -774,11 +805,7 @@ mod tests {
fn test_dynamic_float_range_falls_back_to_independent_sampling() {
let storage = InMemoryStorage::new();
let directions = vec![Direction::Minimize];
let sampler = TpeSampler::from_config(TpeConfig {
multivariate: None,
n_startup_trials: 2,
seed: Some(42),
});
let sampler = TpeBuilder::new().n_startup_trials(2).seed(42).build();
let study = create_study("dynamic-float-range", storage, sampler, directions).unwrap();

study
Expand Down Expand Up @@ -1257,11 +1284,7 @@ mod tests {
vec![Direction::Minimize],
)
.unwrap();
let probe = TpeSampler::from_config(TpeConfig {
multivariate: None,
n_startup_trials: 1,
seed: Some(0),
});
let probe = TpeBuilder::new().n_startup_trials(1).seed(0).build();
let ctx = |trial_id: u32| Context {
study_id: study.id,
directions: vec![Direction::Minimize],
Expand Down Expand Up @@ -1328,11 +1351,7 @@ mod tests {
use rustuna_core::attr::{AttrKey, Attrs};

let storage = InMemoryStorage::new();
let sampler = TpeSampler::from_config(TpeConfig {
multivariate: None,
n_startup_trials: 1,
seed: Some(0),
});
let sampler = TpeBuilder::new().n_startup_trials(1).seed(0).build();
let study = create_study(
"bad-constraints",
storage,
Expand Down
Loading