421 lines
14 KiB
Rust
421 lines
14 KiB
Rust
// 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<u64>,
|
|
) -> Result<SegmentationResult> {
|
|
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<u8> {
|
|
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<f32> {
|
|
(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<u64>,
|
|
) -> Result<SegmentationResult> {
|
|
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<u64>,
|
|
) -> Result<SegmentationResult> {
|
|
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<f32> = 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<f32> = 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<usize> = (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);
|
|
}
|
|
}
|
|
}
|