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.2.tar.gz (4.8 kB view details)

Uploaded Source

File details

Details for the file torchiteration-0.0.2.tar.gz.

File metadata

  • Download URL: torchiteration-0.0.2.tar.gz
  • Upload date:
  • Size: 4.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/4.0.2 CPython/3.11.4

File hashes

Hashes for torchiteration-0.0.2.tar.gz
Algorithm Hash digest
SHA256 c98f4285582930583810e8f13fd3da74c05917957c348d05888bc7ea5dd41116
MD5 f104952c0b3b7d3e09a706c9de664f0d
BLAKE2b-256 d9f2ebd95ec9c508ba8e4e2368b509b12220c04935d98e64749611bd06a78f1b

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