Skip to main content

scirs2_vision/segmentation/
unified.rs

1//! Unified image segmentation API
2//!
3//! Provides a single entry point for segmenting images using multiple methods:
4//! Otsu thresholding, adaptive thresholding, K-means color segmentation,
5//! GrabCut foreground extraction, and watershed segmentation.
6
7use crate::error::{Result, VisionError};
8use image::DynamicImage;
9use scirs2_core::ndarray::Array2;
10
11/// Available segmentation methods
12#[derive(Debug, Clone)]
13pub enum SegmentMethod {
14    /// Otsu's automatic thresholding (binary segmentation)
15    Otsu,
16    /// Adaptive thresholding with local neighborhood
17    Adaptive {
18        /// Block size for local thresholding (must be odd >= 3)
19        block_size: usize,
20        /// Constant subtracted from the mean
21        c: f32,
22        /// Thresholding method (Mean or Gaussian)
23        method: super::AdaptiveMethod,
24    },
25    /// K-means color segmentation
26    KMeans {
27        /// Number of clusters
28        k: usize,
29        /// Maximum iterations
30        max_iterations: usize,
31    },
32    /// GrabCut foreground extraction with bounding box
33    GrabCut {
34        /// Bounding box (x, y, width, height) containing the foreground
35        rect: (u32, u32, u32, u32),
36        /// Number of GMM components
37        n_components: usize,
38    },
39    /// Watershed segmentation
40    Watershed {
41        /// Number of initial markers (None for auto)
42        n_markers: Option<usize>,
43        /// Connectivity (4 or 8)
44        connectivity: u8,
45    },
46}
47
48/// Result of unified segmentation
49#[derive(Debug, Clone)]
50pub struct SegmentResult {
51    /// Label map (height x width). Each pixel has an integer label.
52    /// For binary methods (Otsu, Adaptive): 0 = background, 1 = foreground
53    /// For multi-label methods: 0..n_segments
54    pub labels: Array2<u32>,
55    /// Number of unique segments found
56    pub n_segments: usize,
57    /// Method that was used
58    pub method_name: String,
59}
60
61/// Segment an image using the specified method
62///
63/// This is the unified API for image segmentation. It dispatches to the
64/// appropriate algorithm based on the `method` parameter.
65///
66/// # Arguments
67///
68/// * `img` - Input image
69/// * `method` - Segmentation method and its parameters
70///
71/// # Returns
72///
73/// * Segmentation result containing the label map
74///
75/// # Example
76///
77/// ```rust
78/// use scirs2_vision::segmentation::unified::{segment, SegmentMethod};
79/// use image::{DynamicImage, GrayImage, Luma};
80///
81/// # fn main() -> scirs2_vision::error::Result<()> {
82/// let mut buf = GrayImage::new(32, 32);
83/// for y in 0..32u32 {
84///     for x in 0..32u32 {
85///         let val = if x < 16 { 200u8 } else { 50u8 };
86///         buf.put_pixel(x, y, Luma([val]));
87///     }
88/// }
89/// let img = DynamicImage::ImageLuma8(buf);
90///
91/// let result = segment(&img, SegmentMethod::Otsu)?;
92/// assert_eq!(result.n_segments, 2);
93/// # Ok(())
94/// # }
95/// ```
96pub fn segment(img: &DynamicImage, method: SegmentMethod) -> Result<SegmentResult> {
97    match method {
98        SegmentMethod::Otsu => segment_otsu(img),
99        SegmentMethod::Adaptive {
100            block_size,
101            c,
102            method: adaptive_method,
103        } => segment_adaptive(img, block_size, c, adaptive_method),
104        SegmentMethod::KMeans { k, max_iterations } => segment_kmeans(img, k, max_iterations),
105        SegmentMethod::GrabCut { rect, n_components } => segment_grabcut(img, rect, n_components),
106        SegmentMethod::Watershed {
107            n_markers,
108            connectivity,
109        } => segment_watershed(img, n_markers, connectivity),
110    }
111}
112
113/// Otsu segmentation
114fn segment_otsu(img: &DynamicImage) -> Result<SegmentResult> {
115    let (binary, _threshold) = super::otsu_threshold(img)?;
116    let (width, height) = binary.dimensions();
117    let h = height as usize;
118    let w = width as usize;
119
120    let mut labels = Array2::zeros((h, w));
121    for y in 0..h {
122        for x in 0..w {
123            labels[[y, x]] = if binary.get_pixel(x as u32, y as u32)[0] > 0 {
124                1
125            } else {
126                0
127            };
128        }
129    }
130
131    Ok(SegmentResult {
132        labels,
133        n_segments: 2,
134        method_name: "Otsu".to_string(),
135    })
136}
137
138/// Adaptive thresholding segmentation
139fn segment_adaptive(
140    img: &DynamicImage,
141    block_size: usize,
142    c: f32,
143    method: super::AdaptiveMethod,
144) -> Result<SegmentResult> {
145    let binary = super::adaptive_threshold(img, block_size, c, method)?;
146    let (width, height) = binary.dimensions();
147    let h = height as usize;
148    let w = width as usize;
149
150    let mut labels = Array2::zeros((h, w));
151    for y in 0..h {
152        for x in 0..w {
153            labels[[y, x]] = if binary.get_pixel(x as u32, y as u32)[0] > 0 {
154                1
155            } else {
156                0
157            };
158        }
159    }
160
161    Ok(SegmentResult {
162        labels,
163        n_segments: 2,
164        method_name: "Adaptive".to_string(),
165    })
166}
167
168/// K-means color segmentation
169fn segment_kmeans(img: &DynamicImage, k: usize, max_iterations: usize) -> Result<SegmentResult> {
170    let params = super::kmeans_seg::KMeansSegParams {
171        k,
172        max_iterations,
173        epsilon: 1e-4,
174        n_init: 3,
175        use_color: true,
176    };
177
178    let result = super::kmeans_seg::kmeans_segment(img, &params)?;
179
180    Ok(SegmentResult {
181        labels: result.labels,
182        n_segments: k,
183        method_name: format!("KMeans(k={})", k),
184    })
185}
186
187/// GrabCut foreground extraction
188fn segment_grabcut(
189    img: &DynamicImage,
190    rect: (u32, u32, u32, u32),
191    n_components: usize,
192) -> Result<SegmentResult> {
193    let params = super::grabcut::GrabCutParams {
194        n_components,
195        max_iterations: 10,
196        epsilon: 1e-3,
197        smoothness: 50.0,
198    };
199
200    let result = super::grabcut::grabcut_rect(img, rect, &params)?;
201    let (h, w) = result.foreground_mask.dim();
202
203    let mut labels = Array2::zeros((h, w));
204    for y in 0..h {
205        for x in 0..w {
206            labels[[y, x]] = if result.foreground_mask[[y, x]] { 1 } else { 0 };
207        }
208    }
209
210    Ok(SegmentResult {
211        labels,
212        n_segments: 2,
213        method_name: "GrabCut".to_string(),
214    })
215}
216
217/// Watershed segmentation
218fn segment_watershed(
219    img: &DynamicImage,
220    _n_markers: Option<usize>,
221    connectivity: u8,
222) -> Result<SegmentResult> {
223    let conn = if connectivity == 4 { 4 } else { 8 };
224
225    let labels = super::watershed::watershed(img, None, conn)?;
226
227    // Count unique labels
228    let mut unique_labels = std::collections::HashSet::new();
229    for &label in labels.iter() {
230        unique_labels.insert(label);
231    }
232
233    Ok(SegmentResult {
234        labels,
235        n_segments: unique_labels.len(),
236        method_name: "Watershed".to_string(),
237    })
238}
239
240#[cfg(test)]
241mod tests {
242    use super::*;
243    use image::{GrayImage, Luma};
244
245    fn create_bimodal_image() -> DynamicImage {
246        // Create a bimodal image with a gradient band in the middle.
247        // Left half is bright (220), right half is dark (20), with a smooth
248        // transition zone at the boundary. This ensures Otsu picks a threshold
249        // between the two modes rather than at either mode value.
250        let mut buf = GrayImage::new(40, 32);
251        for y in 0..32u32 {
252            for x in 0..40u32 {
253                let val = if x < 15 {
254                    220u8
255                } else if x > 24 {
256                    20u8
257                } else {
258                    // Transition zone: linearly interpolate
259                    let t = (x - 15) as f32 / 10.0;
260                    (220.0 * (1.0 - t) + 20.0 * t) as u8
261                };
262                buf.put_pixel(x, y, Luma([val]));
263            }
264        }
265        DynamicImage::ImageLuma8(buf)
266    }
267
268    #[test]
269    fn test_segment_otsu() {
270        let img = create_bimodal_image();
271        let result = segment(&img, SegmentMethod::Otsu).expect("Otsu failed");
272        assert_eq!(result.n_segments, 2);
273        assert_eq!(result.labels.dim(), (32, 40));
274        assert_eq!(result.method_name, "Otsu");
275
276        // Bright side should be foreground (label 1), dark side background (label 0)
277        let bright = result.labels[[16, 5]]; // x=5, value 220
278        let dark = result.labels[[16, 35]]; // x=35, value 20
279        assert_ne!(bright, dark, "Otsu should separate bright and dark regions");
280    }
281
282    #[test]
283    fn test_segment_adaptive_mean() {
284        let img = create_bimodal_image();
285        let result = segment(
286            &img,
287            SegmentMethod::Adaptive {
288                block_size: 7,
289                c: 0.0,
290                method: super::super::AdaptiveMethod::Mean,
291            },
292        )
293        .expect("Adaptive mean failed");
294
295        assert_eq!(result.n_segments, 2);
296        assert_eq!(result.labels.dim(), (32, 40));
297    }
298
299    #[test]
300    fn test_segment_adaptive_gaussian() {
301        let img = create_bimodal_image();
302        let result = segment(
303            &img,
304            SegmentMethod::Adaptive {
305                block_size: 7,
306                c: 0.0,
307                method: super::super::AdaptiveMethod::Gaussian,
308            },
309        )
310        .expect("Adaptive gaussian failed");
311
312        assert_eq!(result.n_segments, 2);
313    }
314
315    #[test]
316    fn test_segment_kmeans() {
317        let mut buf = image::RgbImage::new(20, 20);
318        for y in 0..20u32 {
319            for x in 0..20u32 {
320                let color = if x < 10 {
321                    [200u8, 50, 50]
322                } else {
323                    [50u8, 50, 200]
324                };
325                buf.put_pixel(x, y, image::Rgb(color));
326            }
327        }
328        let img = DynamicImage::ImageRgb8(buf);
329
330        let result = segment(
331            &img,
332            SegmentMethod::KMeans {
333                k: 2,
334                max_iterations: 100,
335            },
336        )
337        .expect("KMeans failed");
338
339        assert_eq!(result.n_segments, 2);
340        assert_eq!(result.labels.dim(), (20, 20));
341    }
342
343    #[test]
344    fn test_segment_grabcut() {
345        // Create image with bright center
346        let mut buf = image::RgbImage::new(20, 20);
347        for y in 0..20u32 {
348            for x in 0..20u32 {
349                let is_center = (5..15).contains(&x) && (5..15).contains(&y);
350                let color = if is_center {
351                    [220u8, 220, 220]
352                } else {
353                    [20u8, 20, 20]
354                };
355                buf.put_pixel(x, y, image::Rgb(color));
356            }
357        }
358        let img = DynamicImage::ImageRgb8(buf);
359
360        let result = segment(
361            &img,
362            SegmentMethod::GrabCut {
363                rect: (4, 4, 12, 12),
364                n_components: 3,
365            },
366        )
367        .expect("GrabCut failed");
368
369        assert_eq!(result.n_segments, 2);
370        assert_eq!(result.labels.dim(), (20, 20));
371    }
372
373    #[test]
374    fn test_segment_result_labels_range() {
375        // Use a simple 32x32 image for this test
376        let mut buf = GrayImage::new(32, 32);
377        for y in 0..32u32 {
378            for x in 0..32u32 {
379                buf.put_pixel(x, y, Luma([if x < 16 { 240u8 } else { 10u8 }]));
380            }
381        }
382        // Add transition pixels so Otsu threshold is between the modes
383        for y in 0..32u32 {
384            buf.put_pixel(15, y, Luma([125u8]));
385            buf.put_pixel(16, y, Luma([125u8]));
386        }
387        let img = DynamicImage::ImageLuma8(buf);
388        let result = segment(&img, SegmentMethod::Otsu).expect("Otsu failed");
389
390        // All labels should be 0 or 1 for binary segmentation
391        for &label in result.labels.iter() {
392            assert!(label <= 1, "Label should be 0 or 1, got {}", label);
393        }
394    }
395
396    #[test]
397    fn test_segment_kmeans_three_regions() {
398        let mut buf = image::RgbImage::new(30, 10);
399        for y in 0..10u32 {
400            for x in 0..10u32 {
401                buf.put_pixel(x, y, image::Rgb([255, 0, 0]));
402                buf.put_pixel(x + 10, y, image::Rgb([0, 255, 0]));
403                buf.put_pixel(x + 20, y, image::Rgb([0, 0, 255]));
404            }
405        }
406        let img = DynamicImage::ImageRgb8(buf);
407
408        let result = segment(
409            &img,
410            SegmentMethod::KMeans {
411                k: 3,
412                max_iterations: 100,
413            },
414        )
415        .expect("KMeans 3-cluster failed");
416
417        assert_eq!(result.n_segments, 3);
418
419        // Three regions should have different labels
420        let l0 = result.labels[[5, 5]];
421        let l1 = result.labels[[5, 15]];
422        let l2 = result.labels[[5, 25]];
423        assert_ne!(l0, l1);
424        assert_ne!(l1, l2);
425        assert_ne!(l0, l2);
426    }
427}