diff --git a/examples/demo_agents/demo_SAC.py b/examples/demo_agents/demo_SAC.py index 83ad063cb..be90fd61f 100644 --- a/examples/demo_agents/demo_SAC.py +++ b/examples/demo_agents/demo_SAC.py @@ -9,6 +9,7 @@ import time import gymnasium as gym + from rlberry.agents.torch.sac import SACAgent from rlberry.envs import Pendulum from rlberry.manager import AgentManager diff --git a/examples/demo_agents/video_plot_a2c.py b/examples/demo_agents/video_plot_a2c.py index 6e20c537f..60494cc43 100644 --- a/examples/demo_agents/video_plot_a2c.py +++ b/examples/demo_agents/video_plot_a2c.py @@ -11,10 +11,10 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_a2c.jpg' -from rlberry.agents.torch import A2CAgent -from rlberry.envs.benchmarks.ball_exploration import PBall2D from gymnasium.wrappers import TimeLimit +from rlberry.agents.torch import A2CAgent +from rlberry.envs.benchmarks.ball_exploration import PBall2D env = PBall2D() env = TimeLimit(env, max_episode_steps=256) diff --git a/examples/demo_agents/video_plot_dqn.py b/examples/demo_agents/video_plot_dqn.py index fd8e91f36..a7e798426 100644 --- a/examples/demo_agents/video_plot_dqn.py +++ b/examples/demo_agents/video_plot_dqn.py @@ -20,17 +20,16 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_dqn.jpg' -from rlberry.envs import gym_make +import os +import shutil + +from gymnasium.wrappers.record_video import RecordVideo from torch.utils.tensorboard import SummaryWriter from rlberry.agents.torch.dqn import DQNAgent +from rlberry.envs import gym_make from rlberry.utils.logging import configure_logging -from gymnasium.wrappers.record_video import RecordVideo -import shutil -import os - - configure_logging(level="INFO") env = gym_make("CartPole-v1", render_mode="rgb_array") diff --git a/examples/demo_agents/video_plot_mdqn.py b/examples/demo_agents/video_plot_mdqn.py index d1f84f449..5631bd835 100644 --- a/examples/demo_agents/video_plot_mdqn.py +++ b/examples/demo_agents/video_plot_mdqn.py @@ -20,17 +20,16 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_dqn.jpg' -from rlberry.envs import gym_make +import os +import shutil + +from gymnasium.wrappers.record_video import RecordVideo from torch.utils.tensorboard import SummaryWriter from rlberry.agents.torch.dqn import MunchausenDQNAgent +from rlberry.envs import gym_make from rlberry.utils.logging import configure_logging -from gymnasium.wrappers.record_video import RecordVideo -import shutil -import os - - configure_logging(level="INFO") env = gym_make("CartPole-v1", render_mode="rgb_array") diff --git a/examples/demo_agents/video_plot_ppo.py b/examples/demo_agents/video_plot_ppo.py index 47e4c6629..1b48bfcb2 100644 --- a/examples/demo_agents/video_plot_ppo.py +++ b/examples/demo_agents/video_plot_ppo.py @@ -14,7 +14,6 @@ from rlberry.agents.torch import PPOAgent from rlberry.envs.benchmarks.ball_exploration import PBall2D - env = PBall2D() n_steps = 3e3 diff --git a/examples/demo_agents/video_plot_rs_kernel_ucbvi.py b/examples/demo_agents/video_plot_rs_kernel_ucbvi.py index 8a012d274..54db5e011 100644 --- a/examples/demo_agents/video_plot_rs_kernel_ucbvi.py +++ b/examples/demo_agents/video_plot_rs_kernel_ucbvi.py @@ -11,8 +11,8 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_rs_kernel_ucbvi.jpg' -from rlberry.envs import Acrobot from rlberry.agents import RSKernelUCBVIAgent +from rlberry.envs import Acrobot from rlberry.wrappers import RescaleRewardWrapper env = Acrobot() diff --git a/examples/demo_bandits/plot_TS_bandit.py b/examples/demo_bandits/plot_TS_bandit.py index 599033dbc..ae6f89157 100644 --- a/examples/demo_bandits/plot_TS_bandit.py +++ b/examples/demo_bandits/plot_TS_bandit.py @@ -11,19 +11,19 @@ """ import numpy as np -from rlberry.envs.bandits import BernoulliBandit, NormalBandit + from rlberry.agents.bandits import ( IndexAgent, TSAgent, - makeBoundedUCBIndex, - makeSubgaussianUCBIndex, makeBetaPrior, + makeBoundedUCBIndex, makeGaussianPrior, + makeSubgaussianUCBIndex, ) +from rlberry.envs.bandits import BernoulliBandit, NormalBandit from rlberry.manager import ExperimentManager, plot_writer_data from rlberry.wrappers import WriterWrapper - # Bernoulli # Agents definition diff --git a/examples/demo_bandits/plot_compare_index_bandits.py b/examples/demo_bandits/plot_compare_index_bandits.py index f089c5ac3..bf4d3c649 100644 --- a/examples/demo_bandits/plot_compare_index_bandits.py +++ b/examples/demo_bandits/plot_compare_index_bandits.py @@ -6,11 +6,9 @@ This script Compare several bandits agents and as a sub-product also shows how to use subplots in with `plot_writer_data` """ -import numpy as np import matplotlib.pyplot as plt -from rlberry.envs.bandits import BernoulliBandit -from rlberry.manager import ExperimentManager, plot_writer_data -from rlberry.wrappers import WriterWrapper +import numpy as np + from rlberry.agents.bandits import ( IndexAgent, RandomizedAgent, @@ -22,6 +20,9 @@ makeETCIndex, makeEXP3Index, ) +from rlberry.envs.bandits import BernoulliBandit +from rlberry.manager import ExperimentManager, plot_writer_data +from rlberry.wrappers import WriterWrapper # Agents definition # sphinx_gallery_thumbnail_number = 2 diff --git a/examples/demo_bandits/plot_exp3_bandit.py b/examples/demo_bandits/plot_exp3_bandit.py index f4716a219..3d15e4562 100644 --- a/examples/demo_bandits/plot_exp3_bandit.py +++ b/examples/demo_bandits/plot_exp3_bandit.py @@ -8,17 +8,17 @@ """ import numpy as np -from rlberry.envs.bandits import AdversarialBandit + from rlberry.agents.bandits import ( RandomizedAgent, TSAgent, - makeEXP3Index, makeBetaPrior, + makeEXP3Index, ) +from rlberry.envs.bandits import AdversarialBandit from rlberry.manager import ExperimentManager, plot_writer_data from rlberry.wrappers import WriterWrapper - # Agents definition diff --git a/examples/demo_bandits/plot_mirror_bandit.py b/examples/demo_bandits/plot_mirror_bandit.py index 4e9b9757d..d6ba60a26 100644 --- a/examples/demo_bandits/plot_mirror_bandit.py +++ b/examples/demo_bandits/plot_mirror_bandit.py @@ -12,19 +12,16 @@ The code is in three parts: definition of environment, definition of agent, and finally definition of the experiment. """ +import matplotlib.pyplot as plt import numpy as np - -from rlberry.manager import ExperimentManager, read_writer_data -from rlberry.envs.interface import Model -from rlberry.agents.bandits import BanditWithSimplePolicy -from rlberry.wrappers import WriterWrapper -import rlberry.spaces as spaces - import requests -import matplotlib.pyplot as plt - import rlberry +import rlberry.spaces as spaces +from rlberry.agents.bandits import BanditWithSimplePolicy +from rlberry.envs.interface import Model +from rlberry.manager import ExperimentManager, read_writer_data +from rlberry.wrappers import WriterWrapper logger = rlberry.logger diff --git a/examples/demo_bandits/plot_ucb_bandit.py b/examples/demo_bandits/plot_ucb_bandit.py index 92b9d8ae2..c7e358591 100644 --- a/examples/demo_bandits/plot_ucb_bandit.py +++ b/examples/demo_bandits/plot_ucb_bandit.py @@ -6,14 +6,14 @@ This script shows how to define a bandit environment and an UCB Index-based algorithm. """ +import matplotlib.pyplot as plt import numpy as np -from rlberry.envs.bandits import NormalBandit + from rlberry.agents.bandits import IndexAgent, makeSubgaussianUCBIndex +from rlberry.envs.bandits import NormalBandit from rlberry.manager import ExperimentManager, plot_writer_data -import matplotlib.pyplot as plt from rlberry.wrappers import WriterWrapper - # Agents definition diff --git a/examples/demo_env/example_atari_atlantis_vectorized_ppo.py b/examples/demo_env/example_atari_atlantis_vectorized_ppo.py index 6fc1c187b..701f7faee 100644 --- a/examples/demo_env/example_atari_atlantis_vectorized_ppo.py +++ b/examples/demo_env/example_atari_atlantis_vectorized_ppo.py @@ -14,15 +14,16 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/example_plot_atari_atlantis_vectorized_ppo.jpg' -from rlberry.manager import ExperimentManager +import os +import shutil from datetime import datetime -from rlberry.agents.torch import PPOAgent + from gymnasium.wrappers.record_video import RecordVideo -import shutil -import os -from rlberry.envs.gym_make import atari_make -from rlberry.agents.torch.utils.training import model_factory_from_env +from rlberry.agents.torch import PPOAgent +from rlberry.agents.torch.utils.training import model_factory_from_env +from rlberry.envs.gym_make import atari_make +from rlberry.manager import ExperimentManager initial_time = datetime.now() print("-------- init agent --------") diff --git a/examples/demo_env/example_atari_breakout_vectorized_ppo.py b/examples/demo_env/example_atari_breakout_vectorized_ppo.py index d96221dc3..3de9f6084 100644 --- a/examples/demo_env/example_atari_breakout_vectorized_ppo.py +++ b/examples/demo_env/example_atari_breakout_vectorized_ppo.py @@ -14,15 +14,16 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/example_plot_atari_breakout_vectorized_ppo.jpg' -from rlberry.manager import ExperimentManager +import os +import shutil from datetime import datetime -from rlberry.agents.torch import PPOAgent + from gymnasium.wrappers.record_video import RecordVideo -import shutil -import os -from rlberry.envs.gym_make import atari_make -from rlberry.agents.torch.utils.training import model_factory_from_env +from rlberry.agents.torch import PPOAgent +from rlberry.agents.torch.utils.training import model_factory_from_env +from rlberry.envs.gym_make import atari_make +from rlberry.manager import ExperimentManager initial_time = datetime.now() print("-------- init agent --------") diff --git a/examples/demo_env/video_plot_acrobot.py b/examples/demo_env/video_plot_acrobot.py index 45aa18a4d..44690bbdd 100644 --- a/examples/demo_env/video_plot_acrobot.py +++ b/examples/demo_env/video_plot_acrobot.py @@ -11,8 +11,8 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_acrobot.jpg' -from rlberry.envs import Acrobot from rlberry.agents import RSUCBVIAgent +from rlberry.envs import Acrobot from rlberry.wrappers import RescaleRewardWrapper env = Acrobot() diff --git a/examples/demo_env/video_plot_apple_gold.py b/examples/demo_env/video_plot_apple_gold.py index 74282cca4..c342b865c 100644 --- a/examples/demo_env/video_plot_apple_gold.py +++ b/examples/demo_env/video_plot_apple_gold.py @@ -9,9 +9,10 @@ :width: 600 """ +from rlberry.agents.dynprog import ValueIterationAgent + # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_apple_gold.jpg' from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold -from rlberry.agents.dynprog import ValueIterationAgent env = AppleGold(reward_free=False, array_observation=False) diff --git a/examples/demo_env/video_plot_atari_freeway.py b/examples/demo_env/video_plot_atari_freeway.py index f8e22f2f9..5c42bd01d 100644 --- a/examples/demo_env/video_plot_atari_freeway.py +++ b/examples/demo_env/video_plot_atari_freeway.py @@ -14,14 +14,15 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_atari_freeway.jpg' -from rlberry.manager import ExperimentManager +import os +import shutil from datetime import datetime -from rlberry.agents.torch.dqn.dqn import DQNAgent + from gymnasium.wrappers.record_video import RecordVideo -import shutil -import os -from rlberry.envs.gym_make import atari_make +from rlberry.agents.torch.dqn.dqn import DQNAgent +from rlberry.envs.gym_make import atari_make +from rlberry.manager import ExperimentManager initial_time = datetime.now() print("-------- init agent --------") diff --git a/examples/demo_env/video_plot_gridworld.py b/examples/demo_env/video_plot_gridworld.py index 129e5a7e6..20a34f548 100644 --- a/examples/demo_env/video_plot_gridworld.py +++ b/examples/demo_env/video_plot_gridworld.py @@ -15,7 +15,6 @@ from rlberry.agents.dynprog import ValueIterationAgent from rlberry.envs.finite import GridWorld - env = GridWorld(7, 10, walls=((2, 2), (3, 3))) agent = ValueIterationAgent(env, gamma=0.95) diff --git a/examples/demo_env/video_plot_old_gym_compatibility_wrapper_old_acrobot.py b/examples/demo_env/video_plot_old_gym_compatibility_wrapper_old_acrobot.py index 90f3eb11b..f0f417cf2 100644 --- a/examples/demo_env/video_plot_old_gym_compatibility_wrapper_old_acrobot.py +++ b/examples/demo_env/video_plot_old_gym_compatibility_wrapper_old_acrobot.py @@ -11,10 +11,10 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_old_gym_acrobot.jpg' -from rlberry.wrappers.tests.old_env.old_acrobot import Old_Acrobot from rlberry.agents import RSUCBVIAgent from rlberry.wrappers import RescaleRewardWrapper from rlberry.wrappers.gym_utils import OldGymCompatibilityWrapper +from rlberry.wrappers.tests.old_env.old_acrobot import Old_Acrobot env = Old_Acrobot() env = OldGymCompatibilityWrapper(env) diff --git a/examples/demo_env/video_plot_pball.py b/examples/demo_env/video_plot_pball.py index af6c7c637..6918eea94 100644 --- a/examples/demo_env/video_plot_pball.py +++ b/examples/demo_env/video_plot_pball.py @@ -11,6 +11,7 @@ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_pball.jpg' import numpy as np + from rlberry.envs.benchmarks.ball_exploration import PBall2D p = 5 diff --git a/examples/demo_env/video_plot_rooms.py b/examples/demo_env/video_plot_rooms.py index 9cee6bf6f..858fec279 100644 --- a/examples/demo_env/video_plot_rooms.py +++ b/examples/demo_env/video_plot_rooms.py @@ -10,8 +10,8 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_rooms.jpg' -from rlberry.envs.benchmarks.grid_exploration.nroom import NRoom from rlberry.agents.dynprog import ValueIterationAgent +from rlberry.envs.benchmarks.grid_exploration.nroom import NRoom env = NRoom( nrooms=9, diff --git a/examples/demo_env/video_plot_springcartpole.py b/examples/demo_env/video_plot_springcartpole.py index 7669ef62a..a30c20d41 100644 --- a/examples/demo_env/video_plot_springcartpole.py +++ b/examples/demo_env/video_plot_springcartpole.py @@ -13,10 +13,11 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_springcartpole.jpg' -from rlberry.envs.classic_control import SpringCartPole -from rlberry.agents.torch import DQNAgent from gymnasium.wrappers.time_limit import TimeLimit +from rlberry.agents.torch import DQNAgent +from rlberry.envs.classic_control import SpringCartPole + model_configs = { "type": "MultiLayerPerceptron", "layer_sizes": (256, 256), diff --git a/examples/demo_env/video_plot_twinrooms.py b/examples/demo_env/video_plot_twinrooms.py index 22c36683a..7723a17fa 100644 --- a/examples/demo_env/video_plot_twinrooms.py +++ b/examples/demo_env/video_plot_twinrooms.py @@ -10,10 +10,10 @@ """ # sphinx_gallery_thumbnail_path = 'thumbnails/video_plot_twinrooms.jpg' -from rlberry.envs.benchmarks.generalization.twinrooms import TwinRooms from rlberry.agents.mbqvi import MBQVIAgent -from rlberry.wrappers.discretize_state import DiscretizeStateWrapper +from rlberry.envs.benchmarks.generalization.twinrooms import TwinRooms from rlberry.seeding import Seeder +from rlberry.wrappers.discretize_state import DiscretizeStateWrapper seeder = Seeder(123) diff --git a/examples/demo_experiment/run.py b/examples/demo_experiment/run.py index 87741dd66..592f5d63a 100644 --- a/examples/demo_experiment/run.py +++ b/examples/demo_experiment/run.py @@ -11,11 +11,9 @@ $ python examples/demo_examples/demo_experiment/run.py """ -from rlberry.experiment import load_experiment_results -from rlberry.experiment import experiment_generator +from rlberry.experiment import experiment_generator, load_experiment_results from rlberry.manager.multiple_managers import MultipleManagers - if __name__ == "__main__": multimanagers = MultipleManagers(parallelization="thread") diff --git a/examples/demo_network/run_client.py b/examples/demo_network/run_client.py index 5098f5311..0b6bab6df 100644 --- a/examples/demo_network/run_client.py +++ b/examples/demo_network/run_client.py @@ -3,11 +3,11 @@ Demo: run_client ===================== """ -from rlberry.network.client import BerryClient -from rlberry.network import interface -from rlberry.network.interface import Message, ResourceRequest import numpy as np +from rlberry.network import interface +from rlberry.network.client import BerryClient +from rlberry.network.interface import Message, ResourceRequest port = int(input("Select server port: ")) client = BerryClient(port=port) diff --git a/examples/demo_network/run_remote_manager.py b/examples/demo_network/run_remote_manager.py index 83df52486..de7cfb5a4 100644 --- a/examples/demo_network/run_remote_manager.py +++ b/examples/demo_network/run_remote_manager.py @@ -3,15 +3,12 @@ Demo: run_remote_manager ===================== """ -from rlberry.envs.gym_make import gym_make -from rlberry.network.client import BerryClient -from rlberry.network.interface import ResourceRequest - from rlberry.agents.torch import REINFORCEAgent - +from rlberry.envs.gym_make import gym_make from rlberry.manager import ExperimentManager, MultipleManagers, RemoteExperimentManager from rlberry.manager.evaluation import evaluate_agents, plot_writer_data - +from rlberry.network.client import BerryClient +from rlberry.network.interface import ResourceRequest if __name__ == "__main__": port = int(input("Select server port: ")) diff --git a/examples/demo_network/run_server.py b/examples/demo_network/run_server.py index c1b6a15b5..0684838db 100644 --- a/examples/demo_network/run_server.py +++ b/examples/demo_network/run_server.py @@ -3,11 +3,11 @@ Demo: run_server ===================== """ -from rlberry.network.interface import ResourceItem -from rlberry.network.server import BerryServer from rlberry.agents import ValueIterationAgent -from rlberry.agents.torch import REINFORCEAgent, A2CAgent +from rlberry.agents.torch import A2CAgent, REINFORCEAgent from rlberry.envs import GridWorld, gym_make +from rlberry.network.interface import ResourceItem +from rlberry.network.server import BerryServer from rlberry.utils.writers import DefaultWriter if __name__ == "__main__": diff --git a/examples/plot_agent_manager.py b/examples/plot_agent_manager.py index 338ee417d..4e81f314d 100644 --- a/examples/plot_agent_manager.py +++ b/examples/plot_agent_manager.py @@ -31,6 +31,7 @@ env = env_ctor(**env_kwargs) import numpy as np + from rlberry.agents import AgentWithSimplePolicy diff --git a/examples/plot_checkpointing.py b/examples/plot_checkpointing.py index 3de689cb6..517f96494 100644 --- a/examples/plot_checkpointing.py +++ b/examples/plot_checkpointing.py @@ -9,8 +9,7 @@ your agents, and how to restore from a previous checkpoint. """ from rlberry.agents import Agent -from rlberry.manager import ExperimentManager -from rlberry.manager import plot_writer_data +from rlberry.manager import ExperimentManager, plot_writer_data class MyAgent(Agent): diff --git a/examples/plot_kernels.py b/examples/plot_kernels.py index 84b2b2cfc..1007c4e3e 100644 --- a/examples/plot_kernels.py +++ b/examples/plot_kernels.py @@ -8,6 +8,7 @@ import matplotlib.pyplot as plt import numpy as np + from rlberry.agents.kernel_based.kernels import kernel_func kernel_types = [ diff --git a/examples/plot_writer_wrapper.py b/examples/plot_writer_wrapper.py index 069d4de00..ce52611e8 100644 --- a/examples/plot_writer_wrapper.py +++ b/examples/plot_writer_wrapper.py @@ -20,13 +20,13 @@ """ +import matplotlib.pyplot as plt import numpy as np -from rlberry.wrappers import WriterWrapper -from rlberry.envs import GridWorld -from rlberry.manager import plot_writer_data, ExperimentManager from rlberry.agents import UCBVIAgent -import matplotlib.pyplot as plt +from rlberry.envs import GridWorld +from rlberry.manager import ExperimentManager, plot_writer_data +from rlberry.wrappers import WriterWrapper # We wrape the default writer of the agent in a WriterWrapper # to record rewards. diff --git a/long_tests/rl_agent/ltest_mbqvi_applegold.py b/long_tests/rl_agent/ltest_mbqvi_applegold.py index 5aa85aa25..5267f1be6 100644 --- a/long_tests/rl_agent/ltest_mbqvi_applegold.py +++ b/long_tests/rl_agent/ltest_mbqvi_applegold.py @@ -1,7 +1,8 @@ -from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold +import numpy as np + from rlberry.agents.mbqvi import MBQVIAgent +from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold from rlberry.manager import ExperimentManager, evaluate_agents -import numpy as np params = {} params["n_samples"] = 8 # samples per state-action pair diff --git a/long_tests/torch_agent/ltest_a2c_cartpole.py b/long_tests/torch_agent/ltest_a2c_cartpole.py index ce22dc0f3..ac6f65a24 100644 --- a/long_tests/torch_agent/ltest_a2c_cartpole.py +++ b/long_tests/torch_agent/ltest_a2c_cartpole.py @@ -1,8 +1,9 @@ -from rlberry.envs import gym_make +import numpy as np + from rlberry.agents.torch import A2CAgent -from rlberry.manager import ExperimentManager from rlberry.agents.torch.utils.training import model_factory_from_env -import numpy as np +from rlberry.envs import gym_make +from rlberry.manager import ExperimentManager # Using parameters from deeprl quick start policy_configs = { diff --git a/long_tests/torch_agent/ltest_ctn_ppo_a2c_pendulum.py b/long_tests/torch_agent/ltest_ctn_ppo_a2c_pendulum.py index 01b10b29f..4b72038ae 100644 --- a/long_tests/torch_agent/ltest_ctn_ppo_a2c_pendulum.py +++ b/long_tests/torch_agent/ltest_ctn_ppo_a2c_pendulum.py @@ -1,8 +1,9 @@ -from rlberry.envs import gym_make -from rlberry.agents.torch import A2CAgent, PPOAgent -from rlberry.manager import ExperimentManager, plot_writer_data, evaluate_agents -import seaborn as sns import matplotlib.pyplot as plt +import seaborn as sns + +from rlberry.agents.torch import A2CAgent, PPOAgent +from rlberry.envs import gym_make +from rlberry.manager import ExperimentManager, evaluate_agents, plot_writer_data def test_a2c_vs_ppo_pendul(): diff --git a/long_tests/torch_agent/ltest_dqn_montaincar.py b/long_tests/torch_agent/ltest_dqn_montaincar.py index 553fc9580..89feb5db8 100644 --- a/long_tests/torch_agent/ltest_dqn_montaincar.py +++ b/long_tests/torch_agent/ltest_dqn_montaincar.py @@ -1,7 +1,8 @@ -from rlberry.envs import gym_make +import numpy as np + from rlberry.agents.torch import DQNAgent +from rlberry.envs import gym_make from rlberry.manager import ExperimentManager, evaluate_agents -import numpy as np model_configs = { "type": "MultiLayerPerceptron", diff --git a/long_tests/torch_agent/ltest_dqn_vs_mdqn_acrobot.py b/long_tests/torch_agent/ltest_dqn_vs_mdqn_acrobot.py index 9d17ea936..00b33ee72 100644 --- a/long_tests/torch_agent/ltest_dqn_vs_mdqn_acrobot.py +++ b/long_tests/torch_agent/ltest_dqn_vs_mdqn_acrobot.py @@ -1,9 +1,10 @@ -from rlberry.envs import gym_make +import matplotlib.pyplot as plt +import seaborn as sns + from rlberry.agents.torch import DQNAgent from rlberry.agents.torch import MunchausenDQNAgent as MDQNAgent +from rlberry.envs import gym_make from rlberry.manager import ExperimentManager, evaluate_agents, plot_writer_data -import matplotlib.pyplot as plt -import seaborn as sns def test_dqn_vs_mdqn_acro(): diff --git a/rlberry/__init__.py b/rlberry/__init__.py index c7dd604b3..27e793cc2 100644 --- a/rlberry/__init__.py +++ b/rlberry/__init__.py @@ -1,11 +1,11 @@ -from ._version import __version__ import logging +from ._version import __version__ + logger = logging.getLogger("rlberry_logger") from rlberry.utils.logging import configure_logging - __path__ = __import__("pkgutil").extend_path(__path__, __name__) # Initialize logging level diff --git a/rlberry/agents/__init__.py b/rlberry/agents/__init__.py index 60fd5b8a4..e8c91dd9d 100644 --- a/rlberry/agents/__init__.py +++ b/rlberry/agents/__init__.py @@ -1,17 +1,14 @@ # Interfaces -from .agent import Agent -from .agent import AgentWithSimplePolicy -from .agent import AgentTorch - # Basic agents (in alphabetical order) # basic = does not require torch, jax, etc... from .adaptiveql import AdaptiveQLAgent +from .agent import Agent, AgentTorch, AgentWithSimplePolicy from .dynprog import ValueIterationAgent -from .kernel_based import RSUCBVIAgent, RSKernelUCBVIAgent +from .kernel_based import RSKernelUCBVIAgent, RSUCBVIAgent from .linear import LSVIUCBAgent from .mbqvi import MBQVIAgent from .optql import OptQLAgent from .psrl import PSRLAgent from .rlsvi import RLSVIAgent -from .ucbvi import UCBVIAgent from .tabular_rl import QLAgent, SARSAAgent +from .ucbvi import UCBVIAgent diff --git a/rlberry/agents/adaptiveql/adaptiveql.py b/rlberry/agents/adaptiveql/adaptiveql.py index 667ed54e0..e06cb96a7 100644 --- a/rlberry/agents/adaptiveql/adaptiveql.py +++ b/rlberry/agents/adaptiveql/adaptiveql.py @@ -1,9 +1,9 @@ import gymnasium.spaces as spaces import numpy as np -from rlberry.agents import AgentWithSimplePolicy -from rlberry.agents.adaptiveql.tree import MDPTreePartition import rlberry +from rlberry.agents import AgentWithSimplePolicy +from rlberry.agents.adaptiveql.tree import MDPTreePartition logger = rlberry.logger diff --git a/rlberry/agents/adaptiveql/tree.py b/rlberry/agents/adaptiveql/tree.py index 4aaeb7948..9e4395ac8 100644 --- a/rlberry/agents/adaptiveql/tree.py +++ b/rlberry/agents/adaptiveql/tree.py @@ -1,6 +1,7 @@ import gymnasium.spaces as spaces -import numpy as np import matplotlib.pyplot as plt +import numpy as np + from rlberry.agents.adaptiveql.utils import bounds_contains, split_bounds diff --git a/rlberry/agents/agent.py b/rlberry/agents/agent.py index a7c81e5fb..80c4a8521 100644 --- a/rlberry/agents/agent.py +++ b/rlberry/agents/agent.py @@ -1,21 +1,21 @@ -from abc import ABC, abstractmethod -import dill -import pickle import bz2 -import _pickle as cPickle -import numpy as np +import inspect +import pickle +from abc import ABC, abstractmethod from inspect import signature from pathlib import Path -from rlberry import metadata_utils -from rlberry import types -from rlberry.seeding.seeder import Seeder -from rlberry.seeding import safe_reseed -from rlberry.envs.utils import process_env -from rlberry.utils.writers import DefaultWriter from typing import Optional -import inspect + +import _pickle as cPickle +import dill +import numpy as np import rlberry +from rlberry import metadata_utils, types +from rlberry.envs.utils import process_env +from rlberry.seeding import safe_reseed +from rlberry.seeding.seeder import Seeder +from rlberry.utils.writers import DefaultWriter logger = rlberry.logger @@ -669,9 +669,10 @@ def load(cls, filename, **kwargs): Arguments to required by the __init__ method of the Agent subclass. """ - from rlberry.utils.torch import choose_device import torch + from rlberry.utils.torch import choose_device + device_str = "cuda:best" if "device" in kwargs.keys(): device_str = kwargs.pop("device", None) diff --git a/rlberry/agents/bandits/__init__.py b/rlberry/agents/bandits/__init__.py index b35c171cf..7dd7a0aef 100644 --- a/rlberry/agents/bandits/__init__.py +++ b/rlberry/agents/bandits/__init__.py @@ -11,9 +11,6 @@ makeSubgaussianMOSSIndex, makeSubgaussianUCBIndex, ) -from .priors import ( - makeBetaPrior, - makeGaussianPrior, -) +from .priors import makeBetaPrior, makeGaussianPrior from .randomized_agents import RandomizedAgent from .ts_agents import TSAgent diff --git a/rlberry/agents/bandits/bandit_base.py b/rlberry/agents/bandits/bandit_base.py index 2558b30a2..e2d9a06d7 100644 --- a/rlberry/agents/bandits/bandit_base.py +++ b/rlberry/agents/bandits/bandit_base.py @@ -1,11 +1,12 @@ -import numpy as np -from rlberry.agents import AgentWithSimplePolicy -from .tools import BanditTracker import pickle - from pathlib import Path +import numpy as np + import rlberry +from rlberry.agents import AgentWithSimplePolicy + +from .tools import BanditTracker logger = rlberry.logger diff --git a/rlberry/agents/bandits/index_agents.py b/rlberry/agents/bandits/index_agents.py index c7335aef7..e43287086 100644 --- a/rlberry/agents/bandits/index_agents.py +++ b/rlberry/agents/bandits/index_agents.py @@ -1,8 +1,7 @@ import numpy as np -from rlberry.agents.bandits import BanditWithSimplePolicy - import rlberry +from rlberry.agents.bandits import BanditWithSimplePolicy logger = rlberry.logger diff --git a/rlberry/agents/bandits/indices.py b/rlberry/agents/bandits/indices.py index ebea3ac3f..a367153de 100644 --- a/rlberry/agents/bandits/indices.py +++ b/rlberry/agents/bandits/indices.py @@ -1,6 +1,7 @@ -import numpy as np from typing import Callable +import numpy as np + def makeETCIndex(A: int = 2, m: int = 1): """ diff --git a/rlberry/agents/bandits/randomized_agents.py b/rlberry/agents/bandits/randomized_agents.py index 76a82c97f..ee88447c5 100644 --- a/rlberry/agents/bandits/randomized_agents.py +++ b/rlberry/agents/bandits/randomized_agents.py @@ -1,8 +1,7 @@ import numpy as np -from rlberry.agents.bandits import BanditWithSimplePolicy - import rlberry +from rlberry.agents.bandits import BanditWithSimplePolicy logger = rlberry.logger diff --git a/rlberry/agents/bandits/tools/tracker.py b/rlberry/agents/bandits/tools/tracker.py index 0aa6a17b0..819707aa3 100644 --- a/rlberry/agents/bandits/tools/tracker.py +++ b/rlberry/agents/bandits/tools/tracker.py @@ -1,8 +1,7 @@ +import rlberry from rlberry import metadata_utils from rlberry.utils.writers import DefaultWriter -import rlberry - logger = rlberry.logger diff --git a/rlberry/agents/bandits/ts_agents.py b/rlberry/agents/bandits/ts_agents.py index 528fae0c0..37423e124 100644 --- a/rlberry/agents/bandits/ts_agents.py +++ b/rlberry/agents/bandits/ts_agents.py @@ -1,8 +1,7 @@ import numpy as np -from rlberry.agents.bandits import BanditWithSimplePolicy - import rlberry +from rlberry.agents.bandits import BanditWithSimplePolicy logger = rlberry.logger diff --git a/rlberry/agents/dynprog/utils.py b/rlberry/agents/dynprog/utils.py index 0d01c93b4..6de6cf979 100644 --- a/rlberry/agents/dynprog/utils.py +++ b/rlberry/agents/dynprog/utils.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.utils.jit_setup import numba_jit diff --git a/rlberry/agents/kernel_based/__init__.py b/rlberry/agents/kernel_based/__init__.py index 275e51c86..06093b764 100644 --- a/rlberry/agents/kernel_based/__init__.py +++ b/rlberry/agents/kernel_based/__init__.py @@ -1,2 +1,2 @@ -from .rs_ucbvi import RSUCBVIAgent from .rs_kernel_ucbvi import RSKernelUCBVIAgent +from .rs_ucbvi import RSUCBVIAgent diff --git a/rlberry/agents/kernel_based/common.py b/rlberry/agents/kernel_based/common.py index 33757f66a..bf734641e 100644 --- a/rlberry/agents/kernel_based/common.py +++ b/rlberry/agents/kernel_based/common.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.utils.jit_setup import numba_jit from rlberry.utils.metrics import metric_lp diff --git a/rlberry/agents/kernel_based/kernels.py b/rlberry/agents/kernel_based/kernels.py index 88954432c..bd4e3f5b9 100644 --- a/rlberry/agents/kernel_based/kernels.py +++ b/rlberry/agents/kernel_based/kernels.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.utils.jit_setup import numba_jit diff --git a/rlberry/agents/kernel_based/rs_kernel_ucbvi.py b/rlberry/agents/kernel_based/rs_kernel_ucbvi.py index f27449577..6d6e27bf8 100644 --- a/rlberry/agents/kernel_based/rs_kernel_ucbvi.py +++ b/rlberry/agents/kernel_based/rs_kernel_ucbvi.py @@ -1,15 +1,13 @@ +import gymnasium.spaces as spaces import numpy as np -from rlberry.utils.jit_setup import numba_jit -import gymnasium.spaces as spaces +import rlberry from rlberry.agents import AgentWithSimplePolicy -from rlberry.agents.dynprog.utils import backward_induction -from rlberry.agents.dynprog.utils import backward_induction_in_place -from rlberry.utils.metrics import metric_lp -from rlberry.agents.kernel_based.kernels import kernel_func +from rlberry.agents.dynprog.utils import backward_induction, backward_induction_in_place from rlberry.agents.kernel_based.common import map_to_representative - -import rlberry +from rlberry.agents.kernel_based.kernels import kernel_func +from rlberry.utils.jit_setup import numba_jit +from rlberry.utils.metrics import metric_lp logger = rlberry.logger diff --git a/rlberry/agents/kernel_based/rs_ucbvi.py b/rlberry/agents/kernel_based/rs_ucbvi.py index cee45ce56..78c507e46 100644 --- a/rlberry/agents/kernel_based/rs_ucbvi.py +++ b/rlberry/agents/kernel_based/rs_ucbvi.py @@ -1,12 +1,10 @@ -from rlberry.agents.agent import AgentWithSimplePolicy -import numpy as np - import gymnasium.spaces as spaces -from rlberry.agents.dynprog.utils import backward_induction -from rlberry.agents.dynprog.utils import backward_induction_in_place -from rlberry.agents.kernel_based.common import map_to_representative +import numpy as np import rlberry +from rlberry.agents.agent import AgentWithSimplePolicy +from rlberry.agents.dynprog.utils import backward_induction, backward_induction_in_place +from rlberry.agents.kernel_based.common import map_to_representative logger = rlberry.logger diff --git a/rlberry/agents/linear/lsvi_ucb.py b/rlberry/agents/linear/lsvi_ucb.py index e777d05d9..3c4f6b124 100644 --- a/rlberry/agents/linear/lsvi_ucb.py +++ b/rlberry/agents/linear/lsvi_ucb.py @@ -1,9 +1,9 @@ import numpy as np -from rlberry.agents import AgentWithSimplePolicy from gymnasium.spaces import Discrete -from rlberry.utils.jit_setup import numba_jit import rlberry +from rlberry.agents import AgentWithSimplePolicy +from rlberry.utils.jit_setup import numba_jit logger = rlberry.logger diff --git a/rlberry/agents/mbqvi/mbqvi.py b/rlberry/agents/mbqvi/mbqvi.py index 83031a168..bcd93ff16 100644 --- a/rlberry/agents/mbqvi/mbqvi.py +++ b/rlberry/agents/mbqvi/mbqvi.py @@ -1,11 +1,9 @@ import numpy as np - - -from rlberry.agents import AgentWithSimplePolicy -from rlberry.agents.dynprog.utils import backward_induction, value_iteration from gymnasium.spaces import Discrete import rlberry +from rlberry.agents import AgentWithSimplePolicy +from rlberry.agents.dynprog.utils import backward_induction, value_iteration logger = rlberry.logger diff --git a/rlberry/agents/optql/optql.py b/rlberry/agents/optql/optql.py index 951d1d834..0608166c1 100644 --- a/rlberry/agents/optql/optql.py +++ b/rlberry/agents/optql/optql.py @@ -1,11 +1,10 @@ +import gymnasium.spaces as spaces import numpy as np -import gymnasium.spaces as spaces +import rlberry from rlberry.agents import AgentWithSimplePolicy from rlberry.exploration_tools.discrete_counter import DiscreteCounter -import rlberry - logger = rlberry.logger diff --git a/rlberry/agents/psrl/psrl.py b/rlberry/agents/psrl/psrl.py index dac8440b9..3c3850602 100644 --- a/rlberry/agents/psrl/psrl.py +++ b/rlberry/agents/psrl/psrl.py @@ -1,14 +1,13 @@ +import gymnasium.spaces as spaces import numpy as np -import gymnasium.spaces as spaces +import rlberry from rlberry.agents import AgentWithSimplePolicy -from rlberry.exploration_tools.discrete_counter import DiscreteCounter from rlberry.agents.dynprog.utils import ( backward_induction_in_place, backward_induction_sd, ) - -import rlberry +from rlberry.exploration_tools.discrete_counter import DiscreteCounter logger = rlberry.logger diff --git a/rlberry/agents/rlsvi/rlsvi.py b/rlberry/agents/rlsvi/rlsvi.py index 6e3c2c120..6451b2ce9 100644 --- a/rlberry/agents/rlsvi/rlsvi.py +++ b/rlberry/agents/rlsvi/rlsvi.py @@ -1,15 +1,14 @@ +import gymnasium.spaces as spaces import numpy as np -import gymnasium.spaces as spaces +import rlberry from rlberry.agents import AgentWithSimplePolicy -from rlberry.exploration_tools.discrete_counter import DiscreteCounter from rlberry.agents.dynprog.utils import ( backward_induction_in_place, backward_induction_reward_sd, backward_induction_sd, ) - -import rlberry +from rlberry.exploration_tools.discrete_counter import DiscreteCounter logger = rlberry.logger diff --git a/rlberry/agents/stable_baselines/stable_baselines.py b/rlberry/agents/stable_baselines/stable_baselines.py index ffc43d099..d7740bf78 100644 --- a/rlberry/agents/stable_baselines/stable_baselines.py +++ b/rlberry/agents/stable_baselines/stable_baselines.py @@ -2,17 +2,14 @@ from typing import Any, Dict, Optional, Tuple, Type, Union import dill -from stable_baselines3.common import utils import stable_baselines3.common.logger as sb_logging +from stable_baselines3.common import utils from stable_baselines3.common.base_class import BaseAlgorithm as SB3Algorithm from stable_baselines3.common.policies import BasePolicy as SB3Policy -from rlberry import metadata_utils -from rlberry import types -from rlberry.agents import AgentWithSimplePolicy - - import rlberry +from rlberry import metadata_utils, types +from rlberry.agents import AgentWithSimplePolicy logger = rlberry.logger diff --git a/rlberry/agents/tabular_rl/qlearning.py b/rlberry/agents/tabular_rl/qlearning.py index 024147acb..88218121a 100644 --- a/rlberry/agents/tabular_rl/qlearning.py +++ b/rlberry/agents/tabular_rl/qlearning.py @@ -1,4 +1,5 @@ -from typing import Optional, Literal +from typing import Literal, Optional + import numpy as np from gymnasium import spaces from scipy.special import softmax diff --git a/rlberry/agents/tabular_rl/sarsa.py b/rlberry/agents/tabular_rl/sarsa.py index 3f097d4d4..9c2bad642 100644 --- a/rlberry/agents/tabular_rl/sarsa.py +++ b/rlberry/agents/tabular_rl/sarsa.py @@ -1,4 +1,5 @@ -from typing import Optional, Literal +from typing import Literal, Optional + import numpy as np from gymnasium import spaces from scipy.special import softmax diff --git a/rlberry/agents/tests/test_adaptiveql.py b/rlberry/agents/tests/test_adaptiveql.py index 4079dcbf9..0a714d5c8 100644 --- a/rlberry/agents/tests/test_adaptiveql.py +++ b/rlberry/agents/tests/test_adaptiveql.py @@ -1,6 +1,7 @@ +import matplotlib.pyplot as plt + from rlberry.agents import AdaptiveQLAgent from rlberry.envs.benchmarks.ball_exploration.ball2d import get_benchmark_env -import matplotlib.pyplot as plt def test_adaptive_ql(): diff --git a/rlberry/agents/tests/test_bandits.py b/rlberry/agents/tests/test_bandits.py index 441e5a3f0..fa7e243c9 100644 --- a/rlberry/agents/tests/test_bandits.py +++ b/rlberry/agents/tests/test_bandits.py @@ -1,24 +1,23 @@ -from rlberry.envs.bandits import NormalBandit, BernoulliBandit from rlberry.agents.bandits import ( + BanditWithSimplePolicy, IndexAgent, RandomizedAgent, TSAgent, - BanditWithSimplePolicy, makeBetaPrior, makeBoundedIMEDIndex, makeBoundedMOSSIndex, makeBoundedNPTSIndex, makeBoundedUCBIndex, + makeBoundedUCBVIndex, makeETCIndex, - makeGaussianPrior, makeEXP3Index, + makeGaussianPrior, makeSubgaussianMOSSIndex, makeSubgaussianUCBIndex, - makeBoundedUCBVIndex, ) +from rlberry.envs.bandits import BernoulliBandit, NormalBandit from rlberry.utils import check_bandit_agent - TEST_SEED = 42 diff --git a/rlberry/agents/tests/test_dynprog.py b/rlberry/agents/tests/test_dynprog.py index 6d96b8f49..21a55b487 100644 --- a/rlberry/agents/tests/test_dynprog.py +++ b/rlberry/agents/tests/test_dynprog.py @@ -3,12 +3,14 @@ import rlberry.seeding as seeding from rlberry.agents.dynprog import ValueIterationAgent -from rlberry.agents.dynprog.utils import backward_induction -from rlberry.agents.dynprog.utils import backward_induction_in_place -from rlberry.agents.dynprog.utils import backward_induction_sd -from rlberry.agents.dynprog.utils import backward_induction_reward_sd -from rlberry.agents.dynprog.utils import bellman_operator -from rlberry.agents.dynprog.utils import value_iteration +from rlberry.agents.dynprog.utils import ( + backward_induction, + backward_induction_in_place, + backward_induction_reward_sd, + backward_induction_sd, + bellman_operator, + value_iteration, +) from rlberry.envs.finite import FiniteMDP _rng = seeding.Seeder(123).rng diff --git a/rlberry/agents/tests/test_kernel_based.py b/rlberry/agents/tests/test_kernel_based.py index 65abac706..df9e5b10e 100644 --- a/rlberry/agents/tests/test_kernel_based.py +++ b/rlberry/agents/tests/test_kernel_based.py @@ -1,6 +1,6 @@ import pytest -from rlberry.agents.kernel_based import RSKernelUCBVIAgent -from rlberry.agents.kernel_based import RSUCBVIAgent + +from rlberry.agents.kernel_based import RSKernelUCBVIAgent, RSUCBVIAgent from rlberry.agents.kernel_based.kernels import _str_to_int from rlberry.envs.benchmarks.ball_exploration.ball2d import get_benchmark_env diff --git a/rlberry/agents/tests/test_lsvi_ucb.py b/rlberry/agents/tests/test_lsvi_ucb.py index 03299b747..7d55a81bd 100644 --- a/rlberry/agents/tests/test_lsvi_ucb.py +++ b/rlberry/agents/tests/test_lsvi_ucb.py @@ -1,8 +1,9 @@ import numpy as np import pytest + +from rlberry.agents.dynprog import ValueIterationAgent from rlberry.agents.features import FeatureMap from rlberry.agents.linear.lsvi_ucb import LSVIUCBAgent -from rlberry.agents.dynprog import ValueIterationAgent from rlberry.envs.finite import GridWorld diff --git a/rlberry/agents/tests/test_mbqvi.py b/rlberry/agents/tests/test_mbqvi.py index cafdb5566..3459c958e 100644 --- a/rlberry/agents/tests/test_mbqvi.py +++ b/rlberry/agents/tests/test_mbqvi.py @@ -1,9 +1,9 @@ import numpy as np import pytest -from rlberry.seeding import Seeder from rlberry.agents.mbqvi import MBQVIAgent from rlberry.envs.finite import FiniteMDP +from rlberry.seeding import Seeder @pytest.mark.parametrize("S, A", [(5, 2), (10, 4)]) diff --git a/rlberry/agents/tests/test_psrl.py b/rlberry/agents/tests/test_psrl.py index 325777f6d..1db4891e0 100644 --- a/rlberry/agents/tests/test_psrl.py +++ b/rlberry/agents/tests/test_psrl.py @@ -1,4 +1,5 @@ import pytest + from rlberry.agents.psrl import PSRLAgent from rlberry.envs.finite import GridWorld diff --git a/rlberry/agents/tests/test_replay.py b/rlberry/agents/tests/test_replay.py index bad1a297e..b942e2e5b 100644 --- a/rlberry/agents/tests/test_replay.py +++ b/rlberry/agents/tests/test_replay.py @@ -1,8 +1,9 @@ -import pytest import numpy as np +import pytest +from gymnasium.wrappers import TimeLimit + from rlberry.agents.utils import replay from rlberry.envs.finite import GridWorld -from gymnasium.wrappers import TimeLimit def _get_filled_replay(max_replay_size): diff --git a/rlberry/agents/tests/test_rlsvi.py b/rlberry/agents/tests/test_rlsvi.py index 0907d8d33..548a71bd8 100644 --- a/rlberry/agents/tests/test_rlsvi.py +++ b/rlberry/agents/tests/test_rlsvi.py @@ -1,4 +1,5 @@ import pytest + from rlberry.agents.rlsvi import RLSVIAgent from rlberry.envs.finite import GridWorld diff --git a/rlberry/agents/tests/test_stable_baselines.py b/rlberry/agents/tests/test_stable_baselines.py index 9186e03e3..aac0eb721 100644 --- a/rlberry/agents/tests/test_stable_baselines.py +++ b/rlberry/agents/tests/test_stable_baselines.py @@ -2,8 +2,8 @@ from stable_baselines3 import A2C -from rlberry.envs import gym_make from rlberry.agents.stable_baselines import StableBaselinesAgent +from rlberry.envs import gym_make from rlberry.utils.check_agent import check_rl_agent diff --git a/rlberry/agents/tests/test_tabular_rl.py b/rlberry/agents/tests/test_tabular_rl.py index ab7f618a3..ea3e0d972 100644 --- a/rlberry/agents/tests/test_tabular_rl.py +++ b/rlberry/agents/tests/test_tabular_rl.py @@ -1,4 +1,5 @@ import pytest + from rlberry.agents import QLAgent, SARSAAgent from rlberry.envs import GridWorld diff --git a/rlberry/agents/tests/test_ucbvi.py b/rlberry/agents/tests/test_ucbvi.py index 641fe0c02..1235c6bec 100644 --- a/rlberry/agents/tests/test_ucbvi.py +++ b/rlberry/agents/tests/test_ucbvi.py @@ -1,4 +1,5 @@ import pytest + from rlberry.agents.ucbvi import UCBVIAgent from rlberry.envs.finite import GridWorld diff --git a/rlberry/agents/torch/__init__.py b/rlberry/agents/torch/__init__.py index 896403dc6..effc2299e 100644 --- a/rlberry/agents/torch/__init__.py +++ b/rlberry/agents/torch/__init__.py @@ -1,7 +1,6 @@ # Torch agents (in alphabetical order) from .a2c import A2CAgent -from .dqn import DQNAgent -from .dqn import MunchausenDQNAgent +from .dqn import DQNAgent, MunchausenDQNAgent from .ppo import PPOAgent from .reinforce import REINFORCEAgent from .sac import SACAgent diff --git a/rlberry/agents/torch/a2c/a2c.py b/rlberry/agents/torch/a2c/a2c.py index 9907677e1..e0bdfb254 100644 --- a/rlberry/agents/torch/a2c/a2c.py +++ b/rlberry/agents/torch/a2c/a2c.py @@ -1,18 +1,20 @@ -import torch -import torch.nn as nn +from typing import Optional import gymnasium.spaces as spaces import numpy as np -from rlberry.agents import AgentWithSimplePolicy, AgentTorch -from rlberry.agents.utils.replay import ReplayBuffer -from rlberry.agents.torch.utils.training import optimizer_factory -from rlberry.agents.torch.utils.models import default_policy_net_fn -from rlberry.agents.torch.utils.models import default_value_net_fn -from rlberry.utils.torch import choose_device -from rlberry.utils.factory import load -from typing import Optional +import torch +import torch.nn as nn import rlberry +from rlberry.agents import AgentTorch, AgentWithSimplePolicy +from rlberry.agents.torch.utils.models import ( + default_policy_net_fn, + default_value_net_fn, +) +from rlberry.agents.torch.utils.training import optimizer_factory +from rlberry.agents.utils.replay import ReplayBuffer +from rlberry.utils.factory import load +from rlberry.utils.torch import choose_device logger = rlberry.logger diff --git a/rlberry/agents/torch/dqn/dqn.py b/rlberry/agents/torch/dqn/dqn.py index 84219c8c8..a1b8a2c65 100644 --- a/rlberry/agents/torch/dqn/dqn.py +++ b/rlberry/agents/torch/dqn/dqn.py @@ -1,25 +1,23 @@ import inspect from typing import Callable, Optional, Union -from gymnasium import spaces import numpy as np import torch +from gymnasium import spaces +import rlberry from rlberry import types -from rlberry.agents import AgentWithSimplePolicy, AgentTorch +from rlberry.agents import AgentTorch, AgentWithSimplePolicy +from rlberry.agents.torch.dqn.dqn_utils import lambda_returns, polynomial_schedule from rlberry.agents.torch.utils.training import ( loss_function_factory, model_factory, optimizer_factory, size_model_config, ) -from rlberry.agents.torch.dqn.dqn_utils import polynomial_schedule, lambda_returns from rlberry.agents.utils import replay -from rlberry.utils.torch import choose_device from rlberry.utils.factory import load - - -import rlberry +from rlberry.utils.torch import choose_device logger = rlberry.logger diff --git a/rlberry/agents/torch/dqn/dqn_utils.py b/rlberry/agents/torch/dqn/dqn_utils.py index 2d100b218..78a6881e7 100644 --- a/rlberry/agents/torch/dqn/dqn_utils.py +++ b/rlberry/agents/torch/dqn/dqn_utils.py @@ -2,11 +2,8 @@ import torch import torch.nn.functional as F - -from rlberry.utils.jit_setup import numba_jit - - import rlberry +from rlberry.utils.jit_setup import numba_jit logger = rlberry.logger diff --git a/rlberry/agents/torch/dqn/mdqn.py b/rlberry/agents/torch/dqn/mdqn.py index 746e01d24..39db06725 100644 --- a/rlberry/agents/torch/dqn/mdqn.py +++ b/rlberry/agents/torch/dqn/mdqn.py @@ -1,29 +1,28 @@ import inspect +from typing import Callable, Optional, Union import numpy as np import torch from gymnasium import spaces + +import rlberry from rlberry import types -from rlberry.agents import AgentWithSimplePolicy, AgentTorch -from rlberry.agents.torch.utils.training import ( - loss_function_factory, - model_factory, - optimizer_factory, - size_model_config, -) +from rlberry.agents import AgentTorch, AgentWithSimplePolicy from rlberry.agents.torch.dqn.dqn_utils import ( lambda_returns, polynomial_schedule, stable_scaled_log_softmax, stable_softmax, ) +from rlberry.agents.torch.utils.training import ( + loss_function_factory, + model_factory, + optimizer_factory, + size_model_config, +) from rlberry.agents.utils import replay -from rlberry.utils.torch import choose_device from rlberry.utils.factory import load -from typing import Callable, Optional, Union - - -import rlberry +from rlberry.utils.torch import choose_device logger = rlberry.logger diff --git a/rlberry/agents/torch/ppo/ppo.py b/rlberry/agents/torch/ppo/ppo.py index fdd27442b..e114e8e14 100644 --- a/rlberry/agents/torch/ppo/ppo.py +++ b/rlberry/agents/torch/ppo/ppo.py @@ -1,29 +1,29 @@ +import bz2 +import pickle +from pathlib import Path + +import _pickle as cPickle +import dill +import gymnasium.spaces as spaces import numpy as np import torch import torch.nn as nn -import gymnasium.spaces as spaces import rlberry -from rlberry.agents import AgentWithSimplePolicy -from rlberry.agents import AgentTorch -from rlberry.envs.utils import process_env -from rlberry.agents.torch.utils.training import optimizer_factory -from rlberry.agents.torch.utils.models import default_policy_net_fn -from rlberry.agents.torch.utils.models import default_value_net_fn -from rlberry.utils.torch import choose_device -from rlberry.utils.factory import load +from rlberry.agents import AgentTorch, AgentWithSimplePolicy from rlberry.agents.torch.ppo.ppo_utils import ( - process_ppo_env, - lambda_returns, RolloutBuffer, + lambda_returns, + process_ppo_env, ) - -import dill -import pickle -import bz2 -import _pickle as cPickle -from pathlib import Path - +from rlberry.agents.torch.utils.models import ( + default_policy_net_fn, + default_value_net_fn, +) +from rlberry.agents.torch.utils.training import optimizer_factory +from rlberry.envs.utils import process_env +from rlberry.utils.factory import load +from rlberry.utils.torch import choose_device logger = rlberry.logger diff --git a/rlberry/agents/torch/ppo/ppo_utils.py b/rlberry/agents/torch/ppo/ppo_utils.py index ec7f6df2f..955d8728c 100644 --- a/rlberry/agents/torch/ppo/ppo_utils.py +++ b/rlberry/agents/torch/ppo/ppo_utils.py @@ -7,7 +7,6 @@ from rlberry.envs.utils import process_env from rlberry.utils.jit_setup import numba_jit - logger = logging.getLogger(__name__) diff --git a/rlberry/agents/torch/reinforce/reinforce.py b/rlberry/agents/torch/reinforce/reinforce.py index f9f0c2e2f..1ba808220 100644 --- a/rlberry/agents/torch/reinforce/reinforce.py +++ b/rlberry/agents/torch/reinforce/reinforce.py @@ -1,15 +1,15 @@ -import torch import inspect -import numpy as np import gymnasium.spaces as spaces -from rlberry.agents import AgentWithSimplePolicy, AgentTorch -from rlberry.agents.utils.memories import Memory -from rlberry.agents.torch.utils.training import optimizer_factory -from rlberry.agents.torch.utils.models import default_policy_net_fn -from rlberry.utils.torch import choose_device +import numpy as np +import torch import rlberry +from rlberry.agents import AgentTorch, AgentWithSimplePolicy +from rlberry.agents.torch.utils.models import default_policy_net_fn +from rlberry.agents.torch.utils.training import optimizer_factory +from rlberry.agents.utils.memories import Memory +from rlberry.utils.torch import choose_device logger = rlberry.logger diff --git a/rlberry/agents/torch/sac/sac.py b/rlberry/agents/torch/sac/sac.py index 1828dc566..ea16d9ad4 100644 --- a/rlberry/agents/torch/sac/sac.py +++ b/rlberry/agents/torch/sac/sac.py @@ -2,10 +2,11 @@ import gymnasium.spaces as spaces import numpy as np -import rlberry import torch import torch.nn as nn import torch.optim as optim + +import rlberry from rlberry.agents import AgentTorch, AgentWithSimplePolicy from rlberry.agents.torch.sac.sac_utils import default_policy_net_fn, default_q_net_fn from rlberry.agents.torch.utils.training import optimizer_factory diff --git a/rlberry/agents/torch/tests/test_a2c.py b/rlberry/agents/torch/tests/test_a2c.py index 057649705..c16ea8eed 100644 --- a/rlberry/agents/torch/tests/test_a2c.py +++ b/rlberry/agents/torch/tests/test_a2c.py @@ -1,8 +1,9 @@ -from rlberry.envs import Wrapper +from gymnasium import make + from rlberry.agents.torch import A2CAgent -from rlberry.manager import ExperimentManager, evaluate_agents +from rlberry.envs import Wrapper from rlberry.envs.benchmarks.ball_exploration import PBall2D -from gymnasium import make +from rlberry.manager import ExperimentManager, evaluate_agents def test_a2c(): diff --git a/rlberry/agents/torch/tests/test_dqn.py b/rlberry/agents/torch/tests/test_dqn.py index 5b6848fb7..0d1648cd6 100644 --- a/rlberry/agents/torch/tests/test_dqn.py +++ b/rlberry/agents/torch/tests/test_dqn.py @@ -1,12 +1,13 @@ +import os +import pathlib +import tempfile + import pytest -from rlberry.envs import gym_make + from rlberry.agents.torch.dqn import DQNAgent from rlberry.agents.torch.utils.training import model_factory +from rlberry.envs import gym_make from rlberry.manager import ExperimentManager -import os -import pathlib - -import tempfile @pytest.mark.parametrize( diff --git a/rlberry/agents/torch/tests/test_factory.py b/rlberry/agents/torch/tests/test_factory.py index f1dddb92b..a3f096df3 100644 --- a/rlberry/agents/torch/tests/test_factory.py +++ b/rlberry/agents/torch/tests/test_factory.py @@ -1,4 +1,5 @@ import pytest + from rlberry.agents.torch.utils.training import model_factory diff --git a/rlberry/agents/torch/tests/test_mdqn.py b/rlberry/agents/torch/tests/test_mdqn.py index b327b8599..ae14a6fef 100644 --- a/rlberry/agents/torch/tests/test_mdqn.py +++ b/rlberry/agents/torch/tests/test_mdqn.py @@ -1,7 +1,8 @@ import pytest -from rlberry.envs import gym_make + from rlberry.agents.torch.dqn import MunchausenDQNAgent from rlberry.agents.torch.utils.training import model_factory +from rlberry.envs import gym_make @pytest.mark.parametrize("use_prioritized_replay", [(False), (True)]) diff --git a/rlberry/agents/torch/tests/test_ppo.py b/rlberry/agents/torch/tests/test_ppo.py index ed31465cb..88d2e36ce 100644 --- a/rlberry/agents/torch/tests/test_ppo.py +++ b/rlberry/agents/torch/tests/test_ppo.py @@ -7,14 +7,16 @@ # ppo = PPOAgent(env) # ppo.fit(4096) +import sys + import pytest -from rlberry.envs import Wrapper -from rlberry.agents.torch import PPOAgent -from rlberry.manager import ExperimentManager, evaluate_agents -from rlberry.envs.benchmarks.ball_exploration import PBall2D from gymnasium import make + +from rlberry.agents.torch import PPOAgent from rlberry.agents.torch.utils.training import model_factory_from_env -import sys +from rlberry.envs import Wrapper +from rlberry.envs.benchmarks.ball_exploration import PBall2D +from rlberry.manager import ExperimentManager, evaluate_agents @pytest.mark.timeout(300) diff --git a/rlberry/agents/torch/tests/test_sac.py b/rlberry/agents/torch/tests/test_sac.py index db5f5067f..fdac7e284 100644 --- a/rlberry/agents/torch/tests/test_sac.py +++ b/rlberry/agents/torch/tests/test_sac.py @@ -2,6 +2,7 @@ import pytest from gymnasium import make + from rlberry.agents.torch.sac import SACAgent from rlberry.envs import Wrapper from rlberry.manager import AgentManager, evaluate_agents diff --git a/rlberry/agents/torch/tests/test_torch_atari.py b/rlberry/agents/torch/tests/test_torch_atari.py index bb7629c78..d1f010f21 100644 --- a/rlberry/agents/torch/tests/test_torch_atari.py +++ b/rlberry/agents/torch/tests/test_torch_atari.py @@ -1,15 +1,15 @@ -from rlberry.manager import ExperimentManager -from rlberry.agents.torch.dqn.dqn import DQNAgent -from rlberry.envs.gym_make import atari_make - -from rlberry.agents.torch import PPOAgent -from rlberry.agents.torch.utils.training import model_factory_from_env +import os import pathlib +import tempfile + import numpy as np import pytest -import os -import tempfile +from rlberry.agents.torch import PPOAgent +from rlberry.agents.torch.dqn.dqn import DQNAgent +from rlberry.agents.torch.utils.training import model_factory_from_env +from rlberry.envs.gym_make import atari_make +from rlberry.manager import ExperimentManager def test_forward_dqn(): diff --git a/rlberry/agents/torch/tests/test_torch_models.py b/rlberry/agents/torch/tests/test_torch_models.py index 9bc692294..e720f5785 100644 --- a/rlberry/agents/torch/tests/test_torch_models.py +++ b/rlberry/agents/torch/tests/test_torch_models.py @@ -3,8 +3,12 @@ """ import torch -from rlberry.agents.torch.utils.models import MultiLayerPerceptron -from rlberry.agents.torch.utils.models import ConvolutionalNetwork, DuelingNetwork + +from rlberry.agents.torch.utils.models import ( + ConvolutionalNetwork, + DuelingNetwork, + MultiLayerPerceptron, +) def test_mlp(): diff --git a/rlberry/agents/torch/tests/test_torch_training.py b/rlberry/agents/torch/tests/test_torch_training.py index fe5fb722c..13188ebd8 100644 --- a/rlberry/agents/torch/tests/test_torch_training.py +++ b/rlberry/agents/torch/tests/test_torch_training.py @@ -1,7 +1,8 @@ import torch + +from rlberry.agents.torch.utils.models import default_policy_net_fn from rlberry.agents.torch.utils.training import loss_function_factory, optimizer_factory from rlberry.envs.benchmarks.ball_exploration.ball2d import get_benchmark_env -from rlberry.agents.torch.utils.models import default_policy_net_fn # loss_function_factory assert isinstance(loss_function_factory("l2"), torch.nn.MSELoss) diff --git a/rlberry/agents/torch/utils/models.py b/rlberry/agents/torch/utils/models.py index 709493995..f8a60cdf8 100644 --- a/rlberry/agents/torch/utils/models.py +++ b/rlberry/agents/torch/utils/models.py @@ -3,17 +3,16 @@ # from functools import partial - -from gymnasium import spaces -from gymnasium.vector.sync_vector_env import SyncVectorEnv -from gymnasium.vector.async_vector_env import AsyncVectorEnv import numpy as np import torch import torch.nn as nn import torch.nn.functional as F +from gymnasium import spaces +from gymnasium.vector.async_vector_env import AsyncVectorEnv +from gymnasium.vector.sync_vector_env import SyncVectorEnv from torch.distributions import Categorical, Normal -from rlberry.agents.torch.utils.training import model_factory, activation_factory +from rlberry.agents.torch.utils.training import activation_factory, model_factory def default_twinq_net_fn(env): diff --git a/rlberry/agents/torch/utils/training.py b/rlberry/agents/torch/utils/training.py index ed338b3bb..c2b817d7c 100644 --- a/rlberry/agents/torch/utils/training.py +++ b/rlberry/agents/torch/utils/training.py @@ -63,9 +63,9 @@ def model_factory(type="MultiLayerPerceptron", **kwargs) -> nn.Module: * :class:`~rlberry.agents.torch.utils.models.Table` """ from rlberry.agents.torch.utils.models import ( - MultiLayerPerceptron, - DuelingNetwork, ConvolutionalNetwork, + DuelingNetwork, + MultiLayerPerceptron, Table, ) diff --git a/rlberry/agents/ucbvi/ucbvi.py b/rlberry/agents/ucbvi/ucbvi.py index d5dfc4e67..d150d0c13 100644 --- a/rlberry/agents/ucbvi/ucbvi.py +++ b/rlberry/agents/ucbvi/ucbvi.py @@ -1,19 +1,18 @@ +import gymnasium.spaces as spaces import numpy as np -import gymnasium.spaces as spaces +import rlberry from rlberry.agents import AgentWithSimplePolicy +from rlberry.agents.dynprog.utils import ( + backward_induction_in_place, + backward_induction_reward_sd, + backward_induction_sd, +) from rlberry.agents.ucbvi.utils import ( update_value_and_get_action, update_value_and_get_action_sd, ) from rlberry.exploration_tools.discrete_counter import DiscreteCounter -from rlberry.agents.dynprog.utils import ( - backward_induction_sd, - backward_induction_reward_sd, -) -from rlberry.agents.dynprog.utils import backward_induction_in_place - -import rlberry logger = rlberry.logger diff --git a/rlberry/agents/utils/memories.py b/rlberry/agents/utils/memories.py index 1677efd19..4601bb894 100644 --- a/rlberry/agents/utils/memories.py +++ b/rlberry/agents/utils/memories.py @@ -1,6 +1,7 @@ -import numpy as np from collections import namedtuple +import numpy as np + Transition = namedtuple( "Transition", ("state", "action", "reward", "next_state", "terminal", "info") ) diff --git a/rlberry/agents/utils/replay.py b/rlberry/agents/utils/replay.py index ff4f0e795..2231dd3bc 100644 --- a/rlberry/agents/utils/replay.py +++ b/rlberry/agents/utils/replay.py @@ -2,13 +2,12 @@ New module aiming to replace memories.py """ -import numpy as np - from typing import NamedTuple -from rlberry.agents.utils import replay_utils +import numpy as np import rlberry +from rlberry.agents.utils import replay_utils logger = rlberry.logger diff --git a/rlberry/colab_utils/display_setup.py b/rlberry/colab_utils/display_setup.py index 302e589eb..9bcebc9cc 100644 --- a/rlberry/colab_utils/display_setup.py +++ b/rlberry/colab_utils/display_setup.py @@ -3,12 +3,13 @@ # import base64 -from pyvirtualdisplay import Display -from IPython import display as ipythondisplay # from IPython.display import clear_output from pathlib import Path +from IPython import display as ipythondisplay +from pyvirtualdisplay import Display + def show_video(filename=None, directory="./videos"): """ diff --git a/rlberry/envs/__init__.py b/rlberry/envs/__init__.py index 96d4442f0..6f472b35c 100644 --- a/rlberry/envs/__init__.py +++ b/rlberry/envs/__init__.py @@ -1,6 +1,6 @@ -from .gym_make import gym_make, atari_make from .basewrapper import Wrapper from .classic_control import Acrobot, MountainCar, Pendulum, SpringCartPole from .finite import Chain, FiniteMDP, GridWorld +from .gym_make import atari_make, gym_make from .interface import Model from .pipeline import PipelineEnv diff --git a/rlberry/envs/bandits/bandit_base.py b/rlberry/envs/bandits/bandit_base.py index 95ceeb1a2..69efa12df 100644 --- a/rlberry/envs/bandits/bandit_base.py +++ b/rlberry/envs/bandits/bandit_base.py @@ -1,10 +1,8 @@ from collections import deque - -from rlberry.envs.interface import Model -import rlberry.spaces as spaces - import rlberry +import rlberry.spaces as spaces +from rlberry.envs.interface import Model logger = rlberry.logger diff --git a/rlberry/envs/basewrapper.py b/rlberry/envs/basewrapper.py index f782f96f1..07564ee7f 100644 --- a/rlberry/envs/basewrapper.py +++ b/rlberry/envs/basewrapper.py @@ -1,7 +1,8 @@ import gymnasium as gym -from rlberry.seeding import Seeder, safe_reseed import numpy as np + from rlberry.envs.interface import Model +from rlberry.seeding import Seeder, safe_reseed from rlberry.spaces.from_gym import convert_space_from_gym diff --git a/rlberry/envs/benchmarks/ball_exploration/ball2d.py b/rlberry/envs/benchmarks/ball_exploration/ball2d.py index bb1b43d7e..de6408165 100644 --- a/rlberry/envs/benchmarks/ball_exploration/ball2d.py +++ b/rlberry/envs/benchmarks/ball_exploration/ball2d.py @@ -11,8 +11,8 @@ import numpy as np -from rlberry.wrappers.autoreset import AutoResetWrapper from rlberry.envs.benchmarks.ball_exploration.pball import PBall2D +from rlberry.wrappers.autoreset import AutoResetWrapper def get_benchmark_env(level=1): diff --git a/rlberry/envs/benchmarks/ball_exploration/pball.py b/rlberry/envs/benchmarks/ball_exploration/pball.py index 4f7e9c479..ed935d906 100644 --- a/rlberry/envs/benchmarks/ball_exploration/pball.py +++ b/rlberry/envs/benchmarks/ball_exploration/pball.py @@ -1,11 +1,9 @@ import numpy as np - +import rlberry import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D - -import rlberry +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene logger = rlberry.logger diff --git a/rlberry/envs/benchmarks/generalization/twinrooms.py b/rlberry/envs/benchmarks/generalization/twinrooms.py index f0619e96b..3b486eadf 100644 --- a/rlberry/envs/benchmarks/generalization/twinrooms.py +++ b/rlberry/envs/benchmarks/generalization/twinrooms.py @@ -1,11 +1,11 @@ import numpy as np + +import rlberry import rlberry.spaces as spaces from rlberry.envs import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import circle_shape -import rlberry - logger = rlberry.logger diff --git a/rlberry/envs/benchmarks/grid_exploration/apple_gold.py b/rlberry/envs/benchmarks/grid_exploration/apple_gold.py index 4a4599156..1b9beb472 100644 --- a/rlberry/envs/benchmarks/grid_exploration/apple_gold.py +++ b/rlberry/envs/benchmarks/grid_exploration/apple_gold.py @@ -1,9 +1,9 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.envs.finite import GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.envs.finite import GridWorld +from rlberry.rendering import GeometricPrimitive, Scene logger = rlberry.logger diff --git a/rlberry/envs/benchmarks/grid_exploration/four_room.py b/rlberry/envs/benchmarks/grid_exploration/four_room.py index b4e2d67a5..25c6464d2 100644 --- a/rlberry/envs/benchmarks/grid_exploration/four_room.py +++ b/rlberry/envs/benchmarks/grid_exploration/four_room.py @@ -1,8 +1,8 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.envs.finite import GridWorld import rlberry +import rlberry.spaces as spaces +from rlberry.envs.finite import GridWorld logger = rlberry.logger diff --git a/rlberry/envs/benchmarks/grid_exploration/nroom.py b/rlberry/envs/benchmarks/grid_exploration/nroom.py index 51cc0f279..9dec4937e 100644 --- a/rlberry/envs/benchmarks/grid_exploration/nroom.py +++ b/rlberry/envs/benchmarks/grid_exploration/nroom.py @@ -1,10 +1,11 @@ import math + import numpy as np -import rlberry.spaces as spaces -from rlberry.envs.finite import GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.envs.finite import GridWorld +from rlberry.rendering import GeometricPrimitive, Scene logger = rlberry.logger diff --git a/rlberry/envs/benchmarks/grid_exploration/six_room.py b/rlberry/envs/benchmarks/grid_exploration/six_room.py index 4af6fdb28..07b906953 100644 --- a/rlberry/envs/benchmarks/grid_exploration/six_room.py +++ b/rlberry/envs/benchmarks/grid_exploration/six_room.py @@ -1,9 +1,9 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.envs.finite import GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.envs.finite import GridWorld +from rlberry.rendering import GeometricPrimitive, Scene logger = rlberry.logger diff --git a/rlberry/envs/bullet3/pybullet_envs/__init__.py b/rlberry/envs/bullet3/pybullet_envs/__init__.py index 093f8d9eb..f2a0856ee 100644 --- a/rlberry/envs/bullet3/pybullet_envs/__init__.py +++ b/rlberry/envs/bullet3/pybullet_envs/__init__.py @@ -1,5 +1,5 @@ import gymnasium as gym -from gym.envs.registration import registry, make, spec +from gym.envs.registration import make, registry, spec def register(id, *args, **kvargs): diff --git a/rlberry/envs/bullet3/pybullet_envs/gym_pendulum_envs.py b/rlberry/envs/bullet3/pybullet_envs/gym_pendulum_envs.py index 32ce80c6a..fff013479 100644 --- a/rlberry/envs/bullet3/pybullet_envs/gym_pendulum_envs.py +++ b/rlberry/envs/bullet3/pybullet_envs/gym_pendulum_envs.py @@ -1,10 +1,10 @@ +import numpy as np from gym import spaces from pybullet_envs.env_bases import MJCFBaseBulletEnv from pybullet_envs.gym_pendulum_envs import InvertedPendulumBulletEnv from pybullet_envs.scene_abstract import SingleRobotEmptyScene from rlberry.envs.bullet3.pybullet_envs.robot_pendula import Pendulum, PendulumSwingup -import numpy as np class PendulumBulletEnv(InvertedPendulumBulletEnv): diff --git a/rlberry/envs/bullet3/pybullet_envs/robot_bases.py b/rlberry/envs/bullet3/pybullet_envs/robot_bases.py index d2dc50e75..29f73d9af 100644 --- a/rlberry/envs/bullet3/pybullet_envs/robot_bases.py +++ b/rlberry/envs/bullet3/pybullet_envs/robot_bases.py @@ -1,4 +1,5 @@ import os + import pybullet from pybullet_envs.robot_bases import MJCFBasedRobot, URDFBasedRobot diff --git a/rlberry/envs/classic_control/SpringCartPole.py b/rlberry/envs/classic_control/SpringCartPole.py index 4bfd5f634..73b01b6bd 100644 --- a/rlberry/envs/classic_control/SpringCartPole.py +++ b/rlberry/envs/classic_control/SpringCartPole.py @@ -3,9 +3,10 @@ """ import numpy as np + import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import bar_shape, circle_shape diff --git a/rlberry/envs/classic_control/__init__.py b/rlberry/envs/classic_control/__init__.py index a6cd76c14..056420b3c 100644 --- a/rlberry/envs/classic_control/__init__.py +++ b/rlberry/envs/classic_control/__init__.py @@ -1,4 +1,4 @@ -from .mountain_car import MountainCar from .acrobot import Acrobot +from .mountain_car import MountainCar from .pendulum import Pendulum from .SpringCartPole import SpringCartPole diff --git a/rlberry/envs/classic_control/acrobot.py b/rlberry/envs/classic_control/acrobot.py index 2404b66e1..aa4b5b4dd 100644 --- a/rlberry/envs/classic_control/acrobot.py +++ b/rlberry/envs/classic_control/acrobot.py @@ -11,9 +11,10 @@ """ import numpy as np + import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import bar_shape, circle_shape __copyright__ = "Copyright 2013, RLPy http://acl.mit.edu/RLPy" diff --git a/rlberry/envs/classic_control/mountain_car.py b/rlberry/envs/classic_control/mountain_car.py index ff3cb1335..55c9b0448 100644 --- a/rlberry/envs/classic_control/mountain_car.py +++ b/rlberry/envs/classic_control/mountain_car.py @@ -17,7 +17,7 @@ import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene class MountainCar(RenderInterface2D, Model): diff --git a/rlberry/envs/classic_control/pendulum.py b/rlberry/envs/classic_control/pendulum.py index 972db1ceb..642c87942 100644 --- a/rlberry/envs/classic_control/pendulum.py +++ b/rlberry/envs/classic_control/pendulum.py @@ -10,9 +10,10 @@ """ import numpy as np + import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, RenderInterface2D +from rlberry.rendering import RenderInterface2D, Scene from rlberry.rendering.common_shapes import bar_shape, circle_shape diff --git a/rlberry/envs/finite/__init__.py b/rlberry/envs/finite/__init__.py index 036e4520a..c3cd3305f 100644 --- a/rlberry/envs/finite/__init__.py +++ b/rlberry/envs/finite/__init__.py @@ -1,3 +1,3 @@ +from .chain import Chain from .finite_mdp import FiniteMDP from .gridworld import GridWorld -from .chain import Chain diff --git a/rlberry/envs/finite/chain.py b/rlberry/envs/finite/chain.py index da333d713..56673ae47 100644 --- a/rlberry/envs/finite/chain.py +++ b/rlberry/envs/finite/chain.py @@ -1,7 +1,7 @@ import numpy as np from rlberry.envs.finite import FiniteMDP -from rlberry.rendering import RenderInterface2D, Scene, GeometricPrimitive +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene class Chain(RenderInterface2D, FiniteMDP): diff --git a/rlberry/envs/finite/finite_mdp.py b/rlberry/envs/finite/finite_mdp.py index f10eda7dd..e80bd53bd 100644 --- a/rlberry/envs/finite/finite_mdp.py +++ b/rlberry/envs/finite/finite_mdp.py @@ -1,11 +1,9 @@ import numpy as np - +import rlberry import rlberry.spaces as spaces from rlberry.envs.interface import Model -import rlberry - logger = rlberry.logger diff --git a/rlberry/envs/finite/gridworld.py b/rlberry/envs/finite/gridworld.py index ce585317d..0770989e0 100644 --- a/rlberry/envs/finite/gridworld.py +++ b/rlberry/envs/finite/gridworld.py @@ -1,16 +1,12 @@ import matplotlib -import numpy as np - import matplotlib.pyplot as plt +import numpy as np from matplotlib import cm -from rlberry.envs.finite import FiniteMDP -from rlberry.envs.finite import gridworld_utils -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D -from rlberry.rendering.common_shapes import circle_shape - - import rlberry +from rlberry.envs.finite import FiniteMDP, gridworld_utils +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene +from rlberry.rendering.common_shapes import circle_shape logger = rlberry.logger diff --git a/rlberry/envs/gym_make.py b/rlberry/envs/gym_make.py index 21c6eba54..636ba0093 100644 --- a/rlberry/envs/gym_make.py +++ b/rlberry/envs/gym_make.py @@ -1,9 +1,8 @@ import gymnasium as gym - -from rlberry.envs.basewrapper import Wrapper import numpy as np from numpy import ndarray +from rlberry.envs.basewrapper import Wrapper # VERSION_ORIGINE = True VERSION_ORIGINE = False @@ -70,7 +69,6 @@ def atari_make(id, seed=None, **kwargs): NoopResetEnv, StickyActionEnv, ) - from stable_baselines3.common.monitor import Monitor # Default values for Atari_SB3_wrappers diff --git a/rlberry/envs/interface/model.py b/rlberry/envs/interface/model.py index 065507400..035cbf2e1 100644 --- a/rlberry/envs/interface/model.py +++ b/rlberry/envs/interface/model.py @@ -1,10 +1,10 @@ +import inspect + import gymnasium as gym import numpy as np -import inspect -from rlberry.seeding import Seeder - import rlberry +from rlberry.seeding import Seeder logger = rlberry.logger diff --git a/rlberry/envs/tests/test_bandits.py b/rlberry/envs/tests/test_bandits.py index 0e35fccf4..600f57429 100644 --- a/rlberry/envs/tests/test_bandits.py +++ b/rlberry/envs/tests/test_bandits.py @@ -1,13 +1,12 @@ import numpy as np -from rlberry.seeding import safe_reseed -from rlberry.seeding import Seeder + from rlberry.envs.bandits import ( AdversarialBandit, BernoulliBandit, - NormalBandit, CorruptedNormalBandit, + NormalBandit, ) - +from rlberry.seeding import Seeder, safe_reseed TEST_SEED = 42 diff --git a/rlberry/envs/tests/test_env_seeding.py b/rlberry/envs/tests/test_env_seeding.py index 26682e286..bee725a07 100644 --- a/rlberry/envs/tests/test_env_seeding.py +++ b/rlberry/envs/tests/test_env_seeding.py @@ -1,15 +1,15 @@ +from copy import deepcopy + import numpy as np import pytest -import rlberry.seeding as seeding -from copy import deepcopy -from rlberry.envs.classic_control import MountainCar, Acrobot, Pendulum -from rlberry.envs.finite import Chain -from rlberry.envs.finite import GridWorld +import rlberry.seeding as seeding +from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND +from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold from rlberry.envs.benchmarks.grid_exploration.four_room import FourRoom from rlberry.envs.benchmarks.grid_exploration.six_room import SixRoom -from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold -from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND +from rlberry.envs.classic_control import Acrobot, MountainCar, Pendulum +from rlberry.envs.finite import Chain, GridWorld classes = [ MountainCar, diff --git a/rlberry/envs/tests/test_gym_env_seeding.py b/rlberry/envs/tests/test_gym_env_seeding.py index 9b008e323..3e28e7c74 100644 --- a/rlberry/envs/tests/test_gym_env_seeding.py +++ b/rlberry/envs/tests/test_gym_env_seeding.py @@ -1,11 +1,12 @@ -from rlberry.seeding.seeding import safe_reseed +from copy import deepcopy + import gymnasium as gym import numpy as np import pytest -from rlberry.seeding import Seeder -from rlberry.envs import gym_make -from copy import deepcopy +from rlberry.envs import gym_make +from rlberry.seeding import Seeder +from rlberry.seeding.seeding import safe_reseed gym_envs = [ "Acrobot-v1", diff --git a/rlberry/envs/tests/test_gym_make.py b/rlberry/envs/tests/test_gym_make.py index 9ad80d2ee..b29cf9fcd 100644 --- a/rlberry/envs/tests/test_gym_make.py +++ b/rlberry/envs/tests/test_gym_make.py @@ -22,14 +22,16 @@ def test_atari_make(): def test_rendering_with_atari_make(): - from rlberry.manager import ExperimentManager - from rlberry.agents.torch import PPOAgent - from gymnasium.wrappers.record_video import RecordVideo import os - from rlberry.envs.gym_make import atari_make - from rlberry.agents.torch.utils.training import model_factory_from_env import tempfile + from gymnasium.wrappers.record_video import RecordVideo + + from rlberry.agents.torch import PPOAgent + from rlberry.agents.torch.utils.training import model_factory_from_env + from rlberry.envs.gym_make import atari_make + from rlberry.manager import ExperimentManager + with tempfile.TemporaryDirectory() as tmpdirname: policy_mlp_configs = { "type": "MultiLayerPerceptron", # A network architecture diff --git a/rlberry/envs/tests/test_instantiation.py b/rlberry/envs/tests/test_instantiation.py index d66722484..09b8fd8a6 100644 --- a/rlberry/envs/tests/test_instantiation.py +++ b/rlberry/envs/tests/test_instantiation.py @@ -1,16 +1,15 @@ import numpy as np import pytest -from rlberry.envs import gym_make, PipelineEnv -from rlberry.envs.classic_control import MountainCar, Acrobot, Pendulum -from rlberry.envs.finite import Chain -from rlberry.envs.finite import GridWorld +from rlberry.envs import PipelineEnv, gym_make from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND from rlberry.envs.benchmarks.ball_exploration.ball2d import get_benchmark_env +from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold from rlberry.envs.benchmarks.grid_exploration.four_room import FourRoom -from rlberry.envs.benchmarks.grid_exploration.six_room import SixRoom from rlberry.envs.benchmarks.grid_exploration.nroom import NRoom -from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold +from rlberry.envs.benchmarks.grid_exploration.six_room import SixRoom +from rlberry.envs.classic_control import Acrobot, MountainCar, Pendulum +from rlberry.envs.finite import Chain, GridWorld from rlberry.rendering.render_interface import RenderInterface2D classes = [ diff --git a/rlberry/envs/tests/test_spring_env.py b/rlberry/envs/tests/test_spring_env.py index 9809cd55b..0df034396 100644 --- a/rlberry/envs/tests/test_spring_env.py +++ b/rlberry/envs/tests/test_spring_env.py @@ -1,8 +1,8 @@ import numpy as np + from rlberry.envs import SpringCartPole from rlberry.envs.classic_control.SpringCartPole import rk4 - # # actions # LL = 0 # RR = 1 diff --git a/rlberry/envs/utils.py b/rlberry/envs/utils.py index ba38054ee..1084583e2 100644 --- a/rlberry/envs/utils.py +++ b/rlberry/envs/utils.py @@ -1,9 +1,8 @@ -from typing import Tuple from copy import deepcopy -from rlberry.seeding import safe_reseed - +from typing import Tuple import rlberry +from rlberry.seeding import safe_reseed logger = rlberry.logger diff --git a/rlberry/experiment/__init__.py b/rlberry/experiment/__init__.py index 3205b4c85..7b0280588 100644 --- a/rlberry/experiment/__init__.py +++ b/rlberry/experiment/__init__.py @@ -1,3 +1,3 @@ -from .yaml_utils import parse_experiment_config from .generator import experiment_generator from .load_results import load_experiment_results +from .yaml_utils import parse_experiment_config diff --git a/rlberry/experiment/generator.py b/rlberry/experiment/generator.py index 1dbd6f1f8..d74851977 100644 --- a/rlberry/experiment/generator.py +++ b/rlberry/experiment/generator.py @@ -13,13 +13,14 @@ --max_workers= Number of workers used by ExperimentManager.fit. Set to -1 for the maximum value. [default: -1] """ -from docopt import docopt from pathlib import Path -from rlberry.experiment.yaml_utils import parse_experiment_config -from rlberry.manager import ExperimentManager -from rlberry import check_packages + +from docopt import docopt import rlberry +from rlberry import check_packages +from rlberry.experiment.yaml_utils import parse_experiment_config +from rlberry.manager import ExperimentManager logger = rlberry.logger diff --git a/rlberry/experiment/load_results.py b/rlberry/experiment/load_results.py index f66819f5d..024197802 100644 --- a/rlberry/experiment/load_results.py +++ b/rlberry/experiment/load_results.py @@ -1,9 +1,9 @@ from pathlib import Path -from rlberry.manager import ExperimentManager -import pandas as pd +import pandas as pd import rlberry +from rlberry.manager import ExperimentManager logger = rlberry.logger diff --git a/rlberry/experiment/tests/test_experiment_generator.py b/rlberry/experiment/tests/test_experiment_generator.py index 2c5297198..c9f8b1a5e 100644 --- a/rlberry/experiment/tests/test_experiment_generator.py +++ b/rlberry/experiment/tests/test_experiment_generator.py @@ -1,8 +1,8 @@ -from rlberry.experiment import experiment_generator -from rlberry.agents.kernel_based.rs_ucbvi import RSUCBVIAgent - import numpy as np +from rlberry.agents.kernel_based.rs_ucbvi import RSUCBVIAgent +from rlberry.experiment import experiment_generator + def test_mock_args(monkeypatch): monkeypatch.setattr( diff --git a/rlberry/experiment/tests/test_load_results.py b/rlberry/experiment/tests/test_load_results.py index aca97bf25..0cedb2ae7 100644 --- a/rlberry/experiment/tests/test_load_results.py +++ b/rlberry/experiment/tests/test_load_results.py @@ -1,8 +1,8 @@ -from rlberry.experiment import load_experiment_results -import tempfile -from rlberry.experiment import experiment_generator import os import sys +import tempfile + +from rlberry.experiment import experiment_generator, load_experiment_results TEST_DIR = os.path.dirname(os.path.abspath(__file__)) diff --git a/rlberry/experiment/yaml_utils.py b/rlberry/experiment/yaml_utils.py index 581b254ec..21498425d 100644 --- a/rlberry/experiment/yaml_utils.py +++ b/rlberry/experiment/yaml_utils.py @@ -1,6 +1,8 @@ from pathlib import Path from typing import Generator, Tuple + import yaml + from rlberry.utils.factory import load _AGENT_KEYS = ("init_kwargs", "eval_kwargs", "fit_kwargs") diff --git a/rlberry/exploration_tools/discrete_counter.py b/rlberry/exploration_tools/discrete_counter.py index 549a39955..d17f4917b 100644 --- a/rlberry/exploration_tools/discrete_counter.py +++ b/rlberry/exploration_tools/discrete_counter.py @@ -1,6 +1,7 @@ import numpy as np -from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator + from rlberry.exploration_tools.typing import preprocess_args +from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator from rlberry.spaces import Discrete from rlberry.utils.space_discretizer import Discretizer diff --git a/rlberry/exploration_tools/online_discretization_counter.py b/rlberry/exploration_tools/online_discretization_counter.py index 575114df5..2b35491fb 100644 --- a/rlberry/exploration_tools/online_discretization_counter.py +++ b/rlberry/exploration_tools/online_discretization_counter.py @@ -1,11 +1,11 @@ import numpy as np -from rlberry.utils.jit_setup import numba_jit -from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator -from rlberry.exploration_tools.typing import preprocess_args from gymnasium.spaces import Box, Discrete -from rlberry.utils.metrics import metric_lp import rlberry +from rlberry.exploration_tools.typing import preprocess_args +from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator +from rlberry.utils.jit_setup import numba_jit +from rlberry.utils.metrics import metric_lp logger = rlberry.logger diff --git a/rlberry/exploration_tools/tests/test_discrete_counter.py b/rlberry/exploration_tools/tests/test_discrete_counter.py index ad6c1f2bb..8bc978282 100644 --- a/rlberry/exploration_tools/tests/test_discrete_counter.py +++ b/rlberry/exploration_tools/tests/test_discrete_counter.py @@ -1,7 +1,7 @@ -import pytest import numpy as np -from rlberry.envs import GridWorld -from rlberry.envs import MountainCar +import pytest + +from rlberry.envs import GridWorld, MountainCar from rlberry.envs.benchmarks.grid_exploration.nroom import NRoom from rlberry.exploration_tools.discrete_counter import DiscreteCounter from rlberry.exploration_tools.online_discretization_counter import ( diff --git a/rlberry/exploration_tools/torch/rnd.py b/rlberry/exploration_tools/torch/rnd.py index ac6971c22..9b3e881d7 100644 --- a/rlberry/exploration_tools/torch/rnd.py +++ b/rlberry/exploration_tools/torch/rnd.py @@ -1,14 +1,13 @@ from functools import partial -import torch import gymnasium.spaces as spaces +import torch from torch.nn import functional as F +from rlberry.agents.torch.utils.models import ConvolutionalNetwork, MultiLayerPerceptron from rlberry.agents.utils.memories import ReplayMemory -from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator from rlberry.exploration_tools.typing import preprocess_args -from rlberry.agents.torch.utils.models import ConvolutionalNetwork -from rlberry.agents.torch.utils.models import MultiLayerPerceptron +from rlberry.exploration_tools.uncertainty_estimator import UncertaintyEstimator from rlberry.utils.factory import load from rlberry.utils.torch import choose_device diff --git a/rlberry/exploration_tools/torch/tests/test_rnd.py b/rlberry/exploration_tools/torch/tests/test_rnd.py index 5e8d506fa..78bce4f84 100644 --- a/rlberry/exploration_tools/torch/tests/test_rnd.py +++ b/rlberry/exploration_tools/torch/tests/test_rnd.py @@ -1,5 +1,5 @@ -from rlberry.exploration_tools.torch.rnd import RandomNetworkDistillation from rlberry.envs.benchmarks.ball_exploration.ball2d import get_benchmark_env +from rlberry.exploration_tools.torch.rnd import RandomNetworkDistillation def test_rnd(): diff --git a/rlberry/exploration_tools/uncertainty_estimator.py b/rlberry/exploration_tools/uncertainty_estimator.py index 868b4c90e..b79211f55 100644 --- a/rlberry/exploration_tools/uncertainty_estimator.py +++ b/rlberry/exploration_tools/uncertainty_estimator.py @@ -1,7 +1,9 @@ from abc import ABC, abstractmethod -from rlberry.exploration_tools.typing import _get_type + import numpy as np +from rlberry.exploration_tools.typing import _get_type + class UncertaintyEstimator(ABC): def __init__(self, observation_space, action_space, **kwargs): diff --git a/rlberry/manager/__init__.py b/rlberry/manager/__init__.py index e106d36ea..f761dd68e 100644 --- a/rlberry/manager/__init__.py +++ b/rlberry/manager/__init__.py @@ -1,7 +1,7 @@ +from .evaluation import evaluate_agents, plot_writer_data, read_writer_data from .experiment_manager import ExperimentManager, preset_manager from .multiple_managers import MultipleManagers from .remote_experiment_manager import RemoteExperimentManager -from .evaluation import evaluate_agents, plot_writer_data, read_writer_data # (Remote)AgentManager alias for the (Remote)ExperimentManager class, for backward compatibility AgentManager = ExperimentManager diff --git a/rlberry/manager/evaluation.py b/rlberry/manager/evaluation.py index 1e2e04bb9..133e054b1 100644 --- a/rlberry/manager/evaluation.py +++ b/rlberry/manager/evaluation.py @@ -1,18 +1,19 @@ +import bz2 +import pickle +from datetime import datetime +from distutils.version import LooseVersion +from itertools import cycle +from pathlib import Path + +import _pickle as cPickle +import dill import matplotlib.pyplot as plt import numpy as np import pandas as pd import seaborn as sns -from pathlib import Path -from datetime import datetime -import pickle -import bz2 -import _pickle as cPickle -from itertools import cycle -import dill -from distutils.version import LooseVersion -from rlberry.manager import ExperimentManager import rlberry +from rlberry.manager import ExperimentManager logger = rlberry.logger diff --git a/rlberry/manager/experiment_manager.py b/rlberry/manager/experiment_manager.py index 963fe7f2a..a819f2144 100644 --- a/rlberry/manager/experiment_manager.py +++ b/rlberry/manager/experiment_manager.py @@ -1,35 +1,33 @@ +import bz2 import concurrent.futures -from copy import deepcopy -from pathlib import Path -import cProfile, pstats -from pstats import SortKey +import cProfile import functools +import gc import json import logging -import dill -import gc +import multiprocessing import pickle -import bz2 -import _pickle as cPickle +import pstats import shutil import threading -import multiprocessing +from copy import deepcopy from multiprocessing.spawn import _check_not_importing_main +from pathlib import Path +from pstats import SortKey from typing import List, Optional, Tuple, Union +import _pickle as cPickle +import dill import numpy as np import pandas as pd import rlberry -from rlberry.seeding import safe_reseed, set_external_seed -from rlberry.seeding import Seeder -from rlberry import metadata_utils +from rlberry import metadata_utils, types from rlberry.envs.utils import process_env +from rlberry.manager.utils import create_database +from rlberry.seeding import Seeder, safe_reseed, set_external_seed from rlberry.utils.logging import configure_logging from rlberry.utils.writers import DefaultWriter -from rlberry.manager.utils import create_database -from rlberry import types - _OPTUNA_INSTALLED = True try: diff --git a/rlberry/manager/remote_experiment_manager.py b/rlberry/manager/remote_experiment_manager.py index 38335e2f2..7c9b99854 100644 --- a/rlberry/manager/remote_experiment_manager.py +++ b/rlberry/manager/remote_experiment_manager.py @@ -1,17 +1,16 @@ import base64 -import dill import io - -import pandas as pd import pathlib import pickle import zipfile from typing import Any, Mapping, Optional -from rlberry.network import interface -from rlberry.network.client import BerryClient +import dill +import pandas as pd import rlberry +from rlberry.network import interface +from rlberry.network.client import BerryClient logger = rlberry.logger diff --git a/rlberry/manager/tests/test_experiment_manager.py b/rlberry/manager/tests/test_experiment_manager.py index 63a489623..7cd7f3a8a 100644 --- a/rlberry/manager/tests/test_experiment_manager.py +++ b/rlberry/manager/tests/test_experiment_manager.py @@ -1,13 +1,15 @@ -import pytest -import numpy as np -import sys import os -from rlberry.envs import GridWorld +import sys + +import numpy as np +import pytest + from rlberry.agents import AgentWithSimplePolicy +from rlberry.envs import GridWorld from rlberry.manager import ( ExperimentManager, - plot_writer_data, evaluate_agents, + plot_writer_data, preset_manager, ) diff --git a/rlberry/manager/tests/test_experiment_manager_seeding.py b/rlberry/manager/tests/test_experiment_manager_seeding.py index d0e5a317c..d3a0ccf3e 100644 --- a/rlberry/manager/tests/test_experiment_manager_seeding.py +++ b/rlberry/manager/tests/test_experiment_manager_seeding.py @@ -1,10 +1,11 @@ -from rlberry.envs.tests.test_env_seeding import get_env_trajectory, compare_trajectories +import gymnasium as gym +import pytest + +from rlberry.agents.torch import A2CAgent from rlberry.envs import gym_make from rlberry.envs.classic_control import MountainCar +from rlberry.envs.tests.test_env_seeding import compare_trajectories, get_env_trajectory from rlberry.manager import ExperimentManager, MultipleManagers -from rlberry.agents.torch import A2CAgent -import gymnasium as gym -import pytest @pytest.mark.parametrize( diff --git a/rlberry/manager/tests/test_hyperparam_optim.py b/rlberry/manager/tests/test_hyperparam_optim.py index 2803adbcd..2caee153f 100644 --- a/rlberry/manager/tests/test_hyperparam_optim.py +++ b/rlberry/manager/tests/test_hyperparam_optim.py @@ -1,12 +1,14 @@ -from rlberry.envs import GridWorld +import sys +import tempfile + +import numpy as np +import pytest +from optuna.samplers import TPESampler + from rlberry.agents import AgentWithSimplePolicy from rlberry.agents.dynprog.value_iteration import ValueIterationAgent +from rlberry.envs import GridWorld from rlberry.manager import ExperimentManager -from optuna.samplers import TPESampler -import numpy as np -import pytest -import sys -import tempfile class DummyAgent(AgentWithSimplePolicy): diff --git a/rlberry/manager/tests/test_plot.py b/rlberry/manager/tests/test_plot.py index 70533ac60..b3e2190af 100644 --- a/rlberry/manager/tests/test_plot.py +++ b/rlberry/manager/tests/test_plot.py @@ -1,14 +1,15 @@ -import pytest -import tempfile import os -import numpy as np -from pathlib import Path import sys +import tempfile +from pathlib import Path + +import numpy as np +import pytest -from rlberry.wrappers import WriterWrapper -from rlberry.envs import GridWorld -from rlberry.manager import plot_writer_data, ExperimentManager, read_writer_data from rlberry.agents import UCBVIAgent +from rlberry.envs import GridWorld +from rlberry.manager import ExperimentManager, plot_writer_data, read_writer_data +from rlberry.wrappers import WriterWrapper class VIAgent(UCBVIAgent): diff --git a/rlberry/manager/tests/test_shared_data.py b/rlberry/manager/tests/test_shared_data.py index dc44c7501..d2597559a 100644 --- a/rlberry/manager/tests/test_shared_data.py +++ b/rlberry/manager/tests/test_shared_data.py @@ -1,5 +1,6 @@ -import pytest import numpy as np +import pytest + from rlberry.agents import Agent from rlberry.manager import ExperimentManager diff --git a/rlberry/metadata_utils.py b/rlberry/metadata_utils.py index eecd4be5b..9bece8d69 100644 --- a/rlberry/metadata_utils.py +++ b/rlberry/metadata_utils.py @@ -1,8 +1,7 @@ -from datetime import datetime -import uuid import hashlib -from typing import Optional, NamedTuple - +import uuid +from datetime import datetime +from typing import NamedTuple, Optional # Default output directory used by the library. RLBERRY_DEFAULT_DATA_DIR = "rlberry_data/" diff --git a/rlberry/network/client.py b/rlberry/network/client.py index 32d07177d..3d2563c73 100644 --- a/rlberry/network/client.py +++ b/rlberry/network/client.py @@ -1,7 +1,8 @@ +import json import pprint import socket -import json from typing import List, Union + from rlberry.network import interface from rlberry.network.utils import serialize_message diff --git a/rlberry/network/interface.py b/rlberry/network/interface.py index 929a3f366..8bfb0d3b9 100644 --- a/rlberry/network/interface.py +++ b/rlberry/network/interface.py @@ -1,7 +1,6 @@ import struct from typing import Any, Dict, Mapping, NamedTuple, Optional - REQUEST_PREFIX = "ResourceRequest_" diff --git a/rlberry/network/server.py b/rlberry/network/server.py index e40bd632d..b7679b974 100644 --- a/rlberry/network/server.py +++ b/rlberry/network/server.py @@ -1,20 +1,19 @@ import concurrent.futures +import json import logging import multiprocessing import socket -import json +from typing import Optional + +import rlberry import rlberry.network.server_utils as server_utils +from rlberry.envs import gym_make from rlberry.network import interface from rlberry.network.utils import ( apply_fn_to_tree, map_request_to_obj, serialize_message, ) -from rlberry.envs import gym_make -from typing import Optional - - -import rlberry logger = rlberry.logger diff --git a/rlberry/network/server_utils.py b/rlberry/network/server_utils.py index 75922a83f..50dd80cc7 100644 --- a/rlberry/network/server_utils.py +++ b/rlberry/network/server_utils.py @@ -1,9 +1,10 @@ +import base64 import pathlib -from rlberry.network import interface -from rlberry.manager import ExperimentManager -from rlberry import metadata_utils + import rlberry.utils.io -import base64 +from rlberry import metadata_utils +from rlberry.manager import ExperimentManager +from rlberry.network import interface def execute_message( diff --git a/rlberry/network/tests/conftest.py b/rlberry/network/tests/conftest.py index 91ffaff1f..8ec467d2a 100644 --- a/rlberry/network/tests/conftest.py +++ b/rlberry/network/tests/conftest.py @@ -2,16 +2,15 @@ # This file is used to spawn a server to connect to in the tests from test_server.py import multiprocessing +import sys -from rlberry.network.interface import ResourceItem -from rlberry.network.server import BerryServer from rlberry.agents import ValueIterationAgent from rlberry.agents.torch import REINFORCEAgent from rlberry.envs import GridWorld, gym_make +from rlberry.network.interface import ResourceItem +from rlberry.network.server import BerryServer from rlberry.utils.writers import DefaultWriter -import sys - def print_err(s): sys.stderr.write(s) diff --git a/rlberry/network/tests/test_server.py b/rlberry/network/tests/test_server.py index 8f3bf3ca1..ad4b7cd4a 100644 --- a/rlberry/network/tests/test_server.py +++ b/rlberry/network/tests/test_server.py @@ -1,15 +1,15 @@ import sys +import numpy as np import py import pytest from xprocess import ProcessStarter -import numpy as np -from rlberry.network.client import BerryClient -from rlberry.network import interface -from rlberry.network.interface import Message, ResourceRequest from rlberry.manager import RemoteExperimentManager from rlberry.manager.evaluation import evaluate_agents +from rlberry.network import interface +from rlberry.network.client import BerryClient +from rlberry.network.interface import Message, ResourceRequest server_name = "berry" diff --git a/rlberry/network/utils.py b/rlberry/network/utils.py index 67e2ae1f7..af53aeaa5 100644 --- a/rlberry/network/utils.py +++ b/rlberry/network/utils.py @@ -1,8 +1,8 @@ import json from copy import deepcopy -from rlberry.network import interface from typing import Any, Callable, Mapping, Optional, Tuple, Union +from rlberry.network import interface Tree = Union[Any, Tuple, Mapping[Any, "Tree"]] diff --git a/rlberry/rendering/__init__.py b/rlberry/rendering/__init__.py index 5bdd0e295..fb86f91be 100644 --- a/rlberry/rendering/__init__.py +++ b/rlberry/rendering/__init__.py @@ -1,3 +1,2 @@ -from .core import Scene, GeometricPrimitive -from .render_interface import RenderInterface -from .render_interface import RenderInterface2D +from .core import GeometricPrimitive, Scene +from .render_interface import RenderInterface, RenderInterface2D diff --git a/rlberry/rendering/common_shapes.py b/rlberry/rendering/common_shapes.py index 91f942c14..75e7b8ff1 100644 --- a/rlberry/rendering/common_shapes.py +++ b/rlberry/rendering/common_shapes.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.rendering import GeometricPrimitive diff --git a/rlberry/rendering/opengl_render2d.py b/rlberry/rendering/opengl_render2d.py index 64ec79646..2ce836f88 100644 --- a/rlberry/rendering/opengl_render2d.py +++ b/rlberry/rendering/opengl_render2d.py @@ -2,12 +2,12 @@ OpenGL code for 2D rendering, using pygame. """ -import numpy as np from os import environ -from rlberry.rendering import Scene +import numpy as np import rlberry +from rlberry.rendering import Scene logger = rlberry.logger environ["PYGAME_HIDE_SUPPORT_PROMPT"] = "1" @@ -16,18 +16,36 @@ _IMPORT_ERROR_MSG = "" try: import pygame as pg - from pygame.locals import DOUBLEBUF, OPENGL - + from OpenGL.GL import ( + GL_COLOR_BUFFER_BIT, + GL_FRONT, + GL_LINE_LOOP, + GL_LINE_STRIP, + GL_LINES, + GL_POINTS, + GL_POLYGON, + GL_PROJECTION, + GL_QUAD_STRIP, + GL_QUADS, + GL_RGB, + GL_TRIANGLE_FAN, + GL_TRIANGLE_STRIP, + GL_TRIANGLES, + GL_UNSIGNED_BYTE, + glBegin, + glClear, + glClearColor, + glColor3f, + glEnd, + glFlush, + glLoadIdentity, + glMatrixMode, + glReadBuffer, + glReadPixels, + glVertex2f, + ) from OpenGL.GLU import gluOrtho2D - from OpenGL.GL import glMatrixMode, glLoadIdentity, glClearColor - from OpenGL.GL import glClear, glFlush, glBegin, glEnd - from OpenGL.GL import glColor3f, glVertex2f - from OpenGL.GL import glReadBuffer, glReadPixels - from OpenGL.GL import GL_PROJECTION, GL_COLOR_BUFFER_BIT - from OpenGL.GL import GL_POINTS, GL_LINES, GL_LINE_STRIP, GL_LINE_LOOP - from OpenGL.GL import GL_POLYGON, GL_TRIANGLES, GL_TRIANGLE_STRIP - from OpenGL.GL import GL_TRIANGLE_FAN, GL_QUADS, GL_QUAD_STRIP - from OpenGL.GL import GL_FRONT, GL_RGB, GL_UNSIGNED_BYTE + from pygame.locals import DOUBLEBUF, OPENGL except Exception as ex: _IMPORT_SUCESSFUL = False diff --git a/rlberry/rendering/pygame_render2d.py b/rlberry/rendering/pygame_render2d.py index a8d5b3990..c54bc86bf 100644 --- a/rlberry/rendering/pygame_render2d.py +++ b/rlberry/rendering/pygame_render2d.py @@ -2,12 +2,12 @@ Code for 2D rendering, using pygame (without OpenGL) """ -import numpy as np from os import environ -from rlberry.rendering import Scene +import numpy as np import rlberry +from rlberry.rendering import Scene logger = rlberry.logger diff --git a/rlberry/rendering/render_interface.py b/rlberry/rendering/render_interface.py index af846cf33..f42f9cc5e 100644 --- a/rlberry/rendering/render_interface.py +++ b/rlberry/rendering/render_interface.py @@ -5,12 +5,11 @@ from abc import ABC, abstractmethod +import rlberry from rlberry.rendering.opengl_render2d import OpenGLRender2D from rlberry.rendering.pygame_render2d import PyGameRender2D from rlberry.rendering.utils import video_write -import rlberry - logger = rlberry.logger diff --git a/rlberry/rendering/tests/test_rendering_interface.py b/rlberry/rendering/tests/test_rendering_interface.py index f0c793700..c200d474f 100644 --- a/rlberry/rendering/tests/test_rendering_interface.py +++ b/rlberry/rendering/tests/test_rendering_interface.py @@ -1,23 +1,19 @@ import os -import pytest import sys +import tempfile +import pytest from pyvirtualdisplay import Display -from rlberry.envs.classic_control import MountainCar -from rlberry.envs.classic_control import Acrobot -from rlberry.envs.classic_control import Pendulum -from rlberry.envs.finite import Chain -from rlberry.envs.finite import GridWorld -from rlberry.envs.benchmarks.grid_exploration.four_room import FourRoom -from rlberry.envs.benchmarks.grid_exploration.six_room import SixRoom -from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold + +from rlberry.envs import Wrapper from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND from rlberry.envs.benchmarks.generalization.twinrooms import TwinRooms -from rlberry.rendering import RenderInterface -from rlberry.rendering import RenderInterface2D -from rlberry.envs import Wrapper - -import tempfile +from rlberry.envs.benchmarks.grid_exploration.apple_gold import AppleGold +from rlberry.envs.benchmarks.grid_exploration.four_room import FourRoom +from rlberry.envs.benchmarks.grid_exploration.six_room import SixRoom +from rlberry.envs.classic_control import Acrobot, MountainCar, Pendulum +from rlberry.envs.finite import Chain, GridWorld +from rlberry.rendering import RenderInterface, RenderInterface2D try: display = Display(visible=0, size=(1400, 900)) diff --git a/rlberry/rendering/utils.py b/rlberry/rendering/utils.py index bf09963d3..3ebadd7d0 100644 --- a/rlberry/rendering/utils.py +++ b/rlberry/rendering/utils.py @@ -1,6 +1,5 @@ import numpy as np - _FFMPEG_INSTALLED = True try: import ffmpeg diff --git a/rlberry/seeding/__init__.py b/rlberry/seeding/__init__.py index 6c601ad00..faa3d47f1 100644 --- a/rlberry/seeding/__init__.py +++ b/rlberry/seeding/__init__.py @@ -1,3 +1,2 @@ from .seeder import Seeder -from .seeding import safe_reseed -from .seeding import set_external_seed +from .seeding import safe_reseed, set_external_seed diff --git a/rlberry/seeding/seeding.py b/rlberry/seeding/seeding.py index 560e22477..91ac2c43c 100644 --- a/rlberry/seeding/seeding.py +++ b/rlberry/seeding/seeding.py @@ -1,4 +1,5 @@ import numpy as np + import rlberry.check_packages as check_packages from rlberry.seeding.seeder import Seeder diff --git a/rlberry/seeding/tests/test_threads.py b/rlberry/seeding/tests/test_threads.py index dc4e8f2f7..780c019dd 100644 --- a/rlberry/seeding/tests/test_threads.py +++ b/rlberry/seeding/tests/test_threads.py @@ -1,6 +1,7 @@ -from rlberry.seeding.seeder import Seeder import concurrent.futures +from rlberry.seeding.seeder import Seeder + def get_random_number_setting_seed(seeder): return seeder.rng.integers(2**32) diff --git a/rlberry/seeding/tests/test_threads_torch.py b/rlberry/seeding/tests/test_threads_torch.py index e600a09e6..058666f00 100644 --- a/rlberry/seeding/tests/test_threads_torch.py +++ b/rlberry/seeding/tests/test_threads_torch.py @@ -1,7 +1,8 @@ -from rlberry.seeding.seeder import Seeder -from rlberry.seeding import set_external_seed import concurrent.futures +from rlberry.seeding import set_external_seed +from rlberry.seeding.seeder import Seeder + _TORCH_INSTALLED = True try: import torch diff --git a/rlberry/spaces/__init__.py b/rlberry/spaces/__init__.py index 61b8fc6f8..83d29c7c4 100644 --- a/rlberry/spaces/__init__.py +++ b/rlberry/spaces/__init__.py @@ -1,6 +1,6 @@ -from .discrete import Discrete from .box import Box -from .tuple import Tuple -from .multi_discrete import MultiDiscrete -from .multi_binary import MultiBinary from .dict import Dict +from .discrete import Discrete +from .multi_binary import MultiBinary +from .multi_discrete import MultiDiscrete +from .tuple import Tuple diff --git a/rlberry/spaces/box.py b/rlberry/spaces/box.py index ba507c16a..c719fb00d 100644 --- a/rlberry/spaces/box.py +++ b/rlberry/spaces/box.py @@ -1,5 +1,6 @@ import gymnasium as gym import numpy as np + from rlberry.seeding import Seeder diff --git a/rlberry/spaces/dict.py b/rlberry/spaces/dict.py index dd86cf561..014564eb5 100644 --- a/rlberry/spaces/dict.py +++ b/rlberry/spaces/dict.py @@ -1,4 +1,5 @@ import gymnasium as gym + from rlberry.seeding import Seeder diff --git a/rlberry/spaces/discrete.py b/rlberry/spaces/discrete.py index ebbb1841b..5cd8efc10 100644 --- a/rlberry/spaces/discrete.py +++ b/rlberry/spaces/discrete.py @@ -1,4 +1,5 @@ import gymnasium as gym + from rlberry.seeding import Seeder diff --git a/rlberry/spaces/from_gym.py b/rlberry/spaces/from_gym.py index 081f27ea0..67931f46f 100644 --- a/rlberry/spaces/from_gym.py +++ b/rlberry/spaces/from_gym.py @@ -1,6 +1,7 @@ -import rlberry.spaces import gymnasium.spaces +import rlberry.spaces + def convert_space_from_gym(space): if isinstance(space, gymnasium.spaces.Box) and ( diff --git a/rlberry/spaces/multi_binary.py b/rlberry/spaces/multi_binary.py index 7ea42f3d7..7ccb568c4 100644 --- a/rlberry/spaces/multi_binary.py +++ b/rlberry/spaces/multi_binary.py @@ -1,4 +1,5 @@ import gymnasium as gym + from rlberry.seeding import Seeder diff --git a/rlberry/spaces/multi_discrete.py b/rlberry/spaces/multi_discrete.py index eb34af66e..a7991e0bb 100644 --- a/rlberry/spaces/multi_discrete.py +++ b/rlberry/spaces/multi_discrete.py @@ -1,5 +1,6 @@ import gymnasium as gym import numpy as np + from rlberry.seeding import Seeder diff --git a/rlberry/spaces/tests/test_from_gym.py b/rlberry/spaces/tests/test_from_gym.py index b9b79152d..46f2ebc28 100644 --- a/rlberry/spaces/tests/test_from_gym.py +++ b/rlberry/spaces/tests/test_from_gym.py @@ -1,6 +1,7 @@ +import gymnasium.spaces import numpy as np import pytest -import gymnasium.spaces + import rlberry.spaces from rlberry.spaces.from_gym import convert_space_from_gym diff --git a/rlberry/spaces/tests/test_spaces.py b/rlberry/spaces/tests/test_spaces.py index aac7646ed..8fbdb95ee 100644 --- a/rlberry/spaces/tests/test_spaces.py +++ b/rlberry/spaces/tests/test_spaces.py @@ -1,12 +1,7 @@ import numpy as np import pytest -from rlberry.spaces import Box -from rlberry.spaces import Discrete -from rlberry.spaces import Tuple -from rlberry.spaces import MultiDiscrete -from rlberry.spaces import MultiBinary -from rlberry.spaces import Dict +from rlberry.spaces import Box, Dict, Discrete, MultiBinary, MultiDiscrete, Tuple @pytest.mark.parametrize("n", list(range(1, 10))) diff --git a/rlberry/spaces/tuple.py b/rlberry/spaces/tuple.py index f5b4d2c6e..8e754336d 100644 --- a/rlberry/spaces/tuple.py +++ b/rlberry/spaces/tuple.py @@ -1,4 +1,5 @@ import gymnasium as gym + from rlberry.seeding import Seeder diff --git a/rlberry/tests/test_agent_extra.py b/rlberry/tests/test_agent_extra.py index 61cfcdba6..e415c9214 100644 --- a/rlberry/tests/test_agent_extra.py +++ b/rlberry/tests/test_agent_extra.py @@ -1,15 +1,17 @@ +import sys + +import numpy as np import pytest + import rlberry.agents as agents import rlberry.agents.torch as torch_agents +from rlberry.agents.features import FeatureMap from rlberry.utils.check_agent import ( + check_hyperparam_optimisation_agent, check_rl_agent, check_rlberry_agent, check_vectorized_env_agent, - check_hyperparam_optimisation_agent, ) -from rlberry.agents.features import FeatureMap -import numpy as np -import sys class OneHotFeatureMap(FeatureMap): diff --git a/rlberry/tests/test_agents_base.py b/rlberry/tests/test_agents_base.py index a9c65ee9f..1f75011cb 100644 --- a/rlberry/tests/test_agents_base.py +++ b/rlberry/tests/test_agents_base.py @@ -7,17 +7,14 @@ """ -import pytest -import numpy as np import sys +import numpy as np +import pytest + import rlberry.agents as agents from rlberry.agents.features import FeatureMap - -from rlberry.utils.check_agent import ( - check_rl_agent, - check_rlberry_agent, -) +from rlberry.utils.check_agent import check_rl_agent, check_rlberry_agent class OneHotFeatureMap(FeatureMap): diff --git a/rlberry/tests/test_envs.py b/rlberry/tests/test_envs.py index 9519de04f..91b4eebf4 100644 --- a/rlberry/tests/test_envs.py +++ b/rlberry/tests/test_envs.py @@ -1,4 +1,5 @@ -from rlberry.utils.check_env import check_env, check_rlberry_env +import pytest + from rlberry.envs import Acrobot from rlberry.envs.benchmarks.ball_exploration import PBall2D from rlberry.envs.benchmarks.generalization.twinrooms import TwinRooms @@ -6,7 +7,7 @@ from rlberry.envs.benchmarks.grid_exploration.nroom import NRoom from rlberry.envs.classic_control import MountainCar, SpringCartPole from rlberry.envs.finite import Chain, GridWorld -import pytest +from rlberry.utils.check_env import check_env, check_rlberry_env ALL_ENVS = [ Acrobot, diff --git a/rlberry/types.py b/rlberry/types.py index ac16e8d4f..27308a0c2 100644 --- a/rlberry/types.py +++ b/rlberry/types.py @@ -1,5 +1,7 @@ -import gymnasium as gym from typing import Any, Callable, Mapping, Tuple, Union + +import gymnasium as gym + from rlberry.seeding import Seeder # either a gym.Env or a tuple containing (constructor, kwargs) to build the env diff --git a/rlberry/utils/__init__.py b/rlberry/utils/__init__.py index f70c962c1..dc1557abe 100644 --- a/rlberry/utils/__init__.py +++ b/rlberry/utils/__init__.py @@ -1,9 +1,9 @@ -from .check_bandit_agent import check_bandit_agent from .check_agent import ( + check_experiment_manager, + check_fit_additive, check_rl_agent, check_save_load, - check_fit_additive, check_seeding_agent, - check_experiment_manager, ) +from .check_bandit_agent import check_bandit_agent from .check_env import check_env diff --git a/rlberry/utils/check_agent.py b/rlberry/utils/check_agent.py index f4a02976c..cd393e962 100644 --- a/rlberry/utils/check_agent.py +++ b/rlberry/utils/check_agent.py @@ -1,13 +1,14 @@ +import os +import pathlib +import tempfile + +import numpy as np + from rlberry.envs import Chain, Pendulum from rlberry.envs.benchmarks.ball_exploration import PBall2D +from rlberry.envs.gym_make import gym_make from rlberry.manager import ExperimentManager -import numpy as np from rlberry.seeding import set_external_seed -import tempfile -import os -from rlberry.envs.gym_make import gym_make -import pathlib - SEED = 42 diff --git a/rlberry/utils/check_env.py b/rlberry/utils/check_env.py index 1ffd48d04..c9c49513e 100644 --- a/rlberry/utils/check_env.py +++ b/rlberry/utils/check_env.py @@ -1,6 +1,6 @@ -from rlberry.seeding import safe_reseed -from rlberry.seeding import Seeder import numpy as np + +from rlberry.seeding import Seeder, safe_reseed from rlberry.utils.check_gym_env import check_gym_env seeder = Seeder(42) diff --git a/rlberry/utils/check_gym_env.py b/rlberry/utils/check_gym_env.py index ae38d2e0e..b6d713eca 100644 --- a/rlberry/utils/check_gym_env.py +++ b/rlberry/utils/check_gym_env.py @@ -2,13 +2,12 @@ Based on https://github.com/openai/gym and then modified for our purpose. """ -from typing import Union, Optional import inspect +from typing import Optional, Union import gymnasium as gym import numpy as np -from gymnasium import logger -from gymnasium import spaces +from gymnasium import logger, spaces def _is_numpy_array_space(space: spaces.Space) -> bool: diff --git a/rlberry/utils/io.py b/rlberry/utils/io.py index cb269f29a..b70f1588a 100644 --- a/rlberry/utils/io.py +++ b/rlberry/utils/io.py @@ -1,6 +1,6 @@ import os -import zipfile import pathlib +import zipfile def zipdir(dir_path, ouput_fname): diff --git a/rlberry/utils/logging.py b/rlberry/utils/logging.py index af78bdf2e..242e33283 100644 --- a/rlberry/utils/logging.py +++ b/rlberry/utils/logging.py @@ -1,7 +1,9 @@ import logging import logging.config from pathlib import Path + import gymnasium as gym + import rlberry diff --git a/rlberry/utils/math.py b/rlberry/utils/math.py index 5fcb09841..67cb2cd16 100644 --- a/rlberry/utils/math.py +++ b/rlberry/utils/math.py @@ -1,5 +1,6 @@ +from typing import Tuple, Union + import numpy as np -from typing import Union, Tuple Interval = Union[np.ndarray, Tuple[float, float], Tuple[np.ndarray, np.ndarray]] diff --git a/rlberry/utils/metrics.py b/rlberry/utils/metrics.py index 1f6d48456..fbff0c9e2 100644 --- a/rlberry/utils/metrics.py +++ b/rlberry/utils/metrics.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.utils.jit_setup import numba_jit diff --git a/rlberry/utils/space_discretizer.py b/rlberry/utils/space_discretizer.py index 6c34f4a35..1e66b4cd7 100644 --- a/rlberry/utils/space_discretizer.py +++ b/rlberry/utils/space_discretizer.py @@ -1,7 +1,7 @@ import numpy as np from gymnasium.spaces import Box, Discrete -from rlberry.utils.binsearch import binary_search_nd -from rlberry.utils.binsearch import unravel_index_uniform_bin + +from rlberry.utils.binsearch import binary_search_nd, unravel_index_uniform_bin class Discretizer: diff --git a/rlberry/utils/tests/test_binsearch.py b/rlberry/utils/tests/test_binsearch.py index fd94adde4..741b939b4 100644 --- a/rlberry/utils/tests/test_binsearch.py +++ b/rlberry/utils/tests/test_binsearch.py @@ -1,8 +1,7 @@ import numpy as np import pytest -from rlberry.utils.binsearch import binary_search_nd -from rlberry.utils.binsearch import unravel_index_uniform_bin +from rlberry.utils.binsearch import binary_search_nd, unravel_index_uniform_bin def test_binary_search_nd(): diff --git a/rlberry/utils/tests/test_check.py b/rlberry/utils/tests/test_check.py index 070b580d0..ab86da19e 100644 --- a/rlberry/utils/tests/test_check.py +++ b/rlberry/utils/tests/test_check.py @@ -1,15 +1,16 @@ +import gymnasium as gym import numpy as np import pytest -from rlberry.envs import GridWorld, Chain -from rlberry.utils.check_env import check_env + +from rlberry.agents import UCBVIAgent, ValueIterationAgent +from rlberry.envs import Chain, GridWorld +from rlberry.spaces import Box, Dict, Discrete from rlberry.utils.check_agent import ( - check_rl_agent, _fit_experiment_manager, check_agents_almost_equal, + check_rl_agent, ) -from rlberry.spaces import Box, Dict, Discrete -import gymnasium as gym -from rlberry.agents import ValueIterationAgent, UCBVIAgent +from rlberry.utils.check_env import check_env class ActionDictTestEnv(gym.Env): diff --git a/rlberry/utils/tests/test_metrics.py b/rlberry/utils/tests/test_metrics.py index 487dcbb61..798279aff 100644 --- a/rlberry/utils/tests/test_metrics.py +++ b/rlberry/utils/tests/test_metrics.py @@ -1,5 +1,6 @@ -import pytest import numpy as np +import pytest + from rlberry.utils.metrics import metric_lp diff --git a/rlberry/utils/tests/test_writer.py b/rlberry/utils/tests/test_writer.py index 1345649d2..e9fb17aa4 100644 --- a/rlberry/utils/tests/test_writer.py +++ b/rlberry/utils/tests/test_writer.py @@ -1,6 +1,7 @@ import time -from rlberry.envs import GridWorld + from rlberry.agents import AgentWithSimplePolicy +from rlberry.envs import GridWorld from rlberry.manager import ExperimentManager diff --git a/rlberry/utils/torch.py b/rlberry/utils/torch.py index 4663f9ab2..78d421501 100644 --- a/rlberry/utils/torch.py +++ b/rlberry/utils/torch.py @@ -1,11 +1,11 @@ import os import re import shutil -from subprocess import check_output, run, PIPE +from subprocess import PIPE, check_output, run + import numpy as np import torch - import rlberry logger = rlberry.logger diff --git a/rlberry/utils/writers.py b/rlberry/utils/writers.py index 2b3504df2..d7e548bb2 100644 --- a/rlberry/utils/writers.py +++ b/rlberry/utils/writers.py @@ -1,15 +1,15 @@ -import numpy as np -import pandas as pd +import shutil +import sys from collections import deque -from typing import Optional from timeit import default_timer as timer -from rlberry import check_packages -from rlberry import metadata_utils -import shutil +from typing import Optional + +import numpy as np +import pandas as pd from tqdm import tqdm from tqdm.utils import _screen_shape_wrapper -import sys +from rlberry import check_packages, metadata_utils if check_packages.TENSORBOARD_INSTALLED: from torch.utils.tensorboard import SummaryWriter diff --git a/rlberry/wrappers/__init__.py b/rlberry/wrappers/__init__.py index 9015a9418..4c6908782 100644 --- a/rlberry/wrappers/__init__.py +++ b/rlberry/wrappers/__init__.py @@ -1,11 +1,8 @@ +from .discrete2onehot import DiscreteToOneHotWrapper from .discretize_state import DiscretizeStateWrapper from .rescale_reward import RescaleRewardWrapper -from .writer_utils import WriterWrapper -from .discrete2onehot import DiscreteToOneHotWrapper from .tests import ( old_acrobot, - old_twinrooms, - old_six_room, old_apple_gold, old_ball2d, old_finite_mdp, @@ -15,4 +12,7 @@ old_nroom, old_pball, old_pendulum, + old_six_room, + old_twinrooms, ) +from .writer_utils import WriterWrapper diff --git a/rlberry/wrappers/discrete2onehot.py b/rlberry/wrappers/discrete2onehot.py index 33f424926..41d07ed51 100644 --- a/rlberry/wrappers/discrete2onehot.py +++ b/rlberry/wrappers/discrete2onehot.py @@ -1,7 +1,8 @@ -from rlberry.spaces import Box, Discrete -from rlberry.envs import Wrapper import numpy as np +from rlberry.envs import Wrapper +from rlberry.spaces import Box, Discrete + class DiscreteToOneHotWrapper(Wrapper): """Converts observation spaces from Discrete to Box via one-hot encoding.""" diff --git a/rlberry/wrappers/discretize_state.py b/rlberry/wrappers/discretize_state.py index b437632c7..9f57503f2 100644 --- a/rlberry/wrappers/discretize_state.py +++ b/rlberry/wrappers/discretize_state.py @@ -1,8 +1,8 @@ import numpy as np import rlberry.spaces as spaces -from rlberry.utils.binsearch import binary_search_nd, unravel_index_uniform_bin from rlberry.envs import Wrapper +from rlberry.utils.binsearch import binary_search_nd, unravel_index_uniform_bin class DiscretizeStateWrapper(Wrapper): diff --git a/rlberry/wrappers/gym_utils.py b/rlberry/wrappers/gym_utils.py index aa80e8281..263305c54 100644 --- a/rlberry/wrappers/gym_utils.py +++ b/rlberry/wrappers/gym_utils.py @@ -1,13 +1,8 @@ import gymnasium as gym from gymnasium.utils.step_api_compatibility import step_api_compatibility -from rlberry.spaces import Discrete -from rlberry.spaces import Box -from rlberry.spaces import Tuple -from rlberry.spaces import MultiDiscrete -from rlberry.spaces import MultiBinary -from rlberry.spaces import Dict from rlberry.envs import Wrapper +from rlberry.spaces import Box, Dict, Discrete, MultiBinary, MultiDiscrete, Tuple def convert_space_from_gym(gym_space): diff --git a/rlberry/wrappers/rescale_reward.py b/rlberry/wrappers/rescale_reward.py index 4b6e25686..e3307116c 100644 --- a/rlberry/wrappers/rescale_reward.py +++ b/rlberry/wrappers/rescale_reward.py @@ -1,4 +1,5 @@ import numpy as np + from rlberry.envs import Wrapper diff --git a/rlberry/wrappers/tests/__init__.py b/rlberry/wrappers/tests/__init__.py index dacd0d62a..fa13cf778 100644 --- a/rlberry/wrappers/tests/__init__.py +++ b/rlberry/wrappers/tests/__init__.py @@ -1,7 +1,5 @@ from .old_env import ( old_acrobot, - old_twinrooms, - old_six_room, old_apple_gold, old_ball2d, old_finite_mdp, @@ -11,4 +9,6 @@ old_nroom, old_pball, old_pendulum, + old_six_room, + old_twinrooms, ) diff --git a/rlberry/wrappers/tests/old_env/__init__.py b/rlberry/wrappers/tests/old_env/__init__.py index 0d63274fa..7b631f610 100644 --- a/rlberry/wrappers/tests/old_env/__init__.py +++ b/rlberry/wrappers/tests/old_env/__init__.py @@ -4,7 +4,7 @@ from .old_gridworld import Old_GridWorld from .old_mountain_car import Old_MountainCar from .old_nroom import Old_NRoom -from .old_pendulum import Old_Pendulum from .old_pball import Old_PBall2D, Old_SimplePBallND +from .old_pendulum import Old_Pendulum from .old_six_room import Old_SixRoom from .old_twinrooms import Old_TwinRooms diff --git a/rlberry/wrappers/tests/old_env/old_acrobot.py b/rlberry/wrappers/tests/old_env/old_acrobot.py index 8ee5d24f2..9dfb8d475 100644 --- a/rlberry/wrappers/tests/old_env/old_acrobot.py +++ b/rlberry/wrappers/tests/old_env/old_acrobot.py @@ -10,9 +10,10 @@ """ import numpy as np + import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import bar_shape, circle_shape __copyright__ = "Copyright 2013, RLPy http://acl.mit.edu/RLPy" diff --git a/rlberry/wrappers/tests/old_env/old_apple_gold.py b/rlberry/wrappers/tests/old_env/old_apple_gold.py index 9006c990c..b7f0ace93 100644 --- a/rlberry/wrappers/tests/old_env/old_apple_gold.py +++ b/rlberry/wrappers/tests/old_env/old_apple_gold.py @@ -1,9 +1,9 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.rendering import GeometricPrimitive, Scene +from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_finite_mdp.py b/rlberry/wrappers/tests/old_env/old_finite_mdp.py index ed18b7549..a4d31fcc6 100644 --- a/rlberry/wrappers/tests/old_env/old_finite_mdp.py +++ b/rlberry/wrappers/tests/old_env/old_finite_mdp.py @@ -1,11 +1,9 @@ import numpy as np - +import rlberry import rlberry.spaces as spaces from rlberry.envs.interface import Model -import rlberry - logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_four_room.py b/rlberry/wrappers/tests/old_env/old_four_room.py index a67e0e193..e90a8d04e 100644 --- a/rlberry/wrappers/tests/old_env/old_four_room.py +++ b/rlberry/wrappers/tests/old_env/old_four_room.py @@ -1,8 +1,8 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld import rlberry +import rlberry.spaces as spaces +from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_gridworld.py b/rlberry/wrappers/tests/old_env/old_gridworld.py index 4de564bac..01937e599 100644 --- a/rlberry/wrappers/tests/old_env/old_gridworld.py +++ b/rlberry/wrappers/tests/old_env/old_gridworld.py @@ -1,16 +1,13 @@ import matplotlib -import numpy as np - import matplotlib.pyplot as plt +import numpy as np from matplotlib import cm -from rlberry.wrappers.tests.old_env.old_finite_mdp import Old_FiniteMDP +import rlberry from rlberry.envs.finite import gridworld_utils -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import circle_shape - - -import rlberry +from rlberry.wrappers.tests.old_env.old_finite_mdp import Old_FiniteMDP logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_mountain_car.py b/rlberry/wrappers/tests/old_env/old_mountain_car.py index dc40b31db..608ea29c3 100644 --- a/rlberry/wrappers/tests/old_env/old_mountain_car.py +++ b/rlberry/wrappers/tests/old_env/old_mountain_car.py @@ -16,7 +16,7 @@ import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene class Old_MountainCar(RenderInterface2D, Model): diff --git a/rlberry/wrappers/tests/old_env/old_nroom.py b/rlberry/wrappers/tests/old_env/old_nroom.py index 6820ee780..1ad89edea 100644 --- a/rlberry/wrappers/tests/old_env/old_nroom.py +++ b/rlberry/wrappers/tests/old_env/old_nroom.py @@ -1,10 +1,11 @@ import math + import numpy as np -import rlberry.spaces as spaces -from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.rendering import GeometricPrimitive, Scene +from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_pball.py b/rlberry/wrappers/tests/old_env/old_pball.py index acc7ee29d..1944f8d90 100644 --- a/rlberry/wrappers/tests/old_env/old_pball.py +++ b/rlberry/wrappers/tests/old_env/old_pball.py @@ -1,11 +1,9 @@ import numpy as np - +import rlberry import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D - -import rlberry +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_pendulum.py b/rlberry/wrappers/tests/old_env/old_pendulum.py index e8e93ca01..668dfcaf9 100644 --- a/rlberry/wrappers/tests/old_env/old_pendulum.py +++ b/rlberry/wrappers/tests/old_env/old_pendulum.py @@ -9,9 +9,10 @@ """ import numpy as np + import rlberry.spaces as spaces from rlberry.envs.interface import Model -from rlberry.rendering import Scene, RenderInterface2D +from rlberry.rendering import RenderInterface2D, Scene from rlberry.rendering.common_shapes import bar_shape, circle_shape diff --git a/rlberry/wrappers/tests/old_env/old_six_room.py b/rlberry/wrappers/tests/old_env/old_six_room.py index a51905d2d..299497b2d 100644 --- a/rlberry/wrappers/tests/old_env/old_six_room.py +++ b/rlberry/wrappers/tests/old_env/old_six_room.py @@ -1,9 +1,9 @@ import numpy as np -import rlberry.spaces as spaces -from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld -from rlberry.rendering import Scene, GeometricPrimitive import rlberry +import rlberry.spaces as spaces +from rlberry.rendering import GeometricPrimitive, Scene +from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld logger = rlberry.logger diff --git a/rlberry/wrappers/tests/old_env/old_twinrooms.py b/rlberry/wrappers/tests/old_env/old_twinrooms.py index c9ffa09a5..e32ab8eea 100644 --- a/rlberry/wrappers/tests/old_env/old_twinrooms.py +++ b/rlberry/wrappers/tests/old_env/old_twinrooms.py @@ -1,11 +1,11 @@ import numpy as np + +import rlberry import rlberry.spaces as spaces from rlberry.envs import Model -from rlberry.rendering import Scene, GeometricPrimitive, RenderInterface2D +from rlberry.rendering import GeometricPrimitive, RenderInterface2D, Scene from rlberry.rendering.common_shapes import circle_shape -import rlberry - logger = rlberry.logger diff --git a/rlberry/wrappers/tests/test_basewrapper.py b/rlberry/wrappers/tests/test_basewrapper.py index f624e46f2..8516646f0 100644 --- a/rlberry/wrappers/tests/test_basewrapper.py +++ b/rlberry/wrappers/tests/test_basewrapper.py @@ -1,8 +1,8 @@ -from rlberry.envs.interface import Model -from rlberry.envs import Wrapper -from rlberry.envs import GridWorld import gymnasium as gym +from rlberry.envs import GridWorld, Wrapper +from rlberry.envs.interface import Model + def test_wrapper(): env = GridWorld() diff --git a/rlberry/wrappers/tests/test_common_wrappers.py b/rlberry/wrappers/tests/test_common_wrappers.py index 502d24c2a..a783d160e 100644 --- a/rlberry/wrappers/tests/test_common_wrappers.py +++ b/rlberry/wrappers/tests/test_common_wrappers.py @@ -1,5 +1,6 @@ import numpy as np import pytest + from rlberry import spaces from rlberry.agents import RSUCBVIAgent from rlberry.envs.classic_control import MountainCar @@ -9,23 +10,20 @@ from rlberry.wrappers.autoreset import AutoResetWrapper from rlberry.wrappers.discrete2onehot import DiscreteToOneHotWrapper from rlberry.wrappers.discretize_state import DiscretizeStateWrapper -from rlberry.wrappers.rescale_reward import RescaleRewardWrapper -from rlberry.wrappers.uncertainty_estimator_wrapper import UncertaintyEstimatorWrapper -from rlberry.wrappers.vis2d import Vis2dWrapper from rlberry.wrappers.gym_utils import OldGymCompatibilityWrapper - - +from rlberry.wrappers.rescale_reward import RescaleRewardWrapper from rlberry.wrappers.tests.old_env.old_acrobot import Old_Acrobot from rlberry.wrappers.tests.old_env.old_apple_gold import Old_AppleGold from rlberry.wrappers.tests.old_env.old_four_room import Old_FourRoom from rlberry.wrappers.tests.old_env.old_gridworld import Old_GridWorld from rlberry.wrappers.tests.old_env.old_mountain_car import Old_MountainCar from rlberry.wrappers.tests.old_env.old_nroom import Old_NRoom -from rlberry.wrappers.tests.old_env.old_pendulum import Old_Pendulum from rlberry.wrappers.tests.old_env.old_pball import Old_PBall2D, Old_SimplePBallND +from rlberry.wrappers.tests.old_env.old_pendulum import Old_Pendulum from rlberry.wrappers.tests.old_env.old_six_room import Old_SixRoom from rlberry.wrappers.tests.old_env.old_twinrooms import Old_TwinRooms - +from rlberry.wrappers.uncertainty_estimator_wrapper import UncertaintyEstimatorWrapper +from rlberry.wrappers.vis2d import Vis2dWrapper classes = [ Old_Acrobot, diff --git a/rlberry/wrappers/tests/test_gym_space_conversion.py b/rlberry/wrappers/tests/test_gym_space_conversion.py index 8b4b18f14..586e5c4ad 100644 --- a/rlberry/wrappers/tests/test_gym_space_conversion.py +++ b/rlberry/wrappers/tests/test_gym_space_conversion.py @@ -1,6 +1,7 @@ +import gymnasium as gym import numpy as np import pytest -import gymnasium as gym + import rlberry from rlberry.wrappers.gym_utils import convert_space_from_gym diff --git a/rlberry/wrappers/tests/test_wrapper_seeding.py b/rlberry/wrappers/tests/test_wrapper_seeding.py index 936db0d14..1ae2b87d6 100644 --- a/rlberry/wrappers/tests/test_wrapper_seeding.py +++ b/rlberry/wrappers/tests/test_wrapper_seeding.py @@ -1,13 +1,13 @@ +from copy import deepcopy + import numpy as np import pytest -from rlberry.seeding import Seeder -from copy import deepcopy -from rlberry.envs.classic_control import MountainCar, Acrobot -from rlberry.envs.finite import Chain -from rlberry.envs.finite import GridWorld -from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND from rlberry.envs import Wrapper +from rlberry.envs.benchmarks.ball_exploration import PBall2D, SimplePBallND +from rlberry.envs.classic_control import Acrobot, MountainCar +from rlberry.envs.finite import Chain, GridWorld +from rlberry.seeding import Seeder from rlberry.wrappers import RescaleRewardWrapper _GYM_INSTALLED = True diff --git a/rlberry/wrappers/tests/test_writer_utils.py b/rlberry/wrappers/tests/test_writer_utils.py index da8edad70..5a96b5b5c 100644 --- a/rlberry/wrappers/tests/test_writer_utils.py +++ b/rlberry/wrappers/tests/test_writer_utils.py @@ -1,9 +1,8 @@ import pytest -from rlberry.wrappers import WriterWrapper -from rlberry.envs import GridWorld - from rlberry.agents import UCBVIAgent +from rlberry.envs import GridWorld +from rlberry.wrappers import WriterWrapper @pytest.mark.parametrize("write_scalar", ["action", "reward", "action_and_reward"]) diff --git a/rlberry/wrappers/uncertainty_estimator_wrapper.py b/rlberry/wrappers/uncertainty_estimator_wrapper.py index 383b3e3e5..11604fc49 100644 --- a/rlberry/wrappers/uncertainty_estimator_wrapper.py +++ b/rlberry/wrappers/uncertainty_estimator_wrapper.py @@ -1,12 +1,10 @@ +import numpy as np import torch +import rlberry from rlberry.envs import Wrapper - -import numpy as np from rlberry.utils.factory import load -import rlberry - logger = rlberry.logger diff --git a/rlberry/wrappers/vis2d.py b/rlberry/wrappers/vis2d.py index 808b64bb8..b6a65abce 100644 --- a/rlberry/wrappers/vis2d.py +++ b/rlberry/wrappers/vis2d.py @@ -1,13 +1,12 @@ -from rlberry.envs import Wrapper -from rlberry.exploration_tools.discrete_counter import DiscreteCounter -from matplotlib.backends.backend_agg import FigureCanvasAgg as FigureCanvas -from rlberry.rendering.utils import video_write import gymnasium.spaces as spaces - import matplotlib.pyplot as plt import numpy as np +from matplotlib.backends.backend_agg import FigureCanvasAgg as FigureCanvas import rlberry +from rlberry.envs import Wrapper +from rlberry.exploration_tools.discrete_counter import DiscreteCounter +from rlberry.rendering.utils import video_write logger = rlberry.logger diff --git a/scripts/fetch_contributors.py b/scripts/fetch_contributors.py index 2eac326bf..1737cd086 100644 --- a/scripts/fetch_contributors.py +++ b/scripts/fetch_contributors.py @@ -4,11 +4,11 @@ The table should be updated for each new inclusion in the teams. Generating the table requires admin rights. """ -import requests import time -from pathlib import Path from os import path +from pathlib import Path +import requests LOGO_URL = "https://avatars.githubusercontent.com/u/72948299?v=4" REPO_FOLDER = Path(path.abspath(__file__)).parent.parent diff --git a/setup.py b/setup.py index 60fe1b809..6fb5ad7da 100644 --- a/setup.py +++ b/setup.py @@ -1,6 +1,6 @@ -from setuptools import setup, find_packages import os +from setuptools import find_packages, setup ver_file = os.path.join("rlberry", "_version.py") with open(ver_file) as f: