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 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 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 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 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 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 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 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 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 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 assert!((table.interpolate(&[0.5, 0.5]) - 1.5).abs() < 1e-10);
282 assert!((table.interpolate(&[1.5, 0.5]) - 2.5).abs() < 1e-10);
284 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 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 assert!((table.interpolate(&[0.5, 0.5, 0.5]) - 1.5).abs() < 1e-10);
304 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 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 assert!((table.interpolate(&[0.5, 0.5, 0.5, 0.5]) - 2.0).abs() < 1e-10);
323 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 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 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 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 let axes = vec![
349 vec![0.0, 1.0],
350 vec![5.0], ];
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 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 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 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 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 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 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();
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(); 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])); assert!(result.is_err(), "Should panic on point dimension mismatch");
426 }
427
428 #[test]
429 fn test_large_grid_interpolation() {
430 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 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); }
440}