1use 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#[derive(Debug, Clone, Copy, PartialEq)]
16pub enum SimilarityMetric {
17 SSD,
19 NCC,
21 MI,
23 NMI,
25 CC,
27 MSE,
29}
30
31#[derive(Debug, Clone, Copy, PartialEq)]
33pub enum OptimizationMethod {
34 GradientDescent,
36 Powell,
38 Simplex,
40 ConjugateGradient,
42}
43
44#[derive(Debug, Clone)]
46pub struct IntensityRegistrationConfig {
47 pub metric: SimilarityMetric,
49 pub optimizer: OptimizationMethod,
51 pub params: RegistrationParams,
53 pub step_size: f64,
55 pub learning_rate: f64,
57 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#[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#[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, ¤t_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, ¤t_transform, config)?;
120
121 if (current_cost - new_cost).abs() < config.params.tolerance {
123 converged = true;
124 break;
125 }
126
127 current_transform = match config.optimizer {
129 OptimizationMethod::GradientDescent => update_transform_gradient_descent(
130 ¤t_transform,
131 &gradient,
132 config.learning_rate,
133 ),
134 OptimizationMethod::Powell => {
135 update_transform_gradient_descent(
137 ¤t_transform,
138 &gradient,
139 config.learning_rate,
140 )
141 }
142 OptimizationMethod::Simplex => {
143 update_transform_gradient_descent(
145 ¤t_transform,
146 &gradient,
147 config.learning_rate,
148 )
149 }
150 OptimizationMethod::ConjugateGradient => {
151 update_transform_gradient_descent(
153 ¤t_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(), })
171}
172
173#[allow(dead_code)]
175fn multi_resolution_register(
176 reference: &GrayImage,
177 moving: &GrayImage,
178 initial_transform: &TransformMatrix,
179 config: &IntensityRegistrationConfig,
180) -> Result<RegistrationResult> {
181 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 for level in (0..config.params.pyramid_levels).rev() {
190 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 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 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#[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#[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 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#[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 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#[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 let mut gradient = Array1::zeros(6); for i in 0..6 {
321 let mut perturbed_transform = transform.clone();
322
323 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#[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 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#[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#[allow(dead_code)]
384fn compute_mse(image1: &GrayImage, image2: &GrayImage) -> Result<f64> {
385 compute_ssd(image1, image2) }
387
388#[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 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) } else {
434 Ok(0.0)
435 }
436}
437
438#[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) } else {
460 Ok(0.0)
461 }
462}
463
464#[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) }
487
488#[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 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) }
523
524#[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#[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#[allow(dead_code)]
563pub fn rigid_register_intensity(
564 reference: &GrayImage,
565 moving: &GrayImage,
566 config: &IntensityRegistrationConfig,
567) -> Result<RegistrationResult> {
568 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); let ncc = compute_ncc(&img1, &img2).expect("Operation failed");
606 assert!(ncc <= 0.0); }
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}