diff --git a/src/sampler/mod.rs b/src/sampler/mod.rs index 0f0b600..487b7f4 100644 --- a/src/sampler/mod.rs +++ b/src/sampler/mod.rs @@ -1,16 +1,11 @@ //! Sampler trait and implementations for parameter sampling. pub mod grid; -pub mod multivariate_tpe; pub mod random; pub mod tpe; use std::collections::HashMap; -pub use multivariate_tpe::{ - ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder, -}; - use crate::distribution::Distribution; use crate::param::ParamValue; diff --git a/src/sampler/tpe/mod.rs b/src/sampler/tpe/mod.rs index 407fa9b..3013111 100644 --- a/src/sampler/tpe/mod.rs +++ b/src/sampler/tpe/mod.rs @@ -4,9 +4,13 @@ //! including support for intersection search space calculation. mod gamma; +mod multivariate; mod sampler; pub mod search_space; pub use gamma::{FixedGamma, GammaStrategy, HyperoptGamma, LinearGamma, SqrtGamma}; +pub use multivariate::{ + ConstantLiarStrategy, MultivariateTpeSampler, MultivariateTpeSamplerBuilder, +}; pub use sampler::{TpeSampler, TpeSamplerBuilder}; pub use search_space::{GroupDecomposedSearchSpace, IntersectionSearchSpace}; diff --git a/src/sampler/multivariate_tpe.rs b/src/sampler/tpe/multivariate.rs similarity index 99% rename from src/sampler/multivariate_tpe.rs rename to src/sampler/tpe/multivariate.rs index e040eb2..7967222 100644 --- a/src/sampler/multivariate_tpe.rs +++ b/src/sampler/tpe/multivariate.rs @@ -1,6 +1,6 @@ //! 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 //! parameters. This allows it to capture correlations between parameters, which can //! significantly improve optimization performance on problems where parameters interact. @@ -12,7 +12,7 @@ //! - The optimal value of one parameter depends on another //! - 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 //! - The search space varies dynamically between trials //! - You want simpler, faster optimization @@ -68,7 +68,7 @@ //! ## Basic Usage //! //! ``` -//! use optimizer::sampler::MultivariateTpeSampler; +//! use optimizer::sampler::tpe::MultivariateTpeSampler; //! //! let sampler = MultivariateTpeSampler::builder() //! .gamma(0.15) @@ -82,8 +82,7 @@ //! ## With Custom Gamma Strategy and Group Decomposition //! //! ``` -//! use optimizer::sampler::MultivariateTpeSampler; -//! use optimizer::sampler::tpe::SqrtGamma; +//! use optimizer::sampler::tpe::{MultivariateTpeSampler, SqrtGamma}; //! //! let sampler = MultivariateTpeSampler::builder() //! .gamma_strategy(SqrtGamma::default()) @@ -95,7 +94,7 @@ //! ## Parallel Optimization with Constant Liar //! //! ``` -//! use optimizer::sampler::{ConstantLiarStrategy, MultivariateTpeSampler}; +//! use optimizer::sampler::tpe::{ConstantLiarStrategy, MultivariateTpeSampler}; //! //! // Use mean imputation for parallel workers //! let sampler = MultivariateTpeSampler::builder() @@ -108,7 +107,7 @@ //! ## Suppressing Fallback Warnings //! //! ``` -//! use optimizer::sampler::MultivariateTpeSampler; +//! use optimizer::sampler::tpe::MultivariateTpeSampler; //! //! let sampler = MultivariateTpeSampler::builder().build().unwrap(); //! ``` @@ -120,11 +119,11 @@ use parking_lot::Mutex; use rand::SeedableRng; use rand::rngs::StdRng; -use super::tpe::{FixedGamma, GammaStrategy}; -use super::{CompletedTrial, PendingTrial, Sampler}; +use super::{FixedGamma, GammaStrategy}; use crate::distribution::Distribution; use crate::error::Result; use crate::param::ParamValue; +use crate::sampler::{CompletedTrial, PendingTrial, Sampler}; /// Strategy for imputing objective values for pending/running trials during parallel optimization. /// @@ -143,7 +142,7 @@ use crate::param::ParamValue; /// # Examples /// /// ``` -/// use optimizer::sampler::ConstantLiarStrategy; +/// use optimizer::sampler::tpe::ConstantLiarStrategy; /// /// // Use mean imputation for parallel optimization /// let strategy = ConstantLiarStrategy::Mean; @@ -207,7 +206,7 @@ impl MultivariateTpeSampler { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// /// let sampler = MultivariateTpeSampler::new(); /// ``` @@ -229,7 +228,7 @@ impl MultivariateTpeSampler { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// /// let sampler = MultivariateTpeSampler::builder() /// .gamma(0.15) @@ -446,7 +445,7 @@ impl MultivariateTpeSampler { /// /// ```ignore /// use std::collections::HashMap; - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// use optimizer::sampler::tpe::IntersectionSearchSpace; /// /// let sampler = MultivariateTpeSampler::new(); @@ -473,10 +472,6 @@ impl MultivariateTpeSampler { /// 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`]. /// Trials are sorted by objective value (ascending for minimization), and the /// gamma quantile determines the split point. @@ -499,7 +494,7 @@ impl MultivariateTpeSampler { /// /// ```ignore /// use std::collections::HashMap; - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// use optimizer::sampler::tpe::IntersectionSearchSpace; /// /// let sampler = MultivariateTpeSampler::new(); @@ -593,7 +588,7 @@ impl MultivariateTpeSampler { /// /// ```ignore /// use std::collections::HashMap; - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// /// let sampler = MultivariateTpeSampler::new(); /// let trials = vec![/* ... completed trials ... */]; @@ -661,7 +656,7 @@ impl MultivariateTpeSampler { /// # Examples /// /// ```ignore - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// use optimizer::kde::MultivariateKDE; /// /// let sampler = MultivariateTpeSampler::builder() @@ -859,7 +854,7 @@ impl MultivariateTpeSampler { /// /// ```ignore /// use std::collections::HashMap; - /// use optimizer::sampler::MultivariateTpeSampler; + /// use optimizer::sampler::tpe::MultivariateTpeSampler; /// use optimizer::distribution::{Distribution, FloatDistribution}; /// /// let sampler = MultivariateTpeSampler::builder() @@ -925,7 +920,7 @@ impl MultivariateTpeSampler { ) -> HashMap { use std::collections::HashSet; - use crate::sampler::tpe::GroupDecomposedSearchSpace; + use super::GroupDecomposedSearchSpace; // Decompose the search space into independent parameter groups let groups = GroupDecomposedSearchSpace::calculate(history); @@ -1014,8 +1009,8 @@ impl MultivariateTpeSampler { history: &[CompletedTrial], rng: &mut StdRng, ) -> HashMap { + use super::IntersectionSearchSpace; use crate::kde::MultivariateKDE; - use crate::sampler::tpe::IntersectionSearchSpace; // Early returns for cases requiring random sampling 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 /// 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. #[allow(dead_code)] @@ -1772,7 +1767,7 @@ impl MultivariateTpeSampler { /// Using a fixed gamma value: /// /// ``` -/// use optimizer::sampler::MultivariateTpeSamplerBuilder; +/// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .gamma(0.15) @@ -1786,8 +1781,7 @@ impl MultivariateTpeSampler { /// Using a custom gamma strategy: /// /// ``` -/// use optimizer::sampler::MultivariateTpeSamplerBuilder; -/// use optimizer::sampler::tpe::SqrtGamma; +/// use optimizer::sampler::tpe::{MultivariateTpeSamplerBuilder, SqrtGamma}; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .gamma_strategy(SqrtGamma::default()) @@ -1845,7 +1839,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .gamma(0.10) // Use top 10% as "good" trials @@ -1876,8 +1870,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; - /// use optimizer::sampler::tpe::{LinearGamma, SqrtGamma}; + /// use optimizer::sampler::tpe::{LinearGamma, MultivariateTpeSamplerBuilder, SqrtGamma}; /// /// // Square root strategy (Optuna-style) /// let sampler = MultivariateTpeSamplerBuilder::new() @@ -1911,7 +1904,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .n_startup_trials(20) // Random sample first 20 trials @@ -1937,7 +1930,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .n_ei_candidates(48) // Evaluate more candidates @@ -1964,7 +1957,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .group(true) // Enable group decomposition @@ -1990,7 +1983,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::{ConstantLiarStrategy, MultivariateTpeSamplerBuilder}; + /// use optimizer::sampler::tpe::{ConstantLiarStrategy, MultivariateTpeSamplerBuilder}; /// /// // Use mean imputation for pending trials /// let sampler = MultivariateTpeSamplerBuilder::new() @@ -2025,7 +2018,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .seed(42) // Reproducible results @@ -2047,7 +2040,7 @@ impl MultivariateTpeSamplerBuilder { /// # Examples /// /// ``` - /// use optimizer::sampler::MultivariateTpeSamplerBuilder; + /// use optimizer::sampler::tpe::MultivariateTpeSamplerBuilder; /// /// let sampler = MultivariateTpeSamplerBuilder::new() /// .gamma(0.15) diff --git a/src/sampler/tpe/sampler.rs b/src/sampler/tpe/sampler.rs index 59e1e8c..434967a 100644 --- a/src/sampler/tpe/sampler.rs +++ b/src/sampler/tpe/sampler.rs @@ -280,6 +280,7 @@ impl TpeSampler { clippy::cast_possible_truncation, clippy::cast_sign_loss )] + #[must_use] fn split_trials<'a>( &self, history: &'a [CompletedTrial], diff --git a/tests/multivariate_tpe_integration.rs b/tests/multivariate_tpe_integration.rs index c2c6689..21b54f9 100644 --- a/tests/multivariate_tpe_integration.rs +++ b/tests/multivariate_tpe_integration.rs @@ -9,8 +9,7 @@ clippy::cast_possible_truncation )] -use optimizer::sampler::MultivariateTpeSampler; -use optimizer::sampler::tpe::TpeSampler; +use optimizer::sampler::tpe::{MultivariateTpeSampler, TpeSampler}; use optimizer::{Direction, Error, Study}; // =============================================================================