use torsh_core::error::Result;
use torsh_data::dataloader::{simple_dataloader, simple_random_dataloader};
use torsh_data::prelude::*;
use torsh_tensor::creation::{ones, zeros};
#[test]
fn test_tensor_dataset_pipeline() -> Result<()> {
let data = ones::<f32>(&[10, 3])?;
let labels = zeros::<f32>(&[10])?;
let dataset = TensorDataset::from_tensors(vec![data, labels]);
assert_eq!(dataset.len(), 10);
let sampler = SequentialSampler::new(dataset.len());
let mut indices: Vec<usize> = sampler.iter().collect();
assert_eq!(indices, vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9]);
let random_sampler = RandomSampler::simple(dataset.len()).with_generator(42);
indices = random_sampler.iter().collect();
assert_eq!(indices.len(), 10);
indices.sort();
assert_eq!(indices, vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9]);
let batch_sampler = BatchingSampler::new(SequentialSampler::new(dataset.len()), 3, false);
let batches: Vec<Vec<usize>> = batch_sampler.iter().collect();
assert_eq!(batches.len(), 4); assert_eq!(batches[0], vec![0, 1, 2]);
assert_eq!(batches[3], vec![9]);
Ok(())
}
#[test]
fn test_dataloader_with_collation() -> Result<()> {
let data = ones::<f32>(&[8, 3])?;
let dataset = TensorDataset::from_tensor(data);
let dataloader = DataLoaderBuilder::new(dataset).batch_size(2).build()?;
let mut batch_count = 0;
for batch in dataloader.iter() {
let batch = batch?;
assert_eq!(batch.len(), 1); assert_eq!(batch[0].shape().dims()[0], 2); assert_eq!(batch[0].shape().dims()[1], 3); batch_count += 1;
}
assert_eq!(batch_count, 4);
Ok(())
}
#[test]
fn test_weighted_sampling() -> Result<()> {
let data = ones::<f32>(&[5, 2])?;
let _dataset = TensorDataset::from_tensor(data);
let weights = vec![0.1, 0.1, 0.1, 0.1, 0.6];
let sampler = WeightedRandomSampler::new(weights, 100, true).with_generator(42);
let indices: Vec<usize> = sampler.iter().collect();
assert_eq!(indices.len(), 100);
let mut counts = [0; 5];
for &idx in &indices {
counts[idx] += 1;
}
assert!(counts[4] > counts[0]);
assert!(counts[4] > counts[1]);
assert!(counts[4] > counts[2]);
assert!(counts[4] > counts[3]);
Ok(())
}
#[test]
fn test_distributed_sampling() -> Result<()> {
let dataset_len = 100;
let num_replicas = 4;
let sampler0 = DistributedSampler::new(dataset_len, num_replicas, 0, false).with_generator(42);
let sampler1 = DistributedSampler::new(dataset_len, num_replicas, 1, false).with_generator(42);
let sampler2 = DistributedSampler::new(dataset_len, num_replicas, 2, false).with_generator(42);
let sampler3 = DistributedSampler::new(dataset_len, num_replicas, 3, false).with_generator(42);
let indices0: Vec<usize> = sampler0.iter().collect();
let indices1: Vec<usize> = sampler1.iter().collect();
let indices2: Vec<usize> = sampler2.iter().collect();
let indices3: Vec<usize> = sampler3.iter().collect();
assert_eq!(indices0.len(), 25);
assert_eq!(indices1.len(), 25);
assert_eq!(indices2.len(), 25);
assert_eq!(indices3.len(), 25);
let mut all_indices = indices0.clone();
all_indices.extend(indices1);
all_indices.extend(indices2);
all_indices.extend(indices3);
all_indices.sort();
let expected: Vec<usize> = (0..100).collect();
assert_eq!(all_indices, expected);
Ok(())
}
#[test]
fn test_stratified_sampling() -> Result<()> {
let labels = vec![0, 0, 0, 0, 0, 1, 1, 2];
let sampler = StratifiedSampler::new(&labels, 8, false).with_generator(42);
assert_eq!(sampler.len(), 8);
assert_eq!(sampler.num_strata(), 3);
let indices: Vec<usize> = sampler.iter().collect();
assert_eq!(indices.len(), 8);
let mut class_counts = [0; 3];
for &idx in &indices {
class_counts[labels[idx]] += 1;
}
assert_eq!(class_counts[0], 5);
assert_eq!(class_counts[1], 2);
assert_eq!(class_counts[2], 1);
Ok(())
}
#[test]
fn test_concat_dataset() -> Result<()> {
let ds1 = TensorDataset::from_tensor(ones::<f32>(&[5, 3])?);
let ds2 = TensorDataset::from_tensor(zeros::<f32>(&[3, 3])?);
let concat = ConcatDataset::new(vec![ds1, ds2]);
assert_eq!(concat.len(), 8);
let item0 = concat.get(0)?;
let item5 = concat.get(5)?;
assert_eq!(item0.len(), 1);
assert_eq!(item5.len(), 1);
Ok(())
}
#[test]
fn test_subset_dataset() -> Result<()> {
let dataset = TensorDataset::from_tensor(ones::<f32>(&[10, 3])?);
let subset = Subset::new(dataset, vec![0, 2, 4, 6, 8]);
assert_eq!(subset.len(), 5);
for i in 0..5 {
let item = subset.get(i)?;
assert_eq!(item.len(), 1);
}
assert!(subset.get(5).is_err());
Ok(())
}
#[test]
fn test_text_dataset() -> Result<()> {
let texts = vec![
"This is a positive example".to_string(),
"This is a negative example".to_string(),
"Another positive text".to_string(),
"Another negative text".to_string(),
];
let labels = vec![1, 0, 1, 0];
let dataset = TextClassificationDataset::new(texts, labels)?;
assert_eq!(dataset.len(), 4);
assert_eq!(dataset.num_classes(), 2);
let (tensor, label) = dataset.get(0)?;
assert_eq!(label, 1);
assert!(tensor.ndim() > 0);
let vocab = dataset.vocabulary();
assert!(!vocab.is_empty());
let text = "test text";
let ids = vocab.encode(text);
let decoded = vocab.decode(&ids);
assert!(!ids.is_empty());
assert!(!decoded.is_empty());
Ok(())
}
#[test]
fn test_dynamic_batching() -> Result<()> {
let t1 = ones::<f32>(&[3, 4])?; let t2 = ones::<f32>(&[5, 4])?; let t3 = ones::<f32>(&[2, 4])?;
let batch = vec![t1, t2, t3];
let collator = DynamicBatchCollate::new(0.0f32).with_max_length(6);
let (padded_tensor, lengths) = collator.collate(batch)?;
assert_eq!(padded_tensor.shape().dims(), &[3, 5, 4]);
assert_eq!(lengths.shape().dims(), &[3]);
let length_data = lengths.to_vec()?;
assert_eq!(length_data, vec![3, 5, 2]);
Ok(())
}
#[test]
fn test_cached_dataset() -> Result<()> {
let base_dataset = TensorDataset::from_tensor(ones::<f32>(&[10, 3])?);
let cached = CachedDataset::new(base_dataset, 5);
assert_eq!(cached.len(), 10);
for _ in 0..3 {
let _ = cached.get(0)?;
let _ = cached.get(1)?;
let _ = cached.get(2)?;
}
let hit_rate = cached.cache_hit_rate();
assert!(hit_rate > 0.0);
Ok(())
}
#[test]
fn test_curriculum_sampler() -> Result<()> {
let difficulty_fn = |idx: usize| -> f64 { idx as f64 / 10.0 };
let mut sampler = CurriculumSampler::new(
10,
difficulty_fn,
5, CurriculumStrategy::Linear,
)
.with_generator(42);
sampler.set_epoch(0);
let indices_epoch0: Vec<usize> = sampler.iter().collect();
sampler.set_epoch(4);
let indices_epoch4: Vec<usize> = sampler.iter().collect();
assert!(indices_epoch4.len() >= indices_epoch0.len());
Ok(())
}
#[test]
fn test_active_learning_sampler() -> Result<()> {
let mut sampler = ActiveLearningSampler::new(
100, AcquisitionStrategy::UncertaintySampling,
10, )
.with_generator(42);
let uncertainties: Vec<f64> = (0..100).map(|i| (i as f64) / 100.0).collect();
sampler.update_uncertainties(uncertainties);
let selected: Vec<usize> = sampler.iter().collect();
assert_eq!(selected.len(), 10);
sampler.add_labeled_samples(&selected);
let selected2: Vec<usize> = sampler.iter().collect();
assert_eq!(selected2.len(), 10);
for &idx in &selected2 {
assert!(!selected.contains(&idx));
}
Ok(())
}
#[test]
fn test_importance_sampling() -> Result<()> {
let weights: Vec<f64> = (0..10).map(|i| 2.0_f64.powi(-i)).collect();
let sampler = ImportanceSampler::new(weights, 100, true)
.with_temperature(1.0)
.with_generator(42);
let indices: Vec<usize> = sampler.iter().collect();
assert_eq!(indices.len(), 100);
let mut counts = [0; 10];
for &idx in &indices {
counts[idx] += 1;
}
assert!(counts[0] > counts[9]);
assert!(counts[1] > counts[8]);
Ok(())
}
#[test]
fn test_complete_workflow() -> Result<()> {
let data = ones::<f32>(&[20, 5])?;
let labels = zeros::<f32>(&[20])?;
let dataset = TensorDataset::from_tensors(vec![data, labels]);
let splits = random_split(dataset, &[16, 4], Some(42))?;
let train_dataset = splits[0].clone();
let val_dataset = splits[1].clone();
assert_eq!(train_dataset.len(), 16);
assert_eq!(val_dataset.len(), 4);
let train_loader = simple_random_dataloader(train_dataset, 4, Some(42))?;
let val_loader = simple_dataloader(val_dataset, 2, false)?;
let mut train_batches = 0;
for batch in train_loader.iter() {
let batch = batch?;
assert_eq!(batch.len(), 2); assert!(batch[0].shape().dims()[0] <= 4); train_batches += 1;
}
assert_eq!(train_batches, 4);
let mut val_batches = 0;
for batch in val_loader.iter() {
let batch = batch?;
assert_eq!(batch.len(), 2); assert!(batch[0].shape().dims()[0] <= 2); val_batches += 1;
}
assert_eq!(val_batches, 2);
Ok(())
}