mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Add image and figure to tensorboard logger (#277)
* Added Image and Figure classes to logger. For now, these objects can only be logged by TensorBoardOutputFormat * Added documentation for figure and image logging into tensorboard * Updated changelog * Minor changes to documentation. Reviewed supported types for logging images and figures * Fix type for np arrays * Added more explicit example for logging figures in the documentation. Added docstrings for parameters in logging auxiliary classes * Added tests for image and figure logging * Applied autoformatting * Update doc * Fix documentation example * Bump version Co-authored-by: Carlos Casas <ccasascuadrado@guidewire.com> Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>
This commit is contained in:
parent
06498e8be7
commit
5993033c73
5 changed files with 181 additions and 3 deletions
|
|
@ -81,6 +81,78 @@ Here is a simple example on how to log both additional tensor or arbitrary scala
|
||||||
|
|
||||||
model.learn(50000, callback=TensorboardCallback())
|
model.learn(50000, callback=TensorboardCallback())
|
||||||
|
|
||||||
|
Logging Images
|
||||||
|
--------------
|
||||||
|
|
||||||
|
TensorBoard supports periodic logging of image data, which helps evaluating agents at various stages during training.
|
||||||
|
|
||||||
|
.. warning::
|
||||||
|
To support image logging `pillow <https://github.com/python-pillow/Pillow>`_ must be installed otherwise, TensorBoard ignores the image and logs a warning.
|
||||||
|
|
||||||
|
Here is an example of how to render an image to TensorBoard at regular intervals:
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
from stable_baselines3 import SAC
|
||||||
|
from stable_baselines3.common.callbacks import BaseCallback
|
||||||
|
from stable_baselines3.common.logger import Image
|
||||||
|
|
||||||
|
model = SAC("MlpPolicy", "Pendulum-v0", tensorboard_log="/tmp/sac/", verbose=1)
|
||||||
|
|
||||||
|
|
||||||
|
class ImageRecorderCallback(BaseCallback):
|
||||||
|
def __init__(self, verbose=0):
|
||||||
|
super(ImageRecorderCallback, self).__init__(verbose)
|
||||||
|
|
||||||
|
def _on_step(self):
|
||||||
|
image = self.training_env.render(mode="rgb_array")
|
||||||
|
# "HWC" specify the dataformat of the image, here channel last
|
||||||
|
# (H for height, W for width, C for channel)
|
||||||
|
# See https://pytorch.org/docs/stable/tensorboard.html
|
||||||
|
# for supported formats
|
||||||
|
self.logger.record("trajectory/image", Image(image, "HWC"), exclude=("stdout", "log", "json", "csv"))
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
model.learn(50000, callback=ImageRecorderCallback())
|
||||||
|
|
||||||
|
Logging Figures/Plots
|
||||||
|
---------------------
|
||||||
|
TensorBoard supports periodic logging of figures/plots created with matplotlib, which helps evaluating agents at various stages during training.
|
||||||
|
|
||||||
|
.. warning::
|
||||||
|
To support figure logging `matplotlib <https://matplotlib.org/>`_ must be installed otherwise, TensorBoard ignores the figure and logs a warning.
|
||||||
|
|
||||||
|
Here is an example of how to store a plot in TensorBoard at regular intervals:
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
|
from stable_baselines3 import SAC
|
||||||
|
from stable_baselines3.common.callbacks import BaseCallback
|
||||||
|
from stable_baselines3.common.logger import Figure
|
||||||
|
|
||||||
|
model = SAC("MlpPolicy", "Pendulum-v0", tensorboard_log="/tmp/sac/", verbose=1)
|
||||||
|
|
||||||
|
|
||||||
|
class FigureRecorderCallback(BaseCallback):
|
||||||
|
def __init__(self, verbose=0):
|
||||||
|
super(FigureRecorderCallback, self).__init__(verbose)
|
||||||
|
|
||||||
|
def _on_step(self):
|
||||||
|
# Plot values (here a random variable)
|
||||||
|
figure = plt.figure()
|
||||||
|
figure.add_subplot().plot(np.random.random(3))
|
||||||
|
# Close the figure after logging it
|
||||||
|
self.logger.record("trajectory/figure", Figure(figure, close=True), exclude=("stdout", "log", "json", "csv"))
|
||||||
|
plt.close()
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
model.learn(50000, callback=FigureRecorderCallback())
|
||||||
|
|
||||||
Logging Videos
|
Logging Videos
|
||||||
--------------
|
--------------
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@
|
||||||
Changelog
|
Changelog
|
||||||
==========
|
==========
|
||||||
|
|
||||||
Pre-Release 0.11.0a4 (WIP)
|
Pre-Release 0.11.0a5 (WIP)
|
||||||
-------------------------------
|
-------------------------------
|
||||||
|
|
||||||
Breaking Changes:
|
Breaking Changes:
|
||||||
|
|
@ -26,6 +26,7 @@ New Features:
|
||||||
- Added ``monitor_kwargs`` parameter to ``make_vec_env`` and ``make_atari_env``
|
- Added ``monitor_kwargs`` parameter to ``make_vec_env`` and ``make_atari_env``
|
||||||
- Wrap the environments automatically with a ``Monitor`` wrapper when possible.
|
- Wrap the environments automatically with a ``Monitor`` wrapper when possible.
|
||||||
- ``EvalCallback`` now logs the success rate when available (``is_success`` must be present in the info dict)
|
- ``EvalCallback`` now logs the success rate when available (``is_success`` must be present in the info dict)
|
||||||
|
- Added new wrappers to log images and matplotlib figures to tensorboard. (@zampanteymedio)
|
||||||
|
|
||||||
Bug Fixes:
|
Bug Fixes:
|
||||||
^^^^^^^^^^
|
^^^^^^^^^^
|
||||||
|
|
@ -534,4 +535,4 @@ And all the contributors:
|
||||||
@flodorner @KuKuXia @NeoExtended @PartiallyTyped @mmcenta @richardwu @kinalmehta @rolandgvc @tkelestemur @mloo3
|
@flodorner @KuKuXia @NeoExtended @PartiallyTyped @mmcenta @richardwu @kinalmehta @rolandgvc @tkelestemur @mloo3
|
||||||
@tirafesi @blurLake @koulakis @joeljosephjin @shwang @rk37 @andyshih12 @RaphaelWag @xicocaio
|
@tirafesi @blurLake @koulakis @joeljosephjin @shwang @rk37 @andyshih12 @RaphaelWag @xicocaio
|
||||||
@diditforlulz273 @liorcohen5 @ManifoldFR @mloo3 @SwamyDev @wmmc88 @megan-klaiber @thisray
|
@diditforlulz273 @liorcohen5 @ManifoldFR @mloo3 @SwamyDev @wmmc88 @megan-klaiber @thisray
|
||||||
@tfederico @hn2 @LucasAlegre @AptX395
|
@tfederico @hn2 @LucasAlegre @AptX395 @zampanteymedio
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ from typing import Any, Dict, List, Optional, Sequence, TextIO, Tuple, Union
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas
|
import pandas
|
||||||
import torch as th
|
import torch as th
|
||||||
|
from matplotlib import pyplot as plt
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
|
|
@ -26,6 +27,9 @@ DISABLED = 50
|
||||||
class Video(object):
|
class Video(object):
|
||||||
"""
|
"""
|
||||||
Video data class storing the video frames and the frame per seconds
|
Video data class storing the video frames and the frame per seconds
|
||||||
|
|
||||||
|
:param frames: frames to create the video from
|
||||||
|
:param fps: frames per second
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, frames: th.Tensor, fps: Union[float, int]):
|
def __init__(self, frames: th.Tensor, fps: Union[float, int]):
|
||||||
|
|
@ -33,6 +37,34 @@ class Video(object):
|
||||||
self.fps = fps
|
self.fps = fps
|
||||||
|
|
||||||
|
|
||||||
|
class Figure(object):
|
||||||
|
"""
|
||||||
|
Figure data class storing a matplotlib figure and whether to close the figure after logging it
|
||||||
|
|
||||||
|
:param figure: figure to log
|
||||||
|
:param close: if true, close the figure after logging it
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, figure: plt.figure, close: bool):
|
||||||
|
self.figure = figure
|
||||||
|
self.close = close
|
||||||
|
|
||||||
|
|
||||||
|
class Image(object):
|
||||||
|
"""
|
||||||
|
Image data class storing an image and data format
|
||||||
|
|
||||||
|
:param image: image to log
|
||||||
|
:param dataformats: Image data format specification of the form NCHW, NHWC, CHW, HWC, HW, WH, etc.
|
||||||
|
More info in add_image method doc at https://pytorch.org/docs/stable/tensorboard.html
|
||||||
|
Gym envs normally use 'HWC' (channel last)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, image: Union[th.Tensor, np.ndarray, str], dataformats: str):
|
||||||
|
self.image = image
|
||||||
|
self.dataformats = dataformats
|
||||||
|
|
||||||
|
|
||||||
class FormatUnsupportedError(NotImplementedError):
|
class FormatUnsupportedError(NotImplementedError):
|
||||||
def __init__(self, unsupported_formats: Sequence[str], value_description: str):
|
def __init__(self, unsupported_formats: Sequence[str], value_description: str):
|
||||||
if len(unsupported_formats) > 1:
|
if len(unsupported_formats) > 1:
|
||||||
|
|
@ -108,6 +140,12 @@ class HumanOutputFormat(KVWriter, SeqWriter):
|
||||||
if isinstance(value, Video):
|
if isinstance(value, Video):
|
||||||
raise FormatUnsupportedError(["stdout", "log"], "video")
|
raise FormatUnsupportedError(["stdout", "log"], "video")
|
||||||
|
|
||||||
|
if isinstance(value, Figure):
|
||||||
|
raise FormatUnsupportedError(["stdout", "log"], "figure")
|
||||||
|
|
||||||
|
if isinstance(value, Image):
|
||||||
|
raise FormatUnsupportedError(["stdout", "log"], "image")
|
||||||
|
|
||||||
if isinstance(value, float):
|
if isinstance(value, float):
|
||||||
# Align left
|
# Align left
|
||||||
value_str = f"{value:<8.3g}"
|
value_str = f"{value:<8.3g}"
|
||||||
|
|
@ -196,6 +234,10 @@ class JSONOutputFormat(KVWriter):
|
||||||
def cast_to_json_serializable(value: Any):
|
def cast_to_json_serializable(value: Any):
|
||||||
if isinstance(value, Video):
|
if isinstance(value, Video):
|
||||||
raise FormatUnsupportedError(["json"], "video")
|
raise FormatUnsupportedError(["json"], "video")
|
||||||
|
if isinstance(value, Figure):
|
||||||
|
raise FormatUnsupportedError(["json"], "figure")
|
||||||
|
if isinstance(value, Image):
|
||||||
|
raise FormatUnsupportedError(["json"], "image")
|
||||||
if hasattr(value, "dtype"):
|
if hasattr(value, "dtype"):
|
||||||
if value.shape == () or len(value) == 1:
|
if value.shape == () or len(value) == 1:
|
||||||
# if value is a dimensionless numpy array or of length 1, serialize as a float
|
# if value is a dimensionless numpy array or of length 1, serialize as a float
|
||||||
|
|
@ -258,6 +300,12 @@ class CSVOutputFormat(KVWriter):
|
||||||
if isinstance(value, Video):
|
if isinstance(value, Video):
|
||||||
raise FormatUnsupportedError(["csv"], "video")
|
raise FormatUnsupportedError(["csv"], "video")
|
||||||
|
|
||||||
|
if isinstance(value, Figure):
|
||||||
|
raise FormatUnsupportedError(["csv"], "figure")
|
||||||
|
|
||||||
|
if isinstance(value, Image):
|
||||||
|
raise FormatUnsupportedError(["csv"], "image")
|
||||||
|
|
||||||
if value is not None:
|
if value is not None:
|
||||||
self.file.write(str(value))
|
self.file.write(str(value))
|
||||||
self.file.write("\n")
|
self.file.write("\n")
|
||||||
|
|
@ -296,6 +344,12 @@ class TensorBoardOutputFormat(KVWriter):
|
||||||
if isinstance(value, Video):
|
if isinstance(value, Video):
|
||||||
self.writer.add_video(key, value.frames, step, value.fps)
|
self.writer.add_video(key, value.frames, step, value.fps)
|
||||||
|
|
||||||
|
if isinstance(value, Figure):
|
||||||
|
self.writer.add_figure(key, value.figure, step, close=value.close)
|
||||||
|
|
||||||
|
if isinstance(value, Image):
|
||||||
|
self.writer.add_image(key, value.image, step, dataformats=value.dataformats)
|
||||||
|
|
||||||
# Flush the output to the file
|
# Flush the output to the file
|
||||||
self.writer.flush()
|
self.writer.flush()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1 +1 @@
|
||||||
0.11.0a4
|
0.11.0a5
|
||||||
|
|
|
||||||
|
|
@ -3,11 +3,14 @@ from typing import Sequence
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
import torch as th
|
import torch as th
|
||||||
|
from matplotlib import pyplot as plt
|
||||||
from pandas.errors import EmptyDataError
|
from pandas.errors import EmptyDataError
|
||||||
|
|
||||||
from stable_baselines3.common.logger import (
|
from stable_baselines3.common.logger import (
|
||||||
DEBUG,
|
DEBUG,
|
||||||
|
Figure,
|
||||||
FormatUnsupportedError,
|
FormatUnsupportedError,
|
||||||
|
Image,
|
||||||
ScopedConfigure,
|
ScopedConfigure,
|
||||||
Video,
|
Video,
|
||||||
configure,
|
configure,
|
||||||
|
|
@ -198,3 +201,51 @@ def test_report_video_to_unsupported_format_raises_error(tmp_path, unsupported_f
|
||||||
writer.write({"video": video}, key_excluded={"video": ()})
|
writer.write({"video": video}, key_excluded={"video": ()})
|
||||||
assert unsupported_format in str(exec_info.value)
|
assert unsupported_format in str(exec_info.value)
|
||||||
writer.close()
|
writer.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_report_image_to_tensorboard(tmp_path, read_log):
|
||||||
|
pytest.importorskip("tensorboard")
|
||||||
|
|
||||||
|
image = Image(image=th.rand(16, 16, 3), dataformats="HWC")
|
||||||
|
writer = make_output_format("tensorboard", tmp_path)
|
||||||
|
writer.write({"image": image}, key_excluded={"image": ()})
|
||||||
|
|
||||||
|
assert not read_log("tensorboard").empty
|
||||||
|
writer.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("unsupported_format", ["stdout", "log", "json", "csv"])
|
||||||
|
def test_report_image_to_unsupported_format_raises_error(tmp_path, unsupported_format):
|
||||||
|
writer = make_output_format(unsupported_format, tmp_path)
|
||||||
|
|
||||||
|
with pytest.raises(FormatUnsupportedError) as exec_info:
|
||||||
|
image = Image(image=th.rand(16, 16, 3), dataformats="HWC")
|
||||||
|
writer.write({"image": image}, key_excluded={"image": ()})
|
||||||
|
assert unsupported_format in str(exec_info.value)
|
||||||
|
writer.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_report_figure_to_tensorboard(tmp_path, read_log):
|
||||||
|
pytest.importorskip("tensorboard")
|
||||||
|
|
||||||
|
fig = plt.figure()
|
||||||
|
fig.add_subplot().plot(np.random.random(3))
|
||||||
|
figure = Figure(figure=fig, close=True)
|
||||||
|
writer = make_output_format("tensorboard", tmp_path)
|
||||||
|
writer.write({"figure": figure}, key_excluded={"figure": ()})
|
||||||
|
|
||||||
|
assert not read_log("tensorboard").empty
|
||||||
|
writer.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("unsupported_format", ["stdout", "log", "json", "csv"])
|
||||||
|
def test_report_figure_to_unsupported_format_raises_error(tmp_path, unsupported_format):
|
||||||
|
writer = make_output_format(unsupported_format, tmp_path)
|
||||||
|
|
||||||
|
with pytest.raises(FormatUnsupportedError) as exec_info:
|
||||||
|
fig = plt.figure()
|
||||||
|
fig.add_subplot().plot(np.random.random(3))
|
||||||
|
figure = Figure(figure=fig, close=True)
|
||||||
|
writer.write({"figure": figure}, key_excluded={"figure": ()})
|
||||||
|
assert unsupported_format in str(exec_info.value)
|
||||||
|
writer.close()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue