1use num_complex::Complex64;
17use crate::Grid;
18use crate::fft::{fft3d, ifft3d, fft_real_kernel};
19use crate::kernels::smv::{smv_kernel, erode_mask_smv};
20use crate::solvers::cg_solve_with_progress;
21
22#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
24#[derive(Clone, Debug)]
25pub struct ResharpParams {
26 pub radius: f64,
28 pub tik_reg: f64,
30 pub tol: f64,
32 pub max_iter: usize,
34}
35
36impl Default for ResharpParams {
37 fn default() -> Self {
38 Self {
39 radius: 6.0,
40 tik_reg: 1e-4,
41 tol: 1e-6,
42 max_iter: 30,
43 }
44 }
45}
46
47pub fn resharp(
59 field: &[f64],
60 mask: &[u8],
61 grid: &Grid,
62 params: &ResharpParams,
63 mut progress: impl FnMut(usize, usize),
64) -> (Vec<f64>, Vec<u8>) {
65 let (nx, ny, nz) = grid.dims;
66 let n_total = nx * ny * nz;
67
68 let s_kernel = smv_kernel(grid, params.radius);
70 let s_fft = fft_real_kernel(&s_kernel, nx, ny, nz);
71
72 let dker: Vec<f64> = s_fft.iter().map(|&s| 1.0 - s).collect();
74
75 let eroded_mask = erode_mask_smv(mask, &s_fft, grid, 1.0 - 1e-7_f64.sqrt());
77 let eroded_mask_f64: Vec<f64> = eroded_mask.iter()
78 .map(|&m| m as f64)
79 .collect();
80
81 let b = apply_ht_h(&dker, &eroded_mask_f64, field, nx, ny, nz);
87
88 let tik_reg = params.tik_reg;
90 let x0 = vec![0.0; n_total];
91 let x = cg_solve_with_progress(
92 |x_vec| {
93 let mut result = apply_ht_h(&dker, &eroded_mask_f64, x_vec, nx, ny, nz);
94 for i in 0..n_total {
96 result[i] += tik_reg * x_vec[i];
97 }
98 result
99 },
100 &b,
101 &x0,
102 params.tol,
103 params.max_iter,
104 &mut progress,
105 );
106
107 let local_field: Vec<f64> = x.iter()
109 .enumerate()
110 .map(|(i, &v)| if eroded_mask[i] == 1 { v } else { 0.0 })
111 .collect();
112
113 (local_field, eroded_mask)
114}
115
116fn apply_ht_h(
121 dker: &[f64],
122 mask: &[f64],
123 x: &[f64],
124 nx: usize, ny: usize, nz: usize,
125) -> Vec<f64> {
126 let n_total = nx * ny * nz;
127
128 let mut tmp: Vec<Complex64> = x.iter()
130 .map(|&v| Complex64::new(v, 0.0))
131 .collect();
132 fft3d(&mut tmp, nx, ny, nz);
133
134 for i in 0..n_total {
135 tmp[i] *= dker[i];
136 }
137 ifft3d(&mut tmp, nx, ny, nz);
138
139 for i in 0..n_total {
141 tmp[i] *= mask[i];
142 }
143
144 fft3d(&mut tmp, nx, ny, nz);
146
147 for i in 0..n_total {
148 tmp[i] *= dker[i];
149 }
150 ifft3d(&mut tmp, nx, ny, nz);
151
152 tmp.iter().map(|c| c.re).collect()
153}
154
155#[cfg(test)]
156mod tests {
157 use super::*;
158
159 #[test]
160 fn test_resharp_zero_field() {
161 let n = 16;
162 let field = vec![0.0; n * n * n];
163 let mask = vec![1u8; n * n * n];
164 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
165 let params = ResharpParams { radius: 2.0, tik_reg: 1e-4, tol: 1e-6, max_iter: 50 };
166
167 let (local, _) = resharp(&field, &mask, &grid, ¶ms, |_, _| {});
168
169 for &val in local.iter() {
170 assert!(val.abs() < 1e-8, "Zero field should give zero local field, got {}", val);
171 }
172 }
173
174 #[test]
175 fn test_resharp_finite() {
176 let n = 16;
177 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.01).collect();
178 let mask = vec![1u8; n * n * n];
179 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
180 let params = ResharpParams { radius: 2.0, tik_reg: 1e-4, tol: 1e-6, max_iter: 50 };
181
182 let (local, eroded) = resharp(&field, &mask, &grid, ¶ms, |_, _| {});
183
184 for (i, &val) in local.iter().enumerate() {
185 assert!(val.is_finite(), "Local field should be finite at index {}", i);
186 }
187
188 let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
189 assert!(eroded_count > 0, "Eroded mask should have some voxels");
190 }
191}