Skip to main content

solver/lookup/
table.rs

1use ndarray::ArrayD;
2use serde::{Deserialize, Serialize};
3use std::fs::File;
4use std::io::{BufReader, BufWriter};
5use std::path::Path;
6
7#[derive(Clone, Serialize, Deserialize, Debug)]
8pub struct LookupTable {
9    /// The grid coordinates for each dimension.
10    /// axes\[0] corresponds to the first dimension of `data`, etc.
11    pub axes: Vec<Vec<f64>>,
12    /// The data stored in a dynamic-dimensional array.
13    pub data: ArrayD<f64>,
14    /// Cached flat strides for each dimension (C-order).
15    /// strides[i] = product of axis lengths for dimensions i+1 .. ndim.
16    /// Precomputed in `new` so `interpolate` avoids repeated multiplications.
17    #[serde(skip)]
18    strides: Vec<usize>,
19}
20
21impl LookupTable {
22    pub fn new(axes: Vec<Vec<f64>>, data: ArrayD<f64>) -> Self {
23        assert_eq!(
24            axes.len(),
25            data.ndim(),
26            "Number of axes must match data dimensions"
27        );
28        for (i, axis) in axes.iter().enumerate() {
29            assert_eq!(
30                axis.len(),
31                data.shape()[i],
32                "Axis length must match data shape at dim {}",
33                i
34            );
35            for j in 0..axis.len() - 1 {
36                assert!(
37                    axis[j] < axis[j + 1],
38                    "Axes must be sorted strictly increasing"
39                );
40            }
41        }
42
43        // Precompute C-order strides
44        let ndim = axes.len();
45        let mut strides = vec![1usize; ndim];
46        for i in (0..ndim.saturating_sub(1)).rev() {
47            strides[i] = strides[i + 1] * axes[i + 1].len();
48        }
49
50        Self {
51            axes,
52            data,
53            strides,
54        }
55    }
56
57    pub fn save(&self, path: &Path) -> bincode::Result<()> {
58        let file = File::create(path)?;
59        let writer = BufWriter::new(file);
60        bincode::serialize_into(writer, self)
61    }
62
63    pub fn load(path: &Path) -> bincode::Result<Self> {
64        let file = File::open(path)?;
65        let reader = BufReader::new(file);
66        let mut t: Self = bincode::deserialize_from(reader)?;
67        // Recompute derived fields that are not serialized
68        let ndim = t.axes.len();
69        let mut strides = vec![1usize; ndim];
70        for i in (0..ndim.saturating_sub(1)).rev() {
71            strides[i] = strides[i + 1] * t.axes[i + 1].len();
72        }
73        t.strides = strides;
74        Ok(t)
75    }
76
77    /// Multilinear interpolation at the given point.
78    /// `point` must have the same number of dimensions as `axes`.
79    ///
80    /// For uniform axes (produced by `linspace`) the bracket index is computed
81    /// in O(1) via direct division instead of binary search, avoiding a
82    /// per-call heap allocation.
83    pub fn interpolate(&self, point: &[f64]) -> f64 {
84        let ndim = self.axes.len();
85        assert_eq!(point.len(), ndim, "Point dimension mismatch");
86        let flat = self
87            .data
88            .as_slice_memory_order()
89            .expect("LookupTable data must be contiguous in memory order");
90
91        // Stack-allocated arrays (ndim <= 16 is always true for this use case)
92        let mut indices = [0usize; 16];
93        let mut weights = [0.0f64; 16];
94
95        for i in 0..ndim {
96            let axis = &self.axes[i];
97            let n = axis.len();
98            let val = point[i];
99
100            if n == 1 {
101                // indices[i] = 0, weights[i] = 0.0 already
102                continue;
103            }
104
105            if val <= axis[0] {
106                // weights[i] = 0.0 already
107            } else if val >= axis[n - 1] {
108                indices[i] = n - 2;
109                weights[i] = 1.0;
110            } else {
111                // Fast path for uniform axes: O(1) direct computation.
112                // A uniform axis has constant spacing; check if step is uniform
113                // by comparing first and last intervals.
114                let step0 = axis[1] - axis[0];
115                let step_last = axis[n - 1] - axis[n - 2];
116                let idx = if (step0 - step_last).abs() < step0 * 1e-9 {
117                    // Uniform: direct index
118                    let raw = (val - axis[0]) / step0;
119                    (raw as usize).min(n - 2)
120                } else {
121                    // Non-uniform: binary search
122                    let idx = match axis.binary_search_by(|x| x.partial_cmp(&val).unwrap()) {
123                        Ok(exact) => exact,
124                        Err(insert) => insert - 1,
125                    };
126                    idx.min(n - 2)
127                };
128                let w = (val - axis[idx]) / (axis[idx + 1] - axis[idx]);
129                indices[i] = idx;
130                weights[i] = w;
131            }
132        }
133
134        // Iterate over 2^ndim corners using precomputed strides and flat buffer
135        let num_corners = 1usize << ndim;
136        let mut result = 0.0;
137
138        for corner in 0..num_corners {
139            let mut w = 1.0f64;
140            let mut flat_idx = 0usize;
141
142            for dim in 0..ndim {
143                let is_upper = (corner >> dim) & 1 == 1;
144                if is_upper {
145                    w *= weights[dim];
146                    flat_idx +=
147                        (indices[dim] + 1).min(self.axes[dim].len() - 1) * self.strides[dim];
148                } else {
149                    w *= 1.0 - weights[dim];
150                    flat_idx += indices[dim] * self.strides[dim];
151                }
152            }
153
154            result += flat[flat_idx] * w;
155        }
156
157        result
158    }
159}
160
161#[cfg(test)]
162mod tests {
163    use super::*;
164    use ndarray::IxDyn;
165
166    /// Helper to create a 1D linear table for testing interpolation
167    fn create_linear_1d() -> LookupTable {
168        let axes = vec![vec![0.0, 1.0, 2.0, 3.0, 4.0]];
169        let data = ArrayD::from_shape_vec(IxDyn(&[5]), vec![0.0, 1.0, 2.0, 3.0, 4.0]).unwrap();
170        LookupTable::new(axes, data)
171    }
172
173    /// Helper to create a 2D table: f(x, y) = x + 2*y
174    fn create_2d_linear() -> LookupTable {
175        let x_axis = vec![0.0, 1.0, 2.0];
176        let y_axis = vec![0.0, 1.0, 2.0];
177        let mut data_vec = Vec::new();
178        // C-order: last index changes fastest
179        // For [nx, ny] shape, data[i*ny + j] = f(x[i], y[j])
180        for &x in &x_axis {
181            for &y in &y_axis {
182                data_vec.push(x + 2.0 * y);
183            }
184        }
185        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3]), data_vec).unwrap();
186        LookupTable::new(vec![x_axis, y_axis], data)
187    }
188
189    /// Helper to create a 3D table: f(x, y, z) = x + y + z
190    fn create_3d_linear() -> LookupTable {
191        let x_axis = vec![0.0, 1.0, 2.0];
192        let y_axis = vec![0.0, 1.0, 2.0];
193        let z_axis = vec![0.0, 1.0, 2.0];
194        let mut data_vec = Vec::new();
195        // C-order: for [nx, ny, nz], data[i*ny*nz + j*nz + k] = f(x[i], y[j], z[k])
196        for &x in &x_axis {
197            for &y in &y_axis {
198                for &z in &z_axis {
199                    data_vec.push(x + y + z);
200                }
201            }
202        }
203        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3, 3]), data_vec).unwrap();
204        LookupTable::new(vec![x_axis, y_axis, z_axis], data)
205    }
206
207    /// Helper to create a 4D table: f(a, b, c, d) = a + b + c + d
208    fn create_4d_linear() -> LookupTable {
209        let axes = vec![
210            vec![0.0, 1.0, 2.0],
211            vec![0.0, 1.0, 2.0],
212            vec![0.0, 1.0, 2.0],
213            vec![0.0, 1.0, 2.0],
214        ];
215        let mut data_vec = Vec::new();
216        // C-order: for [na, nb, nc, nd]
217        for a in 0..3 {
218            for b in 0..3 {
219                for c in 0..3 {
220                    for d in 0..3 {
221                        let val = axes[0][a] + axes[1][b] + axes[2][c] + axes[3][d];
222                        data_vec.push(val);
223                    }
224                }
225            }
226        }
227        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3, 3, 3]), data_vec).unwrap();
228        LookupTable::new(axes, data)
229    }
230
231    #[test]
232    fn test_1d_exact_points() {
233        let table = create_linear_1d();
234        // Test at exact grid points
235        assert_eq!(table.interpolate(&[0.0]), 0.0);
236        assert_eq!(table.interpolate(&[1.0]), 1.0);
237        assert_eq!(table.interpolate(&[2.0]), 2.0);
238        assert_eq!(table.interpolate(&[3.0]), 3.0);
239        assert_eq!(table.interpolate(&[4.0]), 4.0);
240    }
241
242    #[test]
243    fn test_1d_interpolation() {
244        let table = create_linear_1d();
245        // Test interpolation at midpoints
246        assert!((table.interpolate(&[0.5]) - 0.5).abs() < 1e-10);
247        assert!((table.interpolate(&[1.5]) - 1.5).abs() < 1e-10);
248        assert!((table.interpolate(&[2.5]) - 2.5).abs() < 1e-10);
249        assert!((table.interpolate(&[3.5]) - 3.5).abs() < 1e-10);
250    }
251
252    #[test]
253    fn test_1d_boundary_clamp() {
254        let table = create_linear_1d();
255        // Values beyond boundaries should be clamped
256        assert_eq!(table.interpolate(&[-1.0]), 0.0);
257        assert_eq!(table.interpolate(&[5.0]), 4.0);
258    }
259
260    #[test]
261    fn test_2d_exact_points() {
262        let table = create_2d_linear();
263        // f(x, y) = x + 2*y
264        assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0);
265        assert_eq!(table.interpolate(&[1.0, 0.0]), 1.0);
266        assert_eq!(table.interpolate(&[0.0, 1.0]), 2.0);
267        assert_eq!(table.interpolate(&[1.0, 1.0]), 3.0);
268        assert_eq!(table.interpolate(&[2.0, 2.0]), 6.0);
269    }
270
271    #[test]
272    fn test_2d_interpolation() {
273        let table = create_2d_linear();
274        // f(x, y) = x + 2*y
275        // At (0.5, 0.5): should be 0.5 + 2*0.5 = 1.5
276        assert!((table.interpolate(&[0.5, 0.5]) - 1.5).abs() < 1e-10);
277        // At (1.5, 0.5): should be 1.5 + 2*0.5 = 2.5
278        assert!((table.interpolate(&[1.5, 0.5]) - 2.5).abs() < 1e-10);
279        // At (1.0, 1.5): should be 1.0 + 2*1.5 = 4.0
280        assert!((table.interpolate(&[1.0, 1.5]) - 4.0).abs() < 1e-10);
281    }
282
283    #[test]
284    fn test_3d_exact_points() {
285        let table = create_3d_linear();
286        // f(x, y, z) = x + y + z
287        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0]), 0.0);
288        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0]), 3.0);
289        assert_eq!(table.interpolate(&[2.0, 2.0, 2.0]), 6.0);
290        assert_eq!(table.interpolate(&[1.0, 2.0, 0.0]), 3.0);
291    }
292
293    #[test]
294    fn test_3d_interpolation() {
295        let table = create_3d_linear();
296        // f(x, y, z) = x + y + z
297        // At (0.5, 0.5, 0.5): should be 1.5
298        assert!((table.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
299        // At (1.5, 1.5, 1.5): should be 4.5
300        assert!((table.interpolate(&[1.5, 1.5, 1.5]) - 4.5).abs() < 1e-10);
301    }
302
303    #[test]
304    fn test_4d_exact_points() {
305        let table = create_4d_linear();
306        // f(a, b, c, d) = a + b + c + d
307        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0, 0.0]), 0.0);
308        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0, 1.0]), 4.0);
309        assert_eq!(table.interpolate(&[2.0, 2.0, 2.0, 2.0]), 8.0);
310    }
311
312    #[test]
313    fn test_4d_interpolation() {
314        let table = create_4d_linear();
315        // f(a, b, c, d) = a + b + c + d
316        // At (0.5, 0.5, 0.5, 0.5): should be 2.0
317        assert!((table.interpolate(&[0.5, 0.5, 0.5, 0.5]) - 2.0).abs() < 1e-10);
318        // At (1.5, 1.5, 1.5, 1.5): should be 6.0
319        assert!((table.interpolate(&[1.5, 1.5, 1.5, 1.5]) - 6.0).abs() < 1e-10);
320    }
321
322    #[test]
323    fn test_nonuniform_axes() {
324        // Create table with non-uniform axis spacing
325        let x_axis = vec![0.0, 0.1, 0.5, 2.0];
326        let data = ArrayD::from_shape_vec(IxDyn(&[4]), vec![0.0, 0.1, 0.5, 2.0]).unwrap();
327        let table = LookupTable::new(vec![x_axis], data);
328
329        // Test at exact points
330        assert_eq!(table.interpolate(&[0.0]), 0.0);
331        assert_eq!(table.interpolate(&[0.1]), 0.1);
332        assert_eq!(table.interpolate(&[0.5]), 0.5);
333        assert_eq!(table.interpolate(&[2.0]), 2.0);
334
335        // Test interpolation between first two points
336        let mid_01 = table.interpolate(&[0.05]);
337        assert!((mid_01 - 0.05).abs() < 1e-10);
338    }
339
340    #[test]
341    fn test_single_element_axis() {
342        // Test with a table that has only one element in a dimension
343        let axes = vec![
344            vec![0.0, 1.0],
345            vec![5.0], // Single element
346        ];
347        let data = ArrayD::from_shape_vec(IxDyn(&[2, 1]), vec![10.0, 20.0]).unwrap();
348        let table = LookupTable::new(axes, data);
349
350        // Should return the single value for any query in that dimension
351        assert_eq!(table.interpolate(&[0.0, 5.0]), 10.0);
352        assert_eq!(table.interpolate(&[0.5, 5.0]), 15.0);
353        assert_eq!(table.interpolate(&[1.0, 5.0]), 20.0);
354        // Out-of-bounds in single-element dimension should still work
355        assert_eq!(table.interpolate(&[0.5, 100.0]), 15.0);
356    }
357
358    #[test]
359    fn test_boundary_clamping() {
360        let table = create_2d_linear();
361        // Query outside bounds should clamp
362        let result_below = table.interpolate(&[-1.0, -1.0]);
363        assert_eq!(result_below, table.interpolate(&[0.0, 0.0]));
364
365        let result_above = table.interpolate(&[10.0, 10.0]);
366        assert_eq!(result_above, table.interpolate(&[2.0, 2.0]));
367
368        // Mixed: below in one dim, above in another
369        let result_mixed = table.interpolate(&[-1.0, 10.0]);
370        assert_eq!(result_mixed, table.interpolate(&[0.0, 2.0]));
371    }
372
373    #[test]
374    fn test_save_load_round_trip() {
375        use tempfile::TempDir;
376
377        let original = create_3d_linear();
378
379        let temp_dir = TempDir::new().expect("Failed to create temp dir");
380        let temp_file = temp_dir.path().join("test_lookup_table.bin");
381        original.save(&temp_file).expect("Save failed");
382        assert!(temp_file.exists(), "File should exist after save");
383
384        let loaded = LookupTable::load(&temp_file).expect("Load failed");
385        assert_eq!(loaded.axes.len(), original.axes.len());
386        for i in 0..loaded.axes.len() {
387            assert_eq!(loaded.axes[i], original.axes[i]);
388        }
389
390        // Test that loaded table produces same interpolation results
391        assert_eq!(loaded.interpolate(&[0.0, 0.0, 0.0]), 0.0);
392        assert_eq!(loaded.interpolate(&[1.5, 1.5, 1.5]), 4.5);
393        assert!((loaded.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
394    }
395
396    #[test]
397    fn test_axes_validation() {
398        // Test that constructor validates axes are strictly increasing
399        let axes_bad = vec![vec![0.0, 1.0, 1.0, 2.0]]; // Not strictly increasing!
400        let data = ArrayD::from_shape_vec(IxDyn(&[4]), vec![0.0, 1.0, 1.0, 2.0]).unwrap();
401        let result = std::panic::catch_unwind(|| LookupTable::new(axes_bad, data));
402        assert!(
403            result.is_err(),
404            "Should panic on non-strictly-increasing axis"
405        );
406    }
407
408    #[test]
409    fn test_dimension_mismatch() {
410        let axes = vec![vec![0.0, 1.0], vec![0.0, 1.0]];
411        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3]), vec![0.0; 9]).unwrap(); // Wrong shape!
412        let result = std::panic::catch_unwind(|| LookupTable::new(axes, data));
413        assert!(result.is_err(), "Should panic on dimension mismatch");
414    }
415
416    #[test]
417    fn test_interpolate_dimension_mismatch() {
418        let table = create_2d_linear();
419        let result = std::panic::catch_unwind(|| table.interpolate(&[0.5])); // Only 1 arg for 2D table
420        assert!(result.is_err(), "Should panic on point dimension mismatch");
421    }
422
423    #[test]
424    fn test_large_grid_interpolation() {
425        // Test with larger grid to ensure performance isn't issues
426        let x_axis: Vec<f64> = (0..100).map(|i| i as f64 / 99.0).collect();
427        let data = ArrayD::from_shape_vec(IxDyn(&[100]), x_axis.clone()).unwrap();
428        let table = LookupTable::new(vec![x_axis], data);
429
430        // Spot checks
431        assert!((table.interpolate(&[0.0]) - 0.0).abs() < 1e-10);
432        assert!((table.interpolate(&[1.0]) - 1.0).abs() < 1e-10);
433        assert!((table.interpolate(&[0.5]) - 0.5).abs() < 1e-3); // Looser tolerance for large grid
434    }
435}