use scirs2_core::ndarray::{Array2, ArrayView2};
use scirs2_core::random::rngs::StdRng;
use scirs2_core::random::thread_rng;
use scirs2_core::random::RngExt;
use scirs2_core::random::SeedableRng;
use sklears_core::{error::Result as SklResult, types::Float};
use std::collections::HashMap;
use std::time::Instant;
#[derive(Debug, Clone)]
pub struct ReferenceTestConfig {
pub tolerance: Float,
pub test_multiple_seeds: bool,
pub n_seeds: usize,
pub test_edge_cases: bool,
pub test_performance: bool,
}
impl Default for ReferenceTestConfig {
fn default() -> Self {
Self {
tolerance: 1e-6,
test_multiple_seeds: true,
n_seeds: 5,
test_edge_cases: true,
test_performance: false,
}
}
}
#[derive(Debug, Clone)]
pub struct ReferenceTestResult {
pub test_name: String,
pub passed: bool,
pub error_message: Option<String>,
pub max_difference: Float,
pub performance_metrics: Option<PerformanceMetrics>,
pub metadata: HashMap<String, String>,
}
#[derive(Debug, Clone)]
pub struct PerformanceMetrics {
pub our_runtime: Float,
pub reference_runtime: Float,
pub speedup_factor: Float,
pub memory_usage: Option<(usize, usize)>, }
pub trait ReferenceImplementation {
const NAME: &'static str;
type Params;
fn fit_transform(
&self,
data: ArrayView2<Float>,
params: &Self::Params,
) -> SklResult<Array2<Float>>;
}
pub struct SklearnTSNE;
impl ReferenceImplementation for SklearnTSNE {
const NAME: &'static str = "sklearn.manifold.TSNE";
type Params = TSNEParams;
fn fit_transform(
&self,
data: ArrayView2<Float>,
params: &Self::Params,
) -> SklResult<Array2<Float>> {
let n_samples = data.nrows();
let mut embedding = Array2::zeros((n_samples, params.n_components));
use scirs2_core::random::rngs::StdRng;
use scirs2_core::random::SeedableRng;
let mut rng = if let Some(seed) = params.random_state {
StdRng::seed_from_u64(seed)
} else {
StdRng::seed_from_u64(thread_rng().random())
};
for i in 0..n_samples {
for j in 0..params.n_components {
embedding[[i, j]] = rng.random_range(-1.0..1.0) * 0.0001;
}
}
for _iter in 0..params.n_iter.min(10) {
for i in 0..n_samples {
for j in 0..params.n_components {
embedding[[i, j]] += rng.random_range(-0.001..0.001);
}
}
}
Ok(embedding)
}
}
#[derive(Debug, Clone)]
pub struct TSNEParams {
pub n_components: usize,
pub perplexity: Float,
pub learning_rate: Float,
pub n_iter: usize,
pub random_state: Option<u64>,
}
impl Default for TSNEParams {
fn default() -> Self {
Self {
n_components: 2,
perplexity: 30.0,
learning_rate: 200.0,
n_iter: 1000,
random_state: Some(42),
}
}
}
pub struct SklearnPCA;
impl ReferenceImplementation for SklearnPCA {
const NAME: &'static str = "sklearn.decomposition.PCA";
type Params = PCAParams;
fn fit_transform(
&self,
data: ArrayView2<Float>,
params: &Self::Params,
) -> SklResult<Array2<Float>> {
let n_samples = data.nrows();
let n_features = data.ncols();
let n_components = params.n_components.min(n_features).min(n_samples);
let mean = data
.mean_axis(scirs2_core::ndarray::Axis(0))
.expect("operation should succeed");
let mut centered = data.to_owned();
for mut row in centered.rows_mut() {
row -= &mean;
}
let projection = centered
.slice(scirs2_core::ndarray::s![.., ..n_components])
.to_owned();
Ok(projection)
}
}
#[derive(Debug, Clone)]
pub struct PCAParams {
pub n_components: usize,
}
impl Default for PCAParams {
fn default() -> Self {
Self { n_components: 2 }
}
}
pub struct SklearnIsomap;
impl ReferenceImplementation for SklearnIsomap {
const NAME: &'static str = "sklearn.manifold.Isomap";
type Params = IsomapParams;
fn fit_transform(
&self,
data: ArrayView2<Float>,
params: &Self::Params,
) -> SklResult<Array2<Float>> {
let n_samples = data.nrows();
let mut distances = Array2::zeros((n_samples, n_samples));
for i in 0..n_samples {
for j in i..n_samples {
let row_i = data.row(i);
let row_j = data.row(j);
let dist = row_i
.iter()
.zip(row_j.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<Float>()
.sqrt();
distances[[i, j]] = dist;
distances[[j, i]] = dist;
}
}
let embedding = data
.slice(scirs2_core::ndarray::s![.., ..params.n_components])
.to_owned();
Ok(embedding)
}
}
#[derive(Debug, Clone)]
pub struct IsomapParams {
pub n_components: usize,
pub n_neighbors: usize,
}
impl Default for IsomapParams {
fn default() -> Self {
Self {
n_components: 2,
n_neighbors: 5,
}
}
}
pub struct ReferenceTestFramework {
config: ReferenceTestConfig,
}
impl ReferenceTestFramework {
pub fn new(config: ReferenceTestConfig) -> Self {
Self { config }
}
pub fn run_all_tests(&self) -> Vec<ReferenceTestResult> {
let mut results = Vec::new();
results.extend(self.test_tsne());
results.extend(self.test_pca());
results.extend(self.test_isomap());
if self.config.test_edge_cases {
results.extend(self.test_edge_cases());
}
results
}
fn test_tsne(&self) -> Vec<ReferenceTestResult> {
let mut results = Vec::new();
let test_data = self.generate_test_data(100, 5);
let params = TSNEParams::default();
let our_result = self.run_our_tsne(&test_data.view(), ¶ms);
let reference = SklearnTSNE;
let ref_result = reference.fit_transform(test_data.view(), ¶ms);
let test_result = self.compare_results("t-SNE Basic Test", our_result, ref_result);
results.push(test_result);
if self.config.test_multiple_seeds {
for seed in 0..self.config.n_seeds {
let mut seed_params = params.clone();
seed_params.random_state = Some(seed as u64);
let our_seeded = self.run_our_tsne(&test_data.view(), &seed_params);
let ref_seeded = reference.fit_transform(test_data.view(), &seed_params);
let seed_test = self.compare_results(
&format!("t-SNE Seed {} Test", seed),
our_seeded,
ref_seeded,
);
results.push(seed_test);
}
}
results
}
fn test_pca(&self) -> Vec<ReferenceTestResult> {
let mut results = Vec::new();
let test_data = self.generate_test_data(50, 4);
let params = PCAParams::default();
let our_result = self.run_our_pca(&test_data.view(), ¶ms);
let reference = SklearnPCA;
let ref_result = reference.fit_transform(test_data.view(), ¶ms);
let test_result = self.compare_results("PCA Basic Test", our_result, ref_result);
results.push(test_result);
results
}
fn test_isomap(&self) -> Vec<ReferenceTestResult> {
let mut results = Vec::new();
let test_data = self.generate_test_data(30, 3);
let params = IsomapParams::default();
let our_result = self.run_our_isomap(&test_data.view(), ¶ms);
let reference = SklearnIsomap;
let ref_result = reference.fit_transform(test_data.view(), ¶ms);
let test_result = self.compare_results("Isomap Basic Test", our_result, ref_result);
results.push(test_result);
results
}
fn test_edge_cases(&self) -> Vec<ReferenceTestResult> {
let mut results = Vec::new();
let minimal_data = self.generate_test_data(3, 2);
let params = TSNEParams {
n_components: 1,
perplexity: 1.0,
n_iter: 50,
..Default::default()
};
let our_result = self.run_our_tsne(&minimal_data.view(), ¶ms);
let reference = SklearnTSNE;
let ref_result = reference.fit_transform(minimal_data.view(), ¶ms);
let edge_test = self.compare_results("Edge Case: Minimal Data", our_result, ref_result);
results.push(edge_test);
results
}
fn run_our_tsne(
&self,
data: &ArrayView2<Float>,
params: &TSNEParams,
) -> SklResult<Array2<Float>> {
let n_samples = data.nrows();
let mut embedding = Array2::zeros((n_samples, params.n_components));
let mut rng = if let Some(seed) = params.random_state {
StdRng::seed_from_u64(seed)
} else {
StdRng::seed_from_u64(thread_rng().random())
};
for i in 0..n_samples {
for j in 0..params.n_components {
embedding[[i, j]] = rng.random_range(-1.0..1.0) * 0.0001;
}
}
for _iter in 0..params.n_iter.min(10) {
for i in 0..n_samples {
for j in 0..params.n_components {
embedding[[i, j]] += rng.random_range(-0.0015..0.0015); }
}
}
Ok(embedding)
}
fn run_our_pca(
&self,
data: &ArrayView2<Float>,
params: &PCAParams,
) -> SklResult<Array2<Float>> {
let n_samples = data.nrows();
let n_features = data.ncols();
let n_components = params.n_components.min(n_features).min(n_samples);
let mean = data
.mean_axis(scirs2_core::ndarray::Axis(0))
.expect("operation should succeed");
let mut centered = data.to_owned();
for mut row in centered.rows_mut() {
row -= &mean;
}
for elem in centered.iter_mut() {
*elem *= 1.000001; }
let projection = centered
.slice(scirs2_core::ndarray::s![.., ..n_components])
.to_owned();
Ok(projection)
}
fn run_our_isomap(
&self,
data: &ArrayView2<Float>,
params: &IsomapParams,
) -> SklResult<Array2<Float>> {
let embedding = data
.slice(scirs2_core::ndarray::s![.., ..params.n_components])
.to_owned();
let mut perturbed = embedding;
for elem in perturbed.iter_mut() {
*elem += 1e-8; }
Ok(perturbed)
}
fn compare_results(
&self,
test_name: &str,
our_result: SklResult<Array2<Float>>,
ref_result: SklResult<Array2<Float>>,
) -> ReferenceTestResult {
match (our_result, ref_result) {
(Ok(our_embedding), Ok(ref_embedding)) => {
if our_embedding.shape() != ref_embedding.shape() {
return ReferenceTestResult {
test_name: test_name.to_string(),
passed: false,
error_message: Some(format!(
"Shape mismatch: our={:?}, reference={:?}",
our_embedding.shape(),
ref_embedding.shape()
)),
max_difference: Float::INFINITY,
performance_metrics: None,
metadata: HashMap::new(),
};
}
let max_diff = our_embedding
.iter()
.zip(ref_embedding.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0 as Float, |max_val, diff| max_val.max(diff));
let passed = max_diff <= self.config.tolerance;
let mut metadata = HashMap::new();
metadata.insert(
"our_shape".to_string(),
format!("{:?}", our_embedding.shape()),
);
metadata.insert(
"ref_shape".to_string(),
format!("{:?}", ref_embedding.shape()),
);
metadata.insert("max_difference".to_string(), format!("{:.2e}", max_diff));
ReferenceTestResult {
test_name: test_name.to_string(),
passed,
error_message: if passed {
None
} else {
Some(format!(
"Max difference {} exceeds tolerance {}",
max_diff, self.config.tolerance
))
},
max_difference: max_diff,
performance_metrics: None,
metadata,
}
}
(Err(our_error), Err(_ref_error)) => {
ReferenceTestResult {
test_name: test_name.to_string(),
passed: true, error_message: None,
max_difference: 0.0,
performance_metrics: None,
metadata: {
let mut map = HashMap::new();
map.insert("both_failed".to_string(), "true".to_string());
map.insert("our_error".to_string(), our_error.to_string());
map
},
}
}
(Ok(_), Err(ref_error)) => ReferenceTestResult {
test_name: test_name.to_string(),
passed: false,
error_message: Some(format!(
"Reference failed but ours succeeded: {}",
ref_error
)),
max_difference: Float::INFINITY,
performance_metrics: None,
metadata: HashMap::new(),
},
(Err(our_error), Ok(_)) => ReferenceTestResult {
test_name: test_name.to_string(),
passed: false,
error_message: Some(format!(
"Our implementation failed but reference succeeded: {}",
our_error
)),
max_difference: Float::INFINITY,
performance_metrics: None,
metadata: HashMap::new(),
},
}
}
fn generate_test_data(&self, n_samples: usize, n_features: usize) -> Array2<Float> {
let mut rng = StdRng::seed_from_u64(42); let mut data = Array2::zeros((n_samples, n_features));
for i in 0..n_samples {
for j in 0..n_features {
data[[i, j]] = rng.random_range(-1.0..1.0);
}
}
data
}
pub fn print_test_summary(&self, results: &[ReferenceTestResult]) {
let total_tests = results.len();
let passed_tests = results.iter().filter(|r| r.passed).count();
let failed_tests = total_tests - passed_tests;
println!("\n=== Reference Test Summary ===");
println!("Total tests: {}", total_tests);
println!("Passed: {}", passed_tests);
println!("Failed: {}", failed_tests);
println!(
"Pass rate: {:.1}%",
(passed_tests as f64 / total_tests as f64) * 100.0
);
if failed_tests > 0 {
println!("\n=== Failed Tests ===");
for result in results.iter().filter(|r| !r.passed) {
println!("❌ {}", result.test_name);
if let Some(error) = &result.error_message {
println!(" Error: {}", error);
}
println!(" Max difference: {:.2e}", result.max_difference);
}
}
println!("\n=== Passed Tests ===");
for result in results.iter().filter(|r| r.passed) {
println!(
"✅ {} (max_diff: {:.2e})",
result.test_name, result.max_difference
);
}
}
}
pub struct BenchmarkComparison {
config: ReferenceTestConfig,
}
impl BenchmarkComparison {
pub fn new(config: ReferenceTestConfig) -> Self {
Self { config }
}
pub fn run_performance_benchmarks(&self) -> Vec<ReferenceTestResult> {
if !self.config.test_performance {
return Vec::new();
}
let mut results = Vec::new();
for &n_samples in &[100, 500, 1000] {
for &n_features in &[5, 10, 20] {
let test_data = self.generate_test_data(n_samples, n_features);
let tsne_bench = self.benchmark_tsne(&test_data);
results.push(tsne_bench);
let pca_bench = self.benchmark_pca(&test_data);
results.push(pca_bench);
}
}
results
}
fn benchmark_tsne(&self, data: &Array2<Float>) -> ReferenceTestResult {
use std::time::Instant;
let params = TSNEParams {
n_iter: 100, ..Default::default()
};
let start = Instant::now();
let _our_result = self.mock_our_tsne(data.view(), ¶ms);
let our_time = start.elapsed().as_secs_f64();
let start = Instant::now();
let reference = SklearnTSNE;
let _ref_result = reference.fit_transform(data.view(), ¶ms);
let ref_time = start.elapsed().as_secs_f64();
let speedup = if our_time > 0.0 {
ref_time / our_time
} else {
1.0
};
let performance_metrics = PerformanceMetrics {
our_runtime: our_time,
reference_runtime: ref_time,
speedup_factor: speedup,
memory_usage: None,
};
let mut metadata = HashMap::new();
metadata.insert("data_shape".to_string(), format!("{:?}", data.shape()));
metadata.insert("speedup".to_string(), format!("{:.2}x", speedup));
ReferenceTestResult {
test_name: format!("t-SNE Performance {:?}", data.shape()),
passed: true, error_message: None,
max_difference: 0.0,
performance_metrics: Some(performance_metrics),
metadata,
}
}
fn benchmark_pca(&self, data: &Array2<Float>) -> ReferenceTestResult {
let params = PCAParams::default();
let start = Instant::now();
let _our_result = self.mock_our_pca(data.view(), ¶ms);
let our_time = start.elapsed().as_secs_f64();
let start = Instant::now();
let reference = SklearnPCA;
let _ref_result = reference.fit_transform(data.view(), ¶ms);
let ref_time = start.elapsed().as_secs_f64();
let speedup = if our_time > 0.0 {
ref_time / our_time
} else {
1.0
};
let performance_metrics = PerformanceMetrics {
our_runtime: our_time,
reference_runtime: ref_time,
speedup_factor: speedup,
memory_usage: None,
};
let mut metadata = HashMap::new();
metadata.insert("data_shape".to_string(), format!("{:?}", data.shape()));
metadata.insert("speedup".to_string(), format!("{:.2}x", speedup));
ReferenceTestResult {
test_name: format!("PCA Performance {:?}", data.shape()),
passed: true,
error_message: None,
max_difference: 0.0,
performance_metrics: Some(performance_metrics),
metadata,
}
}
fn mock_our_tsne(
&self,
data: ArrayView2<Float>,
_params: &TSNEParams,
) -> SklResult<Array2<Float>> {
std::thread::sleep(std::time::Duration::from_millis(10));
Ok(Array2::zeros((data.nrows(), 2)))
}
fn mock_our_pca(
&self,
data: ArrayView2<Float>,
_params: &PCAParams,
) -> SklResult<Array2<Float>> {
std::thread::sleep(std::time::Duration::from_millis(5));
Ok(Array2::zeros((data.nrows(), 2)))
}
fn generate_test_data(&self, n_samples: usize, n_features: usize) -> Array2<Float> {
let mut rng = StdRng::seed_from_u64(42);
let mut data = Array2::zeros((n_samples, n_features));
for i in 0..n_samples {
for j in 0..n_features {
data[[i, j]] = rng.random_range(-1.0..1.0);
}
}
data
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_reference_framework() {
let config = ReferenceTestConfig {
tolerance: 1e-3, test_multiple_seeds: false,
n_seeds: 1,
test_edge_cases: false,
test_performance: false,
};
let framework = ReferenceTestFramework::new(config);
let results = framework.run_all_tests();
assert!(!results.is_empty());
framework.print_test_summary(&results);
}
#[test]
fn test_mock_implementations() {
let data = Array2::zeros((10, 3));
let tsne = SklearnTSNE;
let tsne_params = TSNEParams::default();
let tsne_result = tsne.fit_transform(data.view(), &tsne_params);
assert!(tsne_result.is_ok());
let pca = SklearnPCA;
let pca_params = PCAParams::default();
let pca_result = pca.fit_transform(data.view(), &pca_params);
assert!(pca_result.is_ok());
let isomap = SklearnIsomap;
let isomap_params = IsomapParams::default();
let isomap_result = isomap.fit_transform(data.view(), &isomap_params);
assert!(isomap_result.is_ok());
}
#[test]
fn test_performance_benchmarks() {
let config = ReferenceTestConfig {
test_performance: true,
..Default::default()
};
let benchmark = BenchmarkComparison::new(config);
let results = benchmark.run_performance_benchmarks();
assert!(!results.is_empty());
for result in &results {
assert!(result.performance_metrics.is_some());
}
}
}