// Mock segmentation generator for realistic region-based masks use crate::error::{Result, SegmentationError}; use rand::rngs::StdRng; use rand::{Rng, SeedableRng}; use segmentation_shared::{ SegmentationConfig, SegmentationMask, SegmentationModel, SegmentationResult, }; /// Generate a mock segmentation mask with realistic region-based patterns pub fn generate_mock_segmentation( config: &SegmentationConfig, width: usize, height: usize, seed: Option, ) -> Result { if width == 0 || height == 0 { return Err(SegmentationError::InvalidDimensions( "width and height must be greater than 0".to_string(), )); } if config.num_classes == 0 { return Err(SegmentationError::InvalidConfig( "num_classes must be greater than 0".to_string(), )); } let mut rng = if let Some(s) = seed { StdRng::seed_from_u64(s) } else { StdRng::from_entropy() }; let start_time = std::time::Instant::now(); let class_indices = generate_region_based_mask(width, height, config.num_classes, &mut rng); let inference_time_ms = start_time.elapsed().as_secs_f64() * 1000.0 + 10.0 + rng.r#gen_range(5.0..20.0); let mask = SegmentationMask::new(width, height, class_indices); let class_confidences = generate_class_confidences(config.num_classes, &mut rng); Ok(SegmentationResult::new(mask, inference_time_ms).with_confidences(class_confidences)) } /// Generate region-based mask using simple blob generation fn generate_region_based_mask( width: usize, height: usize, num_classes: usize, rng: &mut StdRng, ) -> Vec { let mut mask = vec![0u8; width * height]; let num_blobs = (num_classes - 1).clamp(1, 10); for _ in 0..num_blobs { let class_id = rng.r#gen_range(1..num_classes) as u8; let center_x = rng.r#gen_range(0..width); let center_y = rng.r#gen_range(0..height); let radius = rng.r#gen_range(10..50).min(width.min(height) / 4); for y in 0..height { for x in 0..width { let dx = x as i32 - center_x as i32; let dy = y as i32 - center_y as i32; let dist_sq = (dx * dx + dy * dy) as f64; let radius_sq = (radius * radius) as f64; if dist_sq < radius_sq { let idx = y * width + x; mask[idx] = class_id; } } } } mask } /// Generate realistic confidence scores per class fn generate_class_confidences(num_classes: usize, rng: &mut StdRng) -> Vec { (0..num_classes) .map(|_| rng.r#gen_range(0.6..0.95)) .collect() } /// Generate mock segmentation for specific model type pub fn generate_for_model( model: SegmentationModel, width: usize, height: usize, seed: Option, ) -> Result { let (num_classes, device) = match model { SegmentationModel::DeepLabV3 => (21, "cpu"), SegmentationModel::SegFormer => (150, "cpu"), SegmentationModel::UNet => (21, "cpu"), SegmentationModel::FCN => (21, "cpu"), }; let config = SegmentationConfig::new(model, device.to_string(), num_classes, 512); generate_mock_segmentation(&config, width, height, seed) } /// Generate mock segmentation with specific class distribution pub fn generate_with_distribution( width: usize, height: usize, class_weights: &[f32], seed: Option, ) -> Result { if class_weights.is_empty() { return Err(SegmentationError::InvalidConfig( "class_weights cannot be empty".to_string(), )); } if width == 0 || height == 0 { return Err(SegmentationError::InvalidDimensions( "width and height must be greater than 0".to_string(), )); } let total_weight: f32 = class_weights.iter().sum(); if total_weight <= 0.0 { return Err(SegmentationError::InvalidConfig( "total class weight must be greater than 0".to_string(), )); } let normalized_weights: Vec = class_weights.iter().map(|w| w / total_weight).collect(); let mut rng = if let Some(s) = seed { StdRng::seed_from_u64(s) } else { StdRng::from_entropy() }; let start_time = std::time::Instant::now(); let mut mask = vec![0u8; width * height]; for idx in 0..mask.len() { let rand_val: f32 = rng.r#gen(); let mut cumulative = 0.0; for (class_id, &weight) in normalized_weights.iter().enumerate() { cumulative += weight; if rand_val < cumulative { mask[idx] = class_id as u8; break; } } } let inference_time_ms = start_time.elapsed().as_secs_f64() * 1000.0 + 10.0 + rng.r#gen_range(5.0..20.0); let mask_obj = SegmentationMask::new(width, height, mask); let confidences = generate_class_confidences(class_weights.len(), &mut rng); Ok(SegmentationResult::new(mask_obj, inference_time_ms).with_confidences(confidences)) } #[cfg(test)] mod tests { use super::*; use segmentation_shared::SegmentationModel; #[test] fn test_generate_mock_segmentation_valid() { let config = SegmentationConfig::new(SegmentationModel::DeepLabV3, "cpu".to_string(), 21, 512); let result = generate_mock_segmentation(&config, 100, 100, Some(42)).unwrap(); assert_eq!(result.mask.width, 100); assert_eq!(result.mask.height, 100); assert_eq!(result.mask.class_indices.len(), 10000); assert!(result.inference_time_ms > 0.0); assert_eq!(result.class_confidences.len(), 21); } #[test] fn test_generate_mock_segmentation_deterministic() { let config = SegmentationConfig::new(SegmentationModel::UNet, "cpu".to_string(), 10, 256); let result1 = generate_mock_segmentation(&config, 50, 50, Some(123)).unwrap(); let result2 = generate_mock_segmentation(&config, 50, 50, Some(123)).unwrap(); assert_eq!(result1.mask.class_indices, result2.mask.class_indices); } #[test] fn test_generate_mock_segmentation_different_seeds() { let config = SegmentationConfig::new(SegmentationModel::FCN, "cpu".to_string(), 5, 256); let result1 = generate_mock_segmentation(&config, 50, 50, Some(1)).unwrap(); let result2 = generate_mock_segmentation(&config, 50, 50, Some(2)).unwrap(); assert_ne!(result1.mask.class_indices, result2.mask.class_indices); } #[test] fn test_generate_mock_segmentation_zero_width() { let config = SegmentationConfig::new(SegmentationModel::DeepLabV3, "cpu".to_string(), 21, 512); let result = generate_mock_segmentation(&config, 0, 100, Some(42)); assert!(result.is_err()); assert!( result .unwrap_err() .to_string() .contains("must be greater than 0") ); } #[test] fn test_generate_mock_segmentation_zero_height() { let config = SegmentationConfig::new(SegmentationModel::DeepLabV3, "cpu".to_string(), 21, 512); let result = generate_mock_segmentation(&config, 100, 0, Some(42)); assert!(result.is_err()); assert!( result .unwrap_err() .to_string() .contains("must be greater than 0") ); } #[test] fn test_generate_mock_segmentation_zero_classes() { let config = SegmentationConfig::new(SegmentationModel::DeepLabV3, "cpu".to_string(), 0, 512); let result = generate_mock_segmentation(&config, 100, 100, Some(42)); assert!(result.is_err()); assert!(result.unwrap_err().to_string().contains("num_classes")); } #[test] fn test_generate_mock_segmentation_class_indices_in_range() { let config = SegmentationConfig::new(SegmentationModel::DeepLabV3, "cpu".to_string(), 5, 512); let result = generate_mock_segmentation(&config, 50, 50, Some(42)).unwrap(); for &idx in &result.mask.class_indices { assert!(idx < 5, "class index {} is out of range", idx); } } #[test] fn test_generate_mock_segmentation_confidences_in_range() { let config = SegmentationConfig::new(SegmentationModel::SegFormer, "cpu".to_string(), 10, 512); let result = generate_mock_segmentation(&config, 50, 50, Some(42)).unwrap(); for &conf in &result.class_confidences { assert!( conf >= 0.0 && conf <= 1.0, "confidence {} is out of range", conf ); } } // generate_for_model tests #[test] fn test_generate_for_model_deeplabv3() { let result = generate_for_model(SegmentationModel::DeepLabV3, 100, 100, Some(42)).unwrap(); assert_eq!(result.class_confidences.len(), 21); } #[test] fn test_generate_for_model_segformer() { let result = generate_for_model(SegmentationModel::SegFormer, 100, 100, Some(42)).unwrap(); assert_eq!(result.class_confidences.len(), 150); } #[test] fn test_generate_for_model_unet() { let result = generate_for_model(SegmentationModel::UNet, 100, 100, Some(42)).unwrap(); assert_eq!(result.class_confidences.len(), 21); } #[test] fn test_generate_for_model_fcn() { let result = generate_for_model(SegmentationModel::FCN, 100, 100, Some(42)).unwrap(); assert_eq!(result.class_confidences.len(), 21); } // generate_with_distribution tests #[test] fn test_generate_with_distribution_uniform() { let weights = vec![1.0, 1.0, 1.0]; let result = generate_with_distribution(100, 100, &weights, Some(42)).unwrap(); assert_eq!(result.mask.width, 100); assert_eq!(result.mask.height, 100); assert_eq!(result.class_confidences.len(), 3); } #[test] fn test_generate_with_distribution_skewed() { let weights = vec![0.8, 0.1, 0.1]; let result = generate_with_distribution(1000, 1000, &weights, Some(42)).unwrap(); let histogram = result.mask.class_histogram(); assert!(histogram[0] > histogram[1], "background should dominate"); assert!(histogram[0] > histogram[2], "background should dominate"); } #[test] fn test_generate_with_distribution_empty_weights() { let weights: Vec = vec![]; let result = generate_with_distribution(100, 100, &weights, Some(42)); assert!(result.is_err()); assert!(result.unwrap_err().to_string().contains("cannot be empty")); } #[test] fn test_generate_with_distribution_zero_weight() { let weights = vec![0.0, 0.0, 0.0]; let result = generate_with_distribution(100, 100, &weights, Some(42)); assert!(result.is_err()); assert!( result .unwrap_err() .to_string() .contains("must be greater than 0") ); } #[test] fn test_generate_with_distribution_single_class() { let weights = vec![1.0]; let result = generate_with_distribution(50, 50, &weights, Some(42)).unwrap(); for &idx in &result.mask.class_indices { assert_eq!(idx, 0); } } #[test] fn test_generate_with_distribution_deterministic() { let weights = vec![0.5, 0.3, 0.2]; let result1 = generate_with_distribution(50, 50, &weights, Some(999)).unwrap(); let result2 = generate_with_distribution(50, 50, &weights, Some(999)).unwrap(); assert_eq!(result1.mask.class_indices, result2.mask.class_indices); } #[test] fn test_generate_with_distribution_zero_dimensions() { let weights = vec![1.0, 1.0]; let result = generate_with_distribution(0, 100, &weights, Some(42)); assert!(result.is_err()); assert!( result .unwrap_err() .to_string() .contains("must be greater than 0") ); } // generate_region_based_mask tests #[test] fn test_generate_region_based_mask_size() { let mut rng = StdRng::seed_from_u64(42); let mask = generate_region_based_mask(100, 100, 5, &mut rng); assert_eq!(mask.len(), 10000); } #[test] fn test_generate_region_based_mask_class_range() { let mut rng = StdRng::seed_from_u64(42); let mask = generate_region_based_mask(100, 100, 5, &mut rng); for &idx in &mask { assert!(idx < 5); } } #[test] fn test_generate_region_based_mask_has_regions() { let mut rng = StdRng::seed_from_u64(42); let mask = generate_region_based_mask(100, 100, 5, &mut rng); let histogram: Vec = (0..5) .map(|class| mask.iter().filter(|&&c| c == class).count()) .collect(); let non_zero_classes = histogram.iter().filter(|&&count| count > 0).count(); assert!( non_zero_classes >= 2, "should have at least 2 classes present" ); } // generate_class_confidences tests #[test] fn test_generate_class_confidences_count() { let mut rng = StdRng::seed_from_u64(42); let confidences = generate_class_confidences(10, &mut rng); assert_eq!(confidences.len(), 10); } #[test] fn test_generate_class_confidences_range() { let mut rng = StdRng::seed_from_u64(42); let confidences = generate_class_confidences(20, &mut rng); for &conf in &confidences { assert!(conf >= 0.6 && conf <= 0.95); } } }