fix: return error instead of panicking on out-of-bounds categorical index
This commit is contained in:
+29
-2
@@ -467,7 +467,11 @@ impl<T: Clone + Debug> Parameter for CategoricalParam<T> {
|
|||||||
|
|
||||||
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
|
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
|
||||||
match param_value {
|
match param_value {
|
||||||
ParamValue::Categorical(index) => Ok(self.choices[*index].clone()),
|
ParamValue::Categorical(index) => self
|
||||||
|
.choices
|
||||||
|
.get(*index)
|
||||||
|
.cloned()
|
||||||
|
.ok_or(Error::Internal("categorical index out of bounds")),
|
||||||
_ => Err(Error::Internal(
|
_ => Err(Error::Internal(
|
||||||
"Categorical distribution should return Categorical value",
|
"Categorical distribution should return Categorical value",
|
||||||
)),
|
)),
|
||||||
@@ -708,7 +712,8 @@ impl<T: Categorical + Debug> Parameter for EnumParam<T> {
|
|||||||
|
|
||||||
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
|
fn cast_param_value(&self, param_value: &ParamValue) -> Result<T> {
|
||||||
match param_value {
|
match param_value {
|
||||||
ParamValue::Categorical(index) => Ok(T::from_index(*index)),
|
ParamValue::Categorical(index) if *index < T::N_CHOICES => Ok(T::from_index(*index)),
|
||||||
|
ParamValue::Categorical(_) => Err(Error::Internal("categorical index out of bounds")),
|
||||||
_ => Err(Error::Internal(
|
_ => Err(Error::Internal(
|
||||||
"Categorical distribution should return Categorical value",
|
"Categorical distribution should return Categorical value",
|
||||||
)),
|
)),
|
||||||
@@ -887,6 +892,17 @@ mod tests {
|
|||||||
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
|
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn categorical_param_cast_out_of_bounds() {
|
||||||
|
let param = CategoricalParam::new(vec!["sgd", "adam", "rmsprop"]);
|
||||||
|
assert!(param.cast_param_value(&ParamValue::Categorical(3)).is_err());
|
||||||
|
assert!(
|
||||||
|
param
|
||||||
|
.cast_param_value(&ParamValue::Categorical(usize::MAX))
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn bool_param_distribution() {
|
fn bool_param_distribution() {
|
||||||
let param = BoolParam::new();
|
let param = BoolParam::new();
|
||||||
@@ -955,6 +971,17 @@ mod tests {
|
|||||||
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
|
assert!(param.cast_param_value(&ParamValue::Float(1.0)).is_err());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn enum_param_cast_out_of_bounds() {
|
||||||
|
let param = EnumParam::<TestEnum>::new();
|
||||||
|
assert!(param.cast_param_value(&ParamValue::Categorical(3)).is_err());
|
||||||
|
assert!(
|
||||||
|
param
|
||||||
|
.cast_param_value(&ParamValue::Categorical(usize::MAX))
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn float_param_suggest_via_trial() {
|
fn float_param_suggest_via_trial() {
|
||||||
let param = FloatParam::new(0.0, 1.0);
|
let param = FloatParam::new(0.0, 1.0);
|
||||||
|
|||||||
Reference in New Issue
Block a user