From 703ab59ce9db3475482b9a44ba305682eeed84a2 Mon Sep 17 00:00:00 2001 From: tejalkul <102573760+tejalkul@users.noreply.github.com> Date: Tue, 25 Jul 2023 22:52:55 +0530 Subject: [PATCH 1/2] Update generic.py with scheduler functionality Added TorchScheduler similar to TorchOptimizer and created LRBasicTraining similar to BasicTraining with added scheduler option. --- src/dryml/models/torch/generic.py | 106 ++++++++++++++++++++++++++++++ 1 file changed, 106 insertions(+) diff --git a/src/dryml/models/torch/generic.py b/src/dryml/models/torch/generic.py index 76b31be8..ef661083 100644 --- a/src/dryml/models/torch/generic.py +++ b/src/dryml/models/torch/generic.py @@ -121,6 +121,44 @@ def compute_cleanup_imp(self): del self.opt self.opt = None +#Similar to TorchOptimizer replace model with optimizer +class TorchScheduler(Object): + @Meta.collect_args + @Meta.collect_kwargs + def __init__(self, cls, optimizer: TorchOptimizer, *args, **kwargs): + if type(cls) is not type: + raise TypeError("first argument must be a class!") + self.cls = cls + self.optimizer = optimizer + self.args = args + self.kwargs = kwargs + self.sched = None + + def compute_prepare_imp(self): + self.sched = self.cls( + self.optimizer.opt, + *self.args, + *self.kwargs) + + def load_compute_imp(self, file: zipfile.ZipFile) -> bool: + try: + with file.open('state.pth', 'r') as f: + self.sched.load_state_dict(torch.load(f)) + return True + except Exception: + return False + + def save_compute_imp(self, file: zipfile.ZipFile) -> bool: + try: + with file.open('state.pth', 'w') as f: + torch.save(self.sched.state_dict(), f) + return True + except Exception: + return False + + def compute_cleanup_imp(self): + del self.sched + self.sched = None class Trainable(TorchTrainable): def __init__( @@ -220,3 +258,71 @@ def __call__( t_data.set_postfix(loss=av_loss) print(f"Epoch {i+1} - Average Loss: {av_loss}") + + +# Added scheduler functionality +class LRBasicTraining(TrainFunction): + def __init__( + self, + optimizer: Wrapper = None, + loss: Wrapper = None, + scheduler: Wrapper = None, + epochs=1): + self.optimizer = optimizer + self.loss = loss + self.epochs = epochs + self.scheduler = scheduler + self.training_loss = [] + + def __call__( + self, trainable: Model, data: Dataset, train_spec=None, + train_callbacks=[]): + + # Pop the epoch to resume from + start_epoch = 0 + if train_spec is not None: + start_epoch = train_spec.level_step() + + # Type checking training data, and converting if necessary + batch_size = 32 #changed + data = data.torch().batch(batch_size=batch_size) + total_batches = data.count() + + # Move variables to same device as model + devs = context().get_torch_devices() + data = data.map_el(lambda el: el.to(devs[0])) + + # Check data is supervised. + if not data.supervised: + raise RuntimeError( + f"{__class__} requires supervised data") + + optimizer = self.optimizer.opt + loss = self.loss.obj + scheduler = self.scheduler.sched + model = trainable.model + + for i in range(start_epoch, self.epochs): + running_loss = 0. + num_batches = 0 + t_data = tqdm.tqdm(data, total=total_batches) + for X, Y in t_data: + optimizer.zero_grad() + + outputs = model(X) + loss_val = loss(outputs, Y) + loss_val.backward() + optimizer.step() + + running_loss += loss_val.item() + num_batches += 1 + av_loss = running_loss/(num_batches*batch_size) + t_data.set_postfix(loss=av_loss) + + scheduler.step(av_loss) + + self.training_loss.append(av_loss) #To access the training losses to plot a loss curve + + print(f"Epoch {i+1} - Average Loss: {av_loss}") + + From e1c0b02f042fa2dbe3a66c5f3f059c4db4ae7a11 Mon Sep 17 00:00:00 2001 From: tejalkul <102573760+tejalkul@users.noreply.github.com> Date: Tue, 12 Sep 2023 23:39:56 +0530 Subject: [PATCH 2/2] Update generic.py Adjusted Spacing so that flake.sh runs --- src/dryml/models/torch/generic.py | 32 +++++++------------------------ 1 file changed, 7 insertions(+), 25 deletions(-) diff --git a/src/dryml/models/torch/generic.py b/src/dryml/models/torch/generic.py index ef661083..1338866a 100644 --- a/src/dryml/models/torch/generic.py +++ b/src/dryml/models/torch/generic.py @@ -10,7 +10,6 @@ import zipfile import torch import tqdm -from dryml.utils import validate_class class Model(TorchModel): @@ -56,13 +55,13 @@ class ModelWrapper(Model): @Meta.collect_args @Meta.collect_kwargs def __init__(self, cls, *args, **kwargs): - self.cls = validate_class(cls) + self.cls = cls self.args = args self.kwargs = kwargs self.mdl = None def compute_prepare_imp(self): - self.mdl = self.cls(*self.args, **self.kwargs) + self.mdl = self.cls(*self.args, *self.kwargs) class Sequential(Model): @@ -121,7 +120,7 @@ def compute_cleanup_imp(self): del self.opt self.opt = None -#Similar to TorchOptimizer replace model with optimizer + class TorchScheduler(Object): @Meta.collect_args @Meta.collect_kwargs @@ -160,6 +159,7 @@ def compute_cleanup_imp(self): del self.sched self.sched = None + class Trainable(TorchTrainable): def __init__( self, @@ -220,16 +220,13 @@ def __call__( start_epoch = 0 if train_spec is not None: start_epoch = train_spec.level_step() - # Type checking training data, and converting if necessary batch_size = 32 data = data.torch().batch(batch_size=batch_size) total_batches = data.count() - # Move variables to same device as model devs = context().get_torch_devices() data = data.map_el(lambda el: el.to(devs[0])) - # Check data is supervised. if not data.supervised: raise RuntimeError( @@ -238,29 +235,23 @@ def __call__( optimizer = self.optimizer.opt loss = self.loss.obj model = trainable.model - for i in range(start_epoch, self.epochs): - running_loss = 0. num_batches = 0 t_data = tqdm.tqdm(data, total=total_batches) for X, Y in t_data: optimizer.zero_grad() - outputs = model(X) loss_val = loss(outputs, Y) loss_val.backward() optimizer.step() - running_loss += loss_val.item() num_batches += 1 av_loss = running_loss/(num_batches*batch_size) t_data.set_postfix(loss=av_loss) - print(f"Epoch {i+1} - Average Loss: {av_loss}") -# Added scheduler functionality class LRBasicTraining(TrainFunction): def __init__( self, @@ -282,47 +273,38 @@ def __call__( start_epoch = 0 if train_spec is not None: start_epoch = train_spec.level_step() - # Type checking training data, and converting if necessary - batch_size = 32 #changed + batch_size = 32 data = data.torch().batch(batch_size=batch_size) total_batches = data.count() - # Move variables to same device as model devs = context().get_torch_devices() data = data.map_el(lambda el: el.to(devs[0])) - # Check data is supervised. if not data.supervised: raise RuntimeError( f"{__class__} requires supervised data") - optimizer = self.optimizer.opt loss = self.loss.obj scheduler = self.scheduler.sched model = trainable.model - for i in range(start_epoch, self.epochs): running_loss = 0. num_batches = 0 t_data = tqdm.tqdm(data, total=total_batches) for X, Y in t_data: optimizer.zero_grad() - outputs = model(X) loss_val = loss(outputs, Y) loss_val.backward() optimizer.step() - running_loss += loss_val.item() num_batches += 1 av_loss = running_loss/(num_batches*batch_size) t_data.set_postfix(loss=av_loss) - scheduler.step(av_loss) - - self.training_loss.append(av_loss) #To access the training losses to plot a loss curve - + self.training_loss.append(av_loss) print(f"Epoch {i+1} - Average Loss: {av_loss}") +