refactor: reorganize sampler module and update imports

This commit is contained in:
Manuel Raimann
2026-02-04 13:14:22 +01:00
parent 183a78c09e
commit 02ec9498c4
5 changed files with 35 additions and 43 deletions
-5
View File
@@ -1,16 +1,11 @@
//! Sampler trait and implementations for parameter sampling. //! Sampler trait and implementations for parameter sampling.
pub mod grid; pub mod grid;
pub mod multivariate_tpe;
pub mod random; pub mod random;
pub mod tpe; pub mod tpe;
use std::collections::HashMap; use std::collections::HashMap;
pub use multivariate_tpe::{
ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder,
};
use crate::distribution::Distribution; use crate::distribution::Distribution;
use crate::param::ParamValue; use crate::param::ParamValue;
+4
View File
@@ -4,9 +4,13 @@
//! including support for intersection search space calculation. //! including support for intersection search space calculation.
mod gamma; mod gamma;
mod multivariate;
mod sampler; mod sampler;
pub mod search_space; pub mod search_space;
pub use gamma::{FixedGamma, GammaStrategy, HyperoptGamma, LinearGamma, SqrtGamma}; pub use gamma::{FixedGamma, GammaStrategy, HyperoptGamma, LinearGamma, SqrtGamma};
pub use multivariate::{
ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder,
};
pub use sampler::{TpeSampler, TpeSamplerBuilder}; pub use sampler::{TpeSampler, TpeSamplerBuilder};
pub use search_space::{GroupDecomposedSearchSpace, IntersectionSearchSpace}; pub use search_space::{GroupDecomposedSearchSpace, IntersectionSearchSpace};
@@ -1,6 +1,6 @@
//! Multivariate Tree-Parzen Estimator (TPE) sampler implementation. //! Multivariate Tree-Parzen Estimator (TPE) sampler implementation.
//! //!
//! Unlike the standard [`TpeSampler`](super::tpe::TpeSampler) which samples parameters //! Unlike the standard [`TpeSampler`] which samples parameters
//! independently, the `MultivariateTpeSampler` models joint distributions over multiple //! independently, the `MultivariateTpeSampler` models joint distributions over multiple
//! parameters. This allows it to capture correlations between parameters, which can //! parameters. This allows it to capture correlations between parameters, which can
//! significantly improve optimization performance on problems where parameters interact. //! significantly improve optimization performance on problems where parameters interact.
@@ -12,7 +12,7 @@
//! - The optimal value of one parameter depends on another //! - The optimal value of one parameter depends on another
//! - You have a fixed search space across all trials //! - You have a fixed search space across all trials
//! //!
//! Use the standard [`TpeSampler`](super::tpe::TpeSampler) when: //! Use the standard [`TpeSampler`] when:
//! - Parameters are independent //! - Parameters are independent
//! - The search space varies dynamically between trials //! - The search space varies dynamically between trials
//! - You want simpler, faster optimization //! - You want simpler, faster optimization
@@ -68,7 +68,7 @@
//! ## Basic Usage //! ## Basic Usage
//! //!
//! ``` //! ```
//! use optimizer::sampler::MultivariateTpeSampler; //! use optimizer::sampler::tpe::MultivariateTpeSampler;
//! //!
//! let sampler = MultivariateTpeSampler::builder() //! let sampler = MultivariateTpeSampler::builder()
//! .gamma(0.15) //! .gamma(0.15)
@@ -82,8 +82,7 @@
//! ## With Custom Gamma Strategy and Group Decomposition //! ## With Custom Gamma Strategy and Group Decomposition
//! //!
//! ``` //! ```
//! use optimizer::sampler::MultivariateTpeSampler; //! use optimizer::sampler::tpe::{MultivariateTpeSampler, SqrtGamma};
//! use optimizer::sampler::tpe::SqrtGamma;
//! //!
//! let sampler = MultivariateTpeSampler::builder() //! let sampler = MultivariateTpeSampler::builder()
//! .gamma_strategy(SqrtGamma::default()) //! .gamma_strategy(SqrtGamma::default())
@@ -95,7 +94,7 @@
//! ## Parallel Optimization with Constant Liar //! ## Parallel Optimization with Constant Liar
//! //!
//! ``` //! ```
//! use optimizer::sampler::{ConstantLiarStrategy, MultivariateTpeSampler}; //! use optimizer::sampler::tpe::{ConstantLiarStrategy, MultivariateTpeSampler};
//! //!
//! // Use mean imputation for parallel workers //! // Use mean imputation for parallel workers
//! let sampler = MultivariateTpeSampler::builder() //! let sampler = MultivariateTpeSampler::builder()
@@ -108,7 +107,7 @@
//! ## Suppressing Fallback Warnings //! ## Suppressing Fallback Warnings
//! //!
//! ``` //! ```
//! use optimizer::sampler::MultivariateTpeSampler; //! use optimizer::sampler::tpe::MultivariateTpeSampler;
//! //!
//! let sampler = MultivariateTpeSampler::builder().build().unwrap(); //! let sampler = MultivariateTpeSampler::builder().build().unwrap();
//! ``` //! ```
@@ -120,11 +119,11 @@ use parking_lot::Mutex;
use rand::SeedableRng; use rand::SeedableRng;
use rand::rngs::StdRng; use rand::rngs::StdRng;
use super::tpe::{FixedGamma, GammaStrategy}; use super::{FixedGamma, GammaStrategy};
use super::{CompletedTrial, PendingTrial, Sampler};
use crate::distribution::Distribution; use crate::distribution::Distribution;
use crate::error::Result; use crate::error::Result;
use crate::param::ParamValue; use crate::param::ParamValue;
use crate::sampler::{CompletedTrial, PendingTrial, Sampler};
/// Strategy for imputing objective values for pending/running trials during parallel optimization. /// Strategy for imputing objective values for pending/running trials during parallel optimization.
/// ///
@@ -143,7 +142,7 @@ use crate::param::ParamValue;
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::ConstantLiarStrategy; /// use optimizer::sampler::tpe::ConstantLiarStrategy;
/// ///
/// // Use mean imputation for parallel optimization /// // Use mean imputation for parallel optimization
/// let strategy = ConstantLiarStrategy::Mean; /// let strategy = ConstantLiarStrategy::Mean;
@@ -207,7 +206,7 @@ impl MultivariateTpeSampler {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// ///
/// let sampler = MultivariateTpeSampler::new(); /// let sampler = MultivariateTpeSampler::new();
/// ``` /// ```
@@ -229,7 +228,7 @@ impl MultivariateTpeSampler {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// ///
/// let sampler = MultivariateTpeSampler::builder() /// let sampler = MultivariateTpeSampler::builder()
/// .gamma(0.15) /// .gamma(0.15)
@@ -446,7 +445,7 @@ impl MultivariateTpeSampler {
/// ///
/// ```ignore /// ```ignore
/// use std::collections::HashMap; /// use std::collections::HashMap;
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// use optimizer::sampler::tpe::IntersectionSearchSpace; /// use optimizer::sampler::tpe::IntersectionSearchSpace;
/// ///
/// let sampler = MultivariateTpeSampler::new(); /// let sampler = MultivariateTpeSampler::new();
@@ -473,10 +472,6 @@ impl MultivariateTpeSampler {
/// Splits filtered trials into good and bad groups based on the gamma quantile. /// Splits filtered trials into good and bad groups based on the gamma quantile.
/// ///
/// This method is similar to [`TpeSampler::split_trials`](super::tpe::TpeSampler)
/// but operates on pre-filtered trials that have been selected for multivariate
/// sampling.
///
/// The gamma value is computed dynamically using the configured [`GammaStrategy`]. /// The gamma value is computed dynamically using the configured [`GammaStrategy`].
/// Trials are sorted by objective value (ascending for minimization), and the /// Trials are sorted by objective value (ascending for minimization), and the
/// gamma quantile determines the split point. /// gamma quantile determines the split point.
@@ -499,7 +494,7 @@ impl MultivariateTpeSampler {
/// ///
/// ```ignore /// ```ignore
/// use std::collections::HashMap; /// use std::collections::HashMap;
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// use optimizer::sampler::tpe::IntersectionSearchSpace; /// use optimizer::sampler::tpe::IntersectionSearchSpace;
/// ///
/// let sampler = MultivariateTpeSampler::new(); /// let sampler = MultivariateTpeSampler::new();
@@ -593,7 +588,7 @@ impl MultivariateTpeSampler {
/// ///
/// ```ignore /// ```ignore
/// use std::collections::HashMap; /// use std::collections::HashMap;
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// ///
/// let sampler = MultivariateTpeSampler::new(); /// let sampler = MultivariateTpeSampler::new();
/// let trials = vec![/* ... completed trials ... */]; /// let trials = vec![/* ... completed trials ... */];
@@ -661,7 +656,7 @@ impl MultivariateTpeSampler {
/// # Examples /// # Examples
/// ///
/// ```ignore /// ```ignore
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// use optimizer::kde::MultivariateKDE; /// use optimizer::kde::MultivariateKDE;
/// ///
/// let sampler = MultivariateTpeSampler::builder() /// let sampler = MultivariateTpeSampler::builder()
@@ -859,7 +854,7 @@ impl MultivariateTpeSampler {
/// ///
/// ```ignore /// ```ignore
/// use std::collections::HashMap; /// use std::collections::HashMap;
/// use optimizer::sampler::MultivariateTpeSampler; /// use optimizer::sampler::tpe::MultivariateTpeSampler;
/// use optimizer::distribution::{Distribution, FloatDistribution}; /// use optimizer::distribution::{Distribution, FloatDistribution};
/// ///
/// let sampler = MultivariateTpeSampler::builder() /// let sampler = MultivariateTpeSampler::builder()
@@ -925,7 +920,7 @@ impl MultivariateTpeSampler {
) -> HashMap<String, ParamValue> { ) -> HashMap<String, ParamValue> {
use std::collections::HashSet; use std::collections::HashSet;
use crate::sampler::tpe::GroupDecomposedSearchSpace; use super::GroupDecomposedSearchSpace;
// Decompose the search space into independent parameter groups // Decompose the search space into independent parameter groups
let groups = GroupDecomposedSearchSpace::calculate(history); let groups = GroupDecomposedSearchSpace::calculate(history);
@@ -1014,8 +1009,8 @@ impl MultivariateTpeSampler {
history: &[CompletedTrial], history: &[CompletedTrial],
rng: &mut StdRng, rng: &mut StdRng,
) -> HashMap<String, ParamValue> { ) -> HashMap<String, ParamValue> {
use super::IntersectionSearchSpace;
use crate::kde::MultivariateKDE; use crate::kde::MultivariateKDE;
use crate::sampler::tpe::IntersectionSearchSpace;
// Early returns for cases requiring random sampling // Early returns for cases requiring random sampling
if history.len() < self.n_startup_trials { if history.len() < self.n_startup_trials {
@@ -1172,7 +1167,7 @@ impl MultivariateTpeSampler {
/// ///
/// This method is used to sample parameters that are not in the intersection /// This method is used to sample parameters that are not in the intersection
/// search space. It uses independent univariate TPE sampling for each parameter, /// search space. It uses independent univariate TPE sampling for each parameter,
/// similar to the standard [`TpeSampler`](super::tpe::TpeSampler). /// similar to the standard [`TpeSampler`].
/// ///
/// When there isn't enough history for a parameter, falls back to uniform sampling. /// When there isn't enough history for a parameter, falls back to uniform sampling.
#[allow(dead_code)] #[allow(dead_code)]
@@ -1772,7 +1767,7 @@ impl MultivariateTpeSampler {
/// Using a fixed gamma value: /// Using a fixed gamma value:
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .gamma(0.15) /// .gamma(0.15)
@@ -1786,8 +1781,7 @@ impl MultivariateTpeSampler {
/// Using a custom gamma strategy: /// Using a custom gamma strategy:
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::{MultivariateTpeSamplerBuilder, SqrtGamma};
/// use optimizer::sampler::tpe::SqrtGamma;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .gamma_strategy(SqrtGamma::default()) /// .gamma_strategy(SqrtGamma::default())
@@ -1845,7 +1839,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .gamma(0.10) // Use top 10% as "good" trials /// .gamma(0.10) // Use top 10% as "good" trials
@@ -1876,8 +1870,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::{LinearGamma, MultivariateTpeSamplerBuilder, SqrtGamma};
/// use optimizer::sampler::tpe::{LinearGamma, SqrtGamma};
/// ///
/// // Square root strategy (Optuna-style) /// // Square root strategy (Optuna-style)
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
@@ -1911,7 +1904,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .n_startup_trials(20) // Random sample first 20 trials /// .n_startup_trials(20) // Random sample first 20 trials
@@ -1937,7 +1930,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .n_ei_candidates(48) // Evaluate more candidates /// .n_ei_candidates(48) // Evaluate more candidates
@@ -1964,7 +1957,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .group(true) // Enable group decomposition /// .group(true) // Enable group decomposition
@@ -1990,7 +1983,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::{ConstantLiarStrategy, MultivariateTpeSamplerBuilder}; /// use optimizer::sampler::tpe::{ConstantLiarStrategy, MultivariateTpeSamplerBuilder};
/// ///
/// // Use mean imputation for pending trials /// // Use mean imputation for pending trials
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
@@ -2025,7 +2018,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .seed(42) // Reproducible results /// .seed(42) // Reproducible results
@@ -2047,7 +2040,7 @@ impl MultivariateTpeSamplerBuilder {
/// # Examples /// # Examples
/// ///
/// ``` /// ```
/// use optimizer::sampler::MultivariateTpeSamplerBuilder; /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder;
/// ///
/// let sampler = MultivariateTpeSamplerBuilder::new() /// let sampler = MultivariateTpeSamplerBuilder::new()
/// .gamma(0.15) /// .gamma(0.15)
+1
View File
@@ -280,6 +280,7 @@ impl TpeSampler {
clippy::cast_possible_truncation, clippy::cast_possible_truncation,
clippy::cast_sign_loss clippy::cast_sign_loss
)] )]
#[must_use]
fn split_trials<'a>( fn split_trials<'a>(
&self, &self,
history: &'a [CompletedTrial], history: &'a [CompletedTrial],
+1 -2
View File
@@ -9,8 +9,7 @@
clippy::cast_possible_truncation clippy::cast_possible_truncation
)] )]
use optimizer::sampler::MultivariateTpeSampler; use optimizer::sampler::tpe::{MultivariateTpeSampler, TpeSampler};
use optimizer::sampler::tpe::TpeSampler;
use optimizer::{Direction, Error, Study}; use optimizer::{Direction, Error, Study};
// ============================================================================= // =============================================================================