From 852d635742e97495d200b104ab8e09f128211e26 Mon Sep 17 00:00:00 2001 From: Zikang Xiong <73256697+ZikangXiong@users.noreply.github.com> Date: Tue, 29 Nov 2022 17:33:46 -0500 Subject: [PATCH] Exposed modules in __init__.py with __all__ (#1195) * Exposed modules in __init__.py with __all__ * Remove flake8 ignore and update root __all__ * Update version Co-authored-by: Antonin Raffin --- docs/misc/changelog.rst | 5 +++-- setup.cfg | 11 ---------- stable_baselines3/__init__.py | 12 ++++++++++ stable_baselines3/a2c/__init__.py | 2 ++ stable_baselines3/common/envs/__init__.py | 11 ++++++++++ stable_baselines3/common/vec_env/__init__.py | 23 +++++++++++++++++++- stable_baselines3/ddpg/__init__.py | 2 ++ stable_baselines3/dqn/__init__.py | 2 ++ stable_baselines3/her/__init__.py | 2 ++ stable_baselines3/ppo/__init__.py | 2 ++ stable_baselines3/sac/__init__.py | 2 ++ stable_baselines3/td3/__init__.py | 2 ++ stable_baselines3/version.txt | 2 +- 13 files changed, 63 insertions(+), 15 deletions(-) diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index 4885d52..01fc719 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -4,7 +4,7 @@ Changelog ========== -Release 1.7.0a4 (WIP) +Release 1.7.0a5 (WIP) -------------------------- Breaking Changes: @@ -43,6 +43,7 @@ Others: - Fixed ``tests/test_distributions.py`` type hint - Fixed ``stable_baselines3/common/type_aliases.py`` type hint - Fixed ``stable_baselines3/common/env_util.py`` type hint +- Exposed modules in ``__init__.py`` with the ``__all__`` attribute (@ZikangXiong) Documentation: ^^^^^^^^^^^^^^ @@ -1126,4 +1127,4 @@ And all the contributors: @simoninithomas @armandpl @manuel-delverme @Gautam-J @gianlucadecola @buoyancy99 @caburu @xy9485 @Gregwar @ycheng517 @quantitative-technologies @bcollazo @git-thor @TibiGG @cool-RR @MWeltevrede @Melanol @qgallouedec @francescoluciano @jlp-ue @burakdmb @timothe-chaumont @honglu2875 -@anand-bala @hughperkins @sidney-tio @AlexPasqua @dominicgkerr @Akhilez @Rocamonde @tobirohrer +@anand-bala @hughperkins @sidney-tio @AlexPasqua @dominicgkerr @Akhilez @Rocamonde @tobirohrer @ZikangXiong diff --git a/setup.cfg b/setup.cfg index 331c8ff..733b5c3 100644 --- a/setup.cfg +++ b/setup.cfg @@ -81,17 +81,6 @@ exclude = (?x)( ignore = W503,W504,E203,E231 # Ignore import not used when aliases are defined per-file-ignores = - ./stable_baselines3/__init__.py:F401 - ./stable_baselines3/common/__init__.py:F401 - ./stable_baselines3/common/envs/__init__.py:F401 - ./stable_baselines3/a2c/__init__.py:F401 - ./stable_baselines3/ddpg/__init__.py:F401 - ./stable_baselines3/dqn/__init__.py:F401 - ./stable_baselines3/her/__init__.py:F401 - ./stable_baselines3/ppo/__init__.py:F401 - ./stable_baselines3/sac/__init__.py:F401 - ./stable_baselines3/td3/__init__.py:F401 - ./stable_baselines3/common/vec_env/__init__.py:F401 # Default implementation in abstract methods ./stable_baselines3/common/callbacks.py:B027 ./stable_baselines3/common/noise.py:B027 diff --git a/stable_baselines3/__init__.py b/stable_baselines3/__init__.py index d73f5f0..0775a8e 100644 --- a/stable_baselines3/__init__.py +++ b/stable_baselines3/__init__.py @@ -20,3 +20,15 @@ def HER(*args, **kwargs): "Since Stable Baselines 2.1.0, `HER` is now a replay buffer class `HerReplayBuffer`.\n " "Please check the documentation for more information: https://stable-baselines3.readthedocs.io/" ) + + +__all__ = [ + "A2C", + "DDPG", + "DQN", + "PPO", + "SAC", + "TD3", + "HerReplayBuffer", + "get_system_info", +] diff --git a/stable_baselines3/a2c/__init__.py b/stable_baselines3/a2c/__init__.py index 7e99964..78fc54f 100644 --- a/stable_baselines3/a2c/__init__.py +++ b/stable_baselines3/a2c/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.a2c.a2c import A2C from stable_baselines3.a2c.policies import CnnPolicy, MlpPolicy, MultiInputPolicy + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "A2C"] diff --git a/stable_baselines3/common/envs/__init__.py b/stable_baselines3/common/envs/__init__.py index 23bd575..3ff0221 100644 --- a/stable_baselines3/common/envs/__init__.py +++ b/stable_baselines3/common/envs/__init__.py @@ -7,3 +7,14 @@ from stable_baselines3.common.envs.identity_env import ( IdentityEnvMultiDiscrete, ) from stable_baselines3.common.envs.multi_input_envs import SimpleMultiObsEnv + +__all__ = [ + "BitFlippingEnv", + "FakeImageEnv", + "IdentityEnv", + "IdentityEnvBox", + "IdentityEnvMultiBinary", + "IdentityEnvMultiDiscrete", + "SimpleMultiObsEnv", + "SimpleMultiObsEnv", +] diff --git a/stable_baselines3/common/vec_env/__init__.py b/stable_baselines3/common/vec_env/__init__.py index 3880fbd..33a103a 100644 --- a/stable_baselines3/common/vec_env/__init__.py +++ b/stable_baselines3/common/vec_env/__init__.py @@ -1,4 +1,3 @@ -# flake8: noqa F401 import typing from copy import deepcopy from typing import Optional, Type, Union @@ -72,3 +71,25 @@ def sync_envs_normalization(env: "GymEnv", eval_env: "GymEnv") -> None: eval_env_tmp.ret_rms = deepcopy(env_tmp.ret_rms) env_tmp = env_tmp.venv eval_env_tmp = eval_env_tmp.venv + + +__all__ = [ + "CloudpickleWrapper", + "VecEnv", + "VecEnvWrapper", + "DummyVecEnv", + "StackedDictObservations", + "StackedObservations", + "SubprocVecEnv", + "VecCheckNan", + "VecExtractDictObs", + "VecFrameStack", + "VecMonitor", + "VecNormalize", + "VecTransposeImage", + "VecVideoRecorder", + "unwrap_vec_wrapper", + "unwrap_vec_normalize", + "is_vecenv_wrapped", + "sync_envs_normalization", +] diff --git a/stable_baselines3/ddpg/__init__.py b/stable_baselines3/ddpg/__init__.py index 262e7f1..257a3e3 100644 --- a/stable_baselines3/ddpg/__init__.py +++ b/stable_baselines3/ddpg/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.ddpg.ddpg import DDPG from stable_baselines3.ddpg.policies import CnnPolicy, MlpPolicy, MultiInputPolicy + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "DDPG"] diff --git a/stable_baselines3/dqn/__init__.py b/stable_baselines3/dqn/__init__.py index f36f96e..2e5e2db 100644 --- a/stable_baselines3/dqn/__init__.py +++ b/stable_baselines3/dqn/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.dqn.dqn import DQN from stable_baselines3.dqn.policies import CnnPolicy, MlpPolicy, MultiInputPolicy + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "DQN"] diff --git a/stable_baselines3/her/__init__.py b/stable_baselines3/her/__init__.py index 1f58921..dc4c8c2 100644 --- a/stable_baselines3/her/__init__.py +++ b/stable_baselines3/her/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.her.goal_selection_strategy import GoalSelectionStrategy from stable_baselines3.her.her_replay_buffer import HerReplayBuffer + +__all__ = ["GoalSelectionStrategy", "HerReplayBuffer"] diff --git a/stable_baselines3/ppo/__init__.py b/stable_baselines3/ppo/__init__.py index e5c23fc..cd91257 100644 --- a/stable_baselines3/ppo/__init__.py +++ b/stable_baselines3/ppo/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.ppo.policies import CnnPolicy, MlpPolicy, MultiInputPolicy from stable_baselines3.ppo.ppo import PPO + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "PPO"] diff --git a/stable_baselines3/sac/__init__.py b/stable_baselines3/sac/__init__.py index 5a84dde..bdf780d 100644 --- a/stable_baselines3/sac/__init__.py +++ b/stable_baselines3/sac/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.sac.policies import CnnPolicy, MlpPolicy, MultiInputPolicy from stable_baselines3.sac.sac import SAC + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "SAC"] diff --git a/stable_baselines3/td3/__init__.py b/stable_baselines3/td3/__init__.py index 0b903cd..428141e 100644 --- a/stable_baselines3/td3/__init__.py +++ b/stable_baselines3/td3/__init__.py @@ -1,2 +1,4 @@ from stable_baselines3.td3.policies import CnnPolicy, MlpPolicy, MultiInputPolicy from stable_baselines3.td3.td3 import TD3 + +__all__ = ["CnnPolicy", "MlpPolicy", "MultiInputPolicy", "TD3"] diff --git a/stable_baselines3/version.txt b/stable_baselines3/version.txt index 0952a4b..5d819d4 100644 --- a/stable_baselines3/version.txt +++ b/stable_baselines3/version.txt @@ -1 +1 @@ -1.7.0a4 +1.7.0a5