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

177 lines
5.3 KiB
Rust

use rtx_auto::{
error::AutoError,
proposal::{Proposal, ProposalStatus, ProposalType, ProposalValidator},
};
use rtx_runtime::Runtime;
use serde_json;
use std::sync::Arc;
use tokio_test;
#[tokio::test]
async fn test_proposal_creation() {
let proposal = Proposal::new(
ProposalType::DataLayout,
"Optimize tensor layout for better cache locality".to_string(),
2.5, // Expected 2.5x performance gain
);
assert_eq!(proposal.proposal_type(), ProposalType::DataLayout);
assert_eq!(proposal.status(), ProposalStatus::Pending);
assert_eq!(proposal.expected_performance_gain(), 2.5);
assert!(proposal.description().contains("cache locality"));
assert!(proposal.id().len() > 0);
assert!(proposal.created_at() > 0);
}
#[tokio::test]
async fn test_proposal_serialization() {
let proposal = Proposal::new(
ProposalType::KernelFusion,
"Fuse relu and bias_add kernels".to_string(),
1.8,
);
let serialized = serde_json::to_string(&proposal);
assert!(serialized.is_ok());
let deserialized: Result<Proposal, _> = serde_json::from_str(&serialized.unwrap());
assert!(deserialized.is_ok());
let deserialized_proposal = deserialized.unwrap();
assert_eq!(deserialized_proposal.id(), proposal.id());
assert_eq!(
deserialized_proposal.proposal_type(),
proposal.proposal_type()
);
assert_eq!(
deserialized_proposal.expected_performance_gain(),
proposal.expected_performance_gain()
);
}
#[tokio::test]
async fn test_proposal_status_transitions() {
let mut proposal = Proposal::new(
ProposalType::Quantization,
"Apply INT8 quantization".to_string(),
3.2,
);
assert_eq!(proposal.status(), ProposalStatus::Pending);
proposal.set_status(ProposalStatus::Evaluating);
assert_eq!(proposal.status(), ProposalStatus::Evaluating);
proposal.set_status(ProposalStatus::Approved);
assert_eq!(proposal.status(), ProposalStatus::Approved);
proposal.set_status(ProposalStatus::Applied);
assert_eq!(proposal.status(), ProposalStatus::Applied);
}
#[tokio::test]
async fn test_proposal_validator_creation() {
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
let validator = ProposalValidator::new(runtime.clone());
assert!(validator.is_ok(), "Failed to create ProposalValidator");
}
#[tokio::test]
async fn test_validate_performance_gain() {
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
let validator = ProposalValidator::new(runtime.clone()).unwrap();
let valid_proposal = Proposal::new(
ProposalType::DataParallel,
"Distribute computation across 4 GPUs".to_string(),
3.8,
);
let is_valid = validator.validate_performance_gain(&valid_proposal).await;
assert!(is_valid.is_ok());
assert!(
is_valid.unwrap(),
"Reasonable performance gain should be valid"
);
let invalid_proposal = Proposal::new(
ProposalType::CacheOptimization,
"Impossible optimization".to_string(),
1000.0, // Unrealistic gain
);
let is_valid = validator.validate_performance_gain(&invalid_proposal).await;
assert!(is_valid.is_ok());
assert!(
!is_valid.unwrap(),
"Unrealistic performance gain should be invalid"
);
}
#[tokio::test]
async fn test_validate_proposal_feasibility() {
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
let validator = ProposalValidator::new(runtime.clone()).unwrap();
let proposal = Proposal::new(
ProposalType::PipelineParallel,
"Split model into 8 pipeline stages".to_string(),
4.2,
);
let is_feasible = validator.validate_feasibility(&proposal).await;
assert!(is_feasible.is_ok());
}
#[tokio::test]
async fn test_score_proposal() {
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
let validator = ProposalValidator::new(runtime.clone()).unwrap();
let proposal = Proposal::new(
ProposalType::MemoryCoalescing,
"Coalesce memory accesses for better bandwidth".to_string(),
2.1,
);
let score = validator.score_proposal(&proposal).await;
assert!(score.is_ok());
let score_value = score.unwrap();
assert!(
score_value >= 0.0 && score_value <= 1.0,
"Score should be between 0 and 1"
);
}
#[tokio::test]
async fn test_rank_proposals() {
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
let validator = ProposalValidator::new(runtime.clone()).unwrap();
let proposals = vec![
Proposal::new(
ProposalType::DataLayout,
"Layout optimization".to_string(),
1.5,
),
Proposal::new(ProposalType::KernelFusion, "Kernel fusion".to_string(), 3.2),
Proposal::new(
ProposalType::Quantization,
"INT8 quantization".to_string(),
2.8,
),
];
let ranked = validator.rank_proposals(&proposals).await;
assert!(ranked.is_ok());
let ranked_proposals = ranked.unwrap();
assert_eq!(ranked_proposals.len(), proposals.len());
// Check that proposals are sorted by score in descending order
for i in 0..ranked_proposals.len() - 1 {
assert!(ranked_proposals[i].1 >= ranked_proposals[i + 1].1);
}
}