1#[derive(Debug, Clone, Copy)]
5pub struct Weights {
6 pub w_nb: f32,
7 pub w_mn: f32,
8 pub w_fp: f32,
9}
10
11pub fn weights_at(t: usize, phase_iters: &[usize; 3]) -> Weights {
18 let p1 = phase_iters[0];
19 let p2 = phase_iters[1];
20
21 if t <= p1 {
22 let progress = if p1 > 1 { (t - 1) as f32 / (p1 - 1) as f32 } else { 1.0 };
24 let w_mn = 1000.0 * (1.0 - progress) + 3.0 * progress;
25 Weights { w_nb: 2.0, w_mn, w_fp: 1.0 }
26 } else if t <= p1 + p2 {
27 Weights { w_nb: 3.0, w_mn: 3.0, w_fp: 1.0 }
29 } else {
30 Weights { w_nb: 1.0, w_mn: 0.0, w_fp: 1.0 }
32 }
33}
34
35#[cfg(test)]
36mod tests {
37 use super::*;
38 use approx::assert_abs_diff_eq;
39
40 #[test]
41 fn phase_1_start() {
42 let w = weights_at(1, &[100, 100, 250]);
43 assert_abs_diff_eq!(w.w_nb, 2.0);
44 assert_abs_diff_eq!(w.w_mn, 1000.0, epsilon = 0.01);
45 assert_abs_diff_eq!(w.w_fp, 1.0);
46 }
47
48 #[test]
49 fn phase_1_end() {
50 let w = weights_at(100, &[100, 100, 250]);
51 assert_abs_diff_eq!(w.w_mn, 3.0, epsilon = 0.01);
52 }
53
54 #[test]
55 fn phase_2() {
56 let w = weights_at(150, &[100, 100, 250]);
57 assert_abs_diff_eq!(w.w_nb, 3.0);
58 assert_abs_diff_eq!(w.w_mn, 3.0);
59 }
60
61 #[test]
62 fn phase_3() {
63 let w = weights_at(300, &[100, 100, 250]);
64 assert_abs_diff_eq!(w.w_nb, 1.0);
65 assert_abs_diff_eq!(w.w_mn, 0.0);
66 assert_abs_diff_eq!(w.w_fp, 1.0);
67 }
68}