Skip to main content

No project description provided

Project description

Torch iteration

A lightweight PyTorch training toolkit.

Installation

To install this small tool from the source code

pip install git+https://github.com/cat-claws/torchiteration/

How to use

Example for MNIST

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

To see the result, run this command in Terminal:

tensorboard --logdir=runs

To extend the usage, you can write your own *_step function in the same input-output format as that in steps.py. Essentially, the format is below. Note that, you'd better let your function end with _step to make things easier.

# 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
	}

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

Uploaded Source

File details

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

File metadata

  • Download URL: torchiteration-0.0.1.tar.gz
  • Upload date:
  • Size: 4.3 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.1.tar.gz
Algorithm Hash digest
SHA256 b4acc640dc7be7b1236576d11ec677f7ed84b626254cd02d582ffb0f5a7c6469
MD5 fce42a39f60d41b16ba8426f2f4e9b84
BLAKE2b-256 100c469c47986da913ffbb3f26503cf5d9f335237dac11ea0e1657cb29235a7e

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