Skip to main content

solver/lookup/
grid.rs

1use super::table::LookupTable;
2use indicatif::ParallelProgressIterator;
3use ndarray::{ArrayD, IxDyn};
4use rayon::prelude::*;
5
6pub struct LookupGrid {
7    pub axes: Vec<Vec<f64>>,
8    pub names: Vec<String>,
9}
10
11impl LookupGrid {
12    pub fn new(axes: Vec<Vec<f64>>, names: Vec<String>) -> Self {
13        assert_eq!(axes.len(), names.len());
14        Self { axes, names }
15    }
16
17    /// Generates a LookupTable by evaluating the function `f` at each grid point.
18    /// Uses rayon for parallel execution.
19    pub fn generate<F>(&self, f: F) -> LookupTable
20    where
21        F: Fn(&[f64]) -> f64 + Sync + Send,
22    {
23        let shape: Vec<usize> = self.axes.iter().map(|axis| axis.len()).collect();
24        let total_points: usize = shape.iter().product();
25
26        // Create a flat list of points to iterate over in parallel
27        // This is a bit memory intensive for very large grids, but simple.
28        // Alternatively, we can map index to coordinates on the fly.
29
30        let data_vec: Vec<f64> = (0..total_points)
31            .into_par_iter()
32            .progress_count(total_points as u64)
33            .map(|idx| {
34                let mut coords = Vec::with_capacity(self.axes.len());
35                let mut temp_idx = idx;
36
37                // Convert flat index to multidimensional coordinates
38                // We need to do this in reverse order of strides if we want standard layout
39                // But ndarray default is C-order (row-major).
40                // Last index changes fastest.
41
42                // Let's pre-calculate strides to be safe or just compute indices.
43                // Actually, let's just compute the indices for each dimension.
44
45                let mut indices = vec![0; self.axes.len()];
46                for i in (0..self.axes.len()).rev() {
47                    indices[i] = temp_idx % shape[i];
48                    temp_idx /= shape[i];
49                }
50
51                for (i, &idx_dim) in indices.iter().enumerate() {
52                    coords.push(self.axes[i][idx_dim]);
53                }
54
55                f(&coords)
56            })
57            .collect();
58
59        let data = ArrayD::from_shape_vec(IxDyn(&shape), data_vec).unwrap();
60
61        LookupTable::new(self.axes.clone(), data)
62    }
63}
64
65pub fn linspace(start: f64, end: f64, n: usize) -> Vec<f64> {
66    if n == 1 {
67        return vec![start];
68    }
69    let step = (end - start) / (n as f64 - 1.0);
70    (0..n).map(|i| start + i as f64 * step).collect()
71}
72
73#[cfg(test)]
74mod tests {
75    use super::*;
76
77    #[test]
78    fn test_linspace_basic() {
79        let result = linspace(0.0, 1.0, 5);
80        assert_eq!(result.len(), 5);
81        assert!((result[0] - 0.0).abs() < 1e-10);
82        assert!((result[4] - 1.0).abs() < 1e-10);
83
84        // Check spacing is uniform
85        let diff01 = result[1] - result[0];
86        let diff12 = result[2] - result[1];
87        assert!((diff01 - diff12).abs() < 1e-10);
88    }
89
90    #[test]
91    fn test_linspace_single_point() {
92        let result = linspace(5.0, 10.0, 1);
93        assert_eq!(result.len(), 1);
94        assert_eq!(result[0], 5.0);
95    }
96
97    #[test]
98    fn test_linspace_two_points() {
99        let result = linspace(0.0, 1.0, 2);
100        assert_eq!(result.len(), 2);
101        assert_eq!(result[0], 0.0);
102        assert_eq!(result[1], 1.0);
103    }
104
105    #[test]
106    fn test_linspace_negative_range() {
107        let result = linspace(-1.0, 1.0, 5);
108        assert_eq!(result.len(), 5);
109        assert!((result[0] - (-1.0)).abs() < 1e-10);
110        assert!((result[4] - 1.0).abs() < 1e-10);
111        assert!((result[2] - 0.0).abs() < 1e-10); // Middle should be 0
112    }
113
114    #[test]
115    fn test_lookup_grid_1d() {
116        let axes = vec![linspace(0.0, 4.0, 5)];
117        let names = vec!["x".to_string()];
118        let grid = LookupGrid::new(axes, names);
119
120        // Identity function: f(x) = x
121        let table = grid.generate(|point| point[0]);
122
123        // Check at exact points
124        assert_eq!(table.interpolate(&[0.0]), 0.0);
125        assert_eq!(table.interpolate(&[1.0]), 1.0);
126        assert_eq!(table.interpolate(&[2.0]), 2.0);
127        assert_eq!(table.interpolate(&[3.0]), 3.0);
128        assert_eq!(table.interpolate(&[4.0]), 4.0);
129
130        // Check interpolation
131        assert!((table.interpolate(&[0.5]) - 0.5).abs() < 1e-10);
132        assert!((table.interpolate(&[2.5]) - 2.5).abs() < 1e-10);
133    }
134
135    #[test]
136    fn test_lookup_grid_2d() {
137        let x_axis = linspace(0.0, 2.0, 3);
138        let y_axis = linspace(0.0, 2.0, 3);
139        let axes = vec![x_axis, y_axis];
140        let names = vec!["x".to_string(), "y".to_string()];
141        let grid = LookupGrid::new(axes, names);
142
143        // Function: f(x, y) = x * y
144        let table = grid.generate(|point| point[0] * point[1]);
145
146        // Check at corners
147        assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0);
148        assert_eq!(table.interpolate(&[2.0, 2.0]), 4.0);
149
150        // Check at some interior points
151        assert!((table.interpolate(&[1.0, 1.0]) - 1.0).abs() < 1e-10);
152
153        // Check interpolation: at (0.5, 0.5) should be 0.25
154        assert!((table.interpolate(&[0.5, 0.5]) - 0.25).abs() < 1e-10);
155    }
156
157    #[test]
158    fn test_lookup_grid_3d() {
159        let x_axis = linspace(0.0, 1.0, 3);
160        let y_axis = linspace(0.0, 1.0, 3);
161        let z_axis = linspace(0.0, 1.0, 3);
162        let axes = vec![x_axis, y_axis, z_axis];
163        let names = vec!["x".to_string(), "y".to_string(), "z".to_string()];
164        let grid = LookupGrid::new(axes, names);
165
166        // Function: f(x, y, z) = x + y + z
167        let table = grid.generate(|point| point[0] + point[1] + point[2]);
168
169        // Check at corners
170        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0]), 0.0);
171        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0]), 3.0);
172
173        // Check interpolation at center
174        assert!((table.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
175    }
176
177    #[test]
178    fn test_lookup_grid_4d() {
179        let axes = vec![
180            linspace(0.0, 1.0, 3),
181            linspace(0.0, 1.0, 3),
182            linspace(0.0, 1.0, 3),
183            linspace(0.0, 1.0, 3),
184        ];
185        let names = vec![
186            "a".to_string(),
187            "b".to_string(),
188            "c".to_string(),
189            "d".to_string(),
190        ];
191        let grid = LookupGrid::new(axes, names);
192
193        // Function: f(a, b, c, d) = a + b + c + d
194        let table = grid.generate(|point| point.iter().sum());
195
196        // Check corners
197        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0, 0.0]), 0.0);
198        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0, 1.0]), 4.0);
199
200        // Check interpolation
201        assert!((table.interpolate(&[0.5, 0.5, 0.5, 0.5]) - 2.0).abs() < 1e-10);
202    }
203
204    #[test]
205    fn test_lookup_grid_complex_function() {
206        let x_axis = linspace(0.0, std::f64::consts::PI, 20);
207        let y_axis = linspace(0.0, std::f64::consts::PI, 20);
208        let axes = vec![x_axis, y_axis];
209        let names = vec!["x".to_string(), "y".to_string()];
210        let grid = LookupGrid::new(axes, names);
211
212        // Function: f(x, y) = sin(x) * cos(y)
213        let table = grid.generate(|point| point[0].sin() * point[1].cos());
214
215        // Check at exact points on the grid
216        assert!((table.interpolate(&[0.0, 0.0]) - 0.0).abs() < 1e-10); // sin(0) * cos(0) = 0
217
218        // Check near pi/2: at the actual grid point closest to pi/2
219        let pi_2 = std::f64::consts::PI / 2.0;
220        let result = table.interpolate(&[pi_2, 0.0]);
221        // sin(pi/2) * cos(0) = 1 * 1 = 1, with some tolerance for numerical precision
222        assert!(
223            (result - 1.0).abs() < 0.05,
224            "At (pi/2, 0) expected ~1.0, got {}",
225            result
226        );
227
228        // Spot check interpolation is reasonable (looser tolerance due to non-linear function)
229        let result = table.interpolate(&[std::f64::consts::PI / 4.0, std::f64::consts::PI / 4.0]);
230        let expected = (std::f64::consts::PI / 4.0).sin() * (std::f64::consts::PI / 4.0).cos();
231        assert!((result - expected).abs() < 0.05);
232    }
233
234    #[test]
235    fn test_lookup_grid_axes_names_mismatch() {
236        let axes = vec![vec![0.0, 1.0], vec![0.0, 1.0]];
237        let names = vec!["x".to_string()]; // Only 1 name for 2 axes
238        let result = std::panic::catch_unwind(|| LookupGrid::new(axes, names));
239        assert!(
240            result.is_err(),
241            "Should panic on axes/names length mismatch"
242        );
243    }
244
245    #[test]
246    fn test_lookup_grid_sparse() {
247        // Test with very small grid (minimal points)
248        let axes = vec![vec![0.0, 1.0], vec![0.0, 1.0]];
249        let names = vec!["x".to_string(), "y".to_string()];
250        let grid = LookupGrid::new(axes, names);
251
252        // f(x, y) = x^2 + y^2
253        let table = grid.generate(|point| point[0].powi(2) + point[1].powi(2));
254
255        // Corners
256        assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0); // 0 + 0
257        assert_eq!(table.interpolate(&[1.0, 0.0]), 1.0); // 1 + 0
258        assert_eq!(table.interpolate(&[0.0, 1.0]), 1.0); // 0 + 1
259        assert_eq!(table.interpolate(&[1.0, 1.0]), 2.0); // 1 + 1
260
261        // Interior: at (0.5, 0.5) with bilinear interpolation of x^2+y^2 on [0,1]^2
262        // The 4 corners are (0,0):0, (1,0):1, (0,1):1, (1,1):2
263        // Linear interpolation gives: 1.0
264        let result = table.interpolate(&[0.5, 0.5]);
265        assert!((result - 1.0).abs() < 0.01, "Expected ~1.0, got {}", result);
266    }
267}