1use num_complex::Complex64;
14use crate::Grid;
15use crate::fft::{fft3d, ifft3d, fft_real_kernel};
16use crate::kernels::smv::{smv_kernel, erode_mask_smv};
17use crate::utils::vec_norm;
18
19#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
21#[derive(Clone, Debug)]
22pub struct IsmvParams {
23 pub tol: f64,
25 pub max_iter: usize,
27 pub radius: f64,
29}
30
31impl Default for IsmvParams {
32 fn default() -> Self {
33 Self { tol: 1e-6, max_iter: 50, radius: 5.0 }
34 }
35}
36
37pub fn ismv(
49 field: &[f64],
50 mask: &[u8],
51 grid: &Grid,
52 params: &IsmvParams,
53 progress: impl FnMut(usize, usize),
54) -> (Vec<f64>, Vec<u8>) {
55 ismv_core(field, mask, grid, params.radius, params.tol, params.max_iter, progress)
56}
57
58pub(crate) fn ismv_core(
63 field: &[f64],
64 mask: &[u8],
65 grid: &Grid,
66 radius: f64,
67 tol: f64,
68 max_iter: usize,
69 mut progress: impl FnMut(usize, usize),
70) -> (Vec<f64>, Vec<u8>) {
71 let (nx, ny, nz) = grid.dims;
72 let n_total = nx * ny * nz;
73
74 let smv = smv_kernel(grid, radius);
76 let smv_fft_real = fft_real_kernel(&smv, nx, ny, nz);
77
78 let mut smv_complex: Vec<Complex64> = smv.iter()
80 .map(|&x| Complex64::new(x, 0.0))
81 .collect();
82 fft3d(&mut smv_complex, nx, ny, nz);
83 let smv_fft = smv_complex;
84
85 let m0: Vec<f64> = mask.iter()
87 .map(|&m| if m != 0 { 1.0 } else { 0.0 })
88 .collect();
89
90 let eroded_mask = erode_mask_f64(mask, &smv_fft_real, grid);
92
93 let boundary: Vec<f64> = m0.iter()
95 .zip(eroded_mask.iter())
96 .map(|(&m, &e)| m - e)
97 .collect();
98
99 let mut f: Vec<f64> = field.to_vec();
101
102 let mut f0: Vec<f64> = field.iter()
104 .zip(eroded_mask.iter())
105 .map(|(&fi, &m)| fi * m)
106 .collect();
107
108 let bc: Vec<f64> = field.iter()
110 .zip(boundary.iter())
111 .map(|(&fi, &b)| fi * b)
112 .collect();
113
114 let mut nr = vec_norm(&f0);
116 let eps = tol * nr;
117
118 for iter in 0..max_iter {
120 progress(iter + 1, max_iter);
122
123 if nr <= eps {
124 progress(iter + 1, iter + 1);
125 break;
126 }
127
128 let mut f_complex: Vec<Complex64> = f.iter()
130 .map(|&x| Complex64::new(x, 0.0))
131 .collect();
132
133 fft3d(&mut f_complex, nx, ny, nz);
134
135 for i in 0..n_total {
136 f_complex[i] *= smv_fft[i];
137 }
138
139 ifft3d(&mut f_complex, nx, ny, nz);
140
141 for i in 0..n_total {
143 f[i] = eroded_mask[i] * f_complex[i].re + bc[i];
144 }
145
146 let mut residual_sq = 0.0;
148 for i in 0..n_total {
149 let diff = f0[i] - f[i];
150 residual_sq += diff * diff;
151 f0[i] = f[i];
152 }
153 nr = residual_sq.sqrt();
154 }
155
156 let mut local_field = vec![0.0; n_total];
158 for i in 0..n_total {
159 if mask[i] != 0 {
160 local_field[i] = field[i] - f[i];
161 }
162 }
163
164 let eroded_mask_u8: Vec<u8> = eroded_mask.iter()
166 .map(|&m| if m > 0.5 { 1 } else { 0 })
167 .collect();
168
169 (local_field, eroded_mask_u8)
170}
171
172fn erode_mask_f64(mask: &[u8], smv_kernel_fft: &[f64], grid: &Grid) -> Vec<f64> {
174 let eroded = erode_mask_smv(mask, smv_kernel_fft, grid, 1.0 - 1e-10);
175 eroded.iter().map(|&m| m as f64).collect()
176}
177
178#[cfg(test)]
179mod tests {
180 use super::*;
181
182 #[test]
183 fn test_ismv_zero_field() {
184 let n = 8;
185 let field = vec![0.0; n * n * n];
186 let mask = vec![1u8; n * n * n];
187 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
188
189 let (local, eroded) = ismv_core(
190 &field, &mask, &grid,
191 2.0, 1e-3, 10, |_, _| {}
192 );
193
194 for &val in local.iter() {
195 assert!(val.abs() < 1e-10, "Zero field should give zero local field");
196 }
197
198 let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
200 assert!(eroded_count > 0, "Eroded mask should have some voxels");
201 }
202
203 #[test]
204 fn test_ismv_finite() {
205 let n = 8;
206 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
207 let mask = vec![1u8; n * n * n];
208 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
209
210 let (local, _eroded) = ismv_core(
211 &field, &mask, &grid,
212 2.0, 1e-3, 20, |_, _| {}
213 );
214
215 for (i, &val) in local.iter().enumerate() {
216 assert!(val.is_finite(), "Local field should be finite at index {}", i);
217 }
218 }
219
220 #[test]
221 fn test_ismv_preserves_interior() {
222 let n = 16;
224 let field = vec![0.1; n * n * n];
225
226 let mut mask = vec![0u8; n * n * n];
228 let center = n / 2;
229 let radius = n / 3;
230
231 for i in 0..n {
232 for j in 0..n {
233 for k in 0..n {
234 let di = (i as i32) - (center as i32);
235 let dj = (j as i32) - (center as i32);
236 let dk = (k as i32) - (center as i32);
237 if di*di + dj*dj + dk*dk <= (radius * radius) as i32 {
238 mask[i * n * n + j * n + k] = 1;
239 }
240 }
241 }
242 }
243
244 let mask_count: usize = mask.iter().map(|&m| m as usize).sum();
245 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
246
247 let (_, eroded) = ismv_core(
249 &field, &mask, &grid,
250 1.5, 1e-3, 50, |_, _| {}
251 );
252
253 let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
254
255 assert!(eroded_count <= mask_count, "Eroded mask should be smaller than original");
257 assert!(eroded_count > 0, "Eroded mask should have some voxels");
259 }
260
261 #[test]
262 fn test_ismv_convergence() {
263 let n = 8;
264 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
265 let mask = vec![1u8; n * n * n];
266 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
267
268 let (local_many, _) = ismv_core(
270 &field, &mask, &grid,
271 2.0, 1e-6, 100, |_, _| {}
272 );
273
274 let (local_few, _) = ismv_core(
276 &field, &mask, &grid,
277 2.0, 1e-6, 5, |_, _| {}
278 );
279
280 for (i, &val) in local_many.iter().enumerate() {
282 assert!(val.is_finite(), "iSMV many iters: finite at index {}", i);
283 }
284 for (i, &val) in local_few.iter().enumerate() {
285 assert!(val.is_finite(), "iSMV few iters: finite at index {}", i);
286 }
287
288 let diff_norm: f64 = local_many.iter()
291 .zip(local_few.iter())
292 .map(|(&a, &b)| (a - b).powi(2))
293 .sum::<f64>()
294 .sqrt();
295
296 assert!(diff_norm.is_finite(), "Difference between runs should be finite");
298 }
299
300 #[test]
301 fn test_ismv_different_radius() {
302 let n = 8;
303 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
304 let mask = vec![1u8; n * n * n];
305 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
306
307 let (local_small, eroded_small) = ismv_core(
309 &field, &mask, &grid,
310 1.5, 1e-3, 20, |_, _| {}
311 );
312
313 let (local_large, eroded_large) = ismv_core(
315 &field, &mask, &grid,
316 3.0, 1e-3, 20, |_, _| {}
317 );
318
319 for (i, &val) in local_small.iter().enumerate() {
321 assert!(val.is_finite(), "iSMV small radius: finite at index {}", i);
322 }
323 for (i, &val) in local_large.iter().enumerate() {
324 assert!(val.is_finite(), "iSMV large radius: finite at index {}", i);
325 }
326
327 let small_count: usize = eroded_small.iter().map(|&m| m as usize).sum();
329 let large_count: usize = eroded_large.iter().map(|&m| m as usize).sum();
330 assert!(
331 large_count <= small_count,
332 "Larger radius should erode more: large={}, small={}",
333 large_count, small_count
334 );
335 }
336
337 #[test]
338 fn test_ismv_larger_volume() {
339 let n = 16;
341
342 let mut field = vec![0.0; n * n * n];
344 for z in 0..n {
345 for y in 0..n {
346 for x in 0..n {
347 field[x + y * n + z * n * n] = (z as f64) * 0.1;
348 }
349 }
350 }
351
352 let mut mask = vec![0u8; n * n * n];
354 let center = n / 2;
355 let radius = n / 3;
356 for z in 0..n {
357 for y in 0..n {
358 for x in 0..n {
359 let dx = (x as i32) - (center as i32);
360 let dy = (y as i32) - (center as i32);
361 let dz = (z as i32) - (center as i32);
362 if dx * dx + dy * dy + dz * dz <= (radius * radius) as i32 {
363 mask[x + y * n + z * n * n] = 1;
364 }
365 }
366 }
367 }
368
369 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
370 let (local, eroded) = ismv_core(
371 &field, &mask, &grid,
372 2.0, 1e-3, 50, |_, _| {}
373 );
374
375 assert_eq!(local.len(), n * n * n);
376 for &val in &local {
377 assert!(val.is_finite());
378 }
379
380 let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
382 assert!(eroded_count > 0, "Eroded mask should have some voxels");
383
384 for i in 0..n * n * n {
386 if mask[i] == 0 {
387 assert_eq!(local[i], 0.0, "Outside mask should be zero");
388 }
389 }
390 }
391
392 #[test]
393 fn test_ismv_with_progress() {
394 let n = 8;
395 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
396 let mask = vec![1u8; n * n * n];
397 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
398
399 let mut progress_calls = Vec::new();
400 let (local, _) = ismv_core(
401 &field, &mask, &grid,
402 2.0, 1e-3, 20,
403 |iter, max| { progress_calls.push((iter, max)); }
404 );
405
406 assert_eq!(local.len(), n * n * n);
407 assert!(!progress_calls.is_empty(), "Progress callback should be called");
408 for &val in &local {
409 assert!(val.is_finite());
410 }
411 }
412
413 #[test]
414 fn test_ismv_anisotropic_voxels() {
415 let n = 8;
416 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
417 let mask = vec![1u8; n * n * n];
418 let grid = Grid::new(n, n, n, 0.5, 1.0, 2.0);
419
420 let (local, eroded) = ismv_core(
422 &field, &mask, &grid,
423 3.0, 1e-3, 20, |_, _| {}
424 );
425
426 for &val in &local {
427 assert!(val.is_finite());
428 }
429 let count: usize = eroded.iter().map(|&m| m as usize).sum();
431 assert!(count <= n * n * n);
432 }
433
434 #[test]
435 fn test_ismv_tight_convergence() {
436 let n = 8;
438 let field = vec![0.5; n * n * n]; let mask = vec![1u8; n * n * n];
440 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
441
442 let (local, _) = ismv_core(
443 &field, &mask, &grid,
444 2.0, 1e-12, 200, |_, _| {} );
446
447 for &val in &local {
448 assert!(val.is_finite());
449 }
450 }
451
452 #[test]
453 fn test_ismv_with_background_mask() {
454 let n = 8;
456 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
457
458 let mut mask = vec![0u8; n * n * n];
460 for z in 1..(n - 1) {
461 for y in 1..(n - 1) {
462 for x in 1..(n - 1) {
463 mask[x + y * n + z * n * n] = 1;
464 }
465 }
466 }
467
468 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
469 let (local, eroded) = ismv_core(
470 &field, &mask, &grid,
471 1.5, 1e-3, 30, |_, _| {}
472 );
473
474 for &val in &local {
475 assert!(val.is_finite());
476 }
477
478 for i in 0..n * n * n {
480 if mask[i] == 0 {
481 assert_eq!(local[i], 0.0, "Outside mask should be zero at index {}", i);
482 }
483 }
484
485 for i in 0..n * n * n {
487 if eroded[i] != 0 {
488 assert_eq!(mask[i], 1, "Eroded voxel must be inside original mask");
489 }
490 }
491 }
492}