Skip to main content

Image Classification Tools

A lightweight PyTorch toolkit for building and training image classification models.

Overview

This package provides utilities for common image classification tasks:

  • Data loading: Unified DataPipeline for automatic data preparation with intelligent splitting
  • Model training: Training loops with progress tracking and validation
  • Evaluation: Accuracy metrics, confusion matrices, and performance analysis
  • Visualization: Learning curves, probability distributions, and evaluation plots
  • Hyperparameter optimization: Optuna integration for automated model tuning

Installation

pip install image-classification-tools

For hyperparameter optimization, install with the optuna extra:

pip install image-classification-tools[optuna]

Quick start

Basic usage

import torch
from torchvision import datasets, transforms
from image_classification_tools.pytorch import DataPipeline
from image_classification_tools.pytorch.training import train_model
from image_classification_tools.pytorch.evaluation import evaluate_model

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

# Create data pipeline (handles everything in one call)
loaders = DataPipeline(
    data_source=datasets.MNIST,
    data_dir='./data/pytorch/mnist',
    split='train/val/test',
    val_size=10000,
    batch_size=64,
    train_transform=transform,
    eval_transform=transform,
    preload='gpu'
).get_loaders()

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

# Display summary
print(loaders.split_info())
print(loaders.memory_estimate())

# Define model, criterion, optimizer
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

model = torch.nn.Sequential(
    torch.nn.Flatten(),
    torch.nn.Linear(784, 128),
    torch.nn.ReLU(),
    torch.nn.Linear(128, 10)
).to(device)

criterion = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

# Train
history = train_model(
    model=model,
    train_loader=train_loader,
    val_loader=val_loader,
    criterion=criterion,
    optimizer=optimizer,
    device=device,
    epochs=10
)

# Evaluate
accuracy, predictions, labels = evaluate_model(model, test_loader)
print(f'Test accuracy: {accuracy:.2f}%')

Data augmentation

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))
])

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

# Create pipeline with augmentation (automatically pregenerated)
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',  # Load augmented data to GPU for fast training
    n_augmentations=5,  # 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 is saved to: ./data/pytorch/augmented_cifar10/strong_aug_v1/
# Subsequent runs with same augmented_dataset_name will load from cache

Hyperparameter optimization

import torch.nn as nn
from image_classification_tools.pytorch.hyperparameter_optimization import (
    create_objective, MockTrial, TrialFailedError
)
import optuna
from torchvision import datasets

# Define your model factory
def create_cnn(trial, num_classes, in_channels):
    '''Model factory that samples architecture from trial.'''
    n_blocks = trial.suggest_int('n_conv_blocks', 1, 3)
    filters = trial.suggest_categorical('initial_filters', [16, 32, 64])
    # Build your model here...
    return model

# Define search space for training hyperparameters
search_space = {
    'batch_size': [32, 64, 128],
    'learning_rate': (1e-4, 1e-2, 'log'),
    'optimizer': ['Adam', 'SGD'],
    'weight_decay': (1e-6, 1e-3, 'log')
}

# Create objective function
objective = create_objective(
    model_factory=create_cnn,
    data_source=datasets.MNIST,
    data_dir='./data',
    train_transform=transform,
    eval_transform=transform,
    n_epochs=20,
    num_classes=10,
    in_channels=1,
    val_size=10000,
    search_space=search_space
)

# Run optimization with multi-GPU support
study = optuna.create_study(direction='maximize')
n_workers = torch.cuda.device_count() if torch.cuda.is_available() else 1
study.optimize(objective, n_trials=50, n_jobs=n_workers, catch=(TrialFailedError,))

# Recreate best model
mock_trial = MockTrial(study.best_params)
best_model = create_cnn(mock_trial, num_classes=10, in_channels=1)

Requirements

  • Python ≥ 3.10
  • PyTorch ≥ 2.0.0
  • torchvision ≥ 0.15.0
  • numpy
  • matplotlib
  • optuna (optional, for hyperparameter optimization — install via pip install image-classification-tools[optuna])

Documentation

Full documentation is available at: https://gperdrizet.github.io/CIFAR10/

Demo project

See a complete example of using this package for CIFAR-10 classification: https://github.com/gperdrizet/CIFAR10

License

GPLv3

Release files for image-classification-tools 0.6.4

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for image-classification-tools 0.6.4
File Size Uploaded
image_classification_tools-0.6.4.tar.gz 22.9 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for image-classification-tools 0.6.4
File Interpreter ABI Platform
image_classification_tools-0.6.4-py3-none-any.whl Python 3 none any Details

Total release size: 50.1 kB

Release files / image_classification_tools-0.6.4.tar.gz

Download URL image_classification_tools-0.6.4.tar.gz
Size 22.9 kB
Tags Source
SHA-256 checksum
How to use checksums
9b8adcf127f3ed0d7ef61732a3be03753f10426940422367061c21b61f994034
BLAKE2b-256 checksum
How to use checksums
81989496b5993c951bffb581b733071c1229f87dab83eb3141fd373afe2768b8
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.7

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 4, 2026.

Transparency log

Release files / image_classification_tools-0.6.4-py3-none-any.whl

Download URL image_classification_tools-0.6.4-py3-none-any.whl
Size 27.2 kB
Tags Python 3
SHA-256 checksum
How to use checksums
2282cb5c3f77e7c7ff3b2aa367e4a9f4f420ad8c06ce1e8b292f0709622b99e9
BLAKE2b-256 checksum
How to use checksums
1b70708caecd49408269845f9d1c7b24266534b583920761106bcabe5fb90c5c
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.7

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 4, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.6.4 This release

2 release files

0.6.3

2 release files

0.5.9

2 release files

0.5.8

2 release files

0.5.7

2 release files

0.5.6

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page