1use crate::error::{Result, VisionError};
8use image::DynamicImage;
9use scirs2_core::ndarray::Array2;
10
11#[derive(Debug, Clone)]
13pub enum SegmentMethod {
14 Otsu,
16 Adaptive {
18 block_size: usize,
20 c: f32,
22 method: super::AdaptiveMethod,
24 },
25 KMeans {
27 k: usize,
29 max_iterations: usize,
31 },
32 GrabCut {
34 rect: (u32, u32, u32, u32),
36 n_components: usize,
38 },
39 Watershed {
41 n_markers: Option<usize>,
43 connectivity: u8,
45 },
46}
47
48#[derive(Debug, Clone)]
50pub struct SegmentResult {
51 pub labels: Array2<u32>,
55 pub n_segments: usize,
57 pub method_name: String,
59}
60
61pub 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
113fn 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
138fn 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
168fn 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, ¶ms)?;
179
180 Ok(SegmentResult {
181 labels: result.labels,
182 n_segments: k,
183 method_name: format!("KMeans(k={})", k),
184 })
185}
186
187fn 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, ¶ms)?;
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
217fn 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 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 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 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 let bright = result.labels[[16, 5]]; let dark = result.labels[[16, 35]]; 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 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 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 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 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 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}