Skip to main content

renegade_ml/
predict.rs

1use crate::neighbor::Neighbor;
2
3/// Result of extrapolating neighbor outputs to distance=0.
4#[derive(Debug, Clone)]
5pub struct ExtrapolatedPrediction {
6    /// Predicted output at distance=0.
7    pub value: f64,
8    /// R² of the linear fit. Low values mean the trend is unreliable.
9    pub r_squared: f64,
10    /// Number of neighbors used.
11    pub k: usize,
12}
13
14impl ExtrapolatedPrediction {
15    pub(crate) fn from_neighbors(neighbors: &[Neighbor]) -> Self {
16        let k = neighbors.len();
17
18        if k == 0 {
19            return ExtrapolatedPrediction {
20                value: f64::NAN,
21                r_squared: 0.0,
22                k: 0,
23            };
24        }
25
26        if k == 1 {
27            return ExtrapolatedPrediction {
28                value: neighbors[0].output,
29                r_squared: 0.0,
30                k: 1,
31            };
32        }
33
34        // Linear regression: output = a + b * distance
35        // Extrapolated prediction = a (intercept at distance=0)
36        let n = k as f64;
37        let sum_x: f64 = neighbors.iter().map(|n| n.distance).sum();
38        let sum_y: f64 = neighbors.iter().map(|n| n.output).sum();
39        let sum_xy: f64 = neighbors.iter().map(|n| n.distance * n.output).sum();
40        let sum_xx: f64 = neighbors.iter().map(|n| n.distance * n.distance).sum();
41
42        let denom = n * sum_xx - sum_x * sum_x;
43
44        if denom.abs() < 1e-15 {
45            // All neighbors at the same distance — just average.
46            return ExtrapolatedPrediction {
47                value: sum_y / n,
48                r_squared: 0.0,
49                k,
50            };
51        }
52
53        let b = (n * sum_xy - sum_x * sum_y) / denom;
54        let a = (sum_y - b * sum_x) / n;
55
56        // R² calculation
57        let mean_y = sum_y / n;
58        let ss_tot: f64 = neighbors.iter().map(|n| (n.output - mean_y).powi(2)).sum();
59        let ss_res: f64 = neighbors
60            .iter()
61            .map(|n| {
62                let predicted = a + b * n.distance;
63                (n.output - predicted).powi(2)
64            })
65            .sum();
66
67        let r_squared = if ss_tot > 1e-15 {
68            1.0 - ss_res / ss_tot
69        } else {
70            1.0 // All outputs identical — perfect "fit"
71        };
72
73        ExtrapolatedPrediction {
74            value: a,
75            r_squared,
76            k,
77        }
78    }
79}