Data loading

Classes

class image_classification_tools.pytorch.data.DataPipeline(data_source, data_dir=None, split='train/val/test', val_size=10000, test_size=10000, batch_size=128, num_workers=0, train_transform=None, eval_transform=None, shuffle_train=True, preload=None, n_augmentations=5, augmented_dataset_name=None, pil_augmentations=None, tensor_augmentations=None, force_regenerate=False, seed=42, **loader_kwargs)[source]

Unified data loading pipeline with auto-detection and intelligent splitting.

Data augmentation is decoupled from training: augmented datasets are pregenerated and saved to disk, then loaded during training. The preload parameter controls whether data is loaded lazily from disk, preloaded to CPU, or preloaded to GPU.

Examples

Basic usage without augmentation, GPU preloading:

>>> loaders = DataPipeline(
...     data_source=datasets.CIFAR10,
...     data_dir='./data/pytorch/cifar10',
...     split='train/val/test',
...     train_transform=my_transform,
...     eval_transform=my_transform,
...     preload='gpu'
... ).get_loaders()
>>> train_loader = loaders.train

With augmentation (pregenerated and saved to disk):

>>> loaders = DataPipeline(
...     data_source=datasets.CIFAR10,
...     data_dir='./data/pytorch/cifar10',
...     split='train/val/test',
...     train_transform=eval_transform,
...     eval_transform=eval_transform,
...     preload='cpu',
...     n_augmentations=5,
...     augmented_dataset_name='strong_aug_v1',  # Saved to ./data/pytorch/augmented_cifar10/strong_aug_v1/
...     pil_augmentations=my_pil_augs
... ).get_loaders()
__init__(data_source, data_dir=None, split='train/val/test', val_size=10000, test_size=10000, batch_size=128, num_workers=0, train_transform=None, eval_transform=None, shuffle_train=True, preload=None, n_augmentations=5, augmented_dataset_name=None, pil_augmentations=None, tensor_augmentations=None, force_regenerate=False, seed=42, **loader_kwargs)[source]

Initialize DataPipeline.

Parameters:
  • data_source (type | str | Path) – PyTorch dataset class (e.g., datasets.CIFAR10) or path to data directory

  • split (str) – Desired split outcome. Must be one of: ‘train’, ‘train/val’, ‘train/test’, ‘train/val/test’

  • data_dir (Union[str, Path, None]) – Root directory for dataset storage (required for PyTorch datasets)

  • train_transform (Optional[Compose]) – Transform for training data (optional, defaults to eval_transform if not provided)

  • eval_transform (Optional[Compose]) – Transform for validation/test data (optional, defaults to train_transform if not provided)

  • preload (Optional[str]) – Loading strategy: ‘gpu’, ‘cpu’, or None for lazy loading from disk

  • pil_augmentations (Optional[Compose]) – PIL augmentation transforms (flip, rotate, etc.). If provided, augmented data will be pregenerated.

  • tensor_augmentations (Optional[Compose]) – Tensor augmentation transforms (blur, erasing, etc.). Applied during pregeneration.

  • n_augmentations (int) – Number of augmented copies per image (default: 5)

  • augmented_dataset_name (Optional[str]) – Name for augmented dataset directory. Defaults to ‘depth_{n_augmentations}’. Saved to {data_dir.parent}/augmented_{data_dir.name}/{augmented_dataset_name}/

  • val_size (int) – Number of validation samples (default: 10000)

  • test_size (int) – Number of test samples for 3-way splits (default: 10000)

  • batch_size (int) – Batch size for all loaders (default: 128)

  • num_workers (int) – Number of workers for DataLoader and augmentation (default: 0)

  • seed (int) – Random seed for reproducible splits (default: 42)

  • shuffle_train (bool) – Whether to shuffle training data (default: True)

  • force_regenerate (bool) – Force regeneration of cached augmented data (default: False)

  • **loader_kwargs – Additional arguments passed to DataLoader

Raises:

ValueError – If split format invalid or incompatible config

get_loaders()[source]

Create and return DataLoaders based on configuration.

Return type:

DataLoaders

Returns:

DataLoaders object with train/val/test loaders and metadata

static compute_dataset_stats(data_source, data_dir, num_samples=5000)[source]

Compute mean and std for dataset normalization.

Parameters:
  • data_source (type) – PyTorch dataset class (e.g., datasets.CIFAR10)

  • data_dir (str | Path) – Root directory for dataset storage

  • num_samples (int) – Number of samples to use for computation (default: 5000)

Return type:

Tuple[Tuple[float, float, float], Tuple[float, float, float]]

Returns:

Tuple of (mean, std) where each is (R, G, B) tuple

Example

>>> mean, std = DataPipeline.compute_dataset_stats(
...     datasets.CIFAR10,
...     data_dir='./data/pytorch/cifar10',
...     num_samples=5000
... )
>>> print(f"Mean: {mean}, Std: {std}")
class image_classification_tools.pytorch.data.DataLoaders(train, val, test, batch_size, train_size, val_size, test_size, device)[source]

Container for train/val/test DataLoaders with convenience methods.

train

Training DataLoader (or None if not in split)

val

Validation DataLoader (or None if not in split)

test

Test DataLoader (or None if not in split)

batch_size

Batch size used for all loaders

train_size

Number of training samples

val_size

Number of validation samples

test_size

Number of test samples

device

Device where data is loaded (‘cpu’, ‘cuda’, or None for lazy)

train: DataLoader | None
val: DataLoader | None
test: DataLoader | None
batch_size: int
train_size: int
val_size: int
test_size: int
device: str | None
get_batch_sizes()[source]

Get batch sizes for each split.

Return type:

Dict[str, int]

Returns:

Dictionary with batch sizes for available splits

total_samples()[source]

Get total sample counts for each split.

Return type:

Dict[str, int]

Returns:

Dictionary with sample counts for available splits

memory_estimate()[source]

Estimate memory usage for preloaded data.

Return type:

float

Returns:

Estimated memory use in GB as float” Returns None if data not preloaded

__init__(train, val, test, batch_size, train_size, val_size, test_size, device)

Overview

The data module provides a unified DataPipeline class that handles all data loading and preparation in a single call. The pipeline automatically detects dataset structure, performs intelligent splitting, and handles augmentation with pregeneration and caching.

Key features:

  • Auto-detection: Automatically detects if data source has pre-made train/test splits

  • Outcome-based: User specifies desired outcome (e.g., split='train/val/test'), pipeline determines the how

  • Intelligent splitting: Performs minimal operations based on source structure and desired outcome

  • Pregenerated augmentation: Augmented data is generated once and saved to disk for reuse

  • Memory optimization: GPU/CPU preloading or lazy loading based on use case

  • Type-safe: Returns frozen DataLoaders object with .train, .val, .test attributes

  • Smart caching: Reuses pregenerated augmentation across training runs

  • Dataset statistics: Built-in method to compute mean/std for normalization

Example usage

Basic workflow (CIFAR-10 with GPU preloading):

from torchvision import datasets, transforms
from image_classification_tools.pytorch import DataPipeline

# Define transform
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# Create pipeline and get loaders in one call
loaders = DataPipeline(
    data_source=datasets.CIFAR10,
    data_dir='./data/pytorch/cifar10',  # Will download if not present
    split='train/val/test',
    val_size=10000,
    batch_size=128,
    train_transform=transform,
    eval_transform=transform,
    preload='gpu',
    seed=42
).get_loaders()

# Access loaders via attributes
train_loader = loaders.train
val_loader = loaders.val
test_loader = loaders.test

# Display pipeline summary
print(loaders.total_samples())
# Output: {'train': 40000, 'val': 10000, 'test': 10000}

print(loaders.memory_estimate())
# Output: 2.3 (GB)

With data augmentation (pregenerated):

from image_classification_tools.pytorch import DataPipeline

# Define augmentation transforms
pil_augmentations = transforms.Compose([
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.RandomRotation(15),
    transforms.ColorJitter(brightness=0.2, contrast=0.2)
])

tensor_augmentations = transforms.Compose([
    transforms.RandomErasing(p=0.2, scale=(0.02, 0.1))
])

# Define base transforms
base_transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# Create pipeline with augmentation
# Augmented data is automatically pregenerated and saved to disk
loaders = DataPipeline(
    data_source=datasets.CIFAR10,
    data_dir='./data/pytorch/cifar10',
    split='train/val/test',
    batch_size=128,
    train_transform=base_transform,
    eval_transform=base_transform,
    preload='gpu',  # Preload augmented data to GPU for fast training
    n_augmentations=5,  # Create 5 augmented copies per image
    augmented_dataset_name='strong_aug_v1',  # Optional: defaults to 'depth_5'
    pil_augmentations=pil_augmentations,
    tensor_augmentations=tensor_augmentations
).get_loaders()

# Augmented data saved to: ./data/pytorch/augmented_cifar10/strong_aug_v1/
# Subsequent runs with same augmented_dataset_name load from cache
# Use force_regenerate=True to regenerate cached data

Computing dataset statistics:

from image_classification_tools.pytorch import DataPipeline

# Compute mean and std for normalization
mean, std = DataPipeline.compute_dataset_stats(
    data_source=datasets.CIFAR10,
    data_dir='./data/pytorch/cifar10',
    num_samples=5000
)
print(f'Mean: {mean}')  # (0.4914, 0.4822, 0.4465)
print(f'Std: {std}')    # (0.2470, 0.2435, 0.2616)

# Use computed values in transform
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean=mean, std=std)
])

Custom datasets (directory-based):

# Pipeline auto-detects directory structure
loaders = DataPipeline(
    data_source='./my_dataset',  # Path to dataset directory
    data_dir='./my_dataset',
    split='train/val/test',
    val_size=5000,
    batch_size=64,
    train_transform=transform,
    eval_transform=transform,
    preload='cpu'
).get_loaders()