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 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 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 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 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); }
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 let table = grid.generate(|point| point[0]);
122
123 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 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 let table = grid.generate(|point| point[0] * point[1]);
145
146 assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0);
148 assert_eq!(table.interpolate(&[2.0, 2.0]), 4.0);
149
150 assert!((table.interpolate(&[1.0, 1.0]) - 1.0).abs() < 1e-10);
152
153 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 let table = grid.generate(|point| point[0] + point[1] + point[2]);
168
169 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 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 let table = grid.generate(|point| point.iter().sum());
195
196 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 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 let table = grid.generate(|point| point[0].sin() * point[1].cos());
214
215 assert!((table.interpolate(&[0.0, 0.0]) - 0.0).abs() < 1e-10); let pi_2 = std::f64::consts::PI / 2.0;
220 let result = table.interpolate(&[pi_2, 0.0]);
221 assert!(
223 (result - 1.0).abs() < 0.05,
224 "At (pi/2, 0) expected ~1.0, got {}",
225 result
226 );
227
228 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()]; 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 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 let table = grid.generate(|point| point[0].powi(2) + point[1].powi(2));
254
255 assert_eq!(table.interpolate(&[0.0, 0.0]), 0.0); assert_eq!(table.interpolate(&[1.0, 0.0]), 1.0); assert_eq!(table.interpolate(&[0.0, 1.0]), 1.0); assert_eq!(table.interpolate(&[1.0, 1.0]), 2.0); let result = table.interpolate(&[0.5, 0.5]);
265 assert!((result - 1.0).abs() < 0.01, "Expected ~1.0, got {}", result);
266 }
267}