Files
rustytorch/crates/training/rtx-nas/tests/integration_tests.rs
T
2026-03-04 00:08:42 +00:00

597 lines
18 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Integration tests for rtx-nas
use rtx_nas::{
algorithms::{DARTS, DARTSConfig, RandomSearch, RandomSearchConfig},
search_space::{CellConfig, DARTSSearchSpace, SearchSpace},
};
use rtx_tensor::Device;
#[test]
fn test_random_search_workflow() {
// Create search space
let search_space = DARTSSearchSpace::default().expect("Failed to create search space");
// Configure random search
let config = RandomSearchConfig::new(5);
let mut search = RandomSearch::new(config).expect("Failed to create random search");
// Sample architectures
search
.sample(&search_space)
.expect("Failed to sample architectures");
// Verify samples
assert_eq!(search.num_samples(), 5);
assert!(!search.samples().is_empty());
// Test get_best
let scores = vec![0.5, 0.3, 0.9, 0.4, 0.6];
let best = search.get_best(&scores).expect("Failed to get best");
assert_eq!(best.id, search.samples()[2].id);
// Test get_topk
let topk = search.get_topk(&scores, 3).expect("Failed to get topk");
assert_eq!(topk.len(), 3);
}
#[test]
fn test_darts_workflow() {
let device = Device::cuda(0).unwrap_or(Device::default());
// Configure DARTS
let config = DARTSConfig::default();
let cell_configs = vec![CellConfig::default_darts()];
// Create DARTS instance
let mut darts = DARTS::new(config, cell_configs, &device).expect("Failed to create DARTS");
// Check initial state
assert_eq!(darts.current_epoch(), 0);
assert!(darts.is_warmup());
// Perform warmup step
darts
.step(0.5, 0.6, None, None)
.expect("Failed to perform step");
assert_eq!(darts.current_epoch(), 1);
// Get probabilities
let probs = darts.get_probabilities();
assert_eq!(probs.len(), 1);
// Check all probabilities sum to 1.0
for prob_map in probs {
for prob_vec in prob_map.values() {
let sum: f32 = prob_vec.iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
}
}
// Derive architecture
let arch = darts
.derive_architecture()
.expect("Failed to derive architecture");
assert!(arch.validate().is_ok());
}
#[test]
fn test_search_space_encode_decode() {
let search_space = DARTSSearchSpace::default().expect("Failed to create search space");
// Sample an architecture
let arch = search_space
.sample()
.expect("Failed to sample architecture");
// Encode
let encoding = search_space
.encode(&arch)
.expect("Failed to encode architecture");
assert!(!encoding.is_empty());
// Decode
let decoded = search_space
.decode(&encoding)
.expect("Failed to decode architecture");
// Both should have the same structure
assert_eq!(decoded.num_cells(), arch.num_cells());
}
#[test]
fn test_cell_structure() {
use rtx_nas::search_space::{Cell, Edge, OperationType};
// Create a cell
let config = CellConfig::default_darts();
let mut cell = Cell::new(config).expect("Failed to create cell");
// Set operations on edges
let edge = Edge::new(0, 2);
cell.set_operation(edge, OperationType::Conv3x3)
.expect("Failed to set operation");
// Query cell structure
assert!(cell.total_nodes() > 0);
assert!(!cell.edges().is_empty());
assert_eq!(cell.get_operation(&edge), Some(OperationType::Conv3x3));
}
#[test]
fn test_darts_multi_step() {
let device = Device::cuda(0).unwrap_or(Device::default());
// Create DARTS with no warmup
let mut config = DARTSConfig::default();
config.warmup_epochs = 0;
let cell_configs = vec![CellConfig::default_darts()];
let mut darts = DARTS::new(config, cell_configs, &device).expect("Failed to create DARTS");
// Perform multiple steps
for _ in 0..5 {
darts
.step(0.5, 0.6, None, None)
.expect("Failed to perform step");
}
assert_eq!(darts.current_epoch(), 5);
assert!(!darts.is_warmup());
// Architecture should still be derivable
let arch = darts
.derive_architecture()
.expect("Failed to derive architecture");
assert!(arch.validate().is_ok());
}
#[test]
fn test_architecture_validation() {
use rtx_nas::search_space::{Architecture, Cell};
let cell = Cell::default_darts().expect("Failed to create cell");
// Valid architecture
let arch = Architecture::new("test_arch".into(), vec![cell.clone()], 16, 8);
assert!(arch.validate().is_ok());
// Invalid: no cells
let arch = Architecture::new("test_arch".into(), vec![], 16, 8);
assert!(arch.validate().is_err());
// Invalid: zero channels
let arch = Architecture::new("test_arch".into(), vec![cell.clone()], 0, 8);
assert!(arch.validate().is_err());
// Invalid: zero layers
let arch = Architecture::new("test_arch".into(), vec![cell], 16, 0);
assert!(arch.validate().is_err());
}
#[test]
fn test_search_space_num_choices() {
let space = DARTSSearchSpace::default().expect("Failed to create search space");
let num_choices = space.num_choices();
// With default config: 2 cells × 14 edges × 9 operations = 252
assert_eq!(num_choices, 252);
}
// =============================================================================
// PC-DARTS Integration Tests
// =============================================================================
#[test]
fn test_pcdarts_workflow() {
use rtx_nas::algorithms::{PCDARTS, PCDARTSConfig};
let device = Device::cuda(0).unwrap_or(Device::default());
// Configure PC-DARTS with partial channels
let config = PCDARTSConfig::default();
let cell_configs = vec![CellConfig::default_darts()];
// Create PC-DARTS instance
let mut pcdarts =
PCDARTS::new(config, cell_configs, &device).expect("Failed to create PC-DARTS");
// Verify initial state
assert_eq!(pcdarts.current_epoch(), 0);
assert!(pcdarts.is_warmup());
// Perform warmup step
pcdarts
.step(0.5, 0.6, None, None)
.expect("Failed to perform PC-DARTS step");
assert_eq!(pcdarts.current_epoch(), 1);
// Derive architecture
let arch = pcdarts
.derive_architecture()
.expect("Failed to derive PC-DARTS architecture");
assert!(arch.validate().is_ok());
}
#[test]
fn test_pcdarts_channel_fraction() {
use rtx_nas::algorithms::{DARTSConfig, PCDARTSConfig};
// Test default channel fraction
let config = PCDARTSConfig::default();
assert_eq!(config.channel_fraction, 0.125);
// Test custom channel fraction
let base_config = DARTSConfig::default();
let config = PCDARTSConfig::new(base_config, 0.25);
assert_eq!(config.channel_fraction, 0.25);
// Validate config
assert!(config.validate().is_ok());
}
#[test]
fn test_pcdarts_edge_normalization() {
use rtx_nas::algorithms::{PCDARTS, PCDARTSConfig};
let device = Device::cuda(0).unwrap_or(Device::default());
// Enable edge normalization
let mut config = PCDARTSConfig::default();
config.edge_normalization = true;
let cell_configs = vec![CellConfig::default_darts()];
let mut pcdarts =
PCDARTS::new(config, cell_configs, &device).expect("Failed to create PC-DARTS");
// Perform multiple steps
for _ in 0..3 {
pcdarts
.step(0.5, 0.6, None, None)
.expect("Failed to perform step");
}
// Architecture should be derivable
let arch = pcdarts
.derive_architecture()
.expect("Failed to derive architecture");
assert!(arch.validate().is_ok());
}
// =============================================================================
// Hardware-Aware NAS Integration Tests
// =============================================================================
#[test]
fn test_hardware_device_profiles() {
use rtx_nas::hardware::device::{CommonDevices, DeviceType};
// Test common device profiles
let rtx3090 = CommonDevices::rtx_3090();
assert_eq!(rtx3090.name, "RTX 3090");
assert!(matches!(rtx3090.device_type, DeviceType::CUDA));
assert!(rtx3090.peak_tflops_fp32 > 30.0);
let a100 = CommonDevices::a100_40gb();
assert_eq!(a100.name, "A100 40GB");
assert!(a100.memory_gb >= 40.0);
let mobile = CommonDevices::mobile_arm();
assert!(matches!(mobile.device_type, DeviceType::Mobile));
}
#[test]
fn test_latency_prediction() {
use rtx_nas::hardware::{
device::CommonDevices,
latency::{LatencyPredictor, LookupTablePredictor},
};
use rtx_nas::search_space::{Architecture, Cell};
// Create test architecture
let cell = Cell::default_darts().expect("Failed to create cell");
let arch = Architecture::new("test_arch".to_string(), vec![cell], 16, 8);
// Create predictor
let predictor = LookupTablePredictor::new();
let device = CommonDevices::rtx_3090();
// Predict latency
let latency = predictor
.predict(&arch, &device)
.expect("Failed to predict latency");
assert!(latency >= 0.0);
// Confidence should be medium for uncalibrated predictor
assert_eq!(predictor.confidence(), 0.5);
}
#[test]
fn test_cost_model() {
use rtx_nas::hardware::cost_model::compute_cost;
use rtx_nas::search_space::{Architecture, Cell, Edge, OperationType};
// Create test architecture with operations set
let mut cell = Cell::default_darts().expect("Failed to create cell");
// Set operations on some edges
cell.set_operation(Edge::new(0, 2), OperationType::Conv3x3)
.expect("Failed to set operation");
cell.set_operation(Edge::new(1, 2), OperationType::Conv3x3)
.expect("Failed to set operation");
cell.set_operation(Edge::new(0, 3), OperationType::MaxPool3x3)
.expect("Failed to set operation");
let arch = Architecture::new("test_arch".to_string(), vec![cell.clone(), cell], 32, 8);
// Compute cost
let cost = compute_cost(&arch).expect("Failed to compute cost");
// Verify cost components - with operations set, we should have non-zero costs
assert!(cost.flops > 0, "Expected non-zero FLOPs");
assert!(cost.params > 0, "Expected non-zero params");
assert!(cost.memory_mb >= 0.0);
assert!(cost.model_size_mb >= 0.0);
}
#[test]
fn test_cost_model_fields() {
use rtx_nas::hardware::cost_model::ArchitectureCost;
let cost = ArchitectureCost {
flops: 1000,
params: 100,
memory_mb: 10.0,
model_size_mb: 1.0,
estimated_latency_ms: Some(5.0),
cell_costs: vec![],
};
// Test that cost fields are accessible
assert_eq!(cost.flops, 1000);
assert_eq!(cost.params, 100);
assert_eq!(cost.memory_mb, 10.0);
assert_eq!(cost.model_size_mb, 1.0);
assert_eq!(cost.estimated_latency_ms, Some(5.0));
}
// =============================================================================
// Multi-Objective Search Integration Tests
// =============================================================================
#[test]
fn test_multi_objective_scoring() {
use rtx_nas::hardware::cost_model::ArchitectureCost;
use rtx_nas::search::{MultiObjective, ObjectiveScorer};
// Create multi-objective config
let objectives = MultiObjective::mobile_optimized();
let scorer = ObjectiveScorer::new(objectives).expect("Failed to create scorer");
// Create cost
let cost = ArchitectureCost {
flops: 500_000_000,
params: 5_000_000,
memory_mb: 500.0,
model_size_mb: 20.0,
estimated_latency_ms: Some(50.0),
cell_costs: vec![],
};
// Score architecture (scorer uses internal reference values)
let accuracy = 0.9;
let score = scorer.score(accuracy, &cost);
// Score should be positive
assert!(score > 0.0);
}
#[test]
fn test_pareto_frontier() {
use rtx_nas::hardware::cost_model::ArchitectureCost;
use rtx_nas::search::{ParetoEntry, ParetoFrontier};
use rtx_nas::search_space::{Architecture, Cell};
let mut frontier = ParetoFrontier::new();
// Add entries with different accuracy/cost tradeoffs
let cell = Cell::default_darts().expect("Failed to create cell");
for i in 0..5 {
let arch = Architecture::new(format!("arch_{}", i), vec![cell.clone()], 16, 8);
let cost = ArchitectureCost {
flops: 1000 * (5 - i) as u64,
params: 100 * (5 - i) as u64,
memory_mb: 10.0 * (5 - i) as f32,
model_size_mb: 1.0 * (5 - i) as f32,
estimated_latency_ms: Some(5.0 * (5 - i) as f32),
cell_costs: vec![],
};
let accuracy = 0.5 + 0.1 * i as f32;
let entry = ParetoEntry::new(arch, cost, accuracy);
frontier.add(entry);
}
// Frontier should contain non-dominated solutions
assert!(!frontier.is_empty());
assert!(frontier.len() <= 5);
}
#[test]
fn test_pareto_dominance() {
use rtx_nas::hardware::cost_model::ArchitectureCost;
use rtx_nas::search::ParetoEntry;
use rtx_nas::search_space::{Architecture, Cell};
let cell = Cell::default_darts().expect("Failed to create cell");
// Create two entries
let arch1 = Architecture::new("arch_1".to_string(), vec![cell.clone()], 16, 8);
let cost1 = ArchitectureCost {
flops: 1000,
params: 100,
memory_mb: 10.0,
model_size_mb: 1.0,
estimated_latency_ms: Some(5.0),
cell_costs: vec![],
};
let entry1 = ParetoEntry::new(arch1, cost1, 0.9);
let arch2 = Architecture::new("arch_2".to_string(), vec![cell], 16, 8);
let cost2 = ArchitectureCost {
flops: 2000, // Worse
params: 200, // Worse
memory_mb: 20.0, // Worse
model_size_mb: 2.0,
estimated_latency_ms: Some(10.0),
cell_costs: vec![],
};
let entry2 = ParetoEntry::new(arch2, cost2, 0.8); // Lower accuracy too
// entry1 dominates entry2
assert!(entry1.dominates(&entry2));
assert!(!entry2.dominates(&entry1));
}
// =============================================================================
// FairNAS Integration Tests
// =============================================================================
#[test]
fn test_fairness_tracking() {
use rtx_nas::algorithms::fairness::{FairnessConfig, FairnessTracker};
use rtx_nas::search_space::{Edge, OperationType};
// Create fairness tracker
let config = FairnessConfig::default();
let mut tracker = FairnessTracker::new(config);
// Simulate fair training
let edge = Edge::new(0, 2);
for _ in 0..100 {
tracker.track_optimization(edge, OperationType::Conv3x3);
tracker.track_optimization(edge, OperationType::Conv5x5);
tracker.track_optimization(edge, OperationType::MaxPool3x3);
}
// Check fairness
let score = tracker.compute_fairness_score();
assert!(score > 0.9, "Expected high fairness for balanced training");
assert!(tracker.is_fair());
}
#[test]
fn test_fairness_reweighting() {
use rtx_nas::algorithms::fairness::{FairnessConfig, FairnessTracker};
use rtx_nas::search_space::{Edge, OperationType};
// Create tracker with auto-rebalancing
let mut config = FairnessConfig::default();
config.auto_rebalance = true;
config.rebalance_interval = 10;
let mut tracker = FairnessTracker::new(config);
// Create imbalanced training
let edge = Edge::new(0, 2);
for _ in 0..100 {
tracker.track_optimization(edge, OperationType::Conv3x3);
}
for _ in 0..10 {
tracker.track_optimization(edge, OperationType::Conv5x5);
}
// Get weights - underrepresented ops should have higher weight
let weight_conv3 = tracker.get_sampling_weight(OperationType::Conv3x3);
let weight_conv5 = tracker.get_sampling_weight(OperationType::Conv5x5);
assert!(
weight_conv5 > weight_conv3,
"Conv5x5 should have higher weight"
);
}
#[test]
fn test_fairness_report() {
use rtx_nas::algorithms::fairness::{FairnessConfig, FairnessTracker};
use rtx_nas::search_space::{Edge, OperationType};
let config = FairnessConfig::default();
let mut tracker = FairnessTracker::new(config);
// Add some tracking
let edge = Edge::new(0, 2);
for _ in 0..50 {
tracker.track_optimization(edge, OperationType::Conv3x3);
tracker.track_optimization(edge, OperationType::Conv5x5);
}
// Generate report
let report = tracker.generate_report();
assert_eq!(report.total_iterations, 100);
assert!(report.overall_score > 0.9);
assert!(report.is_fair(0.8));
}
// =============================================================================
// End-to-End Integration Tests
// =============================================================================
#[test]
fn test_hardware_aware_search() {
use rtx_nas::algorithms::{RandomSearch, RandomSearchConfig};
use rtx_nas::hardware::{
cost_model::compute_cost,
device::CommonDevices,
latency::{LatencyPredictor, LookupTablePredictor},
};
use rtx_nas::search::{MultiObjective, ObjectiveScorer, ParetoEntry, ParetoFrontier};
// Create search space
let search_space = DARTSSearchSpace::default().expect("Failed to create search space");
// Configure random search
let config = RandomSearchConfig::new(10);
let mut search = RandomSearch::new(config).expect("Failed to create random search");
// Sample architectures
search.sample(&search_space).expect("Failed to sample");
// Set up hardware evaluation
let device = CommonDevices::rtx_3090();
let predictor = LookupTablePredictor::new();
// Set up multi-objective scoring
let objectives = MultiObjective::default();
let _scorer = ObjectiveScorer::new(objectives).expect("Failed to create scorer");
// Build Pareto frontier
let mut frontier = ParetoFrontier::new();
for arch in search.samples() {
// Compute cost
let mut cost = compute_cost(arch).expect("Failed to compute cost");
// Predict latency
let latency = predictor
.predict(arch, &device)
.expect("Failed to predict latency");
cost.estimated_latency_ms = Some(latency);
// Simulate accuracy (in real use, this would come from validation)
let accuracy = 0.5 + rand::random::<f32>() * 0.4;
// Add to frontier
let entry = ParetoEntry::new(arch.clone(), cost, accuracy);
frontier.add(entry);
}
// Frontier should have some entries
assert!(!frontier.is_empty());
// Get best by accuracy (objective index 0)
let best = frontier.best_for_objective(0);
assert!(best.is_some());
}