mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-29 20:14:17 +00:00
Use higher resolution time_ns() and avoid division by zero (#979)
* Use higher resolution time and round up to eps * Update changelog * Add test case * Fix formatting, time()->time_ns * Bugfix: ns is integer not float * Move test to better place * Divide by 1e9 earlier
This commit is contained in:
parent
fda3d4d748
commit
b1cc15970a
6 changed files with 27 additions and 6 deletions
|
|
@ -18,6 +18,7 @@ SB3-Contrib
|
|||
Bug Fixes:
|
||||
^^^^^^^^^^
|
||||
- Fixed the issue that ``predict`` does not always return action as ``np.ndarray`` (@qgallouedec)
|
||||
- Fixed division by zero error when computing FPS when a small number of time has elapsed in operating systems with low-precision timers.
|
||||
|
||||
Deprecations:
|
||||
^^^^^^^^^^^^^
|
||||
|
|
|
|||
|
|
@ -422,7 +422,7 @@ class BaseAlgorithm(ABC):
|
|||
:param tb_log_name: the name of the run for tensorboard log
|
||||
:return:
|
||||
"""
|
||||
self.start_time = time.time()
|
||||
self.start_time = time.time_ns()
|
||||
|
||||
if self.ep_info_buffer is None or reset_num_timesteps:
|
||||
# Initialize buffers if they don't exist, or reinitialize if resetting counters
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import io
|
||||
import pathlib
|
||||
import sys
|
||||
import time
|
||||
import warnings
|
||||
from copy import deepcopy
|
||||
|
|
@ -427,8 +428,8 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
|||
"""
|
||||
Write log.
|
||||
"""
|
||||
time_elapsed = time.time() - self.start_time
|
||||
fps = int((self.num_timesteps - self._num_timesteps_at_start) / (time_elapsed + 1e-8))
|
||||
time_elapsed = max((time.time_ns() - self.start_time) / 1e9, sys.float_info.epsilon)
|
||||
fps = int((self.num_timesteps - self._num_timesteps_at_start) / time_elapsed)
|
||||
self.logger.record("time/episodes", self._episode_num, exclude="tensorboard")
|
||||
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
|
||||
self.logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer]))
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import sys
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
||||
|
||||
|
|
@ -254,13 +255,14 @@ class OnPolicyAlgorithm(BaseAlgorithm):
|
|||
|
||||
# Display training infos
|
||||
if log_interval is not None and iteration % log_interval == 0:
|
||||
fps = int((self.num_timesteps - self._num_timesteps_at_start) / (time.time() - self.start_time))
|
||||
time_elapsed = max((time.time_ns() - self.start_time) / 1e9, sys.float_info.epsilon)
|
||||
fps = int((self.num_timesteps - self._num_timesteps_at_start) / time_elapsed)
|
||||
self.logger.record("time/iterations", iteration, exclude="tensorboard")
|
||||
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
|
||||
self.logger.record("rollout/ep_rew_mean", safe_mean([ep_info["r"] for ep_info in self.ep_info_buffer]))
|
||||
self.logger.record("rollout/ep_len_mean", safe_mean([ep_info["l"] for ep_info in self.ep_info_buffer]))
|
||||
self.logger.record("time/fps", fps)
|
||||
self.logger.record("time/time_elapsed", int(time.time() - self.start_time), exclude="tensorboard")
|
||||
self.logger.record("time/time_elapsed", int(time_elapsed), exclude="tensorboard")
|
||||
self.logger.record("time/total_timesteps", self.num_timesteps, exclude="tensorboard")
|
||||
self.logger.dump(step=self.num_timesteps)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import os
|
||||
import time
|
||||
from typing import Sequence
|
||||
from unittest import mock
|
||||
|
||||
import gym
|
||||
import numpy as np
|
||||
|
|
@ -381,3 +382,16 @@ def test_fps_logger(tmp_path, algo):
|
|||
# third time, FPS should be the same
|
||||
model.learn(100, log_interval=1, reset_num_timesteps=False)
|
||||
assert max_fps / 10 <= logger.name_to_value["time/fps"] <= max_fps
|
||||
|
||||
|
||||
@pytest.mark.parametrize("algo", [A2C, DQN])
|
||||
def test_fps_no_div_zero(algo):
|
||||
"""Set time to constant and train algorithm to check no division by zero error.
|
||||
|
||||
Time can appear to be constant during short runs on platforms with low-precision
|
||||
timers. We should avoid division by zero errors e.g. when computing FPS in
|
||||
this situation."""
|
||||
with mock.patch("time.time", lambda: 42.0):
|
||||
with mock.patch("time.time_ns", lambda: 42.0):
|
||||
model = algo("MlpPolicy", "CartPole-v1")
|
||||
model.learn(total_timesteps=100)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,10 @@ normal_action_noise = NormalActionNoise(np.zeros(1), 0.1 * np.ones(1))
|
|||
|
||||
|
||||
@pytest.mark.parametrize("model_class", [TD3, DDPG])
|
||||
@pytest.mark.parametrize("action_noise", [normal_action_noise, OrnsteinUhlenbeckActionNoise(np.zeros(1), 0.1 * np.ones(1))])
|
||||
@pytest.mark.parametrize(
|
||||
"action_noise",
|
||||
[normal_action_noise, OrnsteinUhlenbeckActionNoise(np.zeros(1), 0.1 * np.ones(1))],
|
||||
)
|
||||
def test_deterministic_pg(model_class, action_noise):
|
||||
"""
|
||||
Test for DDPG and variants (TD3).
|
||||
|
|
|
|||
Loading…
Reference in a new issue