Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions benchmarl/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ def _load_hydra_schemas():
from hydra.core.config_store import ConfigStore

from benchmarl.algorithms import algorithm_config_registry
from benchmarl.callbacks import callback_config_registry
from benchmarl.environments import _task_class_registry
from benchmarl.experiment import ExperimentConfig

Expand All @@ -33,6 +34,11 @@ def _load_hydra_schemas():
# Load algos schemas
for algo_name, algo_schema in algorithm_config_registry.items():
cs.store(name=f"{algo_name}_config", group="algorithm", node=algo_schema)
# Load callback schemas
for callback_name, callback_schema in callback_config_registry.items():
cs.store(
name=f"{callback_name}_config", group="callback", node=callback_schema
)
# Load task schemas
for task_schema_name, task_schema in _task_class_registry.items():
cs.store(name=f"{task_schema_name}_config", group="task", node=task_schema)
Expand Down
18 changes: 18 additions & 0 deletions benchmarl/callbacks/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#

from .common import CallbackConfig
from .lr_scheduler import LRSchedulerCallback, LRSchedulerConfig

__all__ = [
"CallbackConfig",
"LRSchedulerCallback",
"LRSchedulerConfig",
]

callback_config_registry = {
"lr_scheduler": LRSchedulerConfig,
}
78 changes: 78 additions & 0 deletions benchmarl/callbacks/common.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#

import pathlib

from abc import abstractmethod
from dataclasses import dataclass
from typing import Any, Dict, Optional, Type

from benchmarl.experiment import Callback
from benchmarl.utils import _read_yaml_config


@dataclass
class CallbackConfig:
"""
Dataclass representing a callback configuration.
This should be overridden by implemented callbacks.
Implementors should:

1. add configuration parameters for their callback
2. implement all abstract methods

"""

def get_callback(self) -> Callback:
"""
Main function to turn the config into the associated callback

Returns: the Callback

"""
return self.associated_class()(
**self.__dict__, # Passes all the custom config parameters
)

@staticmethod
def _load_from_yaml(name: str) -> Dict[str, Any]:
yaml_path = (
pathlib.Path(__file__).parent.parent
/ "conf"
/ "callbacks"
/ f"{name.lower()}.yaml"
)
return _read_yaml_config(str(yaml_path.resolve()))

@classmethod
def get_from_yaml(cls, path: Optional[str] = None):
"""
Load the callback configuration from yaml

Args:
path (str, optional): The full path of the yaml file to load from.
If None, it will default to
``benchmarl/conf/callbacks/self.associated_class().__name__``

Returns: the loaded CallbackConfig
"""

if path is None:
config = CallbackConfig._load_from_yaml(
name=cls.associated_class().__name__
)

else:
config = _read_yaml_config(path)
return cls(**config)

@staticmethod
@abstractmethod
def associated_class() -> Type[Callback]:
"""
The callback class associated to the config
"""
raise NotImplementedError
103 changes: 103 additions & 0 deletions benchmarl/callbacks/lr_scheduler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#

from dataclasses import dataclass, MISSING
from typing import Any, Dict

from tensordict import TensorDictBase

from benchmarl.callbacks.common import CallbackConfig
from benchmarl.experiment.callback import Callback
from benchmarl.utils import _class_from_name


@dataclass
class LRSchedulerConfig(CallbackConfig):
"""Configuration for the LR Scheduler Callback."""

scheduler_class: str = MISSING # "torch.optim.lr_scheduler.StepLR"
scheduler_params: Dict[str, Any] = MISSING
log_lr: bool = MISSING

@staticmethod
def associated_class():
return LRSchedulerCallback


class LRSchedulerCallback(Callback):
"""
Callback that applies learning rate scheduling to a single optimizer.

Uses PyTorch's built-in schedulers.
"""

def __init__(
self,
scheduler_class: str,
scheduler_params: Dict[str, Any],
log_lr: bool = True,
):
super().__init__()
self.scheduler_class = scheduler_class
self.scheduler_params = scheduler_params
self.log_lr = log_lr

self.schedulers = None
self.initial_logging = False

def on_setup(self):
"""Setup the scheduler after the experiment is initialized."""

scheduler_class = _class_from_name(self.scheduler_class)
kwargs = {
k: v
for k, v in self.scheduler_params.items()
if k in scheduler_class.__init__.__code__.co_varnames
}

self.schedulers = {}
for group in self.experiment.optimizers:
self.schedulers[group] = {}
for name, optimizer in self.experiment.optimizers[group].items():
scheduler = scheduler_class(optimizer, **kwargs)
self.schedulers[group][name] = scheduler

def on_load_state_dict(self, state_dict: Dict[str, Any]):
for group in self.schedulers:
for name, scheduler in self.schedulers[group].items():
scheduler.load_state_dict(state_dict[f"schedulers_{group}_{name}"])

def on_state_dict(self, state_dict: Dict[str, Any]):
state_dict.update(
{
f"schedulers_{group}_{name}": scheduler.state_dict()
for group in self.schedulers
for name, scheduler in self.schedulers[group].items()
}
)

def on_batch_collected(self, batch: TensorDictBase):
if self.log_lr and not self.initial_logging:
to_log = {
f"train/{group}/lr": next(
iter(self.schedulers[group].values())
).get_last_lr()[0]
for group in self.experiment.group_map.keys()
}
self.experiment.logger.log(to_log, step=self.experiment.n_iters_performed)
self.initial_logging = True

def on_train_end(self, training_td: TensorDictBase, group: str):
"""Step the scheduler after each collection step."""
for scheduler in self.schedulers[group].values():
scheduler.step()

if self.log_lr:
lr = next(iter(self.schedulers[group].values())).get_last_lr()[0]
to_log = {f"train/{group}/lr": lr}
self.experiment.logger.log(to_log, step=self.experiment.n_iters_performed)

return None
12 changes: 12 additions & 0 deletions benchmarl/conf/callback/lr_scheduler.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
defaults:
- lr_scheduler_config
- _self_

# Scheduler to use (e.g., "StepLR", "CosineAnnealingLR", "ExponentialLR")
scheduler_class: torch.optim.lr_scheduler.StepLR

scheduler_params:
step_size: 1000 # For StepLR: step size for learning rate decay
gamma: 0.9 # For StepLR: multiplicative factor for learning rate decay

log_lr: true
1 change: 1 addition & 0 deletions benchmarl/conf/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ defaults:
- task: ???
- model: layers/mlp
- model@critic_model: layers/mlp
- callback@callbacks.c1: lr_scheduler
- _self_

seed: 0
18 changes: 17 additions & 1 deletion benchmarl/experiment/callback.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

from __future__ import annotations

from typing import List
from typing import Any, Dict, List

from tensordict import TensorDictBase

Expand All @@ -28,6 +28,10 @@ def on_setup(self):
"""A callback called atexperiment setup."""
pass

def on_load_state_dict(self, state_dict: Dict[str, Any]):
"""A callback called at state_dict load."""
pass

def on_batch_collected(self, batch: TensorDictBase):
"""
A callback called at the end of every collection step.
Expand Down Expand Up @@ -73,6 +77,10 @@ def on_evaluation_end(self, rollouts: List[TensorDictBase]):
"""
pass

def on_state_dict(self, state_dict: Dict[str, Any]):
"""A callback called at state_dict save."""
pass


class CallbackNotifier:
def __init__(self, experiment, callbacks: List[Callback]):
Expand All @@ -84,6 +92,10 @@ def _on_setup(self):
for callback in self.callbacks:
callback.on_setup()

def _on_load_state_dict(self, state_dict: Dict[str, Any]):
for callback in self.callbacks:
callback.on_load_state_dict(state_dict)

def _on_batch_collected(self, batch: TensorDictBase):
for callback in self.callbacks:
callback.on_batch_collected(batch)
Expand All @@ -106,3 +118,7 @@ def _on_train_end(self, training_td: TensorDictBase, group: str):
def _on_evaluation_end(self, rollouts: List[TensorDictBase]):
for callback in self.callbacks:
callback.on_evaluation_end(rollouts)

def _on_state_dict(self, state_dict: Dict[str, Any]):
for callback in self.callbacks:
callback.on_state_dict(state_dict)
2 changes: 2 additions & 0 deletions benchmarl/experiment/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -969,6 +969,7 @@ def state_dict(self) -> OrderedDict:
)
if not self.config.collect_with_grad:
state_dict.update({"collector": self.collector.state_dict()})
self._on_state_dict(state_dict)
return state_dict

def load_state_dict(self, state_dict: Dict) -> None:
Expand All @@ -990,6 +991,7 @@ def load_state_dict(self, state_dict: Dict) -> None:
self.total_frames = state_dict["state"]["total_frames"]
self.n_iters_performed = state_dict["state"]["n_iters_performed"]
self.mean_return = state_dict["state"]["mean_return"]
self._on_load_state_dict(state_dict)

def _save_experiment(self) -> None:
"""Checkpoint trainer"""
Expand Down
16 changes: 14 additions & 2 deletions benchmarl/hydra_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,12 @@
import importlib
from dataclasses import is_dataclass
from pathlib import Path
from typing import List

from benchmarl.algorithms.common import AlgorithmConfig
from benchmarl.environments import task_config_registry, TaskClass
from benchmarl.environments.common import _type_check_task_config
from benchmarl.experiment import Experiment, ExperimentConfig
from benchmarl.experiment import Callback, Experiment, ExperimentConfig
from benchmarl.models import model_config_registry
from benchmarl.models.common import ModelConfig, parse_model_config, SequenceModelConfig

Expand Down Expand Up @@ -48,6 +49,7 @@ def load_experiment_from_hydra(
task_config = load_task_config_from_hydra(cfg.task, task_name)
model_config = load_model_config_from_hydra(cfg.model)
critic_model_config = load_model_config_from_hydra(cfg.critic_model)
_callbacks = load_callbacks_from_hydra(getattr(cfg, "callbacks", None) or {})

return Experiment(
task=task_config,
Expand All @@ -56,7 +58,7 @@ def load_experiment_from_hydra(
critic_model_config=critic_model_config,
seed=cfg.seed,
config=experiment_config,
callbacks=callbacks,
callbacks=_callbacks + [*callbacks],
)


Expand Down Expand Up @@ -134,6 +136,16 @@ def load_model_config_from_hydra(cfg: DictConfig) -> ModelConfig:
)


def load_callbacks_from_hydra(cfg: DictConfig) -> List[Callback]:
"""Returns a list of :class:`~benchmarl.callbacks.Callback` from hydra config.

Args:
cfg (DictConfig): the callbacks config dictionary from hydra

"""
return [OmegaConf.to_object(callback).get_callback() for callback in cfg.values()]


def _find_hydra_folder(restore_file: str) -> str:
"""Given the restore file, look for the .hydra folder max three levels above it."""
current_folder = Path(restore_file).parent.resolve()
Expand Down
Loading