use crate::data::BatchDataSet;
use crate::tensor::{Device, Result, Tensor, TensorError};
pub const CLASS_NAMES: [&str; 10] = [
"airplane", "automobile", "bird", "cat", "deer",
"dog", "frog", "horse", "ship", "truck",
];
const PIXELS_PER_IMAGE: usize = 3 * 32 * 32; const BYTES_PER_RECORD: usize = 1 + PIXELS_PER_IMAGE; const IMAGES_PER_BATCH: usize = 10_000;
pub struct Cifar10 {
pub images: Tensor,
pub labels: Tensor,
}
impl Cifar10 {
pub fn parse(batches: &[&[u8]]) -> Result<Self> {
if batches.is_empty() {
return Err(TensorError::new("CIFAR-10: no batch data provided"));
}
let mut all_pixels: Vec<f32> = Vec::new();
let mut all_labels: Vec<i64> = Vec::new();
for (batch_idx, &batch) in batches.iter().enumerate() {
let expected = IMAGES_PER_BATCH * BYTES_PER_RECORD;
if batch.len() != expected {
return Err(TensorError::new(&format!(
"CIFAR-10 batch {}: expected {} bytes, got {}",
batch_idx, expected, batch.len()
)));
}
for img_idx in 0..IMAGES_PER_BATCH {
let offset = img_idx * BYTES_PER_RECORD;
let label = batch[offset] as i64;
if label > 9 {
return Err(TensorError::new(&format!(
"CIFAR-10 batch {} image {}: invalid label {}",
batch_idx, img_idx, label
)));
}
all_labels.push(label);
let pixel_start = offset + 1;
let pixel_end = pixel_start + PIXELS_PER_IMAGE;
for &b in &batch[pixel_start..pixel_end] {
all_pixels.push(b as f32 / 255.0);
}
}
}
let n = all_labels.len() as i64;
let images = Tensor::from_f32(&all_pixels, &[n, 3, 32, 32], Device::CPU)?;
let labels = Tensor::from_i64(&all_labels, &[n], Device::CPU)?;
Ok(Cifar10 { images, labels })
}
pub fn len(&self) -> usize {
self.images.shape()[0] as usize
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl BatchDataSet for Cifar10 {
fn len(&self) -> usize {
self.images.shape()[0] as usize
}
fn get_batch(&self, indices: &[usize]) -> Result<Vec<Tensor>> {
let idx: Vec<i64> = indices.iter().map(|&i| (i % self.len()) as i64).collect();
let idx_tensor = Tensor::from_i64(&idx, &[idx.len() as i64], Device::CPU)?;
let images = self.images.index_select(0, &idx_tensor)?;
let labels = self.labels.index_select(0, &idx_tensor)?;
Ok(vec![images, labels])
}
}
pub const TRAIN_BATCH_FILES: [&str; 5] = [
"data_batch_1.bin",
"data_batch_2.bin",
"data_batch_3.bin",
"data_batch_4.bin",
"data_batch_5.bin",
];
pub const TEST_BATCH_FILE: &str = "test_batch.bin";
pub struct Cifar10Disk {
files: Vec<crate::data::records::FixedStrideRecords>,
starts: Vec<usize>,
total: usize,
}
impl Cifar10Disk {
pub fn open<P: AsRef<std::path::Path>>(paths: &[P]) -> Result<Self> {
if paths.is_empty() {
return Err(TensorError::new("Cifar10Disk: no batch files provided"));
}
let mut files = Vec::with_capacity(paths.len());
let mut starts = Vec::with_capacity(paths.len());
let mut total = 0usize;
for path in paths {
let recs =
crate::data::records::FixedStrideRecords::open(path, BYTES_PER_RECORD)?;
starts.push(total);
total += recs.count();
files.push(recs);
}
Ok(Cifar10Disk { files, starts, total })
}
pub fn open_train(dir: impl AsRef<std::path::Path>) -> Result<Self> {
let dir = dir.as_ref();
let paths: Vec<_> = TRAIN_BATCH_FILES.iter().map(|f| dir.join(f)).collect();
Self::open(&paths)
}
pub fn open_test(dir: impl AsRef<std::path::Path>) -> Result<Self> {
Self::open(&[dir.as_ref().join(TEST_BATCH_FILE)])
}
}
impl crate::data::DataSet for Cifar10Disk {
fn len(&self) -> usize {
self.total
}
fn get(&self, index: usize) -> Result<Vec<Tensor>> {
if index >= self.total {
return Err(TensorError::new(&format!(
"Cifar10Disk: sample {index} out of bounds ({} samples)",
self.total
)));
}
let file_idx = self
.starts
.iter()
.rposition(|&start| start <= index)
.expect("starts[0] == 0 covers every index");
let record = self.files[file_idx].record(index - self.starts[file_idx])?;
let label = record[0] as i64;
if label > 9 {
return Err(TensorError::new(&format!(
"Cifar10Disk: sample {index} in {} has invalid label {label}",
self.files[file_idx].path().display()
)));
}
let mut pixels = Vec::with_capacity(PIXELS_PER_IMAGE);
for &b in &record[1..] {
pixels.push(b as f32 / 255.0);
}
let image = Tensor::from_f32(&pixels, &[3, 32, 32], Device::CPU)?;
let label = Tensor::from_i64(&[label], &[], Device::CPU)?;
Ok(vec![image, label])
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_batch(n: usize) -> Vec<u8> {
let mut buf = Vec::with_capacity(n * BYTES_PER_RECORD);
for i in 0..n {
buf.push((i % 10) as u8); for _ in 0..1024 {
buf.push((i % 256) as u8);
}
buf.extend_from_slice(&[0u8; 1024]);
buf.extend_from_slice(&[255u8; 1024]);
}
buf
}
#[test]
fn parse_single_batch() {
let batch = make_batch(IMAGES_PER_BATCH);
let cifar = Cifar10::parse(&[&batch]).unwrap();
assert_eq!(cifar.images.shape(), &[10000, 3, 32, 32]);
assert_eq!(cifar.labels.shape(), &[10000]);
let l = cifar.labels.select(0, 0).unwrap().to_i64_vec().unwrap()[0];
assert_eq!(l, 0);
let l = cifar.labels.select(0, 1).unwrap().to_i64_vec().unwrap()[0];
assert_eq!(l, 1);
}
#[test]
fn parse_multiple_batches() {
let b1 = make_batch(IMAGES_PER_BATCH);
let b2 = make_batch(IMAGES_PER_BATCH);
let cifar = Cifar10::parse(&[&b1, &b2]).unwrap();
assert_eq!(cifar.images.shape(), &[20000, 3, 32, 32]);
}
#[test]
fn wrong_size_rejected() {
let batch = [0u8; 100]; assert!(Cifar10::parse(&[&batch[..]]).is_err());
}
#[test]
fn pixel_normalization() {
let batch = make_batch(IMAGES_PER_BATCH);
let cifar = Cifar10::parse(&[&batch]).unwrap();
let img0 = cifar.images.select(0, 0).unwrap();
let r_pixel: f64 = img0.select(0, 0).unwrap() .select(0, 0).unwrap() .select(0, 0).unwrap() .item().unwrap();
assert!((r_pixel - 0.0).abs() < 1e-6);
let b_pixel: f64 = img0.select(0, 2).unwrap() .select(0, 0).unwrap()
.select(0, 0).unwrap()
.item().unwrap();
assert!((b_pixel - 1.0).abs() < 1e-6);
}
fn write_batches(name: &str, batches: &[Vec<u8>]) -> Vec<std::path::PathBuf> {
let dir = std::env::temp_dir().join("flodl-cifar10-disk-tests");
std::fs::create_dir_all(&dir).unwrap();
batches
.iter()
.enumerate()
.map(|(i, b)| {
let path = dir.join(format!("{name}-{}-{i}.bin", std::process::id()));
std::fs::write(&path, b).unwrap();
path
})
.collect()
}
#[test]
fn disk_matches_parsed() {
use crate::data::DataSet;
let b1 = make_batch(IMAGES_PER_BATCH);
let b2 = make_batch(IMAGES_PER_BATCH);
let parsed = Cifar10::parse(&[&b1, &b2]).unwrap();
let paths = write_batches("match", &[b1, b2]);
let disk = Cifar10Disk::open(&paths).unwrap();
assert_eq!(DataSet::len(&disk), 2 * IMAGES_PER_BATCH);
for &i in &[0usize, 9_999, 10_000, 10_005, 19_999] {
let sample = disk.get(i).unwrap();
assert_eq!(sample[0].shape(), &[3, 32, 32]);
let bulk_img = parsed.images.select(0, i as i64).unwrap();
let diff: f64 = sample[0].sub(&bulk_img).unwrap().abs().unwrap().sum().unwrap().item().unwrap();
assert_eq!(diff, 0.0);
let bulk_label: f64 = parsed.labels.select(0, i as i64).unwrap().item().unwrap();
let disk_label: f64 = sample[1].item().unwrap();
assert_eq!(disk_label, bulk_label);
assert_eq!(sample[1].shape(), &[] as &[i64]);
}
for p in paths {
std::fs::remove_file(p).unwrap();
}
}
#[test]
fn disk_bounds_and_bad_label_error() {
use crate::data::DataSet;
let mut bad = make_batch(2);
bad[0] = 11; let paths = write_batches("bad", &[bad]);
let disk = Cifar10Disk::open(&paths).unwrap();
assert!(disk.get(2).is_err());
let err = disk.get(0).unwrap_err();
assert!(err.to_string().contains("invalid label"));
for p in paths {
std::fs::remove_file(p).unwrap();
}
}
}