use std::ops::{Add, Mul, Sub};
use num::Float;
use crate::core::{Domain, Transformation};
use crate::dom::{AllDomain, InherentNullDomain, VectorDomain, OptionNullDomain};
use crate::error::Fallible;
use crate::dom::InherentNull;
use crate::samplers::SampleUniform;
use crate::trans::{make_row_by_row, make_row_by_row_fallible};
use crate::dist::SymmetricDistance;
use crate::traits::CheckNull;
pub fn make_impute_uniform_float<T>(
bounds: (T, T)
) -> Fallible<Transformation<VectorDomain<InherentNullDomain<AllDomain<T>>>, VectorDomain<AllDomain<T>>, SymmetricDistance, SymmetricDistance>>
where for<'a> T: 'static + Float + SampleUniform + Clone + Sub<Output=T> + Mul<&'a T, Output=T> + Add<&'a T, Output=T> + InherentNull + CheckNull {
let (lower, upper) = bounds;
if lower.is_nan() { return fallible!(MakeTransformation, "lower may not be nan"); }
if upper.is_nan() { return fallible!(MakeTransformation, "upper may not be nan"); }
if lower > upper { return fallible!(MakeTransformation, "lower may not be greater than upper") }
let scale = upper.clone() - lower.clone();
make_row_by_row_fallible(
InherentNullDomain::new(AllDomain::new()),
AllDomain::new(),
move |v| if v.is_null() {
T::sample_standard_uniform(false).map(|v| v * &scale + &lower)
} else { Ok(v.clone()) })
}
pub trait ImputableDomain: Domain {
type Imputed;
fn impute_constant<'a>(default: &'a Self::Carrier, constant: &'a Self::Imputed) -> &'a Self::Imputed;
fn new() -> Self;
}
impl<T: Clone + CheckNull> ImputableDomain for OptionNullDomain<AllDomain<T>> {
type Imputed = T;
fn impute_constant<'a>(default: &'a Self::Carrier, constant: &'a Self::Imputed) -> &'a Self::Imputed {
default.as_ref().unwrap_or(constant)
}
fn new() -> Self { OptionNullDomain::new(AllDomain::new()) }
}
impl<T: InherentNull + CheckNull> ImputableDomain for InherentNullDomain<AllDomain<T>> {
type Imputed = Self::Carrier;
fn impute_constant<'a>(default: &'a Self::Carrier, constant: &'a Self::Imputed) -> &'a Self::Imputed {
if default.is_null() { constant } else { default }
}
fn new() -> Self { InherentNullDomain::new(AllDomain::new()) }
}
pub fn make_impute_constant<DA>(
constant: DA::Imputed
) -> Fallible<Transformation<VectorDomain<DA>, VectorDomain<AllDomain<DA::Imputed>>, SymmetricDistance, SymmetricDistance>>
where DA: ImputableDomain,
DA::Imputed: 'static + Clone + CheckNull,
DA::Carrier: 'static {
if constant.is_null() { return fallible!(MakeTransformation, "Constant may not be null.") }
make_row_by_row(
DA::new(),
AllDomain::new(),
move |v| DA::impute_constant(v, &constant).clone())
}
#[cfg(test)]
mod tests {
use crate::error::ExplainUnwrap;
use crate::trans::{make_impute_constant, make_impute_uniform_float};
use crate::dom::{OptionNullDomain, InherentNullDomain};
#[test]
fn test_impute_uniform() {
let imputer = make_impute_uniform_float::<f64>((2.0, 2.0)).unwrap_test();
let result = imputer.invoke(&vec![1.0, f64::NAN]).unwrap_test();
assert_eq!(result, vec![1., 2.]);
assert!(imputer.stability_relation
.eval(&1, &1).unwrap_test());
}
#[test]
fn test_impute_constant_option() {
let imputer = make_impute_constant::<OptionNullDomain<_>>("IMPUTED".to_string()).unwrap_test();
let result = imputer.invoke(&vec![Some("A".to_string()), None]).unwrap_test();
assert_eq!(result, vec!["A".to_string(), "IMPUTED".to_string()]);
assert!(imputer.stability_relation
.eval(&1, &1).unwrap_test());
}
#[test]
fn test_impute_constant_inherent() {
let imputer = make_impute_constant::<InherentNullDomain<_>>(12.).unwrap_test();
let result = imputer.invoke(&vec![f64::NAN, 23.]).unwrap_test();
assert_eq!(result, vec![12., 23.]);
assert!(imputer.stability_relation
.eval(&1, &1).unwrap_test());
}
}