Skip to content
Open
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
42 changes: 24 additions & 18 deletions finrl/agents/stablebaselines3/models.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
# DRL models from Stable Baselines 3
from __future__ import annotations

import statistics
import time

import numpy as np
Expand Down Expand Up @@ -56,23 +55,30 @@ def _on_step(self) -> bool:
return True

def _on_rollout_end(self) -> bool:
try:
rollout_buffer_rewards = self.locals["rollout_buffer"].rewards.flatten()
self.logger.record(
key="train/reward_min", value=min(rollout_buffer_rewards)
)
self.logger.record(
key="train/reward_mean", value=statistics.mean(rollout_buffer_rewards)
)
self.logger.record(
key="train/reward_max", value=max(rollout_buffer_rewards)
)
except BaseException as error:
# Handle the case where "rewards" is not found
self.logger.record(key="train/reward_min", value=None)
self.logger.record(key="train/reward_mean", value=None)
self.logger.record(key="train/reward_max", value=None)
print("Logging Error:", error)
# On-policy algorithms (A2C, PPO) expose a ``rollout_buffer``; off-policy
# algorithms (DDPG, TD3, SAC) expose a ``replay_buffer``. A model has
# exactly one of them, so read whichever is present instead of assuming
# ``rollout_buffer`` -- the previous code raised a KeyError for off-policy
# algorithms and logged empty (None) reward stats every rollout (#1395).
buffer = getattr(self.model, "rollout_buffer", None) or getattr(
self.model, "replay_buffer", None
)
if buffer is None:
return True

rewards = np.asarray(buffer.rewards)
# A replay buffer is not full early in training; only its first ``pos``
# rows hold real transitions, so drop the zero-initialized tail to keep
# the reward statistics meaningful.
if not getattr(buffer, "full", False):
rewards = rewards[: getattr(buffer, "pos", 0)]
rewards = rewards.flatten()
if rewards.size == 0:
return True

self.logger.record(key="train/reward_min", value=float(rewards.min()))
self.logger.record(key="train/reward_mean", value=float(rewards.mean()))
self.logger.record(key="train/reward_max", value=float(rewards.max()))
return True


Expand Down
94 changes: 94 additions & 0 deletions unit_tests/test_tensorboard_callback.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
from __future__ import annotations

import numpy as np

from finrl.agents.stablebaselines3.models import TensorboardCallback


class _FakeLogger:
"""Capture ``logger.record(key, value)`` calls made by the callback."""

def __init__(self):
self.records = {}

def record(self, key, value):
self.records[key] = value


class _FakeBuffer:
"""Minimal stand-in for an SB3 rollout/replay buffer.

``rewards`` has shape ``(buffer_size, n_envs)`` like the real buffers.
``full``/``pos`` mirror SB3's fill bookkeeping so we can exercise the
partially-filled replay-buffer case.
"""

def __init__(self, rewards, pos=None, full=True):
self.rewards = np.asarray(rewards, dtype=np.float32)
self.full = full
self.pos = self.rewards.shape[0] if pos is None else pos


class _FakeModel:
"""An SB3-like model exposing exactly one buffer plus a logger."""

def __init__(self, rollout_buffer=None, replay_buffer=None):
self.logger = _FakeLogger()
if rollout_buffer is not None:
self.rollout_buffer = rollout_buffer
if replay_buffer is not None:
self.replay_buffer = replay_buffer


def _callback_for(model):
callback = TensorboardCallback()
callback.model = model # `callback.logger` is a property -> model.logger
return callback


def test_on_rollout_end_logs_on_policy_rollout_buffer():
# A2C/PPO expose a rollout_buffer that is full at the end of a rollout.
model = _FakeModel(rollout_buffer=_FakeBuffer([[1.0], [2.0], [3.0]]))
_callback_for(model)._on_rollout_end()

records = model.logger.records
assert records["train/reward_min"] == 1.0
assert records["train/reward_mean"] == 2.0
assert records["train/reward_max"] == 3.0


def test_on_rollout_end_logs_off_policy_replay_buffer():
# Regression for #1395: DDPG/TD3/SAC expose a replay_buffer rather than a
# rollout_buffer. The old code assumed rollout_buffer, raising a KeyError
# and logging None values for every off-policy rollout.
model = _FakeModel(replay_buffer=_FakeBuffer([[10.0], [20.0], [30.0]]))
_callback_for(model)._on_rollout_end()

records = model.logger.records
assert records["train/reward_min"] == 10.0
assert records["train/reward_mean"] == 20.0
assert records["train/reward_max"] == 30.0
assert None not in records.values()


def test_on_rollout_end_ignores_unfilled_replay_buffer_tail():
# A replay buffer is allocated up front; only its first ``pos`` rows hold
# real transitions until it fills, so the zero tail must be ignored.
model = _FakeModel(
replay_buffer=_FakeBuffer([[5.0], [7.0], [0.0], [0.0]], pos=2, full=False)
)
_callback_for(model)._on_rollout_end()

records = model.logger.records
assert records["train/reward_min"] == 5.0
assert records["train/reward_max"] == 7.0
assert records["train/reward_mean"] == 6.0


def test_on_rollout_end_without_buffer_is_noop():
# Defensive: a model exposing neither buffer must not raise or log.
model = _FakeModel()
callback = _callback_for(model)

assert callback._on_rollout_end() is True
assert model.logger.records == {}