Skip to main content

No project description provided

Project description

Torch Iteration - A Lightweight PyTorch Training Toolkit

GitHub license PyPI version

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

torchiteration-0.0.5.tar.gz (5.7 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

torchiteration-0.0.5-py3-none-any.whl (6.4 kB view details)

Uploaded Python 3

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

Hashes for torchiteration-0.0.5.tar.gz
Algorithm Hash digest
SHA256 68697115f17034ed4e43e982026ab8d0f7bd3bc6ab731093bb9b5e3ee1143a35
MD5 1cc021ce44c4a026113bc9abf203696b
BLAKE2b-256 0ba3251cba1c5b997128b7e5b2477d663bf06ee81cdc0f510cdde8c1d1ef7b40

See more details on using hashes here.

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

Hashes for torchiteration-0.0.5-py3-none-any.whl
Algorithm Hash digest
SHA256 6b3e4fc3cc7fbeeafa408eed57240b55fdfbea884312aab7b2a45d9eb9089c51
MD5 240c6488cc88165a1c2b5a8283a63542
BLAKE2b-256 fac5358c983ed9a3bfe56a9fbfd807c86a760e05c200ec1cf7cfb0091f0aaaec

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page