- registration getting better

This commit is contained in:
Wim Pomp
2026-07-29 21:48:02 +02:00
parent bb3d46cc9d
commit 001df705ef
3 changed files with 52 additions and 35 deletions
+4 -4
View File
@@ -212,7 +212,7 @@ where
shape: Vec<f64>, shape: Vec<f64>,
center: Vec<f64>, center: Vec<f64>,
fixed: BSpline<0, D>, fixed: BSpline<0, D>,
moving: BSpline<1, D>, moving: BSpline<3, D>,
fixed_mu: FixedMu, fixed_mu: FixedMu,
minmax: [f64; 2], minmax: [f64; 2],
sampling: Sampling, sampling: Sampling,
@@ -230,7 +230,7 @@ where
{ {
pub fn new( pub fn new(
fixed: BSpline<0, D>, fixed: BSpline<0, D>,
moving: BSpline<1, D>, moving: BSpline<3, D>,
sampling: SamplingArg, sampling: SamplingArg,
n_bins: usize, n_bins: usize,
edge: f64, edge: f64,
@@ -776,7 +776,7 @@ mod tests {
let a = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0]; let a = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0];
let b = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0]; let b = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0];
let fixed = BSpline::<0, _>::new(a.view()); let fixed = BSpline::<0, _>::new(a.view());
let moving = BSpline::<1, _>::new(b.view()); let moving = BSpline::<3, _>::new(b.view());
let mus = vec![0.0, 0.001]; let mus = vec![0.0, 0.001];
let mut npz = NpzWriter::new(File::create( let mut npz = NpzWriter::new(File::create(
@@ -857,7 +857,7 @@ mod tests {
let a = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0]; let a = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0];
let b = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0]; let b = array![0.0, 0.0, 1.0, 0.0, 1.0, 2.0, 1.0, 0.0, 1.0, 0.0, 0.0];
let fixed = BSpline::<0, _>::new(a.view()); let fixed = BSpline::<0, _>::new(a.view());
let moving = BSpline::<1, _>::new(b.view()); let moving = BSpline::<3, _>::new(b.view());
let mus = Array1::linspace(-10.0, 10.0, 500); let mus = Array1::linspace(-10.0, 10.0, 500);
let s = a.len() as f64; let s = a.len() as f64;
+18 -14
View File
@@ -19,7 +19,7 @@ pub enum Optimizer {
impl Default for Optimizer { impl Default for Optimizer {
fn default() -> Self { fn default() -> Self {
Self::ASGD Self::LBFGS
} }
} }
@@ -61,16 +61,16 @@ impl RegistrationStep {
/// ///
/// Uses a multi-resolution pyramid with downsampling matching SimpleElastix: /// Uses a multi-resolution pyramid with downsampling matching SimpleElastix:
/// schedule [8, 4, 2, 1] → σ = [4.0, 2.0, 1.0, 0.5] at spacing=1. /// schedule [8, 4, 2, 1] → σ = [4.0, 2.0, 1.0, 0.5] at spacing=1.
/// 512 iterations per level, 20488192 cached random samples. /// Uses all pixels at all levels for deterministic, precise convergence.
pub fn default_steps(ndim: usize, n: usize) -> Vec<Self> { pub fn default_steps(ndim: usize, n: usize) -> Vec<Self> {
if ndim == 1 { if ndim == 1 {
return vec![Self { return vec![Self {
sigma: Sigma::Absolute(vec![1.0; 1]), sigma: Sigma::Absolute(vec![1.0; 1]),
samples: SamplingArg::Random(2048.min(n)), samples: SamplingArg::Random(n),
n_bins: 32, n_bins: 32,
tolerance: 1e-6, tolerance: 1e-8,
edge: 0.05, edge: 0.05,
max_iterations: 1500, max_iterations: 2048,
learning_rate: 1.0, learning_rate: 1.0,
downsample: 1, downsample: 1,
}]; }];
@@ -78,19 +78,23 @@ impl RegistrationStep {
// Pyramid with downsampling: schedule [8,4,2,1], sigma [4,2,1,0.5] // Pyramid with downsampling: schedule [8,4,2,1], sigma [4,2,1,0.5]
let sigma_schedule: Vec<f64> = vec![4.0, 2.0, 1.0, 0.5]; let sigma_schedule: Vec<f64> = vec![4.0, 2.0, 1.0, 0.5];
let downsample_schedule: Vec<usize> = vec![8, 4, 2, 1]; let downsample_schedule: Vec<usize> = vec![8, 4, 2, 1];
let n_levels = sigma_schedule.len();
sigma_schedule sigma_schedule
.iter() .iter()
.zip(downsample_schedule.iter()) .zip(downsample_schedule.iter())
.map(|(&s, &d)| { .enumerate()
.map(|(level, (&s, &d))| {
let n_pixels = (n / (d * d)).max(4); let n_pixels = (n / (d * d)).max(4);
let n_samples = n_pixels.min(8192).max(2048); let is_finest = level == n_levels - 1;
// Use all pixels at all levels for deterministic results
let (max_iter, tol) = if is_finest { (2048, 1e-8) } else { (512, 1e-6) };
Self { Self {
sigma: Sigma::Absolute(vec![s; ndim]), sigma: Sigma::Absolute(vec![s; ndim]),
samples: SamplingArg::Random(n_samples), samples: SamplingArg::Random(n_pixels),
n_bins: 32, n_bins: 32,
tolerance: 1e-6, tolerance: tol,
edge: 0.05, edge: 0.05,
max_iterations: 512, max_iterations: max_iter,
learning_rate: 1.0, learning_rate: 1.0,
downsample: d, downsample: d,
} }
@@ -318,7 +322,7 @@ impl<D: Dimension> Registration<D> {
continue; continue;
} }
let bf = BSpline::<0, _>::new(f.view()); let bf = BSpline::<0, _>::new(f.view());
let bm = BSpline::<1, _>::new(m.view()); let bm = BSpline::<3, _>::new(m.view());
let metric = MattesMetric::new(bf, bm, samples, n_bins, edge)? let metric = MattesMetric::new(bf, bm, samples, n_bins, edge)?
.with_fixed_mu(self.fixed_mu.clone()); .with_fixed_mu(self.fixed_mu.clone());
@@ -344,7 +348,7 @@ impl<D: Dimension> Registration<D> {
tolerance, tolerance,
maximum_step_length: 1.0, maximum_step_length: 1.0,
sp_a: 20.0, sp_a: 20.0,
sp_alpha: 1.0, sp_alpha: 0.602,
scales: Some(scales), scales: Some(scales),
..Default::default() ..Default::default()
}; };
@@ -418,7 +422,7 @@ impl<D: Dimension> Registration<D> {
} }
let bf = BSpline::<0, _>::new(f.view()); let bf = BSpline::<0, _>::new(f.view());
let bm = BSpline::<1, _>::new(m.view()); let bm = BSpline::<3, _>::new(m.view());
let n_samples = match &samples { let n_samples = match &samples {
SamplingArg::Fixed(n) => *n, SamplingArg::Fixed(n) => *n,
SamplingArg::Random(n) => *n, SamplingArg::Random(n) => *n,
@@ -449,7 +453,7 @@ impl<D: Dimension> Registration<D> {
tolerance, tolerance,
maximum_step_length: 1.0, maximum_step_length: 1.0,
sp_a: 20.0, sp_a: 20.0,
sp_alpha: 1.0, sp_alpha: 0.602,
scales: Some(scales), scales: Some(scales),
..Default::default() ..Default::default()
}; };
+30 -17
View File
@@ -749,7 +749,7 @@ mod tests {
) )
.slice(s![.., 0]) .slice(s![.., 0])
.mapv(|i| i as f64); .mapv(|i| i as f64);
let q = vec![0.85, 4.0]; let q = vec![0.85, 2.0];
let im_b = let im_b =
Transform::new(q.clone(), vec![im_a.shape()[0]]).interpolate::<1, _, _>(&im_a)?; Transform::new(q.clone(), vec![im_a.shape()[0]]).interpolate::<1, _, _>(&im_a)?;
@@ -871,7 +871,7 @@ mod tests {
let f = gaussian_smooth(im_a.view(), &[*sigma_val; 2])?; let f = gaussian_smooth(im_a.view(), &[*sigma_val; 2])?;
let m = gaussian_smooth(im_b.view(), &[*sigma_val; 2])?; let m = gaussian_smooth(im_b.view(), &[*sigma_val; 2])?;
let bf = BSpline::<0, _>::new(f.view()); let bf = BSpline::<0, _>::new(f.view());
let bm = BSpline::<1, _>::new(m.view()); let bm = BSpline::<3, _>::new(m.view());
let metric = MattesMetric::new(bf, bm, SamplingArg::Random(3000), 128, edge)?; let metric = MattesMetric::new(bf, bm, SamplingArg::Random(3000), 128, edge)?;
let mi_id = metric.evaluate(&identity); let mi_id = metric.evaluate(&identity);
@@ -899,7 +899,7 @@ mod tests {
let f4 = gaussian_smooth(im_a.view(), &[4.0, 4.0])?; let f4 = gaussian_smooth(im_a.view(), &[4.0, 4.0])?;
let m4 = gaussian_smooth(im_b.view(), &[4.0, 4.0])?; let m4 = gaussian_smooth(im_b.view(), &[4.0, 4.0])?;
let bf4 = BSpline::<0, _>::new(f4.view()); let bf4 = BSpline::<0, _>::new(f4.view());
let bm4 = BSpline::<1, _>::new(m4.view()); let bm4 = BSpline::<3, _>::new(m4.view());
let metric4 = MattesMetric::new(bf4, bm4, SamplingArg::Random(5000), 128, edge)?; let metric4 = MattesMetric::new(bf4, bm4, SamplingArg::Random(5000), 128, edge)?;
let mi_p_coarse = metric4.evaluate(&p); let mi_p_coarse = metric4.evaluate(&p);
let mi_qinv_coarse = metric4.evaluate(&q_inv); let mi_qinv_coarse = metric4.evaluate(&q_inv);
@@ -943,7 +943,7 @@ mod tests {
let m = gaussian_smooth(im_b.view(), &[4.0, 4.0])?; let m = gaussian_smooth(im_b.view(), &[4.0, 4.0])?;
let metric = MattesMetric::new( let metric = MattesMetric::new(
BSpline::<0, _>::new(f.view()), BSpline::<0, _>::new(f.view()),
BSpline::<1, _>::new(m.view()), BSpline::<3, _>::new(m.view()),
SamplingArg::Random(3000), SamplingArg::Random(3000),
128, 128,
edge, edge,
@@ -1022,13 +1022,24 @@ mod tests {
.fold(0.0f64, f64::max); .fold(0.0f64, f64::max);
println!("Our: {:?} max_err: {:.4} sse: {:.4}", t, max_err, sse); println!("Our: {:?} max_err: {:.4} sse: {:.4}", t, max_err, sse);
let mut tif = IJTiffFile::new(std::env::home_dir().unwrap().join("tmp/register_real_images.tif"))?;
let mut tif = IJTiffFile::new(
std::env::home_dir()
.unwrap()
.join("tmp/register_real_images.tif"),
)?;
tif.save(fixed.mapv(|i| i as u16), 0, 0, 0)?; tif.save(fixed.mapv(|i| i as u16), 0, 0, 0)?;
tif.save(t.interpolate_par::<1, _, _>(moving.view())?.mapv(|i| i as u16), 1, 0, 0)?; tif.save(
t.interpolate_par::<1, _, _>(moving.view())?
.mapv(|i| i as u16),
1,
0,
0,
)?;
tif.save(moving.mapv(|i| i as u16), 2, 0, 0)?; tif.save(moving.mapv(|i| i as u16), 2, 0, 0)?;
assert!(max_err < 0.1); assert!(max_err < 0.02);
assert!(sse < 0.1); assert!(sse < 0.02);
Ok(()) Ok(())
} }
@@ -1036,7 +1047,8 @@ mod tests {
#[test] #[test]
fn register_real_images2() -> Result<(), Box<dyn std::error::Error>> { fn register_real_images2() -> Result<(), Box<dyn std::error::Error>> {
let fixed = read_tiff("test_files/fixed.tif")?; let fixed = read_tiff("test_files/fixed.tif")?;
let e = Transform::<Ix2>::new(vec![0.8, 0.0, 0.0, 1.0, 0.0, 0.0], fixed.shape().to_vec()).inverse()?; let e = Transform::<Ix2>::new(vec![0.8, 0.0, 0.0, 1.0, 0.0, 0.0], fixed.shape().to_vec())
.inverse()?;
let moving = e.interpolate::<3, _, _>(fixed.view())?; let moving = e.interpolate::<3, _, _>(fixed.view())?;
let t = Transform::<Ix2>::register( let t = Transform::<Ix2>::register(
@@ -1047,27 +1059,28 @@ mod tests {
None, None,
)?; )?;
let e_inv = e.inverse()?;
let sse = t let sse = t
.parameters .parameters
.iter() .iter()
.zip(e.parameters.iter()) .zip(e_inv.parameters.iter())
.map(|(a, b)| (a - b).powi(2)) .map(|(a, b)| (a - b).powi(2))
.sum::<f64>(); .sum::<f64>();
let max_err = t let max_err = t
.parameters .parameters
.iter() .iter()
.zip(e.parameters.iter()) .zip(e_inv.parameters.iter())
.map(|(a, b)| (a - b).abs()) .map(|(a, b)| (a - b).abs())
.fold(0.0f64, f64::max); .fold(0.0f64, f64::max);
println!("Our: {:?} max_err: {:.4} sse: {:.4}", t, max_err, sse); println!("Our: {:?} max_err: {:.4} sse: {:.4}", t, max_err, sse);
// let mut tif = IJTiffFile::new(std::env::home_dir().unwrap().join("tmp/register_real_images2.tif"))?; let mut tif = IJTiffFile::new(std::env::home_dir().unwrap().join("tmp/register_real_images2.tif"))?;
// tif.save(fixed.mapv(|i| i as u16), 0, 0, 0)?; tif.save(fixed.mapv(|i| i as u16), 0, 0, 0)?;
// tif.save(t.interpolate_par::<1, _, _>(moving.view())?.mapv(|i| i as u16), 1, 0, 0)?; tif.save(t.interpolate_par::<1, _, _>(moving.view())?.mapv(|i| i as u16), 1, 0, 0)?;
// tif.save(moving.mapv(|i| i as u16), 2, 0, 0)?; tif.save(moving.mapv(|i| i as u16), 2, 0, 0)?;
assert!(max_err < 0.1); assert!(max_err < 0.02);
assert!(sse < 0.1); assert!(sse < 0.02);
Ok(()) Ok(())
} }