diff --git a/docs/guide/tensorboard.rst b/docs/guide/tensorboard.rst index 7ae2ffa..cf5c4cf 100644 --- a/docs/guide/tensorboard.rst +++ b/docs/guide/tensorboard.rst @@ -80,3 +80,76 @@ Here is a simple example on how to log both additional tensor or arbitrary scala model.learn(50000, callback=TensorboardCallback()) + +Logging Videos +-------------- + +TensorBoard supports periodic logging of video data, which helps evaluating agents at various stages during training. + +.. warning:: + To support video logging `moviepy `_ must be installed otherwise, TensorBoard ignores the video and logs a warning. + +Here is an example of how to render an episode and log the resulting video to TensorBoard at regular intervals: + +.. code-block:: python + + from typing import Any, Dict + + import gym + import torch as th + + from stable_baselines3 import A2C + from stable_baselines3.common.callbacks import BaseCallback + from stable_baselines3.common.evaluation import evaluate_policy + from stable_baselines3.common.logger import Video + + + class VideoRecorderCallback(BaseCallback): + def __init__(self, eval_env: gym.Env, render_freq: int, n_eval_episodes: int = 1, deterministic: bool = True): + """ + Records a video of an agent's trajectory traversing ``eval_env`` and logs it to TensorBoard + + :param eval_env: A gym environment from which the trajectory is recorded + :param render_freq: Render the agent's trajectory every eval_freq call of the callback. + :param n_eval_episodes: Number of episodes to render + :param deterministic: Whether to use deterministic or stochastic policy + """ + super().__init__() + self._eval_env = eval_env + self._render_freq = render_freq + self._n_eval_episodes = n_eval_episodes + self._deterministic = deterministic + + def _on_step(self) -> bool: + if self.n_calls % self._render_freq == 0: + screens = [] + + def grab_screens(_locals: Dict[str, Any], _globals: Dict[str, Any]) -> None: + """ + Renders the environment in its current state, recording the screen in the captured `screens` list + + :param _locals: A dictionary containing all local variables of the callback's scope + :param _globals: A dictionary containing all global variables of the callback's scope + """ + screen = self._eval_env.render(mode="rgb_array") + # PyTorch uses CxHxW vs HxWxC gym (and tensorflow) image convention + screens.append(screen.transpose(2, 0, 1)) + + evaluate_policy( + self.model, + self._eval_env, + callback=grab_screens, + n_eval_episodes=self._n_eval_episodes, + deterministic=self._deterministic, + ) + self.logger.record( + "trajectory/video", + Video(th.ByteTensor([screens]), fps=40), + exclude=("stdout", "log", "json", "csv"), + ) + return True + + + model = A2C("MlpPolicy", "CartPole-v1", tensorboard_log="runs/", verbose=1) + video_recorder = VideoRecorderCallback(gym.make("CartPole-v1"), render_freq=5000) + model.learn(total_timesteps=int(5e4), callback=video_recorder) diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index f635712..e896723 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -14,6 +14,7 @@ Breaking Changes: New Features: ^^^^^^^^^^^^^ - Allow custom actor/critic network architectures using ``net_arch=dict(qf=[400, 300], pi=[64, 64])`` for off-policy algorithms (SAC, TD3, DDPG) +- Support logging videos to Tensorboard (@SwamyDev) Bug Fixes: ^^^^^^^^^^ diff --git a/stable_baselines3/common/logger.py b/stable_baselines3/common/logger.py index f97b69a..391c62e 100644 --- a/stable_baselines3/common/logger.py +++ b/stable_baselines3/common/logger.py @@ -5,7 +5,7 @@ import sys import tempfile import warnings from collections import defaultdict -from typing import Any, Dict, List, Optional, TextIO, Tuple, Union +from typing import Any, Dict, List, Optional, Sequence, TextIO, Tuple, Union import numpy as np import pandas @@ -23,6 +23,28 @@ ERROR = 40 DISABLED = 50 +class Video(object): + """ + Video data class storing the video frames and the frame per seconds + """ + + def __init__(self, frames: th.Tensor, fps: Union[float, int]): + self.frames = frames + self.fps = fps + + +class FormatUnsupportedError(NotImplementedError): + def __init__(self, unsupported_formats: Sequence[str], value_description: str): + if len(unsupported_formats) > 1: + format_str = f"formats {', '.join(unsupported_formats)} are" + else: + format_str = f"format {unsupported_formats[0]} is" + super(FormatUnsupportedError, self).__init__( + f"The {format_str} not supported for the {value_description} value logged.\n" + f"You can exclude formats via the `exclude` parameter of the logger's `record` function." + ) + + class KVWriter(object): """ Key Value writer @@ -83,6 +105,9 @@ class HumanOutputFormat(KVWriter, SeqWriter): if excluded is not None and ("stdout" in excluded or "log" in excluded): continue + if isinstance(value, Video): + raise FormatUnsupportedError(["stdout", "log"], "video") + if isinstance(value, float): # Align left value_str = f"{value:<8.3g}" @@ -169,6 +194,8 @@ class JSONOutputFormat(KVWriter): def write(self, key_values: Dict[str, Any], key_excluded: Dict[str, Union[str, Tuple[str, ...]]], step: int = 0) -> None: def cast_to_json_serializable(value: Any): + if isinstance(value, Video): + raise FormatUnsupportedError(["json"], "video") if hasattr(value, "dtype"): if value.shape == () or len(value) == 1: # if value is a dimensionless numpy array or of length 1, serialize as a float @@ -227,6 +254,10 @@ class CSVOutputFormat(KVWriter): if i > 0: self.file.write(",") value = key_values.get(key) + + if isinstance(value, Video): + raise FormatUnsupportedError(["csv"], "video") + if value is not None: self.file.write(str(value)) self.file.write("\n") @@ -262,6 +293,9 @@ class TensorBoardOutputFormat(KVWriter): if isinstance(value, th.Tensor): self.writer.add_histogram(key, value, step) + if isinstance(value, Video): + self.writer.add_video(key, value.frames, step, value.fps) + # Flush the output to the file self.writer.flush() diff --git a/tests/test_logger.py b/tests/test_logger.py index 5ab540f..98f832d 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -2,11 +2,14 @@ from typing import Sequence import numpy as np import pytest +import torch as th from pandas.errors import EmptyDataError from stable_baselines3.common.logger import ( DEBUG, + FormatUnsupportedError, ScopedConfigure, + Video, configure, debug, dump, @@ -162,3 +165,36 @@ def test_exclude_keys(tmp_path, read_log, _format): writer.write(dict(some_tag=42), key_excluded=dict(some_tag=(_format))) writer.close() assert read_log(_format).empty + + +def test_report_video_to_tensorboard(tmp_path, read_log, capsys): + pytest.importorskip("tensorboard") + + video = Video(frames=th.rand(1, 20, 3, 16, 16), fps=20) + writer = make_output_format("tensorboard", tmp_path) + writer.write({"video": video}, key_excluded={"video": ()}) + + if is_moviepy_installed(): + assert not read_log("tensorboard").empty + else: + assert "moviepy" in capsys.readouterr().out + writer.close() + + +def is_moviepy_installed(): + try: + import moviepy # noqa: F401 + except ModuleNotFoundError: + return False + return True + + +@pytest.mark.parametrize("unsupported_format", ["stdout", "log", "json", "csv"]) +def test_report_video_to_unsupported_format_raises_error(tmp_path, unsupported_format): + writer = make_output_format(unsupported_format, tmp_path) + + with pytest.raises(FormatUnsupportedError) as exec_info: + video = Video(frames=th.rand(1, 20, 3, 16, 16), fps=20) + writer.write({"video": video}, key_excluded={"video": ()}) + assert unsupported_format in str(exec_info.value) + writer.close()