This commit is contained in:
Nicola Demo
2024-08-05 17:34:34 +02:00
parent 686b557144
commit 5245a0b68c
19 changed files with 483 additions and 173 deletions

11
pina/optim/__init__.py Normal file
View File

@@ -0,0 +1,11 @@
__all__ = [
"Optimizer",
"TorchOptimizer",
"Scheduler",
"TorchScheduler",
]
from .optimizer_interface import Optimizer
from .torch_optimizer import TorchOptimizer
from .scheduler_interface import Scheduler
from .torch_scheduler import TorchScheduler

View File

@@ -0,0 +1,7 @@
""" Module for PINA Optimizer """
from abc import ABCMeta
class Optimizer(metaclass=ABCMeta): # TODO improve interface
pass

View File

@@ -0,0 +1,7 @@
""" Module for PINA Optimizer """
from abc import ABCMeta
class Scheduler(metaclass=ABCMeta): # TODO improve interface
pass

View File

@@ -0,0 +1,19 @@
""" Module for PINA Torch Optimizer """
import torch
from ..utils import check_consistency
from .optimizer_interface import Optimizer
class TorchOptimizer(Optimizer):
def __init__(self, optimizer_class, **kwargs):
check_consistency(optimizer_class, torch.optim.Optimizer, subclass=True)
self.optimizer_class = optimizer_class
self.kwargs = kwargs
def hook(self, parameters):
self.optimizer_instance = self.optimizer_class(
parameters, **self.kwargs
)

View File

@@ -0,0 +1,27 @@
""" Module for PINA Torch Optimizer """
import torch
try:
from torch.optim.lr_scheduler import LRScheduler # torch >= 2.0
except ImportError:
from torch.optim.lr_scheduler import (
_LRScheduler as LRScheduler,
) # torch < 2.0
from ..utils import check_consistency
from .optimizer_interface import Optimizer
from .scheduler_interface import Scheduler
class TorchScheduler(Scheduler):
def __init__(self, scheduler_class, **kwargs):
check_consistency(scheduler_class, LRScheduler, subclass=True)
self.scheduler_class = scheduler_class
self.kwargs = kwargs
def hook(self, optimizer):
check_consistency(optimizer, Optimizer)
self.scheduler_instance = self.scheduler_class(
optimizer.optimizer_instance, **self.kwargs
)