Skip to main content

torchtrainer

PyTorch model training made simpler without loosing control. Focus on optimizing your model! Concepts are heavily inspired by the awesome project torchsample and Keras. Further, besides applying Epoch Callbacks it also allows to call Callbacks every time after a specific number of batches passed (iterations) for long epoch durations.

Build Status codecov

Features

  • Torchtrainer
  • Logging utilities
  • Metrics
  • Visdom Visualization
  • Learning Rate Scheduler
  • Checkpointing
  • Flexible for muliple data inputs
  • Setup validation after every ... batches

Usage

Installation

pip install torchtrainer

Example

from torch import nn
from torch.optim import SGD
from torchtrainer import TorchTrainer
from torchtrainer.callbacks import VisdomLinePlotter, ProgressBar, VisdomEpoch, Checkpoint, CSVLogger, \
    EarlyStoppingEpoch, ReduceLROnPlateauCallback
from torchtrainer.metrics import BinaryAccuracy


metrics = [BinaryAccuracy()]

train_loader = ...
val_loader = ...

model = ...
loss = nn.BCELoss()
optimizer = SGD(model.parameters(), lr=0.001, momentum=0.9)

# Setup Visdom Environment for your modl
plotter = VisdomLinePlotter(env_name=f'Model {11}')


# Setup the callbacks of your choice

callbacks = [
    ProgressBar(log_every=10),
    VisdomEpoch(plotter, on_iteration_every=10),
    VisdomEpoch(plotter, on_iteration_every=10, monitor='binary_acc'),
    CSVLogger('test.csv'),
    Checkpoint('./model'),
    EarlyStoppingEpoch(min_delta=0.1, monitor='val_running_loss', patience=10),
    ReduceLROnPlateauCallback(factor=0.1, threshold=0.1, patience=2, verbose=True)
]

trainer = TorchTrainer(model)

# function to transform batch into inputs to your model and y_true values
# if your model accepts multiple inputs, just put all inputs into a tuple (input1, input2), y_true
def transform_fn(batch):
    inputs, y_true = batch
    return inputs, y_true.float()

# prepare your trainer for training
trainer.prepare(optimizer,
                loss,
                train_loader,
                val_loader,
                transform_fn=transform_fn,
                callbacks=callbacks,
                metrics=metrics)

# train your model
result = trainer.train(epochs=10, batch_size=10)

Callbacks

Logger

  • CSVLogger
  • CSVLoggerIteration
  • ProgressBar

Visualization and Logging

  • VisdomEpoch

Optimizers

  • ReduceLROnPlateauCallback
  • StepLRCallback

Regularization

  • EarlyStoppingEpoch
  • EarlyStoppingIteration

Checkpointing

  • Checkpoint
  • CheckpointIteration

Metrics

Currently only BinaryAccuracy is implemented. To implement other Metrics use the abstract base metric class torchtrainer.metrics.metric.Metric.

TODO

  • more tests
  • metrics

Metadata

Release files for torchtrainer 0.3.9

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for torchtrainer 0.3.9
File Size Uploaded
torchtrainer-0.3.9.tar.gz 11.0 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for torchtrainer 0.3.9
File Interpreter ABI Platform
torchtrainer-0.3.9-py3-none-any.whl Python 3 none any Details

Total release size: 25.4 kB

Release files / torchtrainer-0.3.9.tar.gz

Download URL torchtrainer-0.3.9.tar.gz
Size 11.0 kB
Tags Source
SHA-256 checksum
How to use checksums
047cbbf7d92b9d7759666dead1e2847ef6c1ffe142fd9d57764bbde74e62ee4a
BLAKE2b-256 checksum
How to use checksums
bff9eb04c322b7d8aaa30083f484b767ddc06594f9f06a02eaeadb71254e3d21
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via poetry/0.12.17 CPython/3.6.8 Darwin/18.6.0

Release files / torchtrainer-0.3.9-py3-none-any.whl

Download URL torchtrainer-0.3.9-py3-none-any.whl
Size 14.5 kB
Tags Python 3
SHA-256 checksum
How to use checksums
73c190c26037e4876c24d9bb20b930423c73f964bf711632861e7b3354a9feaf
BLAKE2b-256 checksum
How to use checksums
0ff2dfd32580bed08a4dd01c4d1d9b6f95fef5d9041f0a975d3bdb1f2f1150e5
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via poetry/0.12.17 CPython/3.6.8 Darwin/18.6.0

Release history Release notifications | RSS feed

This release

0.3.9 This release

2 release files

0.3.8

2 release files

0.3.7

2 release files

0.3.6

2 release files

0.3.5

2 release files

0.3.4

2 release files

0.3.3

2 release files

0.3.2

2 release files

0.3.1

2 release files

0.3

2 release files

0.2.3

2 release files

0.2

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page