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 pub axes: Vec<Vec<f64>>,
12 pub data: ArrayD<f64>,
14 #[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 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 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 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 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 continue;
103 }
104
105 if val <= axis[0] {
106 } else if val >= axis[n - 1] {
108 indices[i] = n - 2;
109 weights[i] = 1.0;
110 } else {
111 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 let raw = (val - axis[0]) / step0;
119 (raw as usize).min(n - 2)
120 } else {
121 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 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 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 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 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 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 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 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 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 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 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 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 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 assert!((table.interpolate(&[0.5, 0.5]) - 1.5).abs() < 1e-10);
277 assert!((table.interpolate(&[1.5, 0.5]) - 2.5).abs() < 1e-10);
279 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 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 assert!((table.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
299 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 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 assert!((table.interpolate(&[0.5, 0.5, 0.5, 0.5]) - 2.0).abs() < 1e-10);
318 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 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 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 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 let axes = vec![
344 vec![0.0, 1.0],
345 vec![5.0], ];
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 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 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 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 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 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 let axes_bad = vec![vec![0.0, 1.0, 1.0, 2.0]]; 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(); 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])); assert!(result.is_err(), "Should panic on point dimension mismatch");
421 }
422
423 #[test]
424 fn test_large_grid_interpolation() {
425 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 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); }
435}