From b9679a8f7e8405f1541e2e23a9b37d8a8d9dbe35 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Wed, 23 Dec 2020 12:20:26 +0100 Subject: [PATCH] fixed trainer api --- pts/dataset/loader.py | 2 +- pts/model/estimator.py | 10 +++++----- pts/trainer.py | 4 ---- 3 files changed, 6 insertions(+), 10 deletions(-) diff --git a/pts/dataset/loader.py b/pts/dataset/loader.py index 92fd76d..c3e3288 100644 --- a/pts/dataset/loader.py +++ b/pts/dataset/loader.py @@ -1,4 +1,4 @@ -from typing import Callable, Iterable, Iterator, List, Optional +from typing import Optional import itertools from torch.utils.data import IterableDataset diff --git a/pts/model/estimator.py b/pts/model/estimator.py index 5ba9b7b..8962993 100644 --- a/pts/model/estimator.py +++ b/pts/model/estimator.py @@ -80,7 +80,7 @@ class PyTorchEstimator(Estimator): training_data: Dataset, validation_data: Optional[Dataset] = None, num_workers: Optional[int] = None, - prefetch_factor: Optional[int] = 2, + prefetch_factor: int = 2, shuffle_buffer_length: Optional[int] = None, **kwargs, ) -> TrainOutput: @@ -139,15 +139,15 @@ class PyTorchEstimator(Estimator): training_data: Dataset, validation_data: Optional[Dataset] = None, num_workers: Optional[int] = None, - num_prefetch: Optional[int] = None, + prefetch_factor: int = 2, shuffle_buffer_length: Optional[int] = None, **kwargs, ) -> PyTorchPredictor: return self.train_model( training_data, validation_data, - num_workers, - num_prefetch, - shuffle_buffer_length, + num_workers=num_workers, + prefetch_factor=prefetch_factor, + shuffle_buffer_length=shuffle_buffer_length, **kwargs, ).predictor diff --git a/pts/trainer.py b/pts/trainer.py index f7431a8..1585de4 100644 --- a/pts/trainer.py +++ b/pts/trainer.py @@ -17,8 +17,6 @@ class Trainer: epochs: int = 100, batch_size: int = 32, num_batches_per_epoch: int = 50, - num_workers: int = 4, - pin_memory: bool = False, learning_rate: float = 1e-3, weight_decay: float = 1e-6, device: Optional[Union[torch.device, str]] = None, @@ -29,8 +27,6 @@ class Trainer: self.learning_rate = learning_rate self.weight_decay = weight_decay self.device = device - self.num_workers = num_workers - self.pin_memory = pin_memory def __call__( self,