Files
rustytorch/demos/rtx-segmentation-demo/src/mock_segmentation.rs
T
2026-03-04 00:08:42 +00:00

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);
}
}
}