fix(sampler): clamp multivariate TPE candidates to parameter bounds
- Clamp KDE candidates to parameter bounds before evaluating l(x)/g(x), matching the univariate TPE behavior; without this, candidates scored well at out-of-bounds locations but became suboptimal when clamped - Sort HashMap iterations by ParamId before consuming the seeded RNG to eliminate non-deterministic sampling caused by global ParamId counter - Replace wall-clock timing assertion in async concurrency test with atomic max-active counter to avoid CI flakiness - Gate unused HashSet import behind cfg(feature = "async")
This commit is contained in:
@@ -197,7 +197,7 @@ impl MultivariateTpeSampler {
|
|||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
param_order.sort_by_key(|id| format!("{id}"));
|
param_order.sort();
|
||||||
|
|
||||||
// Extract observations, validate, and fit KDEs
|
// Extract observations, validate, and fit KDEs
|
||||||
let good_obs = self.extract_observations(&good, ¶m_order);
|
let good_obs = self.extract_observations(&good, ¶m_order);
|
||||||
@@ -215,7 +215,20 @@ impl MultivariateTpeSampler {
|
|||||||
return result;
|
return result;
|
||||||
};
|
};
|
||||||
|
|
||||||
let selected = self.select_candidate_with_rng(&good_kde, &bad_kde, rng);
|
// Compute parameter bounds for each dimension so candidates are clamped
|
||||||
|
let bounds: Vec<(f64, f64)> = param_order
|
||||||
|
.iter()
|
||||||
|
.filter_map(|id| {
|
||||||
|
intersection.get(id).and_then(|dist| match dist {
|
||||||
|
Distribution::Float(d) => Some((d.low, d.high)),
|
||||||
|
#[allow(clippy::cast_precision_loss)]
|
||||||
|
Distribution::Int(d) => Some((d.low as f64, d.high as f64)),
|
||||||
|
Distribution::Categorical(_) => None,
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
let selected = self.select_candidate_with_rng(&good_kde, &bad_kde, &bounds, rng);
|
||||||
|
|
||||||
// Map selected values to parameter ids
|
// Map selected values to parameter ids
|
||||||
for (idx, param_id) in param_order.iter().enumerate() {
|
for (idx, param_id) in param_order.iter().enumerate() {
|
||||||
@@ -273,23 +286,38 @@ impl MultivariateTpeSampler {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Clamps each dimension of a candidate to the corresponding parameter bounds.
|
||||||
|
fn clamp_candidate(candidate: &mut [f64], bounds: &[(f64, f64)]) {
|
||||||
|
for (val, &(lo, hi)) in candidate.iter_mut().zip(bounds.iter()) {
|
||||||
|
*val = val.clamp(lo, hi);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Selects the best candidate from a set of samples using the joint acquisition function.
|
/// Selects the best candidate from a set of samples using the joint acquisition function.
|
||||||
///
|
///
|
||||||
/// This method implements the core TPE selection criterion: it generates candidates
|
/// This method implements the core TPE selection criterion: it generates candidates
|
||||||
/// from the "good" KDE (l(x)) and selects the one that maximizes the ratio l(x)/g(x),
|
/// from the "good" KDE (l(x)) and selects the one that maximizes the ratio l(x)/g(x),
|
||||||
/// which is equivalent to maximizing `log(l(x)) - log(g(x))`.
|
/// which is equivalent to maximizing `log(l(x)) - log(g(x))`.
|
||||||
|
///
|
||||||
|
/// Candidates are clamped to parameter bounds before evaluation so the acquisition
|
||||||
|
/// function scores the values that will actually be used.
|
||||||
#[must_use]
|
#[must_use]
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
pub(crate) fn select_candidate(
|
pub(crate) fn select_candidate(
|
||||||
&self,
|
&self,
|
||||||
good_kde: &crate::kde::MultivariateKDE,
|
good_kde: &crate::kde::MultivariateKDE,
|
||||||
bad_kde: &crate::kde::MultivariateKDE,
|
bad_kde: &crate::kde::MultivariateKDE,
|
||||||
|
bounds: &[(f64, f64)],
|
||||||
) -> Vec<f64> {
|
) -> Vec<f64> {
|
||||||
let mut rng = self.rng.lock();
|
let mut rng = self.rng.lock();
|
||||||
|
|
||||||
// Generate candidates from the good distribution
|
// Generate candidates from the good distribution, clamped to bounds
|
||||||
let candidates: Vec<Vec<f64>> = (0..self.n_ei_candidates)
|
let candidates: Vec<Vec<f64>> = (0..self.n_ei_candidates)
|
||||||
.map(|_| good_kde.sample(&mut rng))
|
.map(|_| {
|
||||||
|
let mut c = good_kde.sample(&mut rng);
|
||||||
|
Self::clamp_candidate(&mut c, bounds);
|
||||||
|
c
|
||||||
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// Compute log(l(x)) - log(g(x)) for each candidate
|
// Compute log(l(x)) - log(g(x)) for each candidate
|
||||||
@@ -321,15 +349,21 @@ impl MultivariateTpeSampler {
|
|||||||
/// Selects the best candidate using an external RNG.
|
/// Selects the best candidate using an external RNG.
|
||||||
///
|
///
|
||||||
/// This variant accepts an external RNG, used when the caller already holds the lock.
|
/// This variant accepts an external RNG, used when the caller already holds the lock.
|
||||||
|
/// Candidates are clamped to parameter bounds before evaluation.
|
||||||
pub(crate) fn select_candidate_with_rng(
|
pub(crate) fn select_candidate_with_rng(
|
||||||
&self,
|
&self,
|
||||||
good_kde: &crate::kde::MultivariateKDE,
|
good_kde: &crate::kde::MultivariateKDE,
|
||||||
bad_kde: &crate::kde::MultivariateKDE,
|
bad_kde: &crate::kde::MultivariateKDE,
|
||||||
|
bounds: &[(f64, f64)],
|
||||||
rng: &mut fastrand::Rng,
|
rng: &mut fastrand::Rng,
|
||||||
) -> Vec<f64> {
|
) -> Vec<f64> {
|
||||||
// Generate candidates from the good distribution
|
// Generate candidates from the good distribution, clamped to bounds
|
||||||
let candidates: Vec<Vec<f64>> = (0..self.n_ei_candidates)
|
let candidates: Vec<Vec<f64>> = (0..self.n_ei_candidates)
|
||||||
.map(|_| good_kde.sample(rng))
|
.map(|_| {
|
||||||
|
let mut c = good_kde.sample(rng);
|
||||||
|
Self::clamp_candidate(&mut c, bounds);
|
||||||
|
c
|
||||||
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// Compute log(l(x)) - log(g(x)) for each candidate
|
// Compute log(l(x)) - log(g(x)) for each candidate
|
||||||
@@ -367,8 +401,8 @@ impl MultivariateTpeSampler {
|
|||||||
result: &mut HashMap<ParamId, ParamValue>,
|
result: &mut HashMap<ParamId, ParamValue>,
|
||||||
rng: &mut fastrand::Rng,
|
rng: &mut fastrand::Rng,
|
||||||
) {
|
) {
|
||||||
// Identify parameters not in result (and not in intersection)
|
// Identify parameters not in result, sorted for deterministic RNG consumption
|
||||||
let missing_params: Vec<(&ParamId, &Distribution)> = search_space
|
let mut missing_params: Vec<(&ParamId, &Distribution)> = search_space
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|(id, _)| !result.contains_key(id))
|
.filter(|(id, _)| !result.contains_key(id))
|
||||||
.collect();
|
.collect();
|
||||||
@@ -377,6 +411,8 @@ impl MultivariateTpeSampler {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
missing_params.sort_by_key(|(id, _)| *id);
|
||||||
|
|
||||||
// Split trials for independent sampling
|
// Split trials for independent sampling
|
||||||
let (good_trials, bad_trials) = self.split_trials(&history.iter().collect::<Vec<_>>());
|
let (good_trials, bad_trials) = self.split_trials(&history.iter().collect::<Vec<_>>());
|
||||||
|
|
||||||
@@ -390,6 +426,7 @@ impl MultivariateTpeSampler {
|
|||||||
/// Samples all parameters using independent TPE sampling.
|
/// Samples all parameters using independent TPE sampling.
|
||||||
///
|
///
|
||||||
/// This is used as a complete fallback when no intersection search space exists.
|
/// This is used as a complete fallback when no intersection search space exists.
|
||||||
|
/// Parameters are sorted by `ParamId` for deterministic RNG consumption order.
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
pub(crate) fn sample_all_independent(
|
pub(crate) fn sample_all_independent(
|
||||||
&self,
|
&self,
|
||||||
@@ -402,7 +439,9 @@ impl MultivariateTpeSampler {
|
|||||||
let mut rng = self.rng.lock();
|
let mut rng = self.rng.lock();
|
||||||
let mut result = HashMap::new();
|
let mut result = HashMap::new();
|
||||||
|
|
||||||
for (param_id, dist) in search_space {
|
let mut sorted: Vec<_> = search_space.iter().collect();
|
||||||
|
sorted.sort_by_key(|(id, _)| *id);
|
||||||
|
for (param_id, dist) in sorted {
|
||||||
let value =
|
let value =
|
||||||
self.sample_independent_tpe(*param_id, dist, &good_trials, &bad_trials, &mut rng);
|
self.sample_independent_tpe(*param_id, dist, &good_trials, &bad_trials, &mut rng);
|
||||||
result.insert(*param_id, value);
|
result.insert(*param_id, value);
|
||||||
@@ -414,6 +453,7 @@ impl MultivariateTpeSampler {
|
|||||||
/// Samples all parameters using independent TPE sampling with an external RNG.
|
/// Samples all parameters using independent TPE sampling with an external RNG.
|
||||||
///
|
///
|
||||||
/// This variant accepts an external RNG, used when the caller already holds the lock.
|
/// This variant accepts an external RNG, used when the caller already holds the lock.
|
||||||
|
/// Parameters are sorted by `ParamId` for deterministic RNG consumption order.
|
||||||
pub(crate) fn sample_all_independent_with_rng(
|
pub(crate) fn sample_all_independent_with_rng(
|
||||||
&self,
|
&self,
|
||||||
search_space: &HashMap<ParamId, Distribution>,
|
search_space: &HashMap<ParamId, Distribution>,
|
||||||
@@ -425,7 +465,9 @@ impl MultivariateTpeSampler {
|
|||||||
|
|
||||||
let mut result = HashMap::new();
|
let mut result = HashMap::new();
|
||||||
|
|
||||||
for (param_id, dist) in search_space {
|
let mut sorted: Vec<_> = search_space.iter().collect();
|
||||||
|
sorted.sort_by_key(|(id, _)| *id);
|
||||||
|
for (param_id, dist) in sorted {
|
||||||
let value =
|
let value =
|
||||||
self.sample_independent_tpe(*param_id, dist, &good_trials, &bad_trials, rng);
|
self.sample_independent_tpe(*param_id, dist, &good_trials, &bad_trials, rng);
|
||||||
result.insert(*param_id, value);
|
result.insert(*param_id, value);
|
||||||
|
|||||||
@@ -322,14 +322,18 @@ impl MultivariateTpeSampler {
|
|||||||
/// Samples all parameters uniformly at random.
|
/// Samples all parameters uniformly at random.
|
||||||
///
|
///
|
||||||
/// This is a fallback method used when multivariate TPE cannot be applied.
|
/// This is a fallback method used when multivariate TPE cannot be applied.
|
||||||
|
/// Parameters are sorted by `ParamId` to ensure deterministic RNG consumption
|
||||||
|
/// order when using a seeded sampler.
|
||||||
#[allow(clippy::unused_self)]
|
#[allow(clippy::unused_self)]
|
||||||
fn sample_all_uniform(
|
fn sample_all_uniform(
|
||||||
&self,
|
&self,
|
||||||
search_space: &HashMap<ParamId, Distribution>,
|
search_space: &HashMap<ParamId, Distribution>,
|
||||||
rng: &mut fastrand::Rng,
|
rng: &mut fastrand::Rng,
|
||||||
) -> HashMap<ParamId, ParamValue> {
|
) -> HashMap<ParamId, ParamValue> {
|
||||||
search_space
|
let mut sorted: Vec<_> = search_space.iter().collect();
|
||||||
.iter()
|
sorted.sort_by_key(|(id, _)| *id);
|
||||||
|
sorted
|
||||||
|
.into_iter()
|
||||||
.map(|(id, dist)| (*id, crate::sampler::common::sample_random(rng, dist)))
|
.map(|(id, dist)| (*id, crate::sampler::common::sample_random(rng, dist)))
|
||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
@@ -2781,8 +2785,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(0.0, 1.0), (0.0, 1.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
// The selected candidate should have 2 dimensions
|
// The selected candidate should have 2 dimensions
|
||||||
assert_eq!(selected.len(), 2);
|
assert_eq!(selected.len(), 2);
|
||||||
@@ -2822,8 +2827,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(-1.0, 2.0), (-1.0, 2.0), (-1.0, 2.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
assert_eq!(selected.len(), 3);
|
assert_eq!(selected.len(), 3);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2841,8 +2847,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(0.0, 10.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
assert_eq!(selected.len(), 1);
|
assert_eq!(selected.len(), 1);
|
||||||
|
|
||||||
// Selected value should be closer to the good region
|
// Selected value should be closer to the good region
|
||||||
@@ -2873,14 +2880,15 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples.clone()).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples.clone()).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples.clone()).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples.clone()).unwrap();
|
||||||
|
let bounds = &[(0.0, 10.0), (0.0, 10.0)];
|
||||||
|
|
||||||
let selected1 = sampler1.select_candidate(&good_kde, &bad_kde);
|
let selected1 = sampler1.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
// Need to recreate KDEs for second sampler since we consumed them
|
// Need to recreate KDEs for second sampler since we consumed them
|
||||||
let good_kde2 = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde2 = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde2 = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde2 = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
|
||||||
let selected2 = sampler2.select_candidate(&good_kde2, &bad_kde2);
|
let selected2 = sampler2.select_candidate(&good_kde2, &bad_kde2, bounds);
|
||||||
|
|
||||||
// With same seed, should get same result
|
// With same seed, should get same result
|
||||||
assert!(
|
assert!(
|
||||||
@@ -2911,8 +2919,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(-10.0, 10.0), (-10.0, 10.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
assert_eq!(selected.len(), 2);
|
assert_eq!(selected.len(), 2);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2942,8 +2951,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(-5.0, 15.0), (-5.0, 15.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
// With more candidates, should definitely find a point in the good region
|
// With more candidates, should definitely find a point in the good region
|
||||||
assert!(
|
assert!(
|
||||||
@@ -2981,8 +2991,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(-5.0, 5.0), (-5.0, 5.0)];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
// Should still return a valid point
|
// Should still return a valid point
|
||||||
assert_eq!(selected.len(), 2);
|
assert_eq!(selected.len(), 2);
|
||||||
@@ -3017,8 +3028,9 @@ mod tests {
|
|||||||
|
|
||||||
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
let good_kde = MultivariateKDE::new(good_samples).unwrap();
|
||||||
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_samples).unwrap();
|
||||||
|
let bounds = &[(-5.0, 15.0); 5];
|
||||||
|
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
assert_eq!(selected.len(), 5);
|
assert_eq!(selected.len(), 5);
|
||||||
|
|
||||||
@@ -3091,7 +3103,8 @@ mod tests {
|
|||||||
let bad_kde = MultivariateKDE::new(bad_obs).unwrap();
|
let bad_kde = MultivariateKDE::new(bad_obs).unwrap();
|
||||||
|
|
||||||
// Select candidate
|
// Select candidate
|
||||||
let selected = sampler.select_candidate(&good_kde, &bad_kde);
|
let bounds = &[(0.0, 10.0), (0.0, 10.0)];
|
||||||
|
let selected = sampler.select_candidate(&good_kde, &bad_kde, bounds);
|
||||||
|
|
||||||
assert_eq!(selected.len(), 2);
|
assert_eq!(selected.len(), 2);
|
||||||
|
|
||||||
|
|||||||
+18
-6
@@ -197,27 +197,39 @@ async fn test_optimize_parallel_single_concurrency() {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_parallel_executes_concurrently() {
|
async fn test_parallel_executes_concurrently() {
|
||||||
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||||
|
|
||||||
let sampler = RandomSampler::with_seed(42);
|
let sampler = RandomSampler::with_seed(42);
|
||||||
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
let study: Study<f64> = Study::with_sampler(Direction::Minimize, sampler);
|
||||||
|
|
||||||
let x_param = FloatParam::new(0.0, 10.0);
|
let x_param = FloatParam::new(0.0, 10.0);
|
||||||
|
let active = Arc::new(AtomicUsize::new(0));
|
||||||
|
let max_active = Arc::new(AtomicUsize::new(0));
|
||||||
|
|
||||||
|
let active_c = Arc::clone(&active);
|
||||||
|
let max_active_c = Arc::clone(&max_active);
|
||||||
|
|
||||||
let start = tokio::time::Instant::now();
|
|
||||||
study
|
study
|
||||||
.optimize_parallel(4, 4, move |trial: &mut optimizer::Trial| {
|
.optimize_parallel(4, 4, move |trial: &mut optimizer::Trial| {
|
||||||
let x = x_param.suggest(trial)?;
|
let x = x_param.suggest(trial)?;
|
||||||
std::thread::sleep(std::time::Duration::from_millis(100));
|
|
||||||
|
let current = active_c.fetch_add(1, Ordering::SeqCst) + 1;
|
||||||
|
max_active_c.fetch_max(current, Ordering::SeqCst);
|
||||||
|
|
||||||
|
std::thread::sleep(std::time::Duration::from_millis(50));
|
||||||
|
|
||||||
|
active_c.fetch_sub(1, Ordering::SeqCst);
|
||||||
Ok::<_, Error>(x)
|
Ok::<_, Error>(x)
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
.expect("parallel optimization should succeed");
|
.expect("parallel optimization should succeed");
|
||||||
|
|
||||||
let elapsed = start.elapsed();
|
|
||||||
assert_eq!(study.n_trials(), 4);
|
assert_eq!(study.n_trials(), 4);
|
||||||
// Sequential would take ~400ms; parallel with concurrency=4 should be ~100ms
|
let max = max_active.load(Ordering::SeqCst);
|
||||||
|
// With 4 trials and concurrency=4, all should run concurrently
|
||||||
assert!(
|
assert!(
|
||||||
elapsed < std::time::Duration::from_millis(350),
|
max >= 2,
|
||||||
"expected parallel execution under 350ms, took {elapsed:?}"
|
"expected at least 2 concurrent workers, but max was {max}"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ fn test_multivariate_tpe_rosenbrock_finds_good_solution() {
|
|||||||
let sampler = MultivariateTpeSampler::builder()
|
let sampler = MultivariateTpeSampler::builder()
|
||||||
.seed(42)
|
.seed(42)
|
||||||
.n_startup_trials(10)
|
.n_startup_trials(10)
|
||||||
.n_ei_candidates(24)
|
.n_ei_candidates(48)
|
||||||
.build()
|
.build()
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
@@ -55,7 +55,7 @@ fn test_multivariate_tpe_rosenbrock_finds_good_solution() {
|
|||||||
let y_param = FloatParam::new(-2.0, 4.0);
|
let y_param = FloatParam::new(-2.0, 4.0);
|
||||||
|
|
||||||
study
|
study
|
||||||
.optimize(100, |trial: &mut optimizer::Trial| {
|
.optimize(200, |trial: &mut optimizer::Trial| {
|
||||||
let x = x_param.suggest(trial)?;
|
let x = x_param.suggest(trial)?;
|
||||||
let y = y_param.suggest(trial)?;
|
let y = y_param.suggest(trial)?;
|
||||||
Ok::<_, Error>(rosenbrock(x, y))
|
Ok::<_, Error>(rosenbrock(x, y))
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
//! All tests are `#[ignore]`-gated so they don't run in normal CI.
|
//! All tests are `#[ignore]`-gated so they don't run in normal CI.
|
||||||
//! Run with: `cargo test --features async -- --ignored`
|
//! Run with: `cargo test --features async -- --ignored`
|
||||||
|
|
||||||
|
#[cfg(feature = "async")]
|
||||||
use std::collections::HashSet;
|
use std::collections::HashSet;
|
||||||
|
|
||||||
use optimizer::parameter::{FloatParam, Parameter};
|
use optimizer::parameter::{FloatParam, Parameter};
|
||||||
|
|||||||
Reference in New Issue
Block a user