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.
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;
+4
View File
@@ -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};
@@ -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<String, ParamValue> {
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<String, ParamValue> {
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)
+1
View File
@@ -280,6 +280,7 @@ impl TpeSampler {
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
#[must_use]
fn split_trials<'a>(
&self,
history: &'a [CompletedTrial],
+1 -2
View File
@@ -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};
// =============================================================================