use std::str::FromStr;
use strafe_testing::{
r::{RString, RTester},
r_assert_relative_equal, r_assert_relative_equal_result,
};
use strafe_type::FloatConstraint;
use crate::{
distribution::norm::NormalBuilder,
reset_statics,
rng::MarsagliaMulticarry,
traits::{Distribution, RNG},
};
pub fn density_inner(x: f64, mean: f64, standard_deviation: f64) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.density(x);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"dnorm({}, {}, {}, log={})",
RString::from_f64(x),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(false),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn log_density_inner(x: f64, mean: f64, standard_deviation: f64) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.log_density(x);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"dnorm({}, {}, {}, log={})",
RString::from_f64(x),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(true),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn probability_inner(q: f64, mean: f64, standard_deviation: f64, lower_tail: bool) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.probability(q, lower_tail);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"pnorm({}, {}, {}, lower.tail={}, log={})",
RString::from_f64(q),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(lower_tail),
RString::from_bool(false),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn log_probability_inner(q: f64, mean: f64, standard_deviation: f64, lower_tail: bool) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.log_probability(q, lower_tail);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"pnorm({}, {}, {}, lower.tail={}, log={})",
RString::from_f64(q),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(lower_tail),
RString::from_bool(true),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn quantile_inner(p: f64, mean: f64, standard_deviation: f64, lower_tail: bool) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.quantile(p, lower_tail);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"qnorm({}, {}, {}, lower.tail={}, log={})",
RString::from_f64(p),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(lower_tail),
RString::from_bool(false),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn log_quantile_inner(p: f64, mean: f64, standard_deviation: f64, lower_tail: bool) {
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
let rust_ans = norm.log_quantile(p, lower_tail);
let r_ans = RTester::new()
.set_display(&RString::from_string(format!(
"qnorm({}, {}, {}, lower.tail={}, log={})",
RString::from_f64(p),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
RString::from_bool(lower_tail),
RString::from_bool(true),
)))
.run()
.expect("R running error")
.as_f64();
r_assert_relative_equal_result!(r_ans, &rust_ans);
}
pub fn random_inner(seed: u16, mean: f64, standard_deviation: f64) {
let num = 50;
let r_state = RTester::new()
.set_seed(seed)
.set_display(&RString::from_str(".Random.seed").unwrap())
.run()
.unwrap()
.as_f64_vec()
.unwrap()
.into_iter()
.skip(1)
.map(|i| i as i32)
.collect::<Vec<_>>();
let mut rust_rng = MarsagliaMulticarry::new();
rust_rng.set_seed(seed as u32);
let rust_state = rust_rng.get_state();
for (r, rust) in r_state.iter().zip(rust_state.iter()) {
assert_eq!(r, rust, "States not equal!");
}
let r_ret = RTester::new()
.set_seed(seed)
.set_display(&RString::from_string(format!(
"rnorm({}, {}, {})",
RString::from_f64(num),
RString::from_f64(mean),
RString::from_f64(standard_deviation),
)))
.run()
.unwrap()
.as_f64_vec()
.unwrap();
let mut builder = NormalBuilder::new();
builder.with_mean(mean);
builder.with_standard_deviation(standard_deviation);
let norm = builder.build();
reset_statics();
let rust_ret = (0..num)
.map(|_| norm.random_sample(&mut rust_rng))
.collect::<Vec<_>>();
for (r, rust) in r_ret.iter().zip(rust_ret.iter()) {
r_assert_relative_equal!(*r, rust.unwrap());
}
}