From 4a90f84c164fdcbeab95b2de76c7e739234ac5c5 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Sun, 3 Jan 2021 21:36:47 +0100 Subject: [PATCH] try one cycle lr --- pts/trainer.py | 31 +++++++++++++++++++++---------- 1 file changed, 21 insertions(+), 10 deletions(-) diff --git a/pts/trainer.py b/pts/trainer.py index 9486253..f107c0f 100644 --- a/pts/trainer.py +++ b/pts/trainer.py @@ -1,5 +1,6 @@ import time from typing import List, Optional, Union +from torch.optim import lr_scheduler from tqdm import tqdm import wandb @@ -7,7 +8,7 @@ import wandb import torch import torch.nn as nn from torch.optim import AdamW -from torch.optim.lr_scheduler import ReduceLROnPlateau +from torch.optim.lr_scheduler import ReduceLROnPlateau, OneCycleLR from torch.utils.data import DataLoader from gluonts.core.component import validated @@ -25,7 +26,8 @@ class Trainer: learning_rate_decay_factor: float = 0.5, patience: int = 10, minimum_learning_rate: float = 5e-5, - clip_gradient: float = 10.0, + maximum_learning_rate: float = 0.01, + clip_gradient: Optional[float] = None, device: Optional[Union[torch.device, str]] = None, **kwargs, ) -> None: @@ -37,6 +39,7 @@ class Trainer: self.learning_rate_decay_factor = learning_rate_decay_factor self.patience = patience self.minimum_learning_rate = minimum_learning_rate + self.maximum_learning_rate = maximum_learning_rate self.clip_gradient = clip_gradient self.device = device wandb.init(**kwargs) @@ -55,13 +58,19 @@ class Trainer: weight_decay=self.weight_decay ) - lr_scheduler = ReduceLROnPlateau( - optimizer, - mode='min', - factor=self.learning_rate_decay_factor, - patience=self.patience, - min_lr=self.minimum_learning_rate, + lr_scheduler = OneCycleLR( + optimizer, + max_lr=self.maximum_learning_rate, + steps_per_epoch=self.num_batches_per_epoch, + epochs=self.epochs, ) + # lr_scheduler = ReduceLROnPlateau( + # optimizer, + # mode='min', + # factor=self.learning_rate_decay_factor, + # patience=self.patience, + # min_lr=self.minimum_learning_rate, + # ) for epoch_no in range(self.epochs): # mark epoch start time @@ -90,11 +99,13 @@ class Trainer: wandb.log({"loss": loss.item()}) loss.backward() - nn.utils.clip_grad_norm_(net.parameters(), self.clip_gradient) + if self.clip_gradient is not None: + nn.utils.clip_grad_norm_(net.parameters(), self.clip_gradient) optimizer.step() + lr_scheduler.step() if self.num_batches_per_epoch == batch_no: - lr_scheduler.step(avg_epoch_loss / batch_no) + # lr_scheduler.step(avg_epoch_loss / batch_no) break # mark epoch end time and log time cost of current epoch