No project description provided
Project description
Torch Iteration - A Lightweight PyTorch Training Toolkit
Torch Iteration is a versatile PyTorch training toolkit designed to simplify your deep learning projects. With real-time statistics and flexible configuration, it empowers you to train models efficiently and effectively.
Installation
To install Torch Iteration from PyPI, simply run the following command:
pip install torchiteration
To install Torch Iteration from the source code, use the following command:
pip install git+https://github.com/cat-claws/torchiteration/
Table of Contents
How to use
Example: Training a Model on MNIST
Get started with Torch Iteration by training a model on the MNIST dataset. Below is an example script:
import torch
from torch.utils.tensorboard import SummaryWriter
from torchiteration import train, validate, predict, classification_step, predict_classification_step
config = {
'dataset':'mnist',
'training_step':'classification_step',
'batch_size':32,
'optimizer':'Adadelta',
'optimizer_config':{
},
'scheduler':'StepLR',
'scheduler_config':{
'step_size':20,
'gamma':0.1
},
'device':'cuda' if torch.cuda.is_available() else 'cpu',
'validation_step':'classification_step',
}
model = torch.hub.load('cat-claws/nn', 'exampleconvnet', in_channels = 1).to(config['device'])
writer = SummaryWriter(comment = f"_{config['dataset']}_{model._get_name()}_{config['training_step']}", flush_secs=10)
for k, v in config.items():
if k.endswith('_step'):
config[k] = eval(v)
elif k == 'optimizer':
config[k] = vars(torch.optim)[v]([p for p in model.parameters() if p.requires_grad], **config[k+'_config'])
config['scheduler'] = vars(torch.optim.lr_scheduler)[config['scheduler']](config[k], **config['scheduler_config'])
import torchvision
train_set = torchvision.datasets.MNIST('', train=True, download=True, transform=torchvision.transforms.ToTensor())
val_set = torchvision.datasets.MNIST('', train=False, transform=torchvision.transforms.ToTensor())
train_loader = torch.utils.data.DataLoader(train_set, num_workers = 4, batch_size = config['batch_size'])
val_loader = torch.utils.data.DataLoader(val_set, num_workers = 4, batch_size = config['batch_size'])
for epoch in range(10):
if epoch > 0:
train(model, train_loader = train_loader, epoch = epoch, writer = writer, **config)
validate(model, val_loader = val_loader, epoch = epoch, writer = writer, **config)
torch.save(model.state_dict(), writer.log_dir.split('/')[-1] + f"_{epoch:03}.pt")
print(model)
outputs = predict(model, predict_classification_step, val_loader = val_loader, **config)
print(outputs.keys(), outputs['predictions'])
writer.flush()
writer.close()
Visualizing Training Progress
To visualize the training progress, run the following command in your terminal:
tensorboard --logdir=runs
Extending Torch Iteration
You can extend Torch Iteration by creating your own custom *_step function in the same input-output format as those in steps.py. Make sure your function ends with _step for easier integration. Here's a template:
# For now, only one model is accepted, but each step can be very versatile
# net must inherit nn.Module
def customised_step(net, batch, batch_idx, **kw):
# Each dataloader may be defined differently, so you can handle it case by case
_your_data_at_each_batch = batch
# Process your training, testing, etc.
_your_loss_perhaps, _whatever_you_want_to_monitor = _your_process(net, _your_data_at_each_batch, **kw)
return {
'first output':_your_loss_perhaps,
'second output':_whatever_you_want_to_monitor,
'third output, etc.': _just_continue
}
With Torch Iteration, streamline your PyTorch training workflows, and easily customize your training steps for your specific project needs.
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file torchiteration-0.0.5.tar.gz.
File metadata
- Download URL: torchiteration-0.0.5.tar.gz
- Upload date:
- Size: 5.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.12.2
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
68697115f17034ed4e43e982026ab8d0f7bd3bc6ab731093bb9b5e3ee1143a35
|
|
| MD5 |
1cc021ce44c4a026113bc9abf203696b
|
|
| BLAKE2b-256 |
0ba3251cba1c5b997128b7e5b2477d663bf06ee81cdc0f510cdde8c1d1ef7b40
|
File details
Details for the file torchiteration-0.0.5-py3-none-any.whl.
File metadata
- Download URL: torchiteration-0.0.5-py3-none-any.whl
- Upload date:
- Size: 6.4 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.12.2
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
6b3e4fc3cc7fbeeafa408eed57240b55fdfbea884312aab7b2a45d9eb9089c51
|
|
| MD5 |
240c6488cc88165a1c2b5a8283a63542
|
|
| BLAKE2b-256 |
fac5358c983ed9a3bfe56a9fbfd807c86a760e05c200ec1cf7cfb0091f0aaaec
|