Skip to content

Training API v2.1.1

Pre-release
Pre-release

Choose a tag to compare

@codewithdark-git codewithdark-git released this 21 Dec 18:57
· 16 commits to main since this release

Release Notes - v2.1.1 (Pre-Release)

Release Date: December 21, 2025


🎉 What's New in v2.1.1

This pre-release introduces the Training API - a simple way to train and evaluate your models on common datasets like MNIST, CIFAR-10, and CIFAR-100 without writing boilerplate training loops.


✨ New Features

Training API

Train your hybrid models with just a few lines of code:

from torchvision_customizer import HybridBuilder, Trainer

# Build a customized model
model = HybridBuilder().from_torchvision(
    "resnet18",
    weights="IMAGENET1K_V1",
    patches={"layer3": {"wrap": "se"}},
    num_classes=10,
)

# Create trainer and train on CIFAR-10
trainer = Trainer(model, device='auto')
metrics = trainer.fit_cifar10(epochs=10, lr=0.001)

print(metrics.summary())

Trainer Class

Method Description
fit(train_loader, val_loader) Train on custom data loaders
fit_mnist(epochs, batch_size, lr) Train on MNIST dataset
fit_cifar10(epochs, batch_size, lr) Train on CIFAR-10 dataset
fit_cifar100(epochs, batch_size, lr) Train on CIFAR-100 dataset
evaluate(data_loader) Evaluate model on any data loader

Quick Training Function

For even simpler usage, use quick_train():

from torchvision_customizer import HybridBuilder, quick_train

model = HybridBuilder().from_torchvision("resnet18", num_classes=10)
metrics = quick_train(model, dataset='cifar10', epochs=5)

Supported Features

  • Optimizers: Adam, AdamW, SGD (with momentum)
  • Schedulers: Cosine Annealing, Step LR, OneCycle
  • Device: Automatic CPU/CUDA detection
  • Data Augmentation: Built-in for CIFAR datasets
  • Metrics: Training loss, accuracy, validation loss/accuracy, timing

TrainingMetrics Class

Track your training progress with detailed metrics:

metrics = trainer.fit_cifar10(epochs=10)

print(f"Best accuracy: {metrics.best_val_acc:.2%}")
print(f"Best epoch: {metrics.best_epoch}")
print(metrics.summary())

🐛 Bug Fixes

Hybrid Builder - Parameter Detection

Issue: When wrapping blocks like CBAM or ECA, the builder incorrectly passed in_channels to blocks that only accept channels.

Fix: Smart parameter detection using inspect.signature() to determine which parameters each block accepts.

# Now works correctly
model = HybridBuilder().from_torchvision(
    "resnet18",
    patches={"layer3": {"wrap": "cbam_block"}},  # ✅ Fixed
    num_classes=10,
)

Stage Patterns - MBConv Support

Issue: Stage class didn't recognize mbconv or fused_mbconv patterns.

Fix: Added support for EfficientNet-style blocks:

from torchvision_customizer import Stage

# Now works
stage = Stage(channels=64, blocks=3, pattern='mbconv')
stage = Stage(channels=64, blocks=3, pattern='fused_mbconv')

Recipe Parser - Parameter Shortcuts

Issue: Documentation mentioned shortcuts like k and s but they weren't implemented.

Fix: Added parameter shortcuts that expand automatically:

Shortcut Expands To
k kernel_size
s stride
p padding
g groups
e expansion
r reduction
d dropout
# Both work now
recipe = Recipe(stem="conv(64, k=7, s=2)")  # Using shortcuts
recipe = Recipe(stem="conv(64, kernel_size=7, stride=2)")  # Full names

Stem Class - kernel_size Alias

Issue: Stem class only accepted kernel but recipe parser expanded k to kernel_size.

Fix: Added kernel_size as an alias for kernel:

# Both work now
stem = Stem(64, kernel=7)
stem = Stem(64, kernel_size=7)  # ✅ New alias

📁 New Files

File Description
torchvision_customizer/hybrid/trainer.py Training API implementation
tests/test_trainer.py Tests for trainer module
examples/train_mnist.py MNIST training example
examples/train_cifar10.py CIFAR-10 training example
examples/quick_train_example.py Quick training examples

📦 Installation

pip install git+https://github.com/codewithdark-git/torchvision-customizer.git

🔄 Upgrade from v2.1.0

v2.1.1 is fully backward compatible with v2.1.0. Simply update your installation:

pip install --upgrade git+https://github.com/codewithdark-git/torchvision-customizer.git

📋 Full Changelog

See CHANGELOG.md for complete version history.


🙏 Contributors