Initial commit
This commit is contained in:
@@ -0,0 +1,420 @@
|
||||
// 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user