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.
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
CSVLoggerCSVLoggerIterationProgressBar
Visualization and Logging
VisdomEpoch
Optimizers
ReduceLROnPlateauCallbackStepLRCallback
Regularization
EarlyStoppingEpochEarlyStoppingIteration
Checkpointing
CheckpointCheckpointIteration
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)
| File | Size | Uploaded | |
|---|---|---|---|
| torchtrainer-0.3.9.tar.gz | 11.0 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| 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
|