Skip to main content

scirs2_vision/registration/
intensity.rs

1//! Intensity-based image registration
2//!
3//! This module implements registration algorithms that directly use image intensity
4//! values to find optimal alignment between images.
5
6use crate::error::{Result, VisionError};
7use crate::registration::warping::{warp_image, BoundaryMethod, InterpolationMethod};
8use crate::registration::{
9    identity_transform, RegistrationParams, RegistrationResult, TransformMatrix,
10};
11use image::GrayImage;
12use scirs2_core::ndarray::{Array1, Array2};
13
14/// Similarity metric for intensity-based registration
15#[derive(Debug, Clone, Copy, PartialEq)]
16pub enum SimilarityMetric {
17    /// Sum of Squared Differences
18    SSD,
19    /// Normalized Cross-Correlation
20    NCC,
21    /// Mutual Information
22    MI,
23    /// Normalized Mutual Information
24    NMI,
25    /// Cross-Correlation
26    CC,
27    /// Mean Squared Error
28    MSE,
29}
30
31/// Optimization algorithm for registration
32#[derive(Debug, Clone, Copy, PartialEq)]
33pub enum OptimizationMethod {
34    /// Gradient Descent
35    GradientDescent,
36    /// Powell's method
37    Powell,
38    /// Simplex method
39    Simplex,
40    /// Conjugate Gradient
41    ConjugateGradient,
42}
43
44/// Intensity-based registration configuration
45#[derive(Debug, Clone)]
46pub struct IntensityRegistrationConfig {
47    /// Similarity metric to optimize
48    pub metric: SimilarityMetric,
49    /// Optimization method
50    pub optimizer: OptimizationMethod,
51    /// Registration parameters
52    pub params: RegistrationParams,
53    /// Step size for gradient computation
54    pub step_size: f64,
55    /// Learning rate for optimization
56    pub learning_rate: f64,
57    /// Use multi-resolution pyramid
58    pub use_pyramid: bool,
59}
60
61impl Default for IntensityRegistrationConfig {
62    fn default() -> Self {
63        Self {
64            metric: SimilarityMetric::NCC,
65            optimizer: OptimizationMethod::GradientDescent,
66            params: RegistrationParams::default(),
67            step_size: 0.1,
68            learning_rate: 0.01,
69            use_pyramid: true,
70        }
71    }
72}
73
74/// Register images using intensity-based methods
75///
76/// # Arguments
77///
78/// * `reference` - Reference image
79/// * `moving` - Moving image to register
80/// * `initial_transform` - Initial transformation estimate
81/// * `config` - Registration configuration
82///
83/// # Returns
84///
85/// * Result containing registration result
86#[allow(dead_code)]
87pub fn register_images_intensity(
88    reference: &GrayImage,
89    moving: &GrayImage,
90    initial_transform: Option<&TransformMatrix>,
91    config: &IntensityRegistrationConfig,
92) -> Result<RegistrationResult> {
93    let initial_transform = initial_transform
94        .cloned()
95        .unwrap_or_else(identity_transform);
96
97    if config.use_pyramid {
98        multi_resolution_register(reference, moving, &initial_transform, config)
99    } else {
100        single_level_register(reference, moving, &initial_transform, config)
101    }
102}
103
104/// Single-level intensity-based registration
105#[allow(dead_code)]
106fn single_level_register(
107    reference: &GrayImage,
108    moving: &GrayImage,
109    initial_transform: &TransformMatrix,
110    config: &IntensityRegistrationConfig,
111) -> Result<RegistrationResult> {
112    let mut current_transform = initial_transform.clone();
113    let mut current_cost =
114        compute_similarity(reference, moving, &current_transform, config.metric)?;
115    let mut iteration = 0;
116    let mut converged = false;
117
118    while iteration < config.params.max_iterations && !converged {
119        let (gradient, new_cost) = compute_gradient(reference, moving, &current_transform, config)?;
120
121        // Check for convergence
122        if (current_cost - new_cost).abs() < config.params.tolerance {
123            converged = true;
124            break;
125        }
126
127        // Update _transform based on optimization method
128        current_transform = match config.optimizer {
129            OptimizationMethod::GradientDescent => update_transform_gradient_descent(
130                &current_transform,
131                &gradient,
132                config.learning_rate,
133            ),
134            OptimizationMethod::Powell => {
135                // Simplified Powell's method (would need full implementation)
136                update_transform_gradient_descent(
137                    &current_transform,
138                    &gradient,
139                    config.learning_rate,
140                )
141            }
142            OptimizationMethod::Simplex => {
143                // Simplified Simplex method (would need full implementation)
144                update_transform_gradient_descent(
145                    &current_transform,
146                    &gradient,
147                    config.learning_rate,
148                )
149            }
150            OptimizationMethod::ConjugateGradient => {
151                // Simplified Conjugate Gradient (would need full implementation)
152                update_transform_gradient_descent(
153                    &current_transform,
154                    &gradient,
155                    config.learning_rate,
156                )
157            }
158        };
159
160        current_cost = new_cost;
161        iteration += 1;
162    }
163
164    Ok(RegistrationResult {
165        transform: current_transform,
166        final_cost: current_cost,
167        iterations: iteration,
168        converged,
169        inliers: Vec::new(), // Not applicable for intensity-based methods
170    })
171}
172
173/// Multi-resolution pyramid registration
174#[allow(dead_code)]
175fn multi_resolution_register(
176    reference: &GrayImage,
177    moving: &GrayImage,
178    initial_transform: &TransformMatrix,
179    config: &IntensityRegistrationConfig,
180) -> Result<RegistrationResult> {
181    // Build pyramids
182    let ref_pyramid = build_image_pyramid(reference, config.params.pyramid_levels);
183    let moving_pyramid = build_image_pyramid(moving, config.params.pyramid_levels);
184
185    let mut current_transform = initial_transform.clone();
186    let mut final_result = None;
187
188    // Register from coarse to fine
189    for level in (0..config.params.pyramid_levels).rev() {
190        // Scale _transform for current level
191        let scale = 2.0_f64.powi(level as i32);
192        let mut scaled_transform = current_transform.clone();
193        scaled_transform[[0, 2]] /= scale;
194        scaled_transform[[1, 2]] /= scale;
195
196        // Register at current level
197        let result = single_level_register(
198            &ref_pyramid[level],
199            &moving_pyramid[level],
200            &scaled_transform,
201            config,
202        )?;
203
204        current_transform = result.transform.clone();
205
206        // Scale _transform back up for next level
207        if level > 0 {
208            current_transform[[0, 2]] *= 2.0;
209            current_transform[[1, 2]] *= 2.0;
210        }
211
212        final_result = Some(result);
213    }
214
215    final_result.ok_or_else(|| {
216        VisionError::OperationError("Multi-resolution registration failed".to_string())
217    })
218}
219
220/// Build image pyramid by downsampling
221#[allow(dead_code)]
222fn build_image_pyramid(image: &GrayImage, levels: usize) -> Vec<GrayImage> {
223    let mut pyramid = vec![image.clone()];
224
225    for _ in 1..levels {
226        let prev = &pyramid[pyramid.len() - 1];
227        let (width, height) = prev.dimensions();
228
229        if width < 8 || height < 8 {
230            break;
231        }
232
233        let downsampled = downsample_image(prev);
234        pyramid.push(downsampled);
235    }
236
237    pyramid
238}
239
240/// Downsample image by factor of 2
241#[allow(dead_code)]
242fn downsample_image(image: &GrayImage) -> GrayImage {
243    let (width, height) = image.dimensions();
244    let new_width = width / 2;
245    let new_height = height / 2;
246
247    let mut downsampled = GrayImage::new(new_width, new_height);
248
249    for y in 0..new_height {
250        for x in 0..new_width {
251            // Average 2x2 block
252            let x2 = x * 2;
253            let y2 = y * 2;
254
255            let mut sum = image.get_pixel(x2, y2)[0] as u32;
256            let mut count = 1;
257
258            if x2 + 1 < width {
259                sum += image.get_pixel(x2 + 1, y2)[0] as u32;
260                count += 1;
261            }
262            if y2 + 1 < height {
263                sum += image.get_pixel(x2, y2 + 1)[0] as u32;
264                count += 1;
265            }
266            if x2 + 1 < width && y2 + 1 < height {
267                sum += image.get_pixel(x2 + 1, y2 + 1)[0] as u32;
268                count += 1;
269            }
270
271            downsampled.put_pixel(x, y, image::Luma([(sum / count) as u8]));
272        }
273    }
274
275    downsampled
276}
277
278/// Compute similarity metric between images
279#[allow(dead_code)]
280fn compute_similarity(
281    reference: &GrayImage,
282    moving: &GrayImage,
283    transform: &TransformMatrix,
284    metric: SimilarityMetric,
285) -> Result<f64> {
286    let (width, height) = reference.dimensions();
287
288    // Warp moving image
289    let warped = warp_image(
290        moving,
291        transform,
292        (width, height),
293        InterpolationMethod::Bilinear,
294        BoundaryMethod::Zero,
295    )?;
296
297    match metric {
298        SimilarityMetric::SSD => compute_ssd(reference, &warped),
299        SimilarityMetric::NCC => compute_ncc(reference, &warped),
300        SimilarityMetric::MI => compute_mutual_information(reference, &warped),
301        SimilarityMetric::NMI => compute_normalized_mutual_information(reference, &warped),
302        SimilarityMetric::CC => compute_cross_correlation(reference, &warped),
303        SimilarityMetric::MSE => compute_mse(reference, &warped),
304    }
305}
306
307/// Compute gradient of similarity metric
308#[allow(dead_code)]
309fn compute_gradient(
310    reference: &GrayImage,
311    moving: &GrayImage,
312    transform: &TransformMatrix,
313    config: &IntensityRegistrationConfig,
314) -> Result<(Array1<f64>, f64)> {
315    let current_cost = compute_similarity(reference, moving, transform, config.metric)?;
316
317    // Compute gradient using finite differences
318    let mut gradient = Array1::zeros(6); // For affine transform parameters
319
320    for i in 0..6 {
321        let mut perturbed_transform = transform.clone();
322
323        // Perturb parameter
324        match i {
325            0 => perturbed_transform[[0, 0]] += config.step_size,
326            1 => perturbed_transform[[0, 1]] += config.step_size,
327            2 => perturbed_transform[[0, 2]] += config.step_size,
328            3 => perturbed_transform[[1, 0]] += config.step_size,
329            4 => perturbed_transform[[1, 1]] += config.step_size,
330            5 => perturbed_transform[[1, 2]] += config.step_size,
331            _ => {}
332        }
333
334        let perturbed_cost =
335            compute_similarity(reference, moving, &perturbed_transform, config.metric)?;
336        gradient[i] = (perturbed_cost - current_cost) / config.step_size;
337    }
338
339    Ok((gradient, current_cost))
340}
341
342/// Update transformation using gradient descent
343#[allow(dead_code)]
344fn update_transform_gradient_descent(
345    transform: &TransformMatrix,
346    gradient: &Array1<f64>,
347    learning_rate: f64,
348) -> TransformMatrix {
349    let mut updated = transform.clone();
350
351    // Update parameters
352    updated[[0, 0]] -= learning_rate * gradient[0];
353    updated[[0, 1]] -= learning_rate * gradient[1];
354    updated[[0, 2]] -= learning_rate * gradient[2];
355    updated[[1, 0]] -= learning_rate * gradient[3];
356    updated[[1, 1]] -= learning_rate * gradient[4];
357    updated[[1, 2]] -= learning_rate * gradient[5];
358
359    updated
360}
361
362/// Compute Sum of Squared Differences
363#[allow(dead_code)]
364fn compute_ssd(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
365    let (width, height) = image1.dimensions();
366    let mut ssd = 0.0;
367    let mut count = 0;
368
369    for y in 0..height {
370        for x in 0..width {
371            let val1 = image1.get_pixel(x, y)[0] as f64;
372            let val2 = image2.get_pixel(x, y)[0] as f64;
373
374            ssd += (val1 - val2).powi(2);
375            count += 1;
376        }
377    }
378
379    Ok(ssd / count as f64)
380}
381
382/// Compute Mean Squared Error
383#[allow(dead_code)]
384fn compute_mse(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
385    compute_ssd(image1, image2) // MSE is the same as average SSD
386}
387
388/// Compute Normalized Cross-Correlation
389#[allow(dead_code)]
390fn compute_ncc(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
391    let (width, height) = image1.dimensions();
392
393    let mut sum1 = 0.0;
394    let mut sum2 = 0.0;
395    let mut sum1_sq = 0.0;
396    let mut sum2_sq = 0.0;
397    let mut sum_cross = 0.0;
398    let mut count = 0;
399
400    for y in 0..height {
401        for x in 0..width {
402            let val1 = image1.get_pixel(x, y)[0] as f64;
403            let val2 = image2.get_pixel(x, y)[0] as f64;
404
405            // Skip zero pixels (likely from warping)
406            if val2 > 0.0 {
407                sum1 += val1;
408                sum2 += val2;
409                sum1_sq += val1 * val1;
410                sum2_sq += val2 * val2;
411                sum_cross += val1 * val2;
412                count += 1;
413            }
414        }
415    }
416
417    if count == 0 {
418        return Ok(0.0);
419    }
420
421    let n = count as f64;
422    let mean1 = sum1 / n;
423    let mean2 = sum2 / n;
424
425    let numerator = sum_cross - n * mean1 * mean2;
426    let var1 = sum1_sq - n * mean1 * mean1;
427    let var2 = sum2_sq - n * mean2 * mean2;
428
429    let denominator = (var1 * var2).sqrt();
430
431    if denominator > 1e-10 {
432        Ok(-numerator / denominator) // Negative because we want to minimize
433    } else {
434        Ok(0.0)
435    }
436}
437
438/// Compute Cross-Correlation
439#[allow(dead_code)]
440fn compute_cross_correlation(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
441    let (width, height) = image1.dimensions();
442    let mut cc = 0.0;
443    let mut count = 0;
444
445    for y in 0..height {
446        for x in 0..width {
447            let val1 = image1.get_pixel(x, y)[0] as f64;
448            let val2 = image2.get_pixel(x, y)[0] as f64;
449
450            if val2 > 0.0 {
451                cc += val1 * val2;
452                count += 1;
453            }
454        }
455    }
456
457    if count > 0 {
458        Ok(-cc / count as f64) // Negative because we want to maximize CC
459    } else {
460        Ok(0.0)
461    }
462}
463
464/// Compute Mutual Information
465#[allow(dead_code)]
466fn compute_mutual_information(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
467    let _jointhist = compute_joint_histogram(image1, image2, 256);
468    let (hist1, hist2) = compute_marginal_histograms(&_jointhist);
469
470    let mut mi = 0.0;
471    let total = _jointhist.sum();
472
473    for i in 0..256 {
474        for j in 0..256 {
475            let p_xy = _jointhist[[i, j]] / total;
476            let p_x = hist1[i] / total;
477            let p_y = hist2[j] / total;
478
479            if p_xy > 1e-10 && p_x > 1e-10 && p_y > 1e-10 {
480                mi += p_xy * (p_xy / (p_x * p_y)).ln();
481            }
482        }
483    }
484
485    Ok(-mi) // Negative because we want to maximize MI
486}
487
488/// Compute Normalized Mutual Information
489#[allow(dead_code)]
490fn compute_normalized_mutual_information(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
491    let _jointhist = compute_joint_histogram(image1, image2, 256);
492    let (hist1, hist2) = compute_marginal_histograms(&_jointhist);
493
494    let total = _jointhist.sum();
495
496    // Compute entropies
497    let mut h1 = 0.0;
498    let mut h2 = 0.0;
499    let mut h12 = 0.0;
500
501    for i in 0..256 {
502        let p1 = hist1[i] / total;
503        if p1 > 1e-10 {
504            h1 -= p1 * p1.ln();
505        }
506
507        let p2 = hist2[i] / total;
508        if p2 > 1e-10 {
509            h2 -= p2 * p2.ln();
510        }
511
512        for j in 0..256 {
513            let p_xy = _jointhist[[i, j]] / total;
514            if p_xy > 1e-10 {
515                h12 -= p_xy * p_xy.ln();
516            }
517        }
518    }
519
520    let nmi = (h1 + h2) / h12;
521    Ok(-nmi) // Negative because we want to maximize NMI
522}
523
524/// Compute joint histogram of two images
525#[allow(dead_code)]
526fn compute_joint_histogram(image1: &GrayImage, image2: &GrayImage, bins: usize) -> Array2<f64> {
527    let (width, height) = image1.dimensions();
528    let mut hist = Array2::zeros((bins, bins));
529
530    for y in 0..height {
531        for x in 0..width {
532            let val1 = image1.get_pixel(x, y)[0] as usize;
533            let val2 = image2.get_pixel(x, y)[0] as usize;
534
535            if val1 < bins && val2 < bins && val2 > 0 {
536                hist[[val1, val2]] += 1.0;
537            }
538        }
539    }
540
541    hist
542}
543
544/// Compute marginal histograms from joint histogram
545#[allow(dead_code)]
546fn compute_marginal_histograms(_jointhist: &Array2<f64>) -> (Array1<f64>, Array1<f64>) {
547    let (bins1, bins2) = _jointhist.dim();
548    let mut hist1 = Array1::zeros(bins1);
549    let mut hist2 = Array1::zeros(bins2);
550
551    for i in 0..bins1 {
552        for j in 0..bins2 {
553            hist1[i] += _jointhist[[i, j]];
554            hist2[j] += _jointhist[[i, j]];
555        }
556    }
557
558    (hist1, hist2)
559}
560
561/// Rigid registration using intensity-based methods
562#[allow(dead_code)]
563pub fn rigid_register_intensity(
564    reference: &GrayImage,
565    moving: &GrayImage,
566    config: &IntensityRegistrationConfig,
567) -> Result<RegistrationResult> {
568    // For rigid registration, we need to constrain the optimization
569    // This is a simplified version - would need proper rigid constraint handling
570    register_images_intensity(reference, moving, None, config)
571}
572
573#[cfg(test)]
574mod tests {
575    use super::*;
576    use image::{ImageBuffer, Luma};
577
578    fn create_test_image(width: u32, height: u32, pattern: u8) -> GrayImage {
579        ImageBuffer::from_fn(width, height, |x, y| {
580            Luma([((x + y + pattern as u32) % 256) as u8])
581        })
582    }
583
584    #[test]
585    fn test_intensity_config() {
586        let config = IntensityRegistrationConfig::default();
587        assert_eq!(config.metric, SimilarityMetric::NCC);
588        assert_eq!(config.optimizer, OptimizationMethod::GradientDescent);
589    }
590
591    #[test]
592    fn test_ssd_computation() {
593        let img1 = create_test_image(10, 10, 0);
594        let img2 = create_test_image(10, 10, 1);
595
596        let ssd = compute_ssd(&img1, &img2).expect("Operation failed");
597        assert!(ssd > 0.0);
598    }
599
600    #[test]
601    fn test_ncc_computation() {
602        let img1 = create_test_image(10, 10, 0);
603        let img2 = create_test_image(10, 10, 0); // Same image
604
605        let ncc = compute_ncc(&img1, &img2).expect("Operation failed");
606        assert!(ncc <= 0.0); // Should be close to -1 (perfect correlation)
607    }
608
609    #[test]
610    fn test_pyramid_building() {
611        let image = create_test_image(64, 64, 0);
612        let pyramid = build_image_pyramid(&image, 3);
613
614        assert_eq!(pyramid.len(), 3);
615        assert_eq!(pyramid[0].dimensions(), (64, 64));
616        assert_eq!(pyramid[1].dimensions(), (32, 32));
617        assert_eq!(pyramid[2].dimensions(), (16, 16));
618    }
619
620    #[test]
621    fn test_joint_histogram() {
622        let img1 = create_test_image(10, 10, 0);
623        let img2 = create_test_image(10, 10, 0);
624
625        let hist = compute_joint_histogram(&img1, &img2, 256);
626        assert!(hist.sum() > 0.0);
627    }
628
629    #[test]
630    fn test_mutual_information() {
631        let img1 = create_test_image(20, 20, 0);
632        let img2 = create_test_image(20, 20, 1);
633
634        let mi = compute_mutual_information(&img1, &img2).expect("Operation failed");
635        assert!(mi.is_finite());
636    }
637
638    #[test]
639    fn test_gradient_computation() {
640        let img1 = create_test_image(20, 20, 0);
641        let img2 = create_test_image(20, 20, 1);
642        let transform = identity_transform();
643        let config = IntensityRegistrationConfig::default();
644
645        let result = compute_gradient(&img1, &img2, &transform, &config);
646        assert!(result.is_ok());
647
648        let (gradient_cost, _cost_value) = result.expect("Operation failed");
649        assert_eq!(gradient_cost.len(), 6);
650    }
651}