1use crate::utils::{shrink, apply_mask_zero};
17use crate::Grid;
18use super::admm::{AdmmBuffers, admm_step, prepare_admm_spectral};
19
20#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
22#[derive(Clone, Debug)]
23pub struct TvParams {
24 pub lambda: f64,
26 pub rho: f64,
28 pub tol: f64,
30 pub max_iter: usize,
32}
33
34impl Default for TvParams {
35 fn default() -> Self {
36 Self {
37 lambda: 2e-4,
38 rho: 2e-2,
39 tol: 1e-3,
40 max_iter: 250,
41 }
42 }
43}
44
45pub fn tv_admm(
64 local_field: &[f64],
65 mask: &[u8],
66 grid: &Grid,
67 bdir: (f64, f64, f64),
68 params: &TvParams,
69 mut progress: impl FnMut(usize, usize),
70) -> Vec<f64> {
71 let n_total = grid.n_total();
72
73 let (mut fft_ws, inv_a, f_hat) = prepare_admm_spectral(local_field, grid, bdir, params.rho);
75
76 let mut buf = AdmmBuffers::new(n_total);
78 let lambda_over_rho = params.lambda / params.rho;
79
80 for iter in 0..params.max_iter {
81 progress(iter + 1, params.max_iter);
82
83 let converged = admm_step(
84 &mut buf, &mut fft_ws, &f_hat, &inv_a, params.rho, grid, params.tol,
85 |vx, vy, vz, _| (shrink(vx, lambda_over_rho), shrink(vy, lambda_over_rho), shrink(vz, lambda_over_rho)),
86 );
87
88 if converged {
89 progress(iter + 1, iter + 1);
90 break;
91 }
92 }
93
94 apply_mask_zero(&mut buf.x, mask);
95
96 buf.x
97}
98
99#[cfg(test)]
100mod tests {
101 use super::*;
102 use crate::utils::gradient::fgrad;
103
104 #[test]
105 fn test_shrink() {
106 assert!((shrink(1.0, 0.5) - 0.5).abs() < 1e-10);
107 assert!((shrink(-1.0, 0.5) - (-0.5)).abs() < 1e-10);
108 assert!((shrink(0.3, 0.5) - 0.0).abs() < 1e-10);
109 assert!((shrink(-0.3, 0.5) - 0.0).abs() < 1e-10);
110 }
111
112 #[test]
113 fn test_tv_admm_zero_field() {
114 let n = 8;
116 let field = vec![0.0; n * n * n];
117 let mask = vec![1u8; n * n * n];
118 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
119 let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-2, max_iter: 10 };
120
121 let chi = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
122
123 for &val in chi.iter() {
124 assert!(val.abs() < 1e-8, "Zero field should give zero chi, got {}", val);
125 }
126 }
127
128 #[test]
129 fn test_tv_admm_finite() {
130 let n = 8;
132 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
133 let mask = vec![1u8; n * n * n];
134 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
135 let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-2, max_iter: 10 };
136
137 let chi = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
138
139 for (i, &val) in chi.iter().enumerate() {
140 assert!(val.is_finite(), "Chi should be finite at index {}", i);
141 }
142 }
143
144 #[test]
145 fn test_tv_admm_smoother_than_tkd() {
146 let n = 8;
148 let mut field = vec![0.0; n * n * n];
150 for i in 0..n*n*n {
151 field[i] = if i % 2 == 0 { 0.01 } else { -0.01 }; }
153 let mask = vec![1u8; n * n * n];
154 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
155 let params = TvParams { lambda: 1e-2, rho: 1.0, tol: 1e-2, max_iter: 50 };
156
157 let chi_tv = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
158
159 let (gx, gy, gz) = fgrad(&chi_tv, &grid);
161 let tv: f64 = gx.iter().chain(gy.iter()).chain(gz.iter())
162 .map(|&g| g.abs())
163 .sum();
164
165 assert!(tv.is_finite(), "TV should be finite");
168 }
169
170 #[cfg(feature = "parallel")]
172 #[test]
173 fn test_tv_parallel_matches_sequential() {
174 let n = 16;
175 let field: Vec<f64> = (0..n*n*n).map(|i| ((i as f64) * 0.7).sin() * 0.01).collect();
176 let mask = vec![1u8; n * n * n];
177 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
178 let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-3, max_iter: 50 };
179
180 let pool_1 = rayon::ThreadPoolBuilder::new().num_threads(1).build().unwrap();
182 let chi_seq = pool_1.install(|| {
183 tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {})
184 });
185
186 let chi_par = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
188
189 for (i, (s, p)) in chi_seq.iter().zip(chi_par.iter()).enumerate() {
191 assert!(
192 (s - p).abs() < 1e-10,
193 "TV mismatch at voxel {}: seq={} par={} diff={}",
194 i, s, p, (s - p).abs()
195 );
196 }
197 }
198}