refactor: replace TpeError with Error in the optimizer library
This commit is contained in:
+3
-9
@@ -1,10 +1,5 @@
|
||||
//! Error types for the optimizer library.
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
/// The error type for TPE operations.
|
||||
#[derive(Debug, Error)]
|
||||
pub enum TpeError {
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
/// Returned when the lower bound is greater than the upper bound.
|
||||
#[error("invalid bounds: low ({low}) must be less than or equal to high ({high})")]
|
||||
InvalidBounds {
|
||||
@@ -61,5 +56,4 @@ pub enum TpeError {
|
||||
TaskError(String),
|
||||
}
|
||||
|
||||
/// A specialized Result type for TPE operations.
|
||||
pub type Result<T> = core::result::Result<T, TpeError>;
|
||||
pub type Result<T> = core::result::Result<T, Error>;
|
||||
|
||||
+10
-10
@@ -5,7 +5,7 @@
|
||||
|
||||
use rand::Rng;
|
||||
|
||||
use crate::error::{Result, TpeError};
|
||||
use crate::error::{Error, Result};
|
||||
|
||||
/// A Gaussian kernel density estimator for continuous distributions.
|
||||
///
|
||||
@@ -45,10 +45,10 @@ impl KernelDensityEstimator {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::EmptySamples` if `samples` is empty.
|
||||
/// Returns `Error::EmptySamples` if `samples` is empty.
|
||||
pub(crate) fn new(samples: Vec<f64>) -> Result<Self> {
|
||||
if samples.is_empty() {
|
||||
return Err(TpeError::EmptySamples);
|
||||
return Err(Error::EmptySamples);
|
||||
}
|
||||
|
||||
let bandwidth = Self::scotts_rule(&samples);
|
||||
@@ -61,14 +61,14 @@ impl KernelDensityEstimator {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::EmptySamples` if `samples` is empty.
|
||||
/// Returns `TpeError::InvalidBandwidth` if `bandwidth` is not positive.
|
||||
/// Returns `Error::EmptySamples` if `samples` is empty.
|
||||
/// Returns `Error::InvalidBandwidth` if `bandwidth` is not positive.
|
||||
pub(crate) fn with_bandwidth(samples: Vec<f64>, bandwidth: f64) -> Result<Self> {
|
||||
if samples.is_empty() {
|
||||
return Err(TpeError::EmptySamples);
|
||||
return Err(Error::EmptySamples);
|
||||
}
|
||||
if bandwidth <= 0.0 {
|
||||
return Err(TpeError::InvalidBandwidth(bandwidth));
|
||||
return Err(Error::InvalidBandwidth(bandwidth));
|
||||
}
|
||||
|
||||
Ok(Self { samples, bandwidth })
|
||||
@@ -261,20 +261,20 @@ mod tests {
|
||||
fn test_kde_empty_samples() {
|
||||
let samples: Vec<f64> = vec![];
|
||||
let result = KernelDensityEstimator::new(samples);
|
||||
assert!(matches!(result, Err(TpeError::EmptySamples)));
|
||||
assert!(matches!(result, Err(Error::EmptySamples)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_kde_zero_bandwidth() {
|
||||
let samples = vec![1.0, 2.0, 3.0];
|
||||
let result = KernelDensityEstimator::with_bandwidth(samples, 0.0);
|
||||
assert!(matches!(result, Err(TpeError::InvalidBandwidth(_))));
|
||||
assert!(matches!(result, Err(Error::InvalidBandwidth(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_kde_negative_bandwidth() {
|
||||
let samples = vec![1.0, 2.0, 3.0];
|
||||
let result = KernelDensityEstimator::with_bandwidth(samples, -1.0);
|
||||
assert!(matches!(result, Err(TpeError::InvalidBandwidth(_))));
|
||||
assert!(matches!(result, Err(Error::InvalidBandwidth(_))));
|
||||
}
|
||||
}
|
||||
|
||||
+3
-3
@@ -33,7 +33,7 @@
|
||||
//! study
|
||||
//! .optimize_with_sampler(20, |trial| {
|
||||
//! let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
//! Ok::<_, optimizer::TpeError>(x * x)
|
||||
//! Ok::<_, optimizer::Error>(x * x)
|
||||
//! })
|
||||
//! .unwrap();
|
||||
//!
|
||||
@@ -86,7 +86,7 @@
|
||||
//! let optimizer = trial.suggest_categorical("optimizer", &["sgd", "adam", "rmsprop"])?;
|
||||
//!
|
||||
//! // Return objective value
|
||||
//! Ok::<_, optimizer::TpeError>(x * n as f64)
|
||||
//! Ok::<_, optimizer::Error>(x * n as f64)
|
||||
//! })
|
||||
//! .unwrap();
|
||||
//! ```
|
||||
@@ -140,7 +140,7 @@ mod study;
|
||||
mod trial;
|
||||
mod types;
|
||||
|
||||
pub use error::{Result, TpeError};
|
||||
pub use error::{Error, Result};
|
||||
pub use param::ParamValue;
|
||||
pub use study::Study;
|
||||
pub use trial::Trial;
|
||||
|
||||
+12
-12
@@ -9,7 +9,7 @@ use rand::rngs::StdRng;
|
||||
use rand::{Rng, SeedableRng};
|
||||
|
||||
use crate::distribution::Distribution;
|
||||
use crate::error::{Result, TpeError};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::kde::KernelDensityEstimator;
|
||||
use crate::param::ParamValue;
|
||||
use crate::sampler::{CompletedTrial, Sampler};
|
||||
@@ -106,8 +106,8 @@ impl TpeSampler {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::InvalidGamma` if gamma is not in (0.0, 1.0).
|
||||
/// Returns `TpeError::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
|
||||
/// Returns `Error::InvalidGamma` if gamma is not in (0.0, 1.0).
|
||||
/// Returns `Error::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
|
||||
pub fn with_config(
|
||||
gamma: f64,
|
||||
n_startup_trials: usize,
|
||||
@@ -116,12 +116,12 @@ impl TpeSampler {
|
||||
seed: Option<u64>,
|
||||
) -> Result<Self> {
|
||||
if gamma <= 0.0 || gamma >= 1.0 {
|
||||
return Err(TpeError::InvalidGamma(gamma));
|
||||
return Err(Error::InvalidGamma(gamma));
|
||||
}
|
||||
if let Some(bw) = kde_bandwidth
|
||||
&& bw <= 0.0
|
||||
{
|
||||
return Err(TpeError::InvalidBandwidth(bw));
|
||||
return Err(Error::InvalidBandwidth(bw));
|
||||
}
|
||||
|
||||
let rng = match seed {
|
||||
@@ -484,7 +484,7 @@ impl TpeSamplerBuilder {
|
||||
/// # Note
|
||||
///
|
||||
/// Validation happens at `build()` time. If gamma is not in (0.0, 1.0),
|
||||
/// `build()` will return `Err(TpeError::InvalidGamma)`.
|
||||
/// `build()` will return `Err(Error::InvalidGamma)`.
|
||||
#[must_use]
|
||||
pub fn gamma(mut self, gamma: f64) -> Self {
|
||||
self.gamma = gamma;
|
||||
@@ -569,7 +569,7 @@ impl TpeSamplerBuilder {
|
||||
/// # Note
|
||||
///
|
||||
/// Validation happens at `build()` time. If bandwidth is not positive,
|
||||
/// `build()` will return `Err(TpeError::InvalidBandwidth)`.
|
||||
/// `build()` will return `Err(Error::InvalidBandwidth)`.
|
||||
#[must_use]
|
||||
pub fn kde_bandwidth(mut self, bandwidth: f64) -> Self {
|
||||
self.kde_bandwidth = Some(bandwidth);
|
||||
@@ -602,8 +602,8 @@ impl TpeSamplerBuilder {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::InvalidGamma` if gamma is not in (0.0, 1.0).
|
||||
/// Returns `TpeError::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
|
||||
/// Returns `Error::InvalidGamma` if gamma is not in (0.0, 1.0).
|
||||
/// Returns `Error::InvalidBandwidth` if `kde_bandwidth` is Some but not positive.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -817,13 +817,13 @@ mod tests {
|
||||
#[test]
|
||||
fn test_tpe_sampler_invalid_gamma_zero() {
|
||||
let result = TpeSampler::with_config(0.0, 10, 24, None, None);
|
||||
assert!(matches!(result, Err(TpeError::InvalidGamma(_))));
|
||||
assert!(matches!(result, Err(Error::InvalidGamma(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tpe_sampler_invalid_gamma_one() {
|
||||
let result = TpeSampler::with_config(1.0, 10, 24, None, None);
|
||||
assert!(matches!(result, Err(TpeError::InvalidGamma(_))));
|
||||
assert!(matches!(result, Err(Error::InvalidGamma(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1077,7 +1077,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_tpe_sampler_builder_invalid_gamma() {
|
||||
let result = TpeSamplerBuilder::new().gamma(1.5).build();
|
||||
assert!(matches!(result, Err(TpeError::InvalidGamma(_))));
|
||||
assert!(matches!(result, Err(Error::InvalidGamma(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
+38
-38
@@ -275,7 +275,7 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if no trials have been completed.
|
||||
/// Returns `Error::NoCompletedTrials` if no trials have been completed.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -305,7 +305,7 @@ where
|
||||
let trials = self.completed_trials.read();
|
||||
|
||||
if trials.is_empty() {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
let best = trials
|
||||
@@ -325,7 +325,7 @@ where
|
||||
}
|
||||
}
|
||||
})
|
||||
.ok_or(crate::TpeError::NoCompletedTrials)?;
|
||||
.ok_or(crate::Error::NoCompletedTrials)?;
|
||||
|
||||
Ok(best.clone())
|
||||
}
|
||||
@@ -338,7 +338,7 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if no trials have been completed.
|
||||
/// Returns `Error::NoCompletedTrials` if no trials have been completed.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -387,7 +387,7 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -402,7 +402,7 @@ where
|
||||
/// study
|
||||
/// .optimize(10, |trial| {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// Ok::<_, optimizer::TpeError>(x * x)
|
||||
/// Ok::<_, optimizer::Error>(x * x)
|
||||
/// })
|
||||
/// .unwrap();
|
||||
///
|
||||
@@ -431,7 +431,7 @@ where
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -455,7 +455,7 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -474,7 +474,7 @@ where
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// // Simulate async work (e.g., network request)
|
||||
/// let value = x * x;
|
||||
/// Ok::<_, optimizer::TpeError>((trial, value))
|
||||
/// Ok::<_, optimizer::Error>((trial, value))
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
@@ -511,7 +511,7 @@ where
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -536,8 +536,8 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `TpeError::TaskError` if the semaphore is closed or a spawned task panics.
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::TaskError` if the semaphore is closed or a spawned task panics.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -556,7 +556,7 @@ where
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// // Async objective function (e.g., network request)
|
||||
/// let value = x * x;
|
||||
/// Ok::<_, optimizer::TpeError>((trial, value))
|
||||
/// Ok::<_, optimizer::Error>((trial, value))
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
@@ -590,7 +590,7 @@ where
|
||||
.clone()
|
||||
.acquire_owned()
|
||||
.await
|
||||
.map_err(|e| crate::TpeError::TaskError(e.to_string()))?;
|
||||
.map_err(|e| crate::Error::TaskError(e.to_string()))?;
|
||||
let trial = self.create_trial();
|
||||
let objective = Arc::clone(&objective);
|
||||
|
||||
@@ -607,7 +607,7 @@ where
|
||||
for handle in handles {
|
||||
match handle
|
||||
.await
|
||||
.map_err(|e| crate::TpeError::TaskError(e.to_string()))?
|
||||
.map_err(|e| crate::Error::TaskError(e.to_string()))?
|
||||
{
|
||||
Ok((trial, value)) => {
|
||||
self.complete_trial(trial, value);
|
||||
@@ -620,7 +620,7 @@ where
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -643,9 +643,9 @@ where
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if no trials completed successfully
|
||||
/// Returns `Error::NoCompletedTrials` if no trials completed successfully
|
||||
/// before optimization stopped (either by completing all trials or early stopping).
|
||||
/// Returns `TpeError::Internal` if a completed trial is not found after adding (internal invariant violation).
|
||||
/// Returns `Error::Internal` if a completed trial is not found after adding (internal invariant violation).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -664,7 +664,7 @@ where
|
||||
/// 100,
|
||||
/// |trial| {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// Ok::<_, optimizer::TpeError>(x * x)
|
||||
/// Ok::<_, optimizer::Error>(x * x)
|
||||
/// },
|
||||
/// |_study, completed_trial| {
|
||||
/// // Stop early if we find a value less than 1.0
|
||||
@@ -702,7 +702,7 @@ where
|
||||
// Get the just-completed trial for the callback
|
||||
let trials = self.completed_trials.read();
|
||||
let Some(completed) = trials.last() else {
|
||||
return Err(crate::TpeError::Internal(
|
||||
return Err(crate::Error::Internal(
|
||||
"completed trial not found after adding",
|
||||
));
|
||||
};
|
||||
@@ -725,7 +725,7 @@ where
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -783,7 +783,7 @@ impl Study<f64> {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -798,7 +798,7 @@ impl Study<f64> {
|
||||
/// study
|
||||
/// .optimize_with_sampler(10, |trial| {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// Ok::<_, optimizer::TpeError>(x * x)
|
||||
/// Ok::<_, optimizer::Error>(x * x)
|
||||
/// })
|
||||
/// .unwrap();
|
||||
///
|
||||
@@ -829,7 +829,7 @@ impl Study<f64> {
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -851,8 +851,8 @@ impl Study<f64> {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if no trials completed successfully.
|
||||
/// Returns `TpeError::Internal` if a completed trial is not found after adding (internal invariant violation).
|
||||
/// Returns `Error::NoCompletedTrials` if no trials completed successfully.
|
||||
/// Returns `Error::Internal` if a completed trial is not found after adding (internal invariant violation).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -871,7 +871,7 @@ impl Study<f64> {
|
||||
/// 100,
|
||||
/// |trial| {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// Ok::<_, optimizer::TpeError>(x * x)
|
||||
/// Ok::<_, optimizer::Error>(x * x)
|
||||
/// },
|
||||
/// |study, _completed_trial| {
|
||||
/// // Stop after finding 5 good trials
|
||||
@@ -907,7 +907,7 @@ impl Study<f64> {
|
||||
// Get the just-completed trial for the callback
|
||||
let trials = self.completed_trials.read();
|
||||
let Some(completed) = trials.last() else {
|
||||
return Err(crate::TpeError::Internal(
|
||||
return Err(crate::Error::Internal(
|
||||
"completed trial not found after adding",
|
||||
));
|
||||
};
|
||||
@@ -930,7 +930,7 @@ impl Study<f64> {
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -953,7 +953,7 @@ impl Study<f64> {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -972,7 +972,7 @@ impl Study<f64> {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// // Simulate async work (e.g., network request)
|
||||
/// let value = x * x;
|
||||
/// Ok::<_, optimizer::TpeError>((trial, value))
|
||||
/// Ok::<_, optimizer::Error>((trial, value))
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
@@ -1009,7 +1009,7 @@ impl Study<f64> {
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -1034,8 +1034,8 @@ impl Study<f64> {
|
||||
///
|
||||
/// # Errors
|
||||
///
|
||||
/// Returns `TpeError::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `TpeError::TaskError` if the semaphore is closed or a spawned task panics.
|
||||
/// Returns `Error::NoCompletedTrials` if all trials failed (no successful trials).
|
||||
/// Returns `Error::TaskError` if the semaphore is closed or a spawned task panics.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
@@ -1054,7 +1054,7 @@ impl Study<f64> {
|
||||
/// let x = trial.suggest_float("x", -10.0, 10.0)?;
|
||||
/// // Async objective function (e.g., network request)
|
||||
/// let value = x * x;
|
||||
/// Ok::<_, optimizer::TpeError>((trial, value))
|
||||
/// Ok::<_, optimizer::Error>((trial, value))
|
||||
/// })
|
||||
/// .await?;
|
||||
///
|
||||
@@ -1087,7 +1087,7 @@ impl Study<f64> {
|
||||
.clone()
|
||||
.acquire_owned()
|
||||
.await
|
||||
.map_err(|e| crate::TpeError::TaskError(e.to_string()))?;
|
||||
.map_err(|e| crate::Error::TaskError(e.to_string()))?;
|
||||
let trial = self.create_trial_with_sampler();
|
||||
let objective = Arc::clone(&objective);
|
||||
|
||||
@@ -1104,7 +1104,7 @@ impl Study<f64> {
|
||||
for handle in handles {
|
||||
match handle
|
||||
.await
|
||||
.map_err(|e| crate::TpeError::TaskError(e.to_string()))?
|
||||
.map_err(|e| crate::Error::TaskError(e.to_string()))?
|
||||
{
|
||||
Ok((trial, value)) => {
|
||||
self.complete_trial(trial, value);
|
||||
@@ -1117,7 +1117,7 @@ impl Study<f64> {
|
||||
|
||||
// Return error if no trials succeeded
|
||||
if self.n_trials() == 0 {
|
||||
return Err(crate::TpeError::NoCompletedTrials);
|
||||
return Err(crate::Error::NoCompletedTrials);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
|
||||
+26
-32
@@ -8,7 +8,7 @@ use parking_lot::RwLock;
|
||||
use crate::distribution::{
|
||||
CategoricalDistribution, Distribution, FloatDistribution, IntDistribution,
|
||||
};
|
||||
use crate::error::{Result, TpeError};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::param::ParamValue;
|
||||
use crate::sampler::{CompletedTrial, Sampler};
|
||||
use crate::types::TrialState;
|
||||
@@ -190,7 +190,7 @@ impl Trial {
|
||||
/// ```
|
||||
pub fn suggest_float(&mut self, name: impl Into<String>, low: f64, high: f64) -> Result<f64> {
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds { low, high });
|
||||
return Err(Error::InvalidBounds { low, high });
|
||||
}
|
||||
|
||||
let name = name.into();
|
||||
@@ -216,7 +216,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -226,7 +226,7 @@ impl Trial {
|
||||
// Sample using the sampler
|
||||
let dist = Distribution::Float(distribution);
|
||||
let ParamValue::Float(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
return Err(Error::Internal(
|
||||
"Float distribution should return Float value",
|
||||
));
|
||||
};
|
||||
@@ -283,11 +283,11 @@ impl Trial {
|
||||
high: f64,
|
||||
) -> Result<f64> {
|
||||
if low <= 0.0 {
|
||||
return Err(TpeError::InvalidLogBounds);
|
||||
return Err(Error::InvalidLogBounds);
|
||||
}
|
||||
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds { low, high });
|
||||
return Err(Error::InvalidBounds { low, high });
|
||||
}
|
||||
|
||||
let name = name.into();
|
||||
@@ -313,7 +313,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -323,7 +323,7 @@ impl Trial {
|
||||
// Sample using the sampler (sampler handles log-scale transformation)
|
||||
let dist = Distribution::Float(distribution);
|
||||
let ParamValue::Float(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
return Err(Error::Internal(
|
||||
"Float distribution should return Float value",
|
||||
));
|
||||
};
|
||||
@@ -380,11 +380,11 @@ impl Trial {
|
||||
step: f64,
|
||||
) -> Result<f64> {
|
||||
if step <= 0.0 {
|
||||
return Err(TpeError::InvalidStep);
|
||||
return Err(Error::InvalidStep);
|
||||
}
|
||||
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds { low, high });
|
||||
return Err(Error::InvalidBounds { low, high });
|
||||
}
|
||||
|
||||
let name = name.into();
|
||||
@@ -410,7 +410,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -420,7 +420,7 @@ impl Trial {
|
||||
// Sample using the sampler (sampler handles step-grid)
|
||||
let dist = Distribution::Float(distribution);
|
||||
let ParamValue::Float(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
return Err(Error::Internal(
|
||||
"Float distribution should return Float value",
|
||||
));
|
||||
};
|
||||
@@ -466,7 +466,7 @@ impl Trial {
|
||||
#[allow(clippy::cast_precision_loss)]
|
||||
pub fn suggest_int(&mut self, name: impl Into<String>, low: i64, high: i64) -> Result<i64> {
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds {
|
||||
return Err(Error::InvalidBounds {
|
||||
low: low as f64,
|
||||
high: high as f64,
|
||||
});
|
||||
@@ -495,7 +495,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -505,9 +505,7 @@ impl Trial {
|
||||
// Sample using the sampler
|
||||
let dist = Distribution::Int(distribution);
|
||||
let ParamValue::Int(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
"Int distribution should return Int value",
|
||||
));
|
||||
return Err(Error::Internal("Int distribution should return Int value"));
|
||||
};
|
||||
|
||||
// Store distribution and value
|
||||
@@ -554,11 +552,11 @@ impl Trial {
|
||||
#[allow(clippy::cast_precision_loss)]
|
||||
pub fn suggest_int_log(&mut self, name: impl Into<String>, low: i64, high: i64) -> Result<i64> {
|
||||
if low < 1 {
|
||||
return Err(TpeError::InvalidLogBounds);
|
||||
return Err(Error::InvalidLogBounds);
|
||||
}
|
||||
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds {
|
||||
return Err(Error::InvalidBounds {
|
||||
low: low as f64,
|
||||
high: high as f64,
|
||||
});
|
||||
@@ -587,7 +585,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -597,9 +595,7 @@ impl Trial {
|
||||
// Sample using the sampler (sampler handles log-scale transformation)
|
||||
let dist = Distribution::Int(distribution);
|
||||
let ParamValue::Int(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
"Int distribution should return Int value",
|
||||
));
|
||||
return Err(Error::Internal("Int distribution should return Int value"));
|
||||
};
|
||||
|
||||
// Store distribution and value
|
||||
@@ -659,11 +655,11 @@ impl Trial {
|
||||
step: i64,
|
||||
) -> Result<i64> {
|
||||
if step <= 0 {
|
||||
return Err(TpeError::InvalidStep);
|
||||
return Err(Error::InvalidStep);
|
||||
}
|
||||
|
||||
if low > high {
|
||||
return Err(TpeError::InvalidBounds {
|
||||
return Err(Error::InvalidBounds {
|
||||
low: low as f64,
|
||||
high: high as f64,
|
||||
});
|
||||
@@ -692,7 +688,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different bounds or type"
|
||||
.to_string(),
|
||||
@@ -702,9 +698,7 @@ impl Trial {
|
||||
// Sample using the sampler (sampler handles step-grid)
|
||||
let dist = Distribution::Int(distribution);
|
||||
let ParamValue::Int(value) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
"Int distribution should return Int value",
|
||||
));
|
||||
return Err(Error::Internal("Int distribution should return Int value"));
|
||||
};
|
||||
|
||||
// Store distribution and value
|
||||
@@ -760,7 +754,7 @@ impl Trial {
|
||||
choices: &[T],
|
||||
) -> Result<T> {
|
||||
if choices.is_empty() {
|
||||
return Err(TpeError::EmptyChoices);
|
||||
return Err(Error::EmptyChoices);
|
||||
}
|
||||
|
||||
let name = name.into();
|
||||
@@ -779,7 +773,7 @@ impl Trial {
|
||||
}
|
||||
}
|
||||
// Distribution exists but doesn't match
|
||||
return Err(TpeError::ParameterConflict {
|
||||
return Err(Error::ParameterConflict {
|
||||
name,
|
||||
reason: "parameter was previously sampled with different number of choices or type"
|
||||
.to_string(),
|
||||
@@ -789,7 +783,7 @@ impl Trial {
|
||||
// Sample using the sampler
|
||||
let dist = Distribution::Categorical(distribution);
|
||||
let ParamValue::Categorical(index) = self.sample_value(&dist) else {
|
||||
return Err(TpeError::Internal(
|
||||
return Err(Error::Internal(
|
||||
"Categorical distribution should return Categorical value",
|
||||
));
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user