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 i in 0..3 {
181            for j in 0..3 {
182                let x = x_axis[i];
183                let y = y_axis[j];
184                data_vec.push(x + 2.0 * y);
185            }
186        }
187        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3]), data_vec).unwrap();
188        LookupTable::new(vec![x_axis, y_axis], data)
189    }
190
191    /// Helper to create a 3D table: f(x, y, z) = x + y + z
192    fn create_3d_linear() -> LookupTable {
193        let x_axis = vec![0.0, 1.0, 2.0];
194        let y_axis = vec![0.0, 1.0, 2.0];
195        let z_axis = vec![0.0, 1.0, 2.0];
196        let mut data_vec = Vec::new();
197        // C-order: for [nx, ny, nz], data[i*ny*nz + j*nz + k] = f(x[i], y[j], z[k])
198        for i in 0..3 {
199            for j in 0..3 {
200                for k in 0..3 {
201                    let x = x_axis[i];
202                    let y = y_axis[j];
203                    let z = z_axis[k];
204                    data_vec.push(x + y + z);
205                }
206            }
207        }
208        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3, 3]), data_vec).unwrap();
209        LookupTable::new(vec![x_axis, y_axis, z_axis], data)
210    }
211
212    /// Helper to create a 4D table: f(a, b, c, d) = a + b + c + d
213    fn create_4d_linear() -> LookupTable {
214        let axes = vec![
215            vec![0.0, 1.0, 2.0],
216            vec![0.0, 1.0, 2.0],
217            vec![0.0, 1.0, 2.0],
218            vec![0.0, 1.0, 2.0],
219        ];
220        let mut data_vec = Vec::new();
221        // C-order: for [na, nb, nc, nd]
222        for a in 0..3 {
223            for b in 0..3 {
224                for c in 0..3 {
225                    for d in 0..3 {
226                        let val = axes[0][a] + axes[1][b] + axes[2][c] + axes[3][d];
227                        data_vec.push(val);
228                    }
229                }
230            }
231        }
232        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3, 3, 3]), data_vec).unwrap();
233        LookupTable::new(axes, data)
234    }
235
236    #[test]
237    fn test_1d_exact_points() {
238        let table = create_linear_1d();
239        // Test at exact grid points
240        assert_eq!(table.interpolate(&[0.0]), 0.0);
241        assert_eq!(table.interpolate(&[1.0]), 1.0);
242        assert_eq!(table.interpolate(&[2.0]), 2.0);
243        assert_eq!(table.interpolate(&[3.0]), 3.0);
244        assert_eq!(table.interpolate(&[4.0]), 4.0);
245    }
246
247    #[test]
248    fn test_1d_interpolation() {
249        let table = create_linear_1d();
250        // Test interpolation at midpoints
251        assert!((table.interpolate(&[0.5]) - 0.5).abs() < 1e-10);
252        assert!((table.interpolate(&[1.5]) - 1.5).abs() < 1e-10);
253        assert!((table.interpolate(&[2.5]) - 2.5).abs() < 1e-10);
254        assert!((table.interpolate(&[3.5]) - 3.5).abs() < 1e-10);
255    }
256
257    #[test]
258    fn test_1d_boundary_clamp() {
259        let table = create_linear_1d();
260        // Values beyond boundaries should be clamped
261        assert_eq!(table.interpolate(&[-1.0]), 0.0);
262        assert_eq!(table.interpolate(&[5.0]), 4.0);
263    }
264
265    #[test]
266    fn test_2d_exact_points() {
267        let table = create_2d_linear();
268        // f(x, y) = x + 2*y
269        assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0);
270        assert_eq!(table.interpolate(&[1.0, 0.0]), 1.0);
271        assert_eq!(table.interpolate(&[0.0, 1.0]), 2.0);
272        assert_eq!(table.interpolate(&[1.0, 1.0]), 3.0);
273        assert_eq!(table.interpolate(&[2.0, 2.0]), 6.0);
274    }
275
276    #[test]
277    fn test_2d_interpolation() {
278        let table = create_2d_linear();
279        // f(x, y) = x + 2*y
280        // At (0.5, 0.5): should be 0.5 + 2*0.5 = 1.5
281        assert!((table.interpolate(&[0.5, 0.5]) - 1.5).abs() < 1e-10);
282        // At (1.5, 0.5): should be 1.5 + 2*0.5 = 2.5
283        assert!((table.interpolate(&[1.5, 0.5]) - 2.5).abs() < 1e-10);
284        // At (1.0, 1.5): should be 1.0 + 2*1.5 = 4.0
285        assert!((table.interpolate(&[1.0, 1.5]) - 4.0).abs() < 1e-10);
286    }
287
288    #[test]
289    fn test_3d_exact_points() {
290        let table = create_3d_linear();
291        // f(x, y, z) = x + y + z
292        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0]), 0.0);
293        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0]), 3.0);
294        assert_eq!(table.interpolate(&[2.0, 2.0, 2.0]), 6.0);
295        assert_eq!(table.interpolate(&[1.0, 2.0, 0.0]), 3.0);
296    }
297
298    #[test]
299    fn test_3d_interpolation() {
300        let table = create_3d_linear();
301        // f(x, y, z) = x + y + z
302        // At (0.5, 0.5, 0.5): should be 1.5
303        assert!((table.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
304        // At (1.5, 1.5, 1.5): should be 4.5
305        assert!((table.interpolate(&[1.5, 1.5, 1.5]) - 4.5).abs() < 1e-10);
306    }
307
308    #[test]
309    fn test_4d_exact_points() {
310        let table = create_4d_linear();
311        // f(a, b, c, d) = a + b + c + d
312        assert_eq!(table.interpolate(&[0.0, 0.0, 0.0, 0.0]), 0.0);
313        assert_eq!(table.interpolate(&[1.0, 1.0, 1.0, 1.0]), 4.0);
314        assert_eq!(table.interpolate(&[2.0, 2.0, 2.0, 2.0]), 8.0);
315    }
316
317    #[test]
318    fn test_4d_interpolation() {
319        let table = create_4d_linear();
320        // f(a, b, c, d) = a + b + c + d
321        // At (0.5, 0.5, 0.5, 0.5): should be 2.0
322        assert!((table.interpolate(&[0.5, 0.5, 0.5, 0.5]) - 2.0).abs() < 1e-10);
323        // At (1.5, 1.5, 1.5, 1.5): should be 6.0
324        assert!((table.interpolate(&[1.5, 1.5, 1.5, 1.5]) - 6.0).abs() < 1e-10);
325    }
326
327    #[test]
328    fn test_nonuniform_axes() {
329        // Create table with non-uniform axis spacing
330        let x_axis = vec![0.0, 0.1, 0.5, 2.0];
331        let data = ArrayD::from_shape_vec(IxDyn(&[4]), vec![0.0, 0.1, 0.5, 2.0]).unwrap();
332        let table = LookupTable::new(vec![x_axis], data);
333
334        // Test at exact points
335        assert_eq!(table.interpolate(&[0.0]), 0.0);
336        assert_eq!(table.interpolate(&[0.1]), 0.1);
337        assert_eq!(table.interpolate(&[0.5]), 0.5);
338        assert_eq!(table.interpolate(&[2.0]), 2.0);
339
340        // Test interpolation between first two points
341        let mid_01 = table.interpolate(&[0.05]);
342        assert!((mid_01 - 0.05).abs() < 1e-10);
343    }
344
345    #[test]
346    fn test_single_element_axis() {
347        // Test with a table that has only one element in a dimension
348        let axes = vec![
349            vec![0.0, 1.0],
350            vec![5.0], // Single element
351        ];
352        let data = ArrayD::from_shape_vec(IxDyn(&[2, 1]), vec![10.0, 20.0]).unwrap();
353        let table = LookupTable::new(axes, data);
354
355        // Should return the single value for any query in that dimension
356        assert_eq!(table.interpolate(&[0.0, 5.0]), 10.0);
357        assert_eq!(table.interpolate(&[0.5, 5.0]), 15.0);
358        assert_eq!(table.interpolate(&[1.0, 5.0]), 20.0);
359        // Out-of-bounds in single-element dimension should still work
360        assert_eq!(table.interpolate(&[0.5, 100.0]), 15.0);
361    }
362
363    #[test]
364    fn test_boundary_clamping() {
365        let table = create_2d_linear();
366        // Query outside bounds should clamp
367        let result_below = table.interpolate(&[-1.0, -1.0]);
368        assert_eq!(result_below, table.interpolate(&[0.0, 0.0]));
369
370        let result_above = table.interpolate(&[10.0, 10.0]);
371        assert_eq!(result_above, table.interpolate(&[2.0, 2.0]));
372
373        // Mixed: below in one dim, above in another
374        let result_mixed = table.interpolate(&[-1.0, 10.0]);
375        assert_eq!(result_mixed, table.interpolate(&[0.0, 2.0]));
376    }
377
378    #[test]
379    fn test_save_load_round_trip() {
380        use tempfile::TempDir;
381
382        let original = create_3d_linear();
383
384        let temp_dir = TempDir::new().expect("Failed to create temp dir");
385        let temp_file = temp_dir.path().join("test_lookup_table.bin");
386        original.save(&temp_file).expect("Save failed");
387        assert!(temp_file.exists(), "File should exist after save");
388
389        let loaded = LookupTable::load(&temp_file).expect("Load failed");
390        assert_eq!(loaded.axes.len(), original.axes.len());
391        for i in 0..loaded.axes.len() {
392            assert_eq!(loaded.axes[i], original.axes[i]);
393        }
394
395        // Test that loaded table produces same interpolation results
396        assert_eq!(loaded.interpolate(&[0.0, 0.0, 0.0]), 0.0);
397        assert_eq!(loaded.interpolate(&[1.5, 1.5, 1.5]), 4.5);
398        assert!((loaded.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
399    }
400
401    #[test]
402    fn test_axes_validation() {
403        // Test that constructor validates axes are strictly increasing
404        let axes_bad = vec![vec![0.0, 1.0, 1.0, 2.0]]; // Not strictly increasing!
405        let data = ArrayD::from_shape_vec(IxDyn(&[4]), vec![0.0, 1.0, 1.0, 2.0]).unwrap();
406        let result = std::panic::catch_unwind(|| LookupTable::new(axes_bad, data));
407        assert!(
408            result.is_err(),
409            "Should panic on non-strictly-increasing axis"
410        );
411    }
412
413    #[test]
414    fn test_dimension_mismatch() {
415        let axes = vec![vec![0.0, 1.0], vec![0.0, 1.0]];
416        let data = ArrayD::from_shape_vec(IxDyn(&[3, 3]), vec![0.0; 9]).unwrap(); // Wrong shape!
417        let result = std::panic::catch_unwind(|| LookupTable::new(axes, data));
418        assert!(result.is_err(), "Should panic on dimension mismatch");
419    }
420
421    #[test]
422    fn test_interpolate_dimension_mismatch() {
423        let table = create_2d_linear();
424        let result = std::panic::catch_unwind(|| table.interpolate(&[0.5])); // Only 1 arg for 2D table
425        assert!(result.is_err(), "Should panic on point dimension mismatch");
426    }
427
428    #[test]
429    fn test_large_grid_interpolation() {
430        // Test with larger grid to ensure performance isn't issues
431        let x_axis: Vec<f64> = (0..100).map(|i| i as f64 / 99.0).collect();
432        let data = ArrayD::from_shape_vec(IxDyn(&[100]), x_axis.clone()).unwrap();
433        let table = LookupTable::new(vec![x_axis], data);
434
435        // Spot checks
436        assert!((table.interpolate(&[0.0]) - 0.0).abs() < 1e-10);
437        assert!((table.interpolate(&[1.0]) - 1.0).abs() < 1e-10);
438        assert!((table.interpolate(&[0.5]) - 0.5).abs() < 1e-3); // Looser tolerance for large grid
439    }
440}