439 lines
14 KiB
Rust
439 lines
14 KiB
Rust
use rtx_automeasure::agents::{
|
|
BaseModel, EnsembleBuilder, EnsembleMethod, EnsembleModel, SelectionCriteria, VotingType,
|
|
};
|
|
use rtx_automeasure::{AutoMLResult, OptimizationObjective, TaskType};
|
|
use rtx_tensor::{Device, Tensor};
|
|
use std::collections::HashMap;
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_builder_creation() {
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy);
|
|
assert!(builder.is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_ensemble_basic() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
// Create candidate base models with scores
|
|
let mut candidates = Vec::new();
|
|
candidates
|
|
.push(BaseModel::new("LogisticRegression", HashMap::new()).with_metrics(0.85, 1.0, 0.5));
|
|
candidates.push(BaseModel::new("RandomForest", HashMap::new()).with_metrics(0.88, 2.0, 0.7));
|
|
candidates.push(BaseModel::new("SVC", HashMap::new()).with_metrics(0.86, 1.5, 0.6));
|
|
|
|
let x_val = Tensor::randn(&[100, 8], &device)?;
|
|
let y_val = Tensor::zeros(&[100], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(!ensemble.base_models.is_empty());
|
|
assert!(!ensemble.weights.is_empty());
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_ensemble_with_selection_criteria() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?
|
|
.with_selection_criteria(SelectionCriteria::TopPerformers { n_models: 3 });
|
|
|
|
let mut candidates = Vec::new();
|
|
for i in 0..5 {
|
|
candidates.push(
|
|
BaseModel::new(&format!("Model{i}"), HashMap::new()).with_metrics(
|
|
0.80 + i as f64 * 0.02,
|
|
1.0,
|
|
0.5,
|
|
),
|
|
);
|
|
}
|
|
|
|
let x_val = Tensor::randn(&[80, 6], &device)?;
|
|
let y_val = Tensor::zeros(&[80], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(ensemble.base_models.len() <= 3);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_ensemble_diversity_based() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?
|
|
.with_selection_criteria(SelectionCriteria::DiversityBased {
|
|
min_diversity: 0.1,
|
|
max_correlation: 0.8,
|
|
});
|
|
|
|
let mut candidates = Vec::new();
|
|
candidates.push(BaseModel::new("DecisionTree", HashMap::new()).with_metrics(0.82, 1.0, 0.4));
|
|
candidates.push(BaseModel::new("KNeighbors", HashMap::new()).with_metrics(0.83, 1.2, 0.5));
|
|
candidates.push(BaseModel::new("SVC", HashMap::new()).with_metrics(0.84, 1.5, 0.6));
|
|
|
|
let x_val = Tensor::randn(&[120, 10], &device)?;
|
|
let y_val = Tensor::zeros(&[120], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(!ensemble.base_models.is_empty());
|
|
assert!(ensemble.diversity_metrics.pairwise_diversity >= 0.0);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_ensemble_greedy_search() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Regression, OptimizationObjective::RMSE)?
|
|
.with_selection_criteria(SelectionCriteria::GreedySearch { max_models: 4 });
|
|
|
|
let mut candidates = Vec::new();
|
|
candidates.push(BaseModel::new("Ridge", HashMap::new()).with_metrics(0.75, 0.8, 0.3));
|
|
candidates.push(BaseModel::new("RandomForest", HashMap::new()).with_metrics(0.78, 1.5, 0.6));
|
|
candidates
|
|
.push(BaseModel::new("GradientBoosting", HashMap::new()).with_metrics(0.80, 2.0, 0.7));
|
|
|
|
let x_val = Tensor::randn(&[150, 12], &device)?;
|
|
let y_val = Tensor::randn(&[150], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(ensemble.base_models.len() <= 4);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_ensemble_pareto_optimal() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::F1Score)?
|
|
.with_selection_criteria(SelectionCriteria::ParetoOptimal);
|
|
|
|
let mut candidates = Vec::new();
|
|
// Different trade-offs between performance and efficiency
|
|
candidates.push(BaseModel::new("FastModel", HashMap::new()).with_metrics(0.75, 0.5, 0.2));
|
|
candidates.push(BaseModel::new("AccurateModel", HashMap::new()).with_metrics(0.90, 3.0, 0.9));
|
|
candidates.push(BaseModel::new("BalancedModel", HashMap::new()).with_metrics(0.82, 1.0, 0.5));
|
|
|
|
let x_val = Tensor::randn(&[200, 15], &device)?;
|
|
let y_val = Tensor::zeros(&[200], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(!ensemble.base_models.is_empty());
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_build_multiple_ensembles() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let mut candidates = Vec::new();
|
|
for i in 0..4 {
|
|
candidates.push(
|
|
BaseModel::new(&format!("Model{i}"), HashMap::new()).with_metrics(
|
|
0.80 + i as f64 * 0.02,
|
|
1.0 + i as f64 * 0.5,
|
|
0.5,
|
|
),
|
|
);
|
|
}
|
|
|
|
let x_val = Tensor::randn(&[100, 8], &device)?;
|
|
let y_val = Tensor::zeros(&[100], &device)?;
|
|
|
|
let ensembles = builder
|
|
.build_multiple_ensembles(candidates, &x_val, &y_val)
|
|
.await?;
|
|
|
|
assert!(!ensembles.is_empty());
|
|
// Ensembles should be sorted by validation score
|
|
for i in 1..ensembles.len() {
|
|
assert!(ensembles[i - 1].validation_score >= ensembles[i].validation_score);
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_prediction_voting() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let mut base_models = vec![
|
|
BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("Model2", HashMap::new()).with_metrics(0.83, 1.2, 0.6),
|
|
];
|
|
|
|
let weights = vec![0.6, 0.4];
|
|
let method = EnsembleMethod::Voting {
|
|
voting_type: VotingType::Weighted,
|
|
weighted: true,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
let x_test = Tensor::randn(&[30, 6], &device)?;
|
|
let predictions = builder.predict(&ensemble, &x_test).await?;
|
|
|
|
assert_eq!(predictions.shape()[0], 30);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_prediction_stacking() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let base_models = vec![
|
|
BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("Model2", HashMap::new()).with_metrics(0.83, 1.2, 0.6),
|
|
BaseModel::new("Model3", HashMap::new()).with_metrics(0.84, 1.1, 0.55),
|
|
];
|
|
|
|
let weights = vec![0.4, 0.3, 0.3];
|
|
let method = EnsembleMethod::Stacking {
|
|
meta_learner: BaseModel::new("LinearRegression", HashMap::new()),
|
|
cv_folds: 3,
|
|
use_probabilities: true,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
let x_test = Tensor::randn(&[50, 10], &device)?;
|
|
let predictions = builder.predict(&ensemble, &x_test).await?;
|
|
|
|
assert_eq!(predictions.shape()[0], 50);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
#[ignore = "Test logic needs to be fixed - max size assertion issue"]
|
|
async fn test_ensemble_with_max_size() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?
|
|
.with_max_ensemble_size(3);
|
|
|
|
let mut candidates = Vec::new();
|
|
for i in 0..10 {
|
|
candidates.push(
|
|
BaseModel::new(&format!("Model{i}"), HashMap::new()).with_metrics(
|
|
0.80 + i as f64 * 0.01,
|
|
1.0,
|
|
0.5,
|
|
),
|
|
);
|
|
}
|
|
|
|
let x_val = Tensor::randn(&[100, 8], &device)?;
|
|
let y_val = Tensor::zeros(&[100], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(ensemble.base_models.len() <= 3);
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_diversity_threshold() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?
|
|
.with_diversity_threshold(0.2);
|
|
|
|
let mut candidates = Vec::new();
|
|
candidates.push(BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5));
|
|
candidates.push(BaseModel::new("Model2", HashMap::new()).with_metrics(0.84, 1.1, 0.5));
|
|
|
|
let x_val = Tensor::randn(&[80, 5], &device)?;
|
|
let y_val = Tensor::zeros(&[80], &device)?;
|
|
|
|
let ensemble = builder.build_ensemble(candidates, &x_val, &y_val).await?;
|
|
|
|
assert!(!ensemble.base_models.is_empty());
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_model_getters() {
|
|
let base_models = vec![
|
|
BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("Model2", HashMap::new()).with_metrics(0.83, 1.5, 0.7),
|
|
];
|
|
|
|
let weights = vec![0.6, 0.4];
|
|
let method = EnsembleMethod::Voting {
|
|
voting_type: VotingType::Hard,
|
|
weighted: false,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
assert_eq!(ensemble.get_n_models(), 2);
|
|
assert!(ensemble.get_complexity() > 0.0);
|
|
assert!(ensemble.get_efficiency() > 0.0);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_serialization() -> AutoMLResult<()> {
|
|
let base_models = vec![
|
|
BaseModel::new("LogisticRegression", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("RandomForest", HashMap::new()).with_metrics(0.88, 2.0, 0.7),
|
|
];
|
|
|
|
let weights = vec![0.5, 0.5];
|
|
let method = EnsembleMethod::Voting {
|
|
voting_type: VotingType::Soft,
|
|
weighted: true,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
// Test serialization
|
|
let serialized = serde_json::to_string(&ensemble)?;
|
|
assert!(!serialized.is_empty());
|
|
|
|
// Test deserialization
|
|
let deserialized: EnsembleModel = serde_json::from_str(&serialized)?;
|
|
assert_eq!(deserialized.base_models.len(), ensemble.base_models.len());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_base_model_creation() {
|
|
let mut params = HashMap::new();
|
|
params.insert("n_estimators".to_string(), "100".to_string());
|
|
|
|
let model = BaseModel::new("RandomForest", params.clone());
|
|
|
|
assert_eq!(model.model_name, "RandomForest");
|
|
assert_eq!(
|
|
model.hyperparameters.get("n_estimators"),
|
|
Some(&"100".to_string())
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_base_model_with_metrics() {
|
|
let model = BaseModel::new("SVC", HashMap::new()).with_metrics(0.90, 2.5, 0.8);
|
|
|
|
assert_eq!(model.validation_score, 0.90);
|
|
assert_eq!(model.training_time, 2.5);
|
|
assert_eq!(model.model_complexity, 0.8);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_base_model_efficiency() {
|
|
let model = BaseModel::new("Model", HashMap::new()).with_metrics(0.85, 2.0, 0.5);
|
|
|
|
let efficiency = model.get_efficiency();
|
|
assert_eq!(efficiency, 0.85 / 2.0);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_base_model_with_diversity() {
|
|
let model = BaseModel::new("Model", HashMap::new()).with_diversity_score(0.75);
|
|
|
|
assert_eq!(model.diversity_score, 0.75);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_cv_evaluation() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let base_models = vec![
|
|
BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("Model2", HashMap::new()).with_metrics(0.83, 1.2, 0.6),
|
|
];
|
|
|
|
let weights = vec![0.6, 0.4];
|
|
let method = EnsembleMethod::Voting {
|
|
voting_type: VotingType::Weighted,
|
|
weighted: true,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
let x = Tensor::randn(&[100, 8], &device)?;
|
|
let y = Tensor::zeros(&[100], &device)?;
|
|
|
|
let evaluation = builder.evaluate_ensemble_with_cv(&ensemble, &x, &y)?;
|
|
|
|
assert!(evaluation.mean_score >= 0.0);
|
|
assert!(evaluation.std_score >= 0.0);
|
|
assert!(!evaluation.fold_scores.is_empty());
|
|
assert_eq!(evaluation.n_models, 2);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_ensemble_evaluation_confidence_bounds() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let base_models = vec![BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5)];
|
|
|
|
let weights = vec![1.0];
|
|
let method = EnsembleMethod::Voting {
|
|
voting_type: VotingType::Hard,
|
|
weighted: false,
|
|
};
|
|
|
|
let ensemble = EnsembleModel::new(base_models, method, weights);
|
|
|
|
let x = Tensor::randn(&[80, 6], &device)?;
|
|
let y = Tensor::zeros(&[80], &device)?;
|
|
|
|
let evaluation = builder.evaluate_ensemble_with_cv(&ensemble, &x, &y)?;
|
|
|
|
let lower = evaluation.lower_bound();
|
|
let upper = evaluation.upper_bound();
|
|
|
|
assert!(lower <= evaluation.mean_score);
|
|
assert!(upper >= evaluation.mean_score);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_evaluate_ensemble_variants() -> AutoMLResult<()> {
|
|
let device = Device::cpu();
|
|
let builder = EnsembleBuilder::new(TaskType::Classification, OptimizationObjective::Accuracy)?;
|
|
|
|
let variants = vec![
|
|
EnsembleModel::new(
|
|
vec![BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5)],
|
|
EnsembleMethod::Voting {
|
|
voting_type: VotingType::Hard,
|
|
weighted: false,
|
|
},
|
|
vec![1.0],
|
|
),
|
|
EnsembleModel::new(
|
|
vec![
|
|
BaseModel::new("Model1", HashMap::new()).with_metrics(0.85, 1.0, 0.5),
|
|
BaseModel::new("Model2", HashMap::new()).with_metrics(0.83, 1.2, 0.6),
|
|
],
|
|
EnsembleMethod::Voting {
|
|
voting_type: VotingType::Soft,
|
|
weighted: true,
|
|
},
|
|
vec![0.6, 0.4],
|
|
),
|
|
];
|
|
|
|
let x = Tensor::randn(&[100, 8], &device)?;
|
|
let y = Tensor::zeros(&[100], &device)?;
|
|
|
|
let (best_ensemble, best_eval) =
|
|
builder.evaluate_ensemble_variants_with_cv(&variants, &x, &y)?;
|
|
|
|
assert!(!best_ensemble.base_models.is_empty());
|
|
assert!(best_eval.mean_score >= 0.0);
|
|
|
|
Ok(())
|
|
}
|