feat: add optional fixed bandwidth for KDE in TPE sampler
This commit is contained in:
+3
-2
@@ -58,7 +58,7 @@ impl KernelDensityEstimator {
|
|||||||
/// # Panics
|
/// # Panics
|
||||||
///
|
///
|
||||||
/// Panics if `samples` is empty or `bandwidth` is not positive.
|
/// Panics if `samples` is empty or `bandwidth` is not positive.
|
||||||
pub fn with_bandwidth(samples: Vec<f64>, bandwidth: f64) -> Self {
|
pub(crate) fn with_bandwidth(samples: Vec<f64>, bandwidth: f64) -> Self {
|
||||||
assert!(!samples.is_empty(), "KDE requires at least one sample");
|
assert!(!samples.is_empty(), "KDE requires at least one sample");
|
||||||
assert!(bandwidth > 0.0, "Bandwidth must be positive");
|
assert!(bandwidth > 0.0, "Bandwidth must be positive");
|
||||||
|
|
||||||
@@ -135,7 +135,8 @@ impl KernelDensityEstimator {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Returns the bandwidth of this KDE.
|
/// Returns the bandwidth of this KDE.
|
||||||
pub fn bandwidth(&self) -> f64 {
|
#[cfg(test)]
|
||||||
|
pub(crate) fn bandwidth(&self) -> f64 {
|
||||||
self.bandwidth
|
self.bandwidth
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+67
-13
@@ -47,6 +47,8 @@ pub struct TpeSampler {
|
|||||||
n_startup_trials: usize,
|
n_startup_trials: usize,
|
||||||
/// Number of candidate samples to evaluate when selecting the next point.
|
/// Number of candidate samples to evaluate when selecting the next point.
|
||||||
n_ei_candidates: usize,
|
n_ei_candidates: usize,
|
||||||
|
/// Optional fixed bandwidth for KDE. If None, uses Scott's rule.
|
||||||
|
kde_bandwidth: Option<f64>,
|
||||||
/// Thread-safe RNG for sampling.
|
/// Thread-safe RNG for sampling.
|
||||||
rng: Mutex<StdRng>,
|
rng: Mutex<StdRng>,
|
||||||
}
|
}
|
||||||
@@ -58,11 +60,13 @@ impl TpeSampler {
|
|||||||
/// - gamma: 0.25 (top 25% of trials are considered "good")
|
/// - gamma: 0.25 (top 25% of trials are considered "good")
|
||||||
/// - n_startup_trials: 10 (random sampling for first 10 trials)
|
/// - n_startup_trials: 10 (random sampling for first 10 trials)
|
||||||
/// - n_ei_candidates: 24 (evaluate 24 candidates per sample)
|
/// - n_ei_candidates: 24 (evaluate 24 candidates per sample)
|
||||||
|
/// - kde_bandwidth: None (uses Scott's rule for automatic bandwidth)
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
gamma: 0.25,
|
gamma: 0.25,
|
||||||
n_startup_trials: 10,
|
n_startup_trials: 10,
|
||||||
n_ei_candidates: 24,
|
n_ei_candidates: 24,
|
||||||
|
kde_bandwidth: None,
|
||||||
rng: Mutex::new(StdRng::from_os_rng()),
|
rng: Mutex::new(StdRng::from_os_rng()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,21 +96,26 @@ impl TpeSampler {
|
|||||||
/// * `gamma` - Fraction of trials to consider "good" (0.0 to 1.0).
|
/// * `gamma` - Fraction of trials to consider "good" (0.0 to 1.0).
|
||||||
/// * `n_startup_trials` - Number of random trials before TPE sampling.
|
/// * `n_startup_trials` - Number of random trials before TPE sampling.
|
||||||
/// * `n_ei_candidates` - Number of candidates to evaluate per sample.
|
/// * `n_ei_candidates` - Number of candidates to evaluate per sample.
|
||||||
|
/// * `kde_bandwidth` - Optional fixed bandwidth for KDE. If None, uses Scott's rule.
|
||||||
/// * `seed` - Optional seed for reproducibility.
|
/// * `seed` - Optional seed for reproducibility.
|
||||||
///
|
///
|
||||||
/// # Panics
|
/// # Panics
|
||||||
///
|
///
|
||||||
/// Panics if gamma is not in (0.0, 1.0).
|
/// Panics if gamma is not in (0.0, 1.0) or if kde_bandwidth is Some but not positive.
|
||||||
pub fn with_config(
|
pub fn with_config(
|
||||||
gamma: f64,
|
gamma: f64,
|
||||||
n_startup_trials: usize,
|
n_startup_trials: usize,
|
||||||
n_ei_candidates: usize,
|
n_ei_candidates: usize,
|
||||||
|
kde_bandwidth: Option<f64>,
|
||||||
seed: Option<u64>,
|
seed: Option<u64>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
assert!(
|
assert!(
|
||||||
gamma > 0.0 && gamma < 1.0,
|
gamma > 0.0 && gamma < 1.0,
|
||||||
"gamma must be in (0.0, 1.0), got {gamma}"
|
"gamma must be in (0.0, 1.0), got {gamma}"
|
||||||
);
|
);
|
||||||
|
if let Some(bw) = kde_bandwidth {
|
||||||
|
assert!(bw > 0.0, "kde_bandwidth must be positive, got {bw}");
|
||||||
|
}
|
||||||
|
|
||||||
let rng = match seed {
|
let rng = match seed {
|
||||||
Some(s) => StdRng::seed_from_u64(s),
|
Some(s) => StdRng::seed_from_u64(s),
|
||||||
@@ -117,6 +126,7 @@ impl TpeSampler {
|
|||||||
gamma,
|
gamma,
|
||||||
n_startup_trials,
|
n_startup_trials,
|
||||||
n_ei_candidates,
|
n_ei_candidates,
|
||||||
|
kde_bandwidth,
|
||||||
rng: Mutex::new(rng),
|
rng: Mutex::new(rng),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -222,8 +232,14 @@ impl TpeSampler {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Fit KDEs to good and bad groups
|
// Fit KDEs to good and bad groups
|
||||||
let l_kde = KernelDensityEstimator::new(good_internal);
|
let l_kde = match self.kde_bandwidth {
|
||||||
let g_kde = KernelDensityEstimator::new(bad_internal);
|
Some(bw) => KernelDensityEstimator::with_bandwidth(good_internal, bw),
|
||||||
|
None => KernelDensityEstimator::new(good_internal),
|
||||||
|
};
|
||||||
|
let g_kde = match self.kde_bandwidth {
|
||||||
|
Some(bw) => KernelDensityEstimator::with_bandwidth(bad_internal, bw),
|
||||||
|
None => KernelDensityEstimator::new(bad_internal),
|
||||||
|
};
|
||||||
|
|
||||||
// Generate candidates from l(x) and select the one with best l(x)/g(x) ratio
|
// Generate candidates from l(x) and select the one with best l(x)/g(x) ratio
|
||||||
let mut best_candidate = internal_low;
|
let mut best_candidate = internal_low;
|
||||||
@@ -393,6 +409,7 @@ pub struct TpeSamplerBuilder {
|
|||||||
gamma: f64,
|
gamma: f64,
|
||||||
n_startup_trials: usize,
|
n_startup_trials: usize,
|
||||||
n_ei_candidates: usize,
|
n_ei_candidates: usize,
|
||||||
|
kde_bandwidth: Option<f64>,
|
||||||
seed: Option<u64>,
|
seed: Option<u64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -403,12 +420,14 @@ impl TpeSamplerBuilder {
|
|||||||
/// - gamma: 0.25 (top 25% of trials are considered "good")
|
/// - gamma: 0.25 (top 25% of trials are considered "good")
|
||||||
/// - n_startup_trials: 10 (random sampling for first 10 trials)
|
/// - n_startup_trials: 10 (random sampling for first 10 trials)
|
||||||
/// - n_ei_candidates: 24 (evaluate 24 candidates per sample)
|
/// - n_ei_candidates: 24 (evaluate 24 candidates per sample)
|
||||||
|
/// - kde_bandwidth: None (uses Scott's rule for automatic bandwidth)
|
||||||
/// - seed: None (use OS-provided entropy)
|
/// - seed: None (use OS-provided entropy)
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
gamma: 0.25,
|
gamma: 0.25,
|
||||||
n_startup_trials: 10,
|
n_startup_trials: 10,
|
||||||
n_ei_candidates: 24,
|
n_ei_candidates: 24,
|
||||||
|
kde_bandwidth: None,
|
||||||
seed: None,
|
seed: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -492,6 +511,40 @@ impl TpeSamplerBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Sets a fixed bandwidth for the kernel density estimator.
|
||||||
|
///
|
||||||
|
/// By default, TPE uses Scott's rule to automatically select the bandwidth
|
||||||
|
/// based on the sample data. Use this method to override with a fixed value.
|
||||||
|
///
|
||||||
|
/// Smaller bandwidths give more localized, peaky distributions.
|
||||||
|
/// Larger bandwidths give smoother, more spread-out distributions.
|
||||||
|
///
|
||||||
|
/// # Arguments
|
||||||
|
///
|
||||||
|
/// * `bandwidth` - The fixed bandwidth (standard deviation) for Gaussian kernels.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if bandwidth is not positive.
|
||||||
|
///
|
||||||
|
/// # Examples
|
||||||
|
///
|
||||||
|
/// ```
|
||||||
|
/// use optimize::TpeSamplerBuilder;
|
||||||
|
///
|
||||||
|
/// let sampler = TpeSamplerBuilder::new()
|
||||||
|
/// .kde_bandwidth(0.5) // Fixed bandwidth of 0.5
|
||||||
|
/// .build();
|
||||||
|
/// ```
|
||||||
|
pub fn kde_bandwidth(mut self, bandwidth: f64) -> Self {
|
||||||
|
assert!(
|
||||||
|
bandwidth > 0.0,
|
||||||
|
"kde_bandwidth must be positive, got {bandwidth}"
|
||||||
|
);
|
||||||
|
self.kde_bandwidth = Some(bandwidth);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Sets a seed for reproducible sampling.
|
/// Sets a seed for reproducible sampling.
|
||||||
///
|
///
|
||||||
/// # Arguments
|
/// # Arguments
|
||||||
@@ -531,6 +584,7 @@ impl TpeSamplerBuilder {
|
|||||||
self.gamma,
|
self.gamma,
|
||||||
self.n_startup_trials,
|
self.n_startup_trials,
|
||||||
self.n_ei_candidates,
|
self.n_ei_candidates,
|
||||||
|
self.kde_bandwidth,
|
||||||
self.seed,
|
self.seed,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -709,7 +763,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_sampler_with_config() {
|
fn test_tpe_sampler_with_config() {
|
||||||
let sampler = TpeSampler::with_config(0.15, 20, 32, Some(42));
|
let sampler = TpeSampler::with_config(0.15, 20, 32, None, Some(42));
|
||||||
assert_eq!(sampler.gamma, 0.15);
|
assert_eq!(sampler.gamma, 0.15);
|
||||||
assert_eq!(sampler.n_startup_trials, 20);
|
assert_eq!(sampler.n_startup_trials, 20);
|
||||||
assert_eq!(sampler.n_ei_candidates, 32);
|
assert_eq!(sampler.n_ei_candidates, 32);
|
||||||
@@ -718,18 +772,18 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
#[should_panic(expected = "gamma must be in (0.0, 1.0)")]
|
#[should_panic(expected = "gamma must be in (0.0, 1.0)")]
|
||||||
fn test_tpe_sampler_invalid_gamma_zero() {
|
fn test_tpe_sampler_invalid_gamma_zero() {
|
||||||
TpeSampler::with_config(0.0, 10, 24, None);
|
TpeSampler::with_config(0.0, 10, 24, None, None);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
#[should_panic(expected = "gamma must be in (0.0, 1.0)")]
|
#[should_panic(expected = "gamma must be in (0.0, 1.0)")]
|
||||||
fn test_tpe_sampler_invalid_gamma_one() {
|
fn test_tpe_sampler_invalid_gamma_one() {
|
||||||
TpeSampler::with_config(1.0, 10, 24, None);
|
TpeSampler::with_config(1.0, 10, 24, None, None);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_startup_random_sampling() {
|
fn test_tpe_startup_random_sampling() {
|
||||||
let sampler = TpeSampler::with_config(0.25, 10, 24, Some(42));
|
let sampler = TpeSampler::with_config(0.25, 10, 24, None, Some(42));
|
||||||
let dist = Distribution::Float(FloatDistribution {
|
let dist = Distribution::Float(FloatDistribution {
|
||||||
low: 0.0,
|
low: 0.0,
|
||||||
high: 1.0,
|
high: 1.0,
|
||||||
@@ -752,7 +806,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_split_trials() {
|
fn test_tpe_split_trials() {
|
||||||
let sampler = TpeSampler::with_config(0.25, 10, 24, Some(42));
|
let sampler = TpeSampler::with_config(0.25, 10, 24, None, Some(42));
|
||||||
|
|
||||||
let dist = Distribution::Float(FloatDistribution {
|
let dist = Distribution::Float(FloatDistribution {
|
||||||
low: 0.0,
|
low: 0.0,
|
||||||
@@ -786,7 +840,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_samples_float_with_history() {
|
fn test_tpe_samples_float_with_history() {
|
||||||
let sampler = TpeSampler::with_config(0.25, 5, 24, Some(42));
|
let sampler = TpeSampler::with_config(0.25, 5, 24, None, Some(42));
|
||||||
|
|
||||||
let dist = Distribution::Float(FloatDistribution {
|
let dist = Distribution::Float(FloatDistribution {
|
||||||
low: 0.0,
|
low: 0.0,
|
||||||
@@ -828,7 +882,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_categorical_sampling() {
|
fn test_tpe_categorical_sampling() {
|
||||||
let sampler = TpeSampler::with_config(0.25, 5, 24, Some(42));
|
let sampler = TpeSampler::with_config(0.25, 5, 24, None, Some(42));
|
||||||
|
|
||||||
let dist = Distribution::Categorical(CategoricalDistribution { n_choices: 4 });
|
let dist = Distribution::Categorical(CategoricalDistribution { n_choices: 4 });
|
||||||
|
|
||||||
@@ -868,7 +922,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tpe_int_sampling() {
|
fn test_tpe_int_sampling() {
|
||||||
let sampler = TpeSampler::with_config(0.25, 5, 24, Some(42));
|
let sampler = TpeSampler::with_config(0.25, 5, 24, None, Some(42));
|
||||||
|
|
||||||
let dist = Distribution::Int(IntDistribution {
|
let dist = Distribution::Int(IntDistribution {
|
||||||
low: 0,
|
low: 0,
|
||||||
@@ -920,8 +974,8 @@ mod tests {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
let sampler1 = TpeSampler::with_config(0.25, 5, 24, Some(12345));
|
let sampler1 = TpeSampler::with_config(0.25, 5, 24, None, Some(12345));
|
||||||
let sampler2 = TpeSampler::with_config(0.25, 5, 24, Some(12345));
|
let sampler2 = TpeSampler::with_config(0.25, 5, 24, None, Some(12345));
|
||||||
|
|
||||||
for i in 0..10 {
|
for i in 0..10 {
|
||||||
let v1 = sampler1.sample(&dist, i, &history);
|
let v1 = sampler1.sample(&dist, i, &history);
|
||||||
|
|||||||
Reference in New Issue
Block a user