use scirs2_core::ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use scirs2_core::random::{Distribution, RandNormal, Random};
use std::collections::HashMap;
use thiserror::Error;
#[derive(Error, Debug)]
pub enum DatasetTraitError {
#[error("Generation error: {0}")]
Generation(String),
#[error("Validation error: {0}")]
Validation(String),
#[error("Configuration error: {0}")]
Configuration(String),
#[error("IO error: {0}")]
Io(String),
#[error("Dimension mismatch: expected {expected}, got {actual}")]
DimensionMismatch { expected: String, actual: String },
#[error("Unsupported operation: {0}")]
UnsupportedOperation(String),
}
pub type DatasetTraitResult<T> = Result<T, DatasetTraitError>;
pub trait Dataset {
fn n_samples(&self) -> usize;
fn n_features(&self) -> usize;
fn shape(&self) -> (usize, usize) {
(self.n_samples(), self.n_features())
}
fn features(&self) -> DatasetTraitResult<ArrayView2<'_, f64>>;
fn sample(&self, index: usize) -> DatasetTraitResult<ArrayView1<'_, f64>>;
fn has_targets(&self) -> bool;
fn targets(&self) -> DatasetTraitResult<Option<ArrayView1<'_, f64>>>;
fn metadata(&self) -> HashMap<String, String> {
HashMap::new()
}
}
pub trait DatasetGenerator {
type Config: Default + Clone;
type Output: Dataset;
fn generate(&self, config: Self::Config) -> DatasetTraitResult<Self::Output>;
fn name(&self) -> &'static str;
fn description(&self) -> &'static str;
fn validate_config(&self, config: &Self::Config) -> DatasetTraitResult<()> {
let _ = config;
Ok(())
}
}
pub trait DatasetLoader {
type Config: Default + Clone;
type Output: Dataset;
fn load(&self, config: Self::Config) -> DatasetTraitResult<Self::Output>;
fn name(&self) -> &'static str;
fn available_datasets(&self) -> Vec<String>;
fn has_dataset(&self, name: &str) -> bool {
self.available_datasets().contains(&name.to_string())
}
}
pub trait DatasetTransformer {
type Config: Default + Clone;
type Input: Dataset;
type Output: Dataset;
fn transform(
&self,
input: Self::Input,
config: Self::Config,
) -> DatasetTraitResult<Self::Output>;
fn name(&self) -> &'static str;
fn can_transform(&self, input: &Self::Input) -> bool;
}
pub trait DatasetValidator {
type Config: Default + Clone;
type Report: Default;
fn validate(
&self,
dataset: &dyn Dataset,
config: Self::Config,
) -> DatasetTraitResult<Self::Report>;
fn name(&self) -> &'static str;
fn criteria(&self) -> Vec<String>;
}
pub trait StreamingDataset: Dataset {
type Batch;
fn batch(&self, start: usize, size: usize) -> DatasetTraitResult<Self::Batch>;
fn batches(
&self,
batch_size: usize,
) -> Box<dyn Iterator<Item = DatasetTraitResult<Self::Batch>>>;
fn preferred_batch_size(&self) -> usize {
1000
}
}
pub trait MutableDataset: Dataset {
fn set_sample(&mut self, index: usize, sample: ArrayView1<f64>) -> DatasetTraitResult<()>;
fn set_targets(&mut self, targets: ArrayView1<f64>) -> DatasetTraitResult<()>;
fn add_sample(
&mut self,
sample: ArrayView1<f64>,
target: Option<f64>,
) -> DatasetTraitResult<()>;
fn remove_sample(&mut self, index: usize) -> DatasetTraitResult<()>;
}
pub trait GenerationStrategy {
type Config: Default + Clone;
fn apply(&self, config: &mut Self::Config, rng: &mut Random) -> DatasetTraitResult<()>;
fn name(&self) -> &'static str;
fn is_applicable(&self, config: &Self::Config) -> bool;
}
#[derive(Debug, Clone)]
pub struct InMemoryDataset {
features: Array2<f64>,
targets: Option<Array1<f64>>,
metadata: HashMap<String, String>,
}
impl InMemoryDataset {
pub fn new(features: Array2<f64>, targets: Option<Array1<f64>>) -> Self {
Self {
features,
targets,
metadata: HashMap::new(),
}
}
pub fn with_metadata(
features: Array2<f64>,
targets: Option<Array1<f64>>,
metadata: HashMap<String, String>,
) -> Self {
Self {
features,
targets,
metadata,
}
}
pub fn add_metadata(&mut self, key: String, value: String) {
self.metadata.insert(key, value);
}
}
impl Dataset for InMemoryDataset {
fn n_samples(&self) -> usize {
self.features.nrows()
}
fn n_features(&self) -> usize {
self.features.ncols()
}
fn features(&self) -> DatasetTraitResult<ArrayView2<'_, f64>> {
Ok(self.features.view())
}
fn sample(&self, index: usize) -> DatasetTraitResult<ArrayView1<'_, f64>> {
if index >= self.n_samples() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("index < {}", self.n_samples()),
actual: format!("index = {}", index),
});
}
Ok(self.features.row(index))
}
fn has_targets(&self) -> bool {
self.targets.is_some()
}
fn targets(&self) -> DatasetTraitResult<Option<ArrayView1<'_, f64>>> {
Ok(self.targets.as_ref().map(|t| t.view()))
}
fn metadata(&self) -> HashMap<String, String> {
self.metadata.clone()
}
}
impl MutableDataset for InMemoryDataset {
fn set_sample(&mut self, index: usize, sample: ArrayView1<f64>) -> DatasetTraitResult<()> {
if index >= self.n_samples() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("index < {}", self.n_samples()),
actual: format!("index = {}", index),
});
}
if sample.len() != self.n_features() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("{} features", self.n_features()),
actual: format!("{} features", sample.len()),
});
}
self.features.row_mut(index).assign(&sample);
Ok(())
}
fn set_targets(&mut self, targets: ArrayView1<f64>) -> DatasetTraitResult<()> {
if targets.len() != self.n_samples() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("{} targets", self.n_samples()),
actual: format!("{} targets", targets.len()),
});
}
self.targets = Some(targets.to_owned());
Ok(())
}
fn add_sample(
&mut self,
sample: ArrayView1<f64>,
_target: Option<f64>,
) -> DatasetTraitResult<()> {
if sample.len() != self.n_features() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("{} features", self.n_features()),
actual: format!("{} features", sample.len()),
});
}
Err(DatasetTraitError::UnsupportedOperation(
"Adding samples to fixed-size arrays not yet implemented".to_string(),
))
}
fn remove_sample(&mut self, index: usize) -> DatasetTraitResult<()> {
if index >= self.n_samples() {
return Err(DatasetTraitError::DimensionMismatch {
expected: format!("index < {}", self.n_samples()),
actual: format!("index = {}", index),
});
}
Err(DatasetTraitError::UnsupportedOperation(
"Removing samples from fixed-size arrays not yet implemented".to_string(),
))
}
}
pub struct GeneratorRegistry {
generators: HashMap<
String,
Box<dyn DatasetGenerator<Config = GeneratorConfig, Output = InMemoryDataset>>,
>,
}
impl GeneratorRegistry {
pub fn new() -> Self {
Self {
generators: HashMap::new(),
}
}
pub fn register<G>(&mut self, generator: G)
where
G: DatasetGenerator<Config = GeneratorConfig, Output = InMemoryDataset> + 'static,
{
self.generators
.insert(generator.name().to_string(), Box::new(generator));
}
pub fn get(
&self,
name: &str,
) -> Option<&dyn DatasetGenerator<Config = GeneratorConfig, Output = InMemoryDataset>> {
self.generators.get(name).map(|g| g.as_ref())
}
pub fn list(&self) -> Vec<String> {
self.generators.keys().cloned().collect()
}
pub fn generate(
&self,
name: &str,
config: GeneratorConfig,
) -> DatasetTraitResult<InMemoryDataset> {
let generator = self.get(name).ok_or_else(|| {
DatasetTraitError::Configuration(format!("Unknown generator: {}", name))
})?;
generator.generate(config)
}
}
impl Default for GeneratorRegistry {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct GeneratorConfig {
pub n_samples: usize,
pub n_features: usize,
pub random_state: Option<u64>,
pub parameters: HashMap<String, ConfigValue>,
}
impl Default for GeneratorConfig {
fn default() -> Self {
Self {
n_samples: 100,
n_features: 2,
random_state: None,
parameters: HashMap::new(),
}
}
}
impl GeneratorConfig {
pub fn new(n_samples: usize, n_features: usize) -> Self {
Self {
n_samples,
n_features,
random_state: None,
parameters: HashMap::new(),
}
}
pub fn set_parameter<T: Into<ConfigValue>>(&mut self, key: String, value: T) {
self.parameters.insert(key, value.into());
}
pub fn get_parameter(&self, key: &str) -> Option<&ConfigValue> {
self.parameters.get(key)
}
pub fn with_random_state(mut self, seed: u64) -> Self {
self.random_state = Some(seed);
self
}
}
#[derive(Debug, Clone)]
pub enum ConfigValue {
Int(i64),
Float(f64),
String(String),
Bool(bool),
IntArray(Vec<i64>),
FloatArray(Vec<f64>),
}
impl From<i64> for ConfigValue {
fn from(value: i64) -> Self {
ConfigValue::Int(value)
}
}
impl From<f64> for ConfigValue {
fn from(value: f64) -> Self {
ConfigValue::Float(value)
}
}
impl From<String> for ConfigValue {
fn from(value: String) -> Self {
ConfigValue::String(value)
}
}
impl From<bool> for ConfigValue {
fn from(value: bool) -> Self {
ConfigValue::Bool(value)
}
}
impl From<Vec<i64>> for ConfigValue {
fn from(value: Vec<i64>) -> Self {
ConfigValue::IntArray(value)
}
}
impl From<Vec<f64>> for ConfigValue {
fn from(value: Vec<f64>) -> Self {
ConfigValue::FloatArray(value)
}
}
pub struct ClassificationGenerator;
impl DatasetGenerator for ClassificationGenerator {
type Config = GeneratorConfig;
type Output = InMemoryDataset;
fn generate(&self, config: Self::Config) -> DatasetTraitResult<Self::Output> {
let mut rng = match config.random_state {
Some(seed) => Random::seed(seed),
None => Random::seed(42),
};
let n_classes = config
.get_parameter("n_classes")
.and_then(|v| match v {
ConfigValue::Int(n) => Some(*n as usize),
_ => None,
})
.unwrap_or(2);
let mut features = Array2::<f64>::zeros((config.n_samples, config.n_features));
let normal_dist = RandNormal::new(0.0, 1.0).expect("operation should succeed");
for mut row in features.rows_mut() {
for val in row.iter_mut() {
*val = normal_dist.sample(&mut rng);
}
}
let targets: Array1<f64> =
Array1::from_shape_fn(config.n_samples, |_| rng.gen_range(0..n_classes) as f64);
let mut metadata = HashMap::new();
metadata.insert("generator".to_string(), "classification".to_string());
metadata.insert("n_classes".to_string(), n_classes.to_string());
Ok(InMemoryDataset::with_metadata(
features,
Some(targets),
metadata,
))
}
fn name(&self) -> &'static str {
"classification"
}
fn description(&self) -> &'static str {
"Generates a classification dataset with Gaussian features"
}
fn validate_config(&self, config: &Self::Config) -> DatasetTraitResult<()> {
if config.n_samples == 0 {
return Err(DatasetTraitError::Configuration(
"n_samples must be > 0".to_string(),
));
}
if config.n_features == 0 {
return Err(DatasetTraitError::Configuration(
"n_features must be > 0".to_string(),
));
}
if let Some(ConfigValue::Int(n_classes)) = config.get_parameter("n_classes") {
if *n_classes <= 0 {
return Err(DatasetTraitError::Configuration(
"n_classes must be > 0".to_string(),
));
}
}
Ok(())
}
}
pub struct RegressionGenerator;
impl DatasetGenerator for RegressionGenerator {
type Config = GeneratorConfig;
type Output = InMemoryDataset;
fn generate(&self, config: Self::Config) -> DatasetTraitResult<Self::Output> {
let mut rng = match config.random_state {
Some(seed) => Random::seed(seed),
None => Random::seed(42),
};
let noise = config
.get_parameter("noise")
.and_then(|v| match v {
ConfigValue::Float(n) => Some(*n),
_ => None,
})
.unwrap_or(0.1);
let mut features = Array2::<f64>::zeros((config.n_samples, config.n_features));
let normal_dist = RandNormal::new(0.0, 1.0).expect("operation should succeed");
for mut row in features.rows_mut() {
for val in row.iter_mut() {
*val = normal_dist.sample(&mut rng);
}
}
let coefficients: Array1<f64> =
Array1::from_shape_fn(config.n_features, |_| rng.random_range(-1.0..1.0));
let mut targets = Array1::<f64>::zeros(config.n_samples);
for (i, target) in targets.iter_mut().enumerate() {
let feature_row = features.row(i);
let noise_dist = RandNormal::new(0.0, noise).expect("operation should succeed");
*target = feature_row.dot(&coefficients) + noise_dist.sample(&mut rng);
}
let mut metadata = HashMap::new();
metadata.insert("generator".to_string(), "regression".to_string());
metadata.insert("noise".to_string(), noise.to_string());
Ok(InMemoryDataset::with_metadata(
features,
Some(targets),
metadata,
))
}
fn name(&self) -> &'static str {
"regression"
}
fn description(&self) -> &'static str {
"Generates a regression dataset with linear relationship and noise"
}
}
pub fn create_default_registry() -> GeneratorRegistry {
let mut registry = GeneratorRegistry::new();
registry.register(ClassificationGenerator);
registry.register(RegressionGenerator);
registry
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array;
#[test]
fn test_in_memory_dataset() {
let features = Array::from_shape_vec((10, 3), (0..30).map(|x| x as f64).collect())
.expect("shape and data length should match");
let targets = Array1::from_shape_vec(10, (0..10).map(|x| x as f64).collect())
.expect("shape and data length should match");
let dataset = InMemoryDataset::new(features, Some(targets));
assert_eq!(dataset.n_samples(), 10);
assert_eq!(dataset.n_features(), 3);
assert_eq!(dataset.shape(), (10, 3));
assert!(dataset.has_targets());
let features_view = dataset.features().expect("operation should succeed");
assert_eq!(features_view.dim(), (10, 3));
let sample = dataset.sample(5).expect("sampling should succeed");
assert_eq!(sample.len(), 3);
assert_eq!(sample[0], 15.0);
let targets_view = dataset
.targets()
.expect("operation should succeed")
.expect("operation should succeed");
assert_eq!(targets_view.len(), 10);
assert_eq!(targets_view[5], 5.0);
}
#[test]
fn test_generator_registry() {
let mut registry = GeneratorRegistry::new();
registry.register(ClassificationGenerator);
registry.register(RegressionGenerator);
let generators = registry.list();
assert!(generators.contains(&"classification".to_string()));
assert!(generators.contains(&"regression".to_string()));
let config = GeneratorConfig::new(50, 4);
let dataset = registry
.generate("classification", config)
.expect("operation should succeed");
assert_eq!(dataset.n_samples(), 50);
assert_eq!(dataset.n_features(), 4);
assert!(dataset.has_targets());
}
#[test]
fn test_classification_generator() {
let generator = ClassificationGenerator;
let mut config = GeneratorConfig::new(100, 5);
config.set_parameter("n_classes".to_string(), 3i64);
config.random_state = Some(42);
let dataset = generator
.generate(config)
.expect("operation should succeed");
assert_eq!(dataset.n_samples(), 100);
assert_eq!(dataset.n_features(), 5);
assert!(dataset.has_targets());
let targets = dataset
.targets()
.expect("operation should succeed")
.expect("operation should succeed");
assert!(targets.iter().all(|&t| (0.0..3.0).contains(&t)));
let metadata = dataset.metadata();
assert_eq!(
metadata.get("generator"),
Some(&"classification".to_string())
);
assert_eq!(metadata.get("n_classes"), Some(&"3".to_string()));
}
#[test]
fn test_regression_generator() {
let generator = RegressionGenerator;
let mut config = GeneratorConfig::new(100, 3);
config.set_parameter("noise".to_string(), 0.05);
config.random_state = Some(42);
let dataset = generator
.generate(config)
.expect("operation should succeed");
assert_eq!(dataset.n_samples(), 100);
assert_eq!(dataset.n_features(), 3);
assert!(dataset.has_targets());
let metadata = dataset.metadata();
assert_eq!(metadata.get("generator"), Some(&"regression".to_string()));
assert_eq!(metadata.get("noise"), Some(&"0.05".to_string()));
}
#[test]
fn test_config_validation() {
let generator = ClassificationGenerator;
let valid_config = GeneratorConfig::new(100, 5);
assert!(generator.validate_config(&valid_config).is_ok());
let invalid_config = GeneratorConfig::new(0, 5);
assert!(generator.validate_config(&invalid_config).is_err());
let invalid_config = GeneratorConfig::new(100, 0);
assert!(generator.validate_config(&invalid_config).is_err());
}
#[test]
fn test_config_parameters() {
let mut config = GeneratorConfig::new(100, 5);
config.set_parameter("n_classes".to_string(), 3i64);
config.set_parameter("noise".to_string(), 0.1);
config.set_parameter("seed".to_string(), "test".to_string());
config.set_parameter("enabled".to_string(), true);
assert!(matches!(
config.get_parameter("n_classes"),
Some(ConfigValue::Int(3))
));
assert!(matches!(
config.get_parameter("noise"),
Some(ConfigValue::Float(0.1))
));
assert!(matches!(
config.get_parameter("seed"),
Some(ConfigValue::String(_))
));
assert!(matches!(
config.get_parameter("enabled"),
Some(ConfigValue::Bool(true))
));
}
#[test]
fn test_default_registry() {
let registry = create_default_registry();
let generators = registry.list();
assert!(generators.contains(&"classification".to_string()));
assert!(generators.contains(&"regression".to_string()));
assert_eq!(generators.len(), 2);
}
#[test]
fn test_mutable_dataset() {
let features = Array::from_shape_vec((3, 2), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
.expect("shape and data length should match");
let targets = Array1::from_shape_vec(3, vec![10.0, 20.0, 30.0])
.expect("shape and data length should match");
let mut dataset = InMemoryDataset::new(features, Some(targets));
let new_sample = Array1::from_vec(vec![99.0, 88.0]);
assert!(dataset.set_sample(1, new_sample.view()).is_ok());
let updated_sample = dataset.sample(1).expect("sampling should succeed");
assert_eq!(updated_sample[0], 99.0);
assert_eq!(updated_sample[1], 88.0);
let wrong_sample = Array1::from_vec(vec![1.0, 2.0, 3.0]); assert!(dataset.set_sample(0, wrong_sample.view()).is_err());
let sample = Array1::from_vec(vec![1.0, 2.0]);
assert!(dataset.set_sample(10, sample.view()).is_err());
}
}