use scirs2_core::ndarray::{ArrayView1, ArrayView2, ArrayViewMut1, ArrayViewMut2, Axis, Slice, s};
use std::marker::PhantomData;
use std::ops::Range;
use thiserror::Error;
#[derive(Error, Debug)]
pub enum ZeroCopyError {
#[error("Index out of bounds: {index} >= {len}")]
IndexOutOfBounds { index: usize, len: usize },
#[error("Range out of bounds: {start}..{end} exceeds {len}")]
RangeOutOfBounds {
start: usize,
end: usize,
len: usize,
},
#[error("Dimension mismatch: expected {expected}, got {actual}")]
DimensionMismatch { expected: String, actual: String },
#[error("Invalid slice parameters: {0}")]
InvalidSlice(String),
#[error("Alignment error: {0}")]
Alignment(String),
}
pub type ZeroCopyResult<T> = Result<T, ZeroCopyError>;
pub struct DatasetView<'a> {
features: ArrayView2<'a, f64>,
targets: Option<ArrayView1<'a, f64>>,
feature_names: Option<&'a [String]>,
sample_indices: Option<&'a [usize]>,
}
impl<'a> DatasetView<'a> {
pub fn new(features: ArrayView2<'a, f64>, targets: Option<ArrayView1<'a, f64>>) -> Self {
Self {
features,
targets,
feature_names: None,
sample_indices: None,
}
}
pub fn with_feature_names(
features: ArrayView2<'a, f64>,
targets: Option<ArrayView1<'a, f64>>,
feature_names: &'a [String],
) -> ZeroCopyResult<Self> {
if feature_names.len() != features.ncols() {
return Err(ZeroCopyError::DimensionMismatch {
expected: format!("{} feature names", features.ncols()),
actual: format!("{} feature names", feature_names.len()),
});
}
Ok(Self {
features,
targets,
feature_names: Some(feature_names),
sample_indices: None,
})
}
pub fn n_samples(&self) -> usize {
self.features.nrows()
}
pub fn n_features(&self) -> usize {
self.features.ncols()
}
pub fn shape(&self) -> (usize, usize) {
(self.n_samples(), self.n_features())
}
pub fn features(&self) -> ArrayView2<'a, f64> {
self.features
}
pub fn targets(&self) -> Option<ArrayView1<'a, f64>> {
self.targets
}
pub fn feature_names(&self) -> Option<&[String]> {
self.feature_names
}
pub fn sample(&self, index: usize) -> ZeroCopyResult<ArrayView1<'a, f64>> {
if index >= self.n_samples() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_samples(),
});
}
Ok(self.features.row(index))
}
pub fn samples(&self, range: Range<usize>) -> ZeroCopyResult<ArrayView2<'a, f64>> {
if range.end > self.n_samples() {
return Err(ZeroCopyError::RangeOutOfBounds {
start: range.start,
end: range.end,
len: self.n_samples(),
});
}
Ok(self.features.slice(s![range, ..]))
}
pub fn feature(&self, index: usize) -> ZeroCopyResult<ArrayView1<'a, f64>> {
if index >= self.n_features() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_features(),
});
}
Ok(self.features.column(index))
}
pub fn features_range(&self, range: Range<usize>) -> ZeroCopyResult<ArrayView2<'a, f64>> {
if range.end > self.n_features() {
return Err(ZeroCopyError::RangeOutOfBounds {
start: range.start,
end: range.end,
len: self.n_features(),
});
}
Ok(self.features.slice(s![.., range]))
}
pub fn subview(
&self,
sample_range: Range<usize>,
feature_range: Range<usize>,
) -> ZeroCopyResult<ArrayView2<'a, f64>> {
if sample_range.end > self.n_samples() {
return Err(ZeroCopyError::RangeOutOfBounds {
start: sample_range.start,
end: sample_range.end,
len: self.n_samples(),
});
}
if feature_range.end > self.n_features() {
return Err(ZeroCopyError::RangeOutOfBounds {
start: feature_range.start,
end: feature_range.end,
len: self.n_features(),
});
}
Ok(self.features.slice(s![sample_range, feature_range]))
}
pub fn targets_range(
&self,
range: Range<usize>,
) -> ZeroCopyResult<Option<ArrayView1<'a, f64>>> {
if let Some(targets) = self.targets {
if range.end > targets.len() {
return Err(ZeroCopyError::RangeOutOfBounds {
start: range.start,
end: range.end,
len: targets.len(),
});
}
Ok(Some(targets.slice(s![range])))
} else {
Ok(None)
}
}
pub fn filter<F>(&self, predicate: F) -> ZeroCopyResult<FilteredDatasetView<'a>>
where
F: Fn(ArrayView1<f64>) -> bool,
{
let mut selected_indices = Vec::new();
for (i, sample) in self.features.axis_iter(Axis(0)).enumerate() {
if predicate(sample) {
selected_indices.push(i);
}
}
Ok(FilteredDatasetView {
original: self,
indices: selected_indices,
})
}
pub fn select_features(&self, indices: &[usize]) -> ZeroCopyResult<SelectedFeaturesView<'a>> {
for &idx in indices {
if idx >= self.n_features() {
return Err(ZeroCopyError::IndexOutOfBounds {
index: idx,
len: self.n_features(),
});
}
}
Ok(SelectedFeaturesView {
original: self,
feature_indices: indices,
})
}
pub fn strided(&self, stride: usize, offset: usize) -> ZeroCopyResult<StridedDatasetView<'a>> {
if stride == 0 {
return Err(ZeroCopyError::InvalidSlice(
"Stride cannot be zero".to_string(),
));
}
if offset >= self.n_samples() {
return Err(ZeroCopyError::IndexOutOfBounds {
index: offset,
len: self.n_samples(),
});
}
Ok(StridedDatasetView {
original: self,
stride,
offset,
})
}
pub fn has_targets(&self) -> bool {
self.targets.is_some()
}
}
pub struct DatasetViewMut<'a> {
features: ArrayViewMut2<'a, f64>,
targets: Option<ArrayViewMut1<'a, f64>>,
feature_names: Option<&'a [String]>,
}
impl<'a> DatasetViewMut<'a> {
pub fn new(features: ArrayViewMut2<'a, f64>, targets: Option<ArrayViewMut1<'a, f64>>) -> Self {
Self {
features,
targets,
feature_names: None,
}
}
pub fn n_samples(&self) -> usize {
self.features.nrows()
}
pub fn n_features(&self) -> usize {
self.features.ncols()
}
pub fn sample_mut(&mut self, index: usize) -> ZeroCopyResult<ArrayViewMut1<f64>> {
if index >= self.n_samples() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_samples(),
});
}
Ok(self.features.row_mut(index))
}
pub fn feature_mut(&mut self, index: usize) -> ZeroCopyResult<ArrayViewMut1<f64>> {
if index >= self.n_features() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_features(),
});
}
Ok(self.features.column_mut(index))
}
pub fn targets_mut(&mut self) -> Option<ArrayViewMut1<f64>> {
self.targets.as_mut().map(|t| t.view_mut())
}
pub fn as_view(&self) -> DatasetView {
DatasetView::new(
self.features.view(),
self.targets.as_ref().map(|t| t.view()),
)
}
}
pub struct FilteredDatasetView<'a> {
original: &'a DatasetView<'a>,
indices: Vec<usize>,
}
impl<'a> FilteredDatasetView<'a> {
pub fn n_samples(&self) -> usize {
self.indices.len()
}
pub fn n_features(&self) -> usize {
self.original.n_features()
}
pub fn sample(&self, index: usize) -> ZeroCopyResult<ArrayView1<f64>> {
if index >= self.indices.len() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.indices.len(),
});
}
let original_index = self.indices[index];
self.original.sample(original_index)
}
pub fn indices(&self) -> &[usize] {
&self.indices
}
pub fn target(&self, index: usize) -> ZeroCopyResult<Option<f64>> {
if index >= self.indices.len() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.indices.len(),
});
}
let original_index = self.indices[index];
if let Some(targets) = self.original.targets() {
Ok(Some(targets[original_index]))
} else {
Ok(None)
}
}
}
pub struct SelectedFeaturesView<'a> {
original: &'a DatasetView<'a>,
feature_indices: &'a [usize],
}
impl<'a> SelectedFeaturesView<'a> {
pub fn n_samples(&self) -> usize {
self.original.n_samples()
}
pub fn n_features(&self) -> usize {
self.feature_indices.len()
}
pub fn sample(&self, index: usize) -> ZeroCopyResult<Vec<f64>> {
let original_sample = self.original.sample(index)?;
let selected: Vec<f64> = self
.feature_indices
.iter()
.map(|&feat_idx| original_sample[feat_idx])
.collect();
Ok(selected)
}
pub fn feature_indices(&self) -> &[usize] {
self.feature_indices
}
pub fn target(&self, index: usize) -> ZeroCopyResult<Option<f64>> {
if let Some(targets) = self.original.targets() {
if index >= targets.len() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: targets.len(),
});
}
Ok(Some(targets[index]))
} else {
Ok(None)
}
}
}
pub struct StridedDatasetView<'a> {
original: &'a DatasetView<'a>,
stride: usize,
offset: usize,
}
impl<'a> StridedDatasetView<'a> {
pub fn n_samples(&self) -> usize {
if self.offset >= self.original.n_samples() {
0
} else {
(self.original.n_samples() - self.offset + self.stride - 1) / self.stride
}
}
pub fn n_features(&self) -> usize {
self.original.n_features()
}
pub fn sample(&self, index: usize) -> ZeroCopyResult<ArrayView1<f64>> {
if index >= self.n_samples() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_samples(),
});
}
let original_index = self.offset + index * self.stride;
self.original.sample(original_index)
}
pub fn original_index(&self, strided_index: usize) -> Option<usize> {
if strided_index < self.n_samples() {
Some(self.offset + strided_index * self.stride)
} else {
None
}
}
pub fn stride_info(&self) -> (usize, usize) {
(self.stride, self.offset)
}
}
pub struct DatasetSampleIterator<'a> {
view: &'a DatasetView<'a>,
current: usize,
}
impl<'a> DatasetSampleIterator<'a> {
fn new(view: &'a DatasetView<'a>) -> Self {
Self { view, current: 0 }
}
}
impl<'a> Iterator for DatasetSampleIterator<'a> {
type Item = ZeroCopyResult<ArrayView1<'a, f64>>;
fn next(&mut self) -> Option<Self::Item> {
if self.current < self.view.n_samples() {
let result = self.view.sample(self.current);
self.current += 1;
Some(result)
} else {
None
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
let remaining = self.view.n_samples().saturating_sub(self.current);
(remaining, Some(remaining))
}
}
impl<'a> ExactSizeIterator for DatasetSampleIterator<'a> {}
impl<'a> DatasetView<'a> {
pub fn samples_iter(&self) -> DatasetSampleIterator {
DatasetSampleIterator::new(self)
}
}
pub struct BatchIterator<'a> {
view: &'a DatasetView<'a>,
batch_size: usize,
current: usize,
}
impl<'a> BatchIterator<'a> {
fn new(view: &'a DatasetView<'a>, batch_size: usize) -> Self {
Self {
view,
batch_size,
current: 0,
}
}
}
impl<'a> Iterator for BatchIterator<'a> {
type Item = ZeroCopyResult<ArrayView2<'a, f64>>;
fn next(&mut self) -> Option<Self::Item> {
if self.current >= self.view.n_samples() {
return None;
}
let end = (self.current + self.batch_size).min(self.view.n_samples());
let range = self.current..end;
self.current = end;
Some(self.view.samples(range))
}
}
impl<'a> DatasetView<'a> {
pub fn batches(&self, batch_size: usize) -> BatchIterator {
BatchIterator::new(self, batch_size)
}
}
pub struct WindowView<'a> {
view: &'a DatasetView<'a>,
window_size: usize,
step: usize,
}
impl<'a> WindowView<'a> {
pub fn new(view: &'a DatasetView<'a>, window_size: usize, step: usize) -> ZeroCopyResult<Self> {
if window_size == 0 {
return Err(ZeroCopyError::InvalidSlice(
"Window size cannot be zero".to_string(),
));
}
if step == 0 {
return Err(ZeroCopyError::InvalidSlice(
"Step cannot be zero".to_string(),
));
}
Ok(Self {
view,
window_size,
step,
})
}
pub fn n_windows(&self) -> usize {
if self.view.n_samples() < self.window_size {
0
} else {
(self.view.n_samples() - self.window_size) / self.step + 1
}
}
pub fn window(&self, index: usize) -> ZeroCopyResult<ArrayView2<'a, f64>> {
if index >= self.n_windows() {
return Err(ZeroCopyError::IndexOutOfBounds {
index,
len: self.n_windows(),
});
}
let start = index * self.step;
let end = start + self.window_size;
self.view.samples(start..end)
}
}
impl<'a> DatasetView<'a> {
pub fn windows(&self, window_size: usize, step: usize) -> ZeroCopyResult<WindowView<'a>> {
WindowView::new(self, window_size, step)
}
}
macro_rules! s {
($($x:expr),*) => {
($(Slice::from($x),)*)
};
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array;
#[test]
fn test_dataset_view_basic() {
let features = Array::from_shape_vec((6, 3), (0..18).map(|x| x as f64).collect()).expect("shape and data length should match");
let targets = Array::from_vec((0..6).map(|x| x as f64).collect());
let view = DatasetView::new(features.view(), Some(targets.view()));
assert_eq!(view.n_samples(), 6);
assert_eq!(view.n_features(), 3);
assert_eq!(view.shape(), (6, 3));
assert!(view.has_targets());
let sample = view.sample(2).expect("sampling should succeed");
assert_eq!(sample[0], 6.0); assert_eq!(sample[1], 7.0); assert_eq!(sample[2], 8.0);
let feature = view.feature(1).expect("operation should succeed");
assert_eq!(feature[0], 1.0); assert_eq!(feature[2], 7.0);
let targets_view = view.targets().expect("operation should succeed");
assert_eq!(targets_view[2], 2.0);
}
#[test]
fn test_dataset_view_ranges() {
let features = Array::from_shape_vec((10, 4), (0..40).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let samples = view.samples(2..5).expect("sampling should succeed");
assert_eq!(samples.dim(), (3, 4));
assert_eq!(samples[[0, 0]], 8.0);
let features_range = view.features_range(1..3).expect("operation should succeed");
assert_eq!(features_range.dim(), (10, 2));
assert_eq!(features_range[[0, 0]], 1.0);
let subview = view.subview(1..4, 1..3).expect("operation should succeed");
assert_eq!(subview.dim(), (3, 2));
assert_eq!(subview[[0, 0]], 5.0); }
#[test]
fn test_filtered_view() {
let features = Array::from_shape_vec(
(5, 2),
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, ],
)
.expect("operation should succeed");
let view = DatasetView::new(features.view(), None);
let filtered = view.filter(|sample| sample.sum() > 10.0).expect("sampling should succeed");
assert_eq!(filtered.n_samples(), 3); assert_eq!(filtered.n_features(), 2);
let sample0 = filtered.sample(0).expect("sampling should succeed");
assert_eq!(sample0[0], 5.0); assert_eq!(sample0[1], 6.0);
let indices = filtered.indices();
assert_eq!(indices, &[2, 3, 4]);
}
#[test]
fn test_selected_features_view() {
let features = Array::from_shape_vec((3, 4), (0..12).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let selected = view.select_features(&[0, 2]).expect("operation should succeed");
assert_eq!(selected.n_samples(), 3);
assert_eq!(selected.n_features(), 2);
let sample0 = selected.sample(0).expect("sampling should succeed");
assert_eq!(sample0, vec![0.0, 2.0]);
let sample1 = selected.sample(1).expect("sampling should succeed");
assert_eq!(sample1, vec![4.0, 6.0]); }
#[test]
fn test_strided_view() {
let features = Array::from_shape_vec((10, 2), (0..20).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let strided = view.strided(3, 1).expect("operation should succeed");
assert_eq!(strided.n_samples(), 3); assert_eq!(strided.n_features(), 2);
let sample0 = strided.sample(0).expect("sampling should succeed");
assert_eq!(sample0[0], 2.0);
let sample1 = strided.sample(1).expect("sampling should succeed");
assert_eq!(sample1[0], 8.0);
assert_eq!(strided.original_index(0), Some(1));
assert_eq!(strided.original_index(1), Some(4));
assert_eq!(strided.original_index(2), Some(7));
}
#[test]
fn test_sample_iterator() {
let features = Array::from_shape_vec((3, 2), (0..6).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let mut iter = view.samples_iter();
let sample0 = iter.next().expect("sampling should succeed").expect("sampling should succeed");
assert_eq!(sample0[0], 0.0);
assert_eq!(sample0[1], 1.0);
let sample1 = iter.next().expect("sampling should succeed").expect("sampling should succeed");
assert_eq!(sample1[0], 2.0);
assert_eq!(sample1[1], 3.0);
let sample2 = iter.next().expect("sampling should succeed").expect("sampling should succeed");
assert_eq!(sample2[0], 4.0);
assert_eq!(sample2[1], 5.0);
assert!(iter.next().is_none());
}
#[test]
fn test_batch_iterator() {
let features = Array::from_shape_vec((10, 2), (0..20).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let mut iter = view.batches(3);
let batch0 = iter.next().expect("operation should succeed").expect("operation should succeed");
assert_eq!(batch0.dim(), (3, 2));
let batch1 = iter.next().expect("operation should succeed").expect("operation should succeed");
assert_eq!(batch1.dim(), (3, 2));
let batch2 = iter.next().expect("operation should succeed").expect("operation should succeed");
assert_eq!(batch2.dim(), (3, 2));
let batch3 = iter.next().expect("operation should succeed").expect("operation should succeed");
assert_eq!(batch3.dim(), (1, 2));
assert!(iter.next().is_none());
}
#[test]
fn test_window_view() {
let features = Array::from_shape_vec((10, 2), (0..20).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
let windows = view.windows(3, 2).expect("operation should succeed");
assert_eq!(windows.n_windows(), 4);
let window0 = windows.window(0).expect("operation should succeed");
assert_eq!(window0.dim(), (3, 2));
let window1 = windows.window(1).expect("operation should succeed");
assert_eq!(window1.dim(), (3, 2));
assert_eq!(window0[[0, 0]], 0.0); assert_eq!(window1[[0, 0]], 4.0); }
#[test]
fn test_mutable_dataset_view() {
let mut features =
Array::from_shape_vec((3, 2), (0..6).map(|x| x as f64).collect()).expect("shape and data length should match");
let mut targets = Array::from_vec(vec![10.0, 20.0, 30.0]);
let mut view = DatasetViewMut::new(features.view_mut(), Some(targets.view_mut()));
assert_eq!(view.n_samples(), 3);
assert_eq!(view.n_features(), 2);
{
let mut sample = view.sample_mut(1).expect("sampling should succeed");
sample[0] = 99.0;
sample[1] = 88.0;
}
let immutable_view = view.as_view();
let sample = immutable_view.sample(1).expect("sampling should succeed");
assert_eq!(sample[0], 99.0);
assert_eq!(sample[1], 88.0);
if let Some(mut targets_mut) = view.targets_mut() {
targets_mut[1] = 999.0;
}
let targets_view = immutable_view.targets().expect("operation should succeed");
assert_eq!(targets_view[1], 999.0);
}
#[test]
fn test_error_handling() {
let features = Array::from_shape_vec((3, 2), (0..6).map(|x| x as f64).collect()).expect("shape and data length should match");
let view = DatasetView::new(features.view(), None);
assert!(view.sample(10).is_err());
assert!(view.feature(10).is_err());
assert!(view.samples(0..10).is_err());
assert!(view.features_range(0..10).is_err());
assert!(view.strided(0, 0).is_err()); assert!(view.strided(1, 10).is_err()); }
}