use super::traits::{Function, Number};
pub fn solver<T, F>(func: F,
initial_conditions: Vec<T>,
time_interval: &[T; 2],
step: T,
weights: &Vec<T>,
weight_sum: T)
-> (Vec<T>, Vec<Vec<T>>)
where T: Number,
F: Function<T> {
let mut time_stamps: Vec<T> = vec![time_interval[0]];
let mut calculated_vals: Vec<Vec<T>> = Vec::with_capacity(
initial_conditions.len());
for i in 0..initial_conditions.len() {
calculated_vals.push(vec![]);
calculated_vals[i].push(initial_conditions[i]);
}
let mut current_vals: Vec<T> = initial_conditions.clone();
let mut current_time: T = time_stamps[time_stamps.len() - 1];
let t_2 = T::from_str_radix("2", 10).ok().unwrap();
while current_time + step < time_interval[1] {
let mut k1: Vec<T> = func(¤t_time, ¤t_vals);
for i in 0..k1.len() {
k1[i] = k1[i] * step;
}
let mut currvls_k2 = current_vals.clone();
for i in 0..currvls_k2.len() {
currvls_k2[i] = currvls_k2[i] + k1[i]/t_2;
}
let mut k2: Vec<T> = func(&(current_time + step/t_2), &currvls_k2);
for i in 0..k2.len() {
k2[i] = k2[i] * step;
}
let mut currvls_k3 = current_vals.clone();
for i in 0..currvls_k3.len() {
currvls_k3[i] = currvls_k3[i] + k2[i]/t_2;
}
let mut k3: Vec<T> = func(&(current_time + step/t_2), &currvls_k3);
for i in 0..k3.len() {
k3[i] = k3[i] * step;
}
let mut currvls_k4 = current_vals.clone();
for i in 0..currvls_k4.len() {
currvls_k4[i] = currvls_k4[i] + k3[i];
}
let mut k4: Vec<T> = func(&(current_time + step), &currvls_k4);
for i in 0..k4.len() {
k4[i] = k4[i] * step;
}
let mut curr_point = current_vals.clone();
for i in 0..curr_point.len() {
curr_point[i] = curr_point[i] + step*(
k1[i]*weights[0] + k2[i]*weights[1] + k3[i]*weights[2]
+ k4[i]*weights[3]
)/weight_sum;
}
for i in 0..calculated_vals.len() {
current_vals[i] = calculated_vals[i].last().unwrap().clone();
}
for i in 0..calculated_vals.len() {
calculated_vals[i].push(curr_point[i]);
current_vals[i] = calculated_vals[i].last().unwrap().clone();
}
current_time = time_stamps[time_stamps.len() - 1] + step;
time_stamps.push(current_time);
}
(time_stamps, calculated_vals)
}
#[cfg(test)]
mod tests {
use super::solver;
#[test]
fn integrate_2_t() {
let start = 0;
let end = 500;
let (_, num_sol) = solver(|t: &f32, _: &Vec<f32>| vec![2.*t],
vec![0.],
&[start as f32, end as f32],
1.0,
&vec![1., 2., 2., 1.],
6.);
let mut an_sol: Vec<u32> = vec![];
for i in start..end {
an_sol.push(i*i);
}
let mut num_sol_u32: Vec<u32> = Vec::new();
for el in num_sol[0].clone() {
let el_cl = el.clone();
num_sol_u32.push(el_cl as u32);
}
assert_eq!(num_sol[0].len(), an_sol.len());
assert_eq!(num_sol_u32, an_sol);
}
}