use scirs2_core::ndarray::{Array1, Array2, ArrayView2};
use scirs2_linalg::compat::{ArrayLinalgExt, UPLO};
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use sklears_core::{
error::{Result as SklResult, SklearsError},
types::Float,
};
use std::collections::HashMap;
use std::fmt::Debug;
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct PipelineContext {
pub data: Array2<f64>,
pub metadata: PipelineMetadata,
pub parameters: HashMap<String, ParameterValue>,
pub stats: ExecutionStats,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct PipelineMetadata {
pub original_shape: (usize, usize),
pub current_shape: (usize, usize),
pub transformations: Vec<TransformationInfo>,
pub quality_metrics: HashMap<String, f64>,
pub warnings: Vec<String>,
pub custom: HashMap<String, String>,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct TransformationInfo {
pub name: String,
pub parameters: HashMap<String, String>,
pub execution_time: f64,
pub input_shape: (usize, usize),
pub output_shape: (usize, usize),
pub metrics: HashMap<String, f64>,
}
#[derive(Debug, Clone)]
pub struct ExecutionStats {
pub total_time: Duration,
pub memory_usage: MemoryStats,
pub step_times: Vec<Duration>,
pub counters: HashMap<String, u64>,
}
#[derive(Debug, Clone)]
pub struct MemoryStats {
pub peak_memory: u64,
pub current_memory: u64,
pub allocations: u64,
pub deallocations: u64,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum ParameterValue {
Int(i64),
Float(f64),
String(String),
Bool(bool),
IntArray(Vec<i64>),
FloatArray(Vec<f64>),
}
pub trait PipelineMiddleware: Send + Sync + Debug {
fn process(&self, context: PipelineContext) -> SklResult<PipelineContext>;
fn name(&self) -> &str;
fn description(&self) -> &str {
"No description provided"
}
fn can_apply(&self, _context: &PipelineContext) -> bool {
true
}
fn configuration(&self) -> MiddlewareConfig {
MiddlewareConfig::default()
}
}
#[derive(Debug, Clone, Default)]
pub struct MiddlewareConfig {
pub optional: bool,
pub continue_on_error: bool,
pub timeout: Option<Duration>,
pub priority: i32,
}
pub struct ManifoldPipeline {
middleware: Vec<Arc<dyn PipelineMiddleware>>,
config: PipelineConfig,
}
#[derive(Debug, Clone)]
pub struct PipelineConfig {
pub verbose: bool,
pub collect_metrics: bool,
pub parallel: bool,
pub max_execution_time: Option<Duration>,
pub fail_fast: bool,
}
impl Default for PipelineConfig {
fn default() -> Self {
Self {
verbose: false,
collect_metrics: true,
parallel: false,
max_execution_time: None,
fail_fast: true,
}
}
}
impl Default for ManifoldPipeline {
fn default() -> Self {
Self::new()
}
}
impl ManifoldPipeline {
pub fn new() -> Self {
Self {
middleware: Vec::new(),
config: PipelineConfig::default(),
}
}
pub fn add_middleware(mut self, middleware: Arc<dyn PipelineMiddleware>) -> Self {
self.middleware.push(middleware);
self
}
pub fn config(mut self, config: PipelineConfig) -> Self {
self.config = config;
self
}
pub fn execute(&self, data: &ArrayView2<'_, Float>) -> SklResult<PipelineResult> {
let start_time = Instant::now();
let data_f64 = data.mapv(|x| x);
let original_shape = data_f64.dim();
let mut context = PipelineContext {
data: data_f64,
metadata: PipelineMetadata {
original_shape,
current_shape: original_shape,
transformations: Vec::new(),
quality_metrics: HashMap::new(),
warnings: Vec::new(),
custom: HashMap::new(),
},
parameters: HashMap::new(),
stats: ExecutionStats {
total_time: Duration::from_secs(0),
memory_usage: MemoryStats {
peak_memory: 0,
current_memory: 0,
allocations: 0,
deallocations: 0,
},
step_times: Vec::new(),
counters: HashMap::new(),
},
};
let mut sorted_middleware = self.middleware.clone();
sorted_middleware.sort_by_key(|m| std::cmp::Reverse(m.configuration().priority));
for middleware in &sorted_middleware {
if !middleware.can_apply(&context) {
if self.config.verbose {
println!(
"Skipping middleware '{}' - not applicable",
middleware.name()
);
}
continue;
}
let step_start = Instant::now();
let context_for_middleware = context.clone();
let step_result = self.execute_middleware(middleware.as_ref(), context_for_middleware);
match step_result {
Ok(new_context) => {
context = new_context;
let step_time = step_start.elapsed();
context.stats.step_times.push(step_time);
if self.config.verbose {
println!(
"Completed middleware '{}' in {:?}",
middleware.name(),
step_time
);
}
}
Err(e) => {
let middleware_config = middleware.configuration();
if middleware_config.continue_on_error && !self.config.fail_fast {
context.metadata.warnings.push(format!(
"Middleware '{}' failed: {}",
middleware.name(),
e
));
if self.config.verbose {
println!(
"Warning: Middleware '{}' failed but continuing: {}",
middleware.name(),
e
);
}
} else {
return Err(SklearsError::InvalidInput(format!(
"Pipeline failed at middleware '{}': {}",
middleware.name(),
e
)));
}
}
}
if let Some(max_time) = self.config.max_execution_time {
if start_time.elapsed() > max_time {
return Err(SklearsError::InvalidInput(
"Pipeline execution timed out".to_string(),
));
}
}
}
context.stats.total_time = start_time.elapsed();
context.metadata.current_shape = context.data.dim();
Ok(PipelineResult {
data: context.data,
metadata: context.metadata,
execution_stats: context.stats,
})
}
fn execute_middleware(
&self,
middleware: &dyn PipelineMiddleware,
context: PipelineContext,
) -> SklResult<PipelineContext> {
let start_time = Instant::now();
let config = middleware.configuration();
let result = if let Some(timeout) = config.timeout {
let result = middleware.process(context);
if start_time.elapsed() > timeout {
return Err(SklearsError::InvalidInput(format!(
"Middleware '{}' timed out after {:?}",
middleware.name(),
timeout
)));
}
result
} else {
middleware.process(context)
};
result
}
pub fn summary(&self) -> PipelineSummary {
PipelineSummary {
middleware_count: self.middleware.len(),
middleware_names: self
.middleware
.iter()
.map(|m| m.name().to_string())
.collect(),
config: self.config.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct PipelineResult {
pub data: Array2<f64>,
pub metadata: PipelineMetadata,
pub execution_stats: ExecutionStats,
}
#[derive(Debug, Clone)]
pub struct PipelineSummary {
pub middleware_count: usize,
pub middleware_names: Vec<String>,
pub config: PipelineConfig,
}
#[derive(Debug)]
pub struct DataValidationMiddleware {
min_samples: usize,
max_samples: usize,
min_features: usize,
max_features: usize,
check_finite: bool,
check_duplicates: bool,
}
impl Default for DataValidationMiddleware {
fn default() -> Self {
Self::new()
}
}
impl DataValidationMiddleware {
pub fn new() -> Self {
Self {
min_samples: 1,
max_samples: usize::MAX,
min_features: 1,
max_features: usize::MAX,
check_finite: true,
check_duplicates: false,
}
}
pub fn min_samples(mut self, min_samples: usize) -> Self {
self.min_samples = min_samples;
self
}
pub fn max_samples(mut self, max_samples: usize) -> Self {
self.max_samples = max_samples;
self
}
pub fn min_features(mut self, min_features: usize) -> Self {
self.min_features = min_features;
self
}
pub fn max_features(mut self, max_features: usize) -> Self {
self.max_features = max_features;
self
}
pub fn check_finite(mut self, check_finite: bool) -> Self {
self.check_finite = check_finite;
self
}
pub fn check_duplicates(mut self, check_duplicates: bool) -> Self {
self.check_duplicates = check_duplicates;
self
}
}
impl PipelineMiddleware for DataValidationMiddleware {
fn process(&self, mut context: PipelineContext) -> SklResult<PipelineContext> {
let (n_samples, n_features) = context.data.dim();
if n_samples < self.min_samples {
return Err(SklearsError::InvalidInput(format!(
"Too few samples: {} < {}",
n_samples, self.min_samples
)));
}
if n_samples > self.max_samples {
return Err(SklearsError::InvalidInput(format!(
"Too many samples: {} > {}",
n_samples, self.max_samples
)));
}
if n_features < self.min_features {
return Err(SklearsError::InvalidInput(format!(
"Too few features: {} < {}",
n_features, self.min_features
)));
}
if n_features > self.max_features {
return Err(SklearsError::InvalidInput(format!(
"Too many features: {} > {}",
n_features, self.max_features
)));
}
if self.check_finite && !context.data.iter().all(|&x| x.is_finite()) {
return Err(SklearsError::InvalidInput(
"Data contains non-finite values (NaN or infinity)".to_string(),
));
}
if self.check_duplicates {
for i in 0..n_samples {
for j in (i + 1)..n_samples {
let row_i = context.data.row(i);
let row_j = context.data.row(j);
if row_i
.iter()
.zip(row_j.iter())
.all(|(a, b)| (a - b).abs() < 1e-15)
{
context
.metadata
.warnings
.push(format!("Duplicate rows found: {} and {}", i, j));
}
}
}
}
context
.metadata
.quality_metrics
.insert("n_samples".to_string(), n_samples as f64);
context
.metadata
.quality_metrics
.insert("n_features".to_string(), n_features as f64);
Ok(context)
}
fn name(&self) -> &str {
"DataValidation"
}
fn description(&self) -> &str {
"Validates input data for common issues"
}
}
#[derive(Debug)]
pub struct StandardizationMiddleware {
with_mean: bool,
with_std: bool,
}
impl Default for StandardizationMiddleware {
fn default() -> Self {
Self::new()
}
}
impl StandardizationMiddleware {
pub fn new() -> Self {
Self {
with_mean: true,
with_std: true,
}
}
pub fn with_mean(mut self, with_mean: bool) -> Self {
self.with_mean = with_mean;
self
}
pub fn with_std(mut self, with_std: bool) -> Self {
self.with_std = with_std;
self
}
}
impl PipelineMiddleware for StandardizationMiddleware {
fn process(&self, mut context: PipelineContext) -> SklResult<PipelineContext> {
let (n_samples, n_features) = context.data.dim();
if self.with_mean || self.with_std {
let means = if self.with_mean {
context
.data
.mean_axis(scirs2_core::ndarray::Axis(0))
.expect("operation should succeed")
} else {
Array1::zeros(n_features)
};
let stds = if self.with_std {
context.data.std_axis(scirs2_core::ndarray::Axis(0), 0.0)
} else {
Array1::ones(n_features)
};
for i in 0..n_samples {
for j in 0..n_features {
if stds[j] > 1e-15 {
context.data[[i, j]] = (context.data[[i, j]] - means[j]) / stds[j];
} else {
context.data[[i, j]] -= means[j];
}
}
}
context
.metadata
.custom
.insert("means".to_string(), format!("{:?}", means));
context
.metadata
.custom
.insert("stds".to_string(), format!("{:?}", stds));
}
context.metadata.transformations.push(TransformationInfo {
name: "Standardization".to_string(),
parameters: {
let mut params = HashMap::new();
params.insert("with_mean".to_string(), self.with_mean.to_string());
params.insert("with_std".to_string(), self.with_std.to_string());
params
},
execution_time: 0.0, input_shape: (n_samples, n_features),
output_shape: (n_samples, n_features),
metrics: HashMap::new(),
});
Ok(context)
}
fn name(&self) -> &str {
"Standardization"
}
fn description(&self) -> &str {
"Standardizes data by removing mean and scaling to unit variance"
}
}
#[derive(Debug)]
pub struct QualityAssessmentMiddleware {
compute_intrinsic_dim: bool,
compute_clustering_metrics: bool,
}
impl Default for QualityAssessmentMiddleware {
fn default() -> Self {
Self::new()
}
}
impl QualityAssessmentMiddleware {
pub fn new() -> Self {
Self {
compute_intrinsic_dim: true,
compute_clustering_metrics: false,
}
}
pub fn compute_intrinsic_dim(mut self, compute: bool) -> Self {
self.compute_intrinsic_dim = compute;
self
}
pub fn compute_clustering_metrics(mut self, compute: bool) -> Self {
self.compute_clustering_metrics = compute;
self
}
}
impl PipelineMiddleware for QualityAssessmentMiddleware {
fn process(&self, mut context: PipelineContext) -> SklResult<PipelineContext> {
if self.compute_intrinsic_dim {
let intrinsic_dim = estimate_intrinsic_dimensionality(&context.data);
context
.metadata
.quality_metrics
.insert("estimated_intrinsic_dim".to_string(), intrinsic_dim);
}
let variance = context
.data
.var_axis(scirs2_core::ndarray::Axis(0), 0.0)
.sum();
context
.metadata
.quality_metrics
.insert("total_variance".to_string(), variance);
let condition_number = estimate_condition_number(&context.data);
context
.metadata
.quality_metrics
.insert("condition_number".to_string(), condition_number);
Ok(context)
}
fn name(&self) -> &str {
"QualityAssessment"
}
fn description(&self) -> &str {
"Assesses data quality and computes various metrics"
}
}
fn estimate_intrinsic_dimensionality(data: &Array2<f64>) -> f64 {
let (n_samples, n_features) = data.dim();
if n_samples == 0 || n_features == 0 {
return 0.0;
}
let mean = data
.mean_axis(scirs2_core::ndarray::Axis(0))
.expect("operation should succeed");
let centered = data - &mean;
let cov_matrix = centered.t().dot(¢ered) / (n_samples - 1) as f64;
if let Ok((eigenvalues, _)) = cov_matrix.eigh(UPLO::Lower) {
let total_var: f64 = eigenvalues.iter().filter(|&&x| x > 0.0).sum();
let mut cumsum = 0.0;
let mut count = 0;
for &val in eigenvalues.iter().rev() {
if val > 0.0 {
cumsum += val;
count += 1;
if cumsum / total_var >= 0.95 {
break;
}
}
}
count as f64
} else {
data.ncols() as f64
}
}
fn estimate_condition_number(data: &Array2<f64>) -> f64 {
if let Ok((_, singular_values, _)) = data.svd(false) {
if let (Some(&max_sv), Some(&min_sv)) = (
singular_values
.iter()
.filter(|&&x| x > 0.0)
.max_by(|a, b| a.partial_cmp(b).expect("operation should succeed")),
singular_values
.iter()
.filter(|&&x| x > 0.0)
.min_by(|a, b| a.partial_cmp(b).expect("operation should succeed")),
) {
max_sv / min_sv
} else {
f64::INFINITY
}
} else {
f64::INFINITY
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array2;
#[test]
fn test_pipeline_execution() {
let data = Array2::from_shape_vec((10, 3), (0..30).map(|i| i as f64).collect())
.expect("operation should succeed");
let pipeline = ManifoldPipeline::new()
.add_middleware(Arc::new(DataValidationMiddleware::new().min_samples(5)))
.add_middleware(Arc::new(StandardizationMiddleware::new()))
.add_middleware(Arc::new(QualityAssessmentMiddleware::new()));
let result = pipeline
.execute(&data.view())
.expect("operation should succeed");
assert_eq!(result.data.dim(), (10, 3));
assert!(!result.metadata.transformations.is_empty());
assert!(!result.metadata.quality_metrics.is_empty());
}
#[test]
fn test_data_validation_middleware() {
let data = Array2::from_shape_vec((2, 2), vec![1.0, 2.0, 3.0, 4.0])
.expect("operation should succeed");
let context = PipelineContext {
data,
metadata: PipelineMetadata {
original_shape: (2, 2),
current_shape: (2, 2),
transformations: Vec::new(),
quality_metrics: HashMap::new(),
warnings: Vec::new(),
custom: HashMap::new(),
},
parameters: HashMap::new(),
stats: ExecutionStats {
total_time: Duration::from_secs(0),
memory_usage: MemoryStats {
peak_memory: 0,
current_memory: 0,
allocations: 0,
deallocations: 0,
},
step_times: Vec::new(),
counters: HashMap::new(),
},
};
let middleware = DataValidationMiddleware::new()
.min_samples(1)
.max_samples(10);
let result = middleware
.process(context)
.expect("operation should succeed");
assert_eq!(result.metadata.quality_metrics.get("n_samples"), Some(&2.0));
assert_eq!(
result.metadata.quality_metrics.get("n_features"),
Some(&2.0)
);
}
#[test]
fn test_standardization_middleware() {
let data = Array2::from_shape_vec((4, 2), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0])
.expect("operation should succeed");
let context = PipelineContext {
data,
metadata: PipelineMetadata {
original_shape: (4, 2),
current_shape: (4, 2),
transformations: Vec::new(),
quality_metrics: HashMap::new(),
warnings: Vec::new(),
custom: HashMap::new(),
},
parameters: HashMap::new(),
stats: ExecutionStats {
total_time: Duration::from_secs(0),
memory_usage: MemoryStats {
peak_memory: 0,
current_memory: 0,
allocations: 0,
deallocations: 0,
},
step_times: Vec::new(),
counters: HashMap::new(),
},
};
let middleware = StandardizationMiddleware::new();
let result = middleware
.process(context)
.expect("operation should succeed");
let means = result
.data
.mean_axis(scirs2_core::ndarray::Axis(0))
.expect("operation should succeed");
let stds = result.data.std_axis(scirs2_core::ndarray::Axis(0), 0.0);
for &mean in means.iter() {
assert!((mean.abs()) < 1e-10);
}
for &std in stds.iter() {
assert!((std - 1.0).abs() < 1e-10);
}
}
#[test]
fn test_pipeline_summary() {
let pipeline = ManifoldPipeline::new()
.add_middleware(Arc::new(DataValidationMiddleware::new()))
.add_middleware(Arc::new(StandardizationMiddleware::new()));
let summary = pipeline.summary();
assert_eq!(summary.middleware_count, 2);
assert!(summary
.middleware_names
.contains(&"DataValidation".to_string()));
assert!(summary
.middleware_names
.contains(&"Standardization".to_string()));
}
}