diff --git a/doc/source/rllib-algorithms.rst b/doc/source/rllib-algorithms.rst index 13ef655ff..1f5a7f15d 100644 --- a/doc/source/rllib-algorithms.rst +++ b/doc/source/rllib-algorithms.rst @@ -320,7 +320,9 @@ SpaceInvaders 650 1001 1025 Policy Gradients ---------------- |pytorch| |tensorflow| -`[paper] `__ `[implementation] `__ We include a vanilla policy gradients implementation as an example algorithm. +`[paper] `__ +`[implementation] `__ +We include a vanilla policy gradients implementation as an example algorithm. .. figure:: a2c-arch.svg @@ -340,7 +342,8 @@ Tuned examples: `CartPole-v0 `__ `[implementation] `__ +`[paper] `__ +`[implementation] `__ PPO's clipped objective supports multiple SGD passes over the same batch of experiences. RLlib's multi-GPU optimizer pins that data in GPU memory to avoid unnecessary transfers from host memory, substantially improving performance over a naive implementation. PPO scales out using multiple workers for experience collection, and also to multiple GPUs for SGD. .. tip:: @@ -399,15 +402,21 @@ HalfCheetah 9664 ~7700 Soft Actor Critic (SAC) ------------------------ |pytorch| |tensorflow| -`[paper] `__ `[implementation] `__ +`[original paper] `__, `[follow up paper] `__, `[discrete actions paper] `__ +`[implementation] `__ .. figure:: dqn-arch.svg SAC architecture (same as DQN) -RLlib's soft-actor critic implementation is ported from the `official SAC repo `__ to better integrate with RLlib APIs. Note that SAC has two fields to configure for custom models: ``policy_model`` and ``Q_model``. +RLlib's soft-actor critic implementation is ported from the `official SAC repo `__ to better integrate with RLlib APIs. +Note that SAC has two fields to configure for custom models: ``policy_model`` and ``Q_model``, the ``model`` field of the config will be ignored. -Tuned examples: `Pendulum-v0 `__, `HalfCheetah-v3 `__ +Tuned examples (continuous actions): +`Pendulum-v0 `__, +`HalfCheetah-v3 `__, +Tuned examples (discrete actions): +`CartPole-v0 `__ **MuJoCo results @3M steps:** `more details `__ diff --git a/rllib/agents/ddpg/ddpg_tf_model.py b/rllib/agents/ddpg/ddpg_tf_model.py index 84f8a0878..cddc38006 100644 --- a/rllib/agents/ddpg/ddpg_tf_model.py +++ b/rllib/agents/ddpg/ddpg_tf_model.py @@ -130,7 +130,7 @@ class DDPGTFModel(TFModelV2): This implements Q(s, a). - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. actions (Tensor): Actions to return the Q-values for. @@ -149,7 +149,7 @@ class DDPGTFModel(TFModelV2): This implements the twin Q(s, a). - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. actions (Tensor): Actions to return the Q-values for. @@ -169,7 +169,7 @@ class DDPGTFModel(TFModelV2): This outputs the support for pi(s). For continuous action spaces, this is the action directly. - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. diff --git a/rllib/agents/ddpg/ddpg_torch_model.py b/rllib/agents/ddpg/ddpg_torch_model.py index de2d355da..688c6e414 100644 --- a/rllib/agents/ddpg/ddpg_torch_model.py +++ b/rllib/agents/ddpg/ddpg_torch_model.py @@ -135,7 +135,7 @@ class DDPGTorchModel(TorchModelV2, nn.Module): This implements Q(s, a). - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. actions (Tensor): Actions to return the Q-values for. @@ -151,7 +151,7 @@ class DDPGTorchModel(TorchModelV2, nn.Module): This implements the twin Q(s, a). - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. actions (Optional[Tensor]): Actions to return the Q-values for. @@ -168,7 +168,7 @@ class DDPGTorchModel(TorchModelV2, nn.Module): This outputs the support for pi(s). For continuous action spaces, this is the action directly. For discrete, is is the mean / std dev. - Arguments: + Args: model_out (Tensor): obs embeddings from the model layers, of shape [BATCH_SIZE, num_outputs]. diff --git a/rllib/agents/dqn/distributional_q_tf_model.py b/rllib/agents/dqn/distributional_q_tf_model.py index 8bd8b99ef..ed8064999 100644 --- a/rllib/agents/dqn/distributional_q_tf_model.py +++ b/rllib/agents/dqn/distributional_q_tf_model.py @@ -176,7 +176,7 @@ class DistributionalQTFModel(TFModelV2): Override this in your custom model to customize the Q output head. - Arguments: + Args: model_out (Tensor): embedding from the model layers Returns: diff --git a/rllib/agents/impala/vtrace_tf_policy.py b/rllib/agents/impala/vtrace_tf_policy.py index fb68617d3..3482adef0 100644 --- a/rllib/agents/impala/vtrace_tf_policy.py +++ b/rllib/agents/impala/vtrace_tf_policy.py @@ -112,7 +112,7 @@ class VTraceLoss: def _make_time_major(policy, seq_lens, tensor, drop_last=False): """Swaps batch and trajectory axis. - Arguments: + Args: policy: Policy reference seq_lens: Sequence lengths if recurrent or None tensor: A tensor or list of tensors to reshape. diff --git a/rllib/agents/impala/vtrace_torch_policy.py b/rllib/agents/impala/vtrace_torch_policy.py index f1ee1edb3..c8befc7dd 100644 --- a/rllib/agents/impala/vtrace_torch_policy.py +++ b/rllib/agents/impala/vtrace_torch_policy.py @@ -189,7 +189,7 @@ def build_vtrace_loss(policy, model, dist_class, train_batch): def make_time_major(policy, seq_lens, tensor, drop_last=False): """Swaps batch and trajectory axis. - Arguments: + Args: policy: Policy reference seq_lens: Sequence lengths if recurrent or None tensor: A tensor or list of tensors to reshape. diff --git a/rllib/agents/mbmpo/mbmpo.py b/rllib/agents/mbmpo/mbmpo.py index eebbb1dcd..13d215797 100644 --- a/rllib/agents/mbmpo/mbmpo.py +++ b/rllib/agents/mbmpo/mbmpo.py @@ -104,7 +104,7 @@ class MetaUpdate: """Computes the MetaUpdate step in MAML, adapted for MBMPO for multiple MAML Iterations - Arguments: + Args: workers (WorkerSet): Set of Workers num_steps (int): Number of meta-update steps per MAML Iteration maml_steps (int): MAML Iterations per MBMPO Iteration @@ -121,7 +121,7 @@ class MetaUpdate: self.metrics = {} def __call__(self, data_tuple): - """Arguments: + """Args: data_tuple (tuple): 1st element is samples collected from MAML Inner adaptation steps and 2nd element is accumulated metrics """ @@ -186,7 +186,7 @@ class MetaUpdate: def postprocess_metrics(self, metrics, prefix=""): """Appends prefix to current metrics - Arguments: + Args: metrics (dict): Dictionary of current metrics prefix (str): Prefix string to be appended """ @@ -197,7 +197,7 @@ class MetaUpdate: def post_process_metrics(prefix, workers, metrics): """Update Current Dataset Metrics and filter out specific keys - Arguments: + Args: prefix (str): Prefix string to be appended workers (WorkerSet): Set of workers metrics (dict): Current metrics dictionary @@ -221,7 +221,7 @@ def fit_dynamics(policy, pid): def sync_ensemble(workers): """Syncs dynamics ensemble weights from main to workers - Arguments: + Args: workers (WorkerSet): Set of workers, including main """ diff --git a/rllib/agents/ppo/ddppo.py b/rllib/agents/ppo/ddppo.py index faf9e5d02..77dd1768b 100644 --- a/rllib/agents/ppo/ddppo.py +++ b/rllib/agents/ppo/ddppo.py @@ -86,7 +86,7 @@ def validate_config(config): Args: config (TrainerConfigDict): The Trainer's config to check. - Throws: + Raises: ValueError: In case something is wrong with the config. """ diff --git a/rllib/agents/ppo/ppo.py b/rllib/agents/ppo/ppo.py index 2981b5077..d7eee985d 100644 --- a/rllib/agents/ppo/ppo.py +++ b/rllib/agents/ppo/ppo.py @@ -101,7 +101,7 @@ def validate_config(config: TrainerConfigDict) -> None: Args: config (TrainerConfigDict): The Trainer's config to check. - Throws: + Raises: ValueError: In case something is wrong with the config. """ if isinstance(config["entropy_coeff"], int): @@ -284,4 +284,5 @@ PPOTrainer = build_trainer( validate_config=validate_config, default_policy=PPOTFPolicy, get_policy_class=get_policy_class, - execution_plan=execution_plan) + execution_plan=execution_plan, +) diff --git a/rllib/agents/ppo/ppo_tf_policy.py b/rllib/agents/ppo/ppo_tf_policy.py index 8b9202705..900195a72 100644 --- a/rllib/agents/ppo/ppo_tf_policy.py +++ b/rllib/agents/ppo/ppo_tf_policy.py @@ -326,7 +326,7 @@ def setup_config(policy: Policy, obs_space: gym.spaces.Space, def setup_mixins(policy: Policy, obs_space: gym.spaces.Space, action_space: gym.spaces.Space, config: TrainerConfigDict) -> None: - """Call all mixin classes' constructors before PPOPolicy initialization. + """Call mixin classes' constructors before Policy's loss initialization. Args: policy (Policy): The Policy object. diff --git a/rllib/agents/qmix/mixers.py b/rllib/agents/qmix/mixers.py index 7d790d797..5a929889b 100644 --- a/rllib/agents/qmix/mixers.py +++ b/rllib/agents/qmix/mixers.py @@ -36,7 +36,7 @@ class QMixer(nn.Module): def forward(self, agent_qs, states): """Forward pass for the mixer. - Arguments: + Args: agent_qs: Tensor of shape [B, T, n_agents, n_actions] states: Tensor of shape [B, T, state_dim] """ diff --git a/rllib/agents/qmix/qmix_policy.py b/rllib/agents/qmix/qmix_policy.py index cc13ff46b..f32128424 100644 --- a/rllib/agents/qmix/qmix_policy.py +++ b/rllib/agents/qmix/qmix_policy.py @@ -57,7 +57,7 @@ class QMixLoss(nn.Module): next_state=None): """Forward pass of the loss. - Arguments: + Args: rewards: Tensor of shape [B, T, n_agents] actions: Tensor of shape [B, T, n_agents] terminated: Tensor of shape [B, T, n_agents] @@ -530,7 +530,7 @@ def _validate(obs_space, action_space): def _mac(model, obs, h): """Forward pass of the multi-agent controller. - Arguments: + Args: model: TorchModelV2 class obs: Tensor of shape [B, n_agents, obs_size] h: List of tensors of shape [B, n_agents, h_size] diff --git a/rllib/agents/sac/README.md b/rllib/agents/sac/README.md index a85329f13..bacb48a15 100644 --- a/rllib/agents/sac/README.md +++ b/rllib/agents/sac/README.md @@ -1,8 +1,10 @@ -Implementation of the Soft Actor-Critic algorithm: +Soft Actor Critic (SAC) +======================= -[1] Soft Actor-Critic Algorithms and Applications - T. Haarnoja, A. Zhou, K. Hartikainen, et al. -https://arxiv.org/abs/1812.05905.pdf +Implementations of: -For supporting discrete action spaces, we implemented this patch on top of the original algorithm: -[2] Soft Actor-Critic for Discrete Action Settings - Petros Christodoulou -https://arxiv.org/pdf/1910.07207v2.pdf +Soft Actor-Critic Algorithm (SAC) and a discrete action extension. + +**[Detailed Documentation](https://docs.ray.io/en/latest/rllib-algorithms.html#sac)** + +**[Implementation](https://github.com/ray-project/ray/blob/master/rllib/agents/sac/sac.py)** diff --git a/rllib/agents/sac/__init__.py b/rllib/agents/sac/__init__.py index 260b4f781..61dbdd75b 100644 --- a/rllib/agents/sac/__init__.py +++ b/rllib/agents/sac/__init__.py @@ -3,8 +3,8 @@ from ray.rllib.agents.sac.sac_tf_policy import SACTFPolicy from ray.rllib.agents.sac.sac_torch_policy import SACTorchPolicy __all__ = [ + "DEFAULT_CONFIG", "SACTFPolicy", "SACTorchPolicy", "SACTrainer", - "DEFAULT_CONFIG", ] diff --git a/rllib/agents/sac/sac.py b/rllib/agents/sac/sac.py index 864812053..72aac5977 100644 --- a/rllib/agents/sac/sac.py +++ b/rllib/agents/sac/sac.py @@ -1,8 +1,22 @@ +""" +Soft Actor Critic (SAC) +======================= + +This file defines the distributed Trainer class for the soft actor critic +algorithm. +See `sac_[tf|torch]_policy.py` for the definition of the policy loss. + +Detailed documentation: https://docs.ray.io/en/latest/rllib-algorithms.html#sac +""" + import logging +from typing import Optional, Type from ray.rllib.agents.trainer import with_common_config from ray.rllib.agents.dqn.dqn import GenericOffPolicyTrainer from ray.rllib.agents.sac.sac_tf_policy import SACTFPolicy +from ray.rllib.policy.policy import Policy +from ray.rllib.utils.typing import TrainerConfigDict logger = logging.getLogger(__name__) @@ -14,16 +28,24 @@ OPTIMIZER_SHARED_CONFIGS = [ # yapf: disable # __sphinx_doc_begin__ + +# Adds the following updates to the (base) `Trainer` config in +# rllib/agents/trainer.py (`COMMON_CONFIG` dict). DEFAULT_CONFIG = with_common_config({ # === Model === + # Use two Q-networks (instead of one) for action-value estimation. + # Note: Each Q-network will have its own target network. "twin_q": True, + # Use a e.g. conv2D state preprocessing network before concatenating the + # resulting (feature) vector with the action input for the input to + # the Q-networks. "use_state_preprocessor": False, - # RLlib model options for the Q function(s). + # Model options for the Q network(s). "Q_model": { "fcnet_activation": "relu", "fcnet_hiddens": [256, 256], }, - # RLlib model options for the policy function. + # Model options for the policy function. "policy_model": { "fcnet_activation": "relu", "fcnet_hiddens": [256, 256], @@ -43,10 +65,10 @@ DEFAULT_CONFIG = with_common_config({ # Target entropy lower bound. If "auto", will be set to -|A| (e.g. -2.0 for # Discrete(2), -3.0 for Box(shape=(3,))). # This is the inverse of reward scale, and will be optimized automatically. - "target_entropy": "auto", - # N-step target updates. + "target_entropy": None, + # N-step target updates. If >1, sars' tuples in trajectories will be + # postprocessed to become sa[discounted sum of R][s t+n] tuples. "n_step": 1, - # Number of env steps to optimize for before returning. "timesteps_per_iteration": 100, @@ -117,15 +139,15 @@ DEFAULT_CONFIG = with_common_config({ # yapf: enable -def get_policy_class(config): - if config["framework"] == "torch": - from ray.rllib.agents.sac.sac_torch_policy import SACTorchPolicy - return SACTorchPolicy - else: - return SACTFPolicy +def validate_config(config: TrainerConfigDict) -> None: + """Validates the Trainer's config dict. + Args: + config (TrainerConfigDict): The Trainer's config to check. -def validate_config(config): + Raises: + ValueError: In case something is wrong with the config. + """ if config["model"].get("custom_model"): logger.warning( "Setting use_state_preprocessor=True since a custom model " @@ -136,10 +158,28 @@ def validate_config(config): raise ValueError("`grad_clip` value must be > 0.0!") +def get_policy_class(config: TrainerConfigDict) -> Optional[Type[Policy]]: + """Policy class picker function. Class is chosen based on DL-framework. + + Args: + config (TrainerConfigDict): The trainer's configuration dict. + + Returns: + Optional[Type[Policy]]: The Policy class to use with PPOTrainer. + If None, use `default_policy` provided in build_trainer(). + """ + if config["framework"] == "torch": + from ray.rllib.agents.sac.sac_torch_policy import SACTorchPolicy + return SACTorchPolicy + + +# Build a child class of `Trainer` (based on the kwargs used to create the +# GenericOffPolicyTrainer class and the kwargs used in the call below), which +# uses the framework specific Policy determined in `get_policy_class()` above. SACTrainer = GenericOffPolicyTrainer.with_updates( name="SAC", default_config=DEFAULT_CONFIG, + validate_config=validate_config, default_policy=SACTFPolicy, get_policy_class=get_policy_class, - validate_config=validate_config, ) diff --git a/rllib/agents/sac/sac_tf_model.py b/rllib/agents/sac/sac_tf_model.py index e2ac33f0b..f8b2f6f1c 100644 --- a/rllib/agents/sac/sac_tf_model.py +++ b/rllib/agents/sac/sac_tf_model.py @@ -1,47 +1,58 @@ +import gym from gym.spaces import Discrete import numpy as np +from typing import Optional, Tuple from ray.rllib.models.tf.tf_modelv2 import TFModelV2 from ray.rllib.utils.framework import try_import_tf +from ray.rllib.utils.typing import ModelConfigDict, TensorType tf1, tf, tfv = try_import_tf() class SACTFModel(TFModelV2): - """Extension of standard TFModel for SAC. + """Extension of the standard TFModelV2 for SAC. + + Instances of this Model get created via wrapping this class around another + default- or custom model (inside + rllib/agents/sac/sac_tf_policy.py::build_sac_model). Doing so simply adds + this class' methods (`get_q_values`, etc..) to the wrapped model, such that + the wrapped model can be used by the SAC algorithm. Data flow: - obs -> forward() -> model_out - model_out -> get_policy_output() -> pi(s) - model_out, actions -> get_q_values() -> Q(s, a) - model_out, actions -> get_twin_q_values() -> Q_twin(s, a) - - Note that this class by itself is not a valid model unless you - implement forward() in a subclass.""" + `obs` -> forward() -> `model_out` + `model_out` -> get_policy_output() -> pi(actions|obs) + `model_out`, `actions` -> get_q_values() -> Q(s, a) + `model_out`, `actions` -> get_twin_q_values() -> Q_twin(s, a) + """ def __init__(self, - obs_space, - action_space, - num_outputs, - model_config, - name, - actor_hidden_activation="relu", - actor_hiddens=(256, 256), - critic_hidden_activation="relu", - critic_hiddens=(256, 256), - twin_q=False, - initial_alpha=1.0, - target_entropy=None): - """Initialize variables of this model. + obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + num_outputs: Optional[int], + model_config: ModelConfigDict, + name: str, + actor_hidden_activation: str = "relu", + actor_hiddens: Tuple[int] = (256, 256), + critic_hidden_activation: str = "relu", + critic_hiddens: Tuple[int] = (256, 256), + twin_q: bool = False, + initial_alpha: float = 1.0, + target_entropy: Optional[float] = None): + """Initialize a SACTFModel instance. - Extra model kwargs: - actor_hidden_activation (str): activation for actor network - actor_hiddens (list): hidden layers sizes for actor network - critic_hidden_activation (str): activation for critic network - critic_hiddens (list): hidden layers sizes for critic network - twin_q (bool): build twin Q networks. + Args: + actor_hidden_activation (str): Activation for the actor network. + actor_hiddens (list): Hidden layers sizes for the actor network. + critic_hidden_activation (str): Activation for the critic network. + critic_hiddens (list): Hidden layers sizes for the critic network. + twin_q (bool): Build twin Q networks (Q-net and target) for more + stable Q-learning. initial_alpha (float): The initial value for the to-be-optimized alpha parameter (default: 1.0). + target_entropy (Optional[float]): A target entropy value for + the to-be-optimized alpha parameter. If None, will use the + defaults described in the papers for SAC (and discrete SAC). Note that the core layers for forward() are not defined here, this only defines the layers for the output heads. Those layers for @@ -131,58 +142,69 @@ class SACTFModel(TFModelV2): self.register_variables([self.log_alpha]) - def get_q_values(self, model_out, actions=None): - """Return the Q estimates for the most recent forward pass. + def get_q_values(self, + model_out: TensorType, + actions: Optional[TensorType] = None) -> TensorType: + """Returns Q-values, given the output of self.__call__(). - This implements Q(s, a). + This implements Q(s, a) -> [single Q-value] for the continuous case and + Q(s) -> [Q-values for all actions] for the discrete case. - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. - actions (Optional[Tensor]): Actions to return the Q-values for. - Shape: [BATCH_SIZE, action_dim]. If None (discrete action - case), return Q-values for all actions. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). + actions (Optional[TensorType]): Continuous action batch to return + Q-values for. Shape: [BATCH_SIZE, action_dim]. If None + (discrete action case), return Q-values for all actions. Returns: - tensor of shape [BATCH_SIZE]. + TensorType: Q-values tensor of shape [BATCH_SIZE, 1]. """ + # Continuous case -> concat actions to model_out. if actions is not None: return self.q_net([model_out, actions]) + # Discrete case -> return q-vals for all actions. else: return self.q_net(model_out) - def get_twin_q_values(self, model_out, actions=None): + def get_twin_q_values(self, + model_out: TensorType, + actions: Optional[TensorType] = None) -> TensorType: """Same as get_q_values but using the twin Q net. This implements the twin Q(s, a). - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). actions (Optional[Tensor]): Actions to return the Q-values for. Shape: [BATCH_SIZE, action_dim]. If None (discrete action case), return Q-values for all actions. Returns: - tensor of shape [BATCH_SIZE]. + TensorType: Q-values tensor of shape [BATCH_SIZE, 1]. """ + # Continuous case -> concat actions to model_out. if actions is not None: return self.twin_q_net([model_out, actions]) + # Discrete case -> return q-vals for all actions. else: return self.twin_q_net(model_out) - def get_policy_output(self, model_out): - """Return the action output for the most recent forward pass. + def get_policy_output(self, model_out: TensorType) -> TensorType: + """Returns policy outputs, given the output of self.__call__(). - This outputs the support for pi(s). For continuous action spaces, this - is the action directly. For discrete, is is the mean / std dev. + For continuous action spaces, these will be the mean/stddev + distribution inputs for the (SquashedGaussian) action distribution. + For discrete action spaces, these will be the logits for a categorical + distribution. - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). Returns: - tensor of shape [BATCH_SIZE, action_out_size] + TensorType: Distribution inputs for sampling actions. """ return self.action_model(model_out) diff --git a/rllib/agents/sac/sac_tf_policy.py b/rllib/agents/sac/sac_tf_policy.py index 451567547..7e4cece5e 100644 --- a/rllib/agents/sac/sac_tf_policy.py +++ b/rllib/agents/sac/sac_tf_policy.py @@ -1,6 +1,12 @@ +""" +TensorFlow policy class used for SAC. +""" + +import gym from gym.spaces import Box, Discrete from functools import partial import logging +from typing import Dict, List, Optional, Tuple, Type, Union import ray import ray.experimental.tf_utils @@ -9,14 +15,19 @@ from ray.rllib.agents.ddpg.ddpg_tf_policy import ComputeTDErrorMixin, \ from ray.rllib.agents.dqn.dqn_tf_policy import postprocess_nstep_and_prio from ray.rllib.agents.sac.sac_tf_model import SACTFModel from ray.rllib.agents.sac.sac_torch_model import SACTorchModel +from ray.rllib.evaluation.episode import MultiAgentEpisode from ray.rllib.models import ModelCatalog +from ray.rllib.models.modelv2 import ModelV2 from ray.rllib.models.tf.tf_action_dist import Beta, Categorical, \ - DiagGaussian, SquashedGaussian + DiagGaussian, SquashedGaussian, TFActionDistribution +from ray.rllib.policy.policy import Policy from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy_template import build_tf_policy from ray.rllib.utils.error import UnsupportedSpaceException from ray.rllib.utils.framework import get_variable, try_import_tf, \ try_import_tfp +from ray.rllib.utils.typing import AgentID, LocalOptimizer, ModelGradients, \ + TensorType, TrainerConfigDict tf1, tf, tfv = try_import_tf() tfp = try_import_tfp() @@ -24,12 +35,26 @@ tfp = try_import_tfp() logger = logging.getLogger(__name__) -def build_sac_model(policy, obs_space, action_space, config): - # 2 cases: - # 1) with separate state-preprocessor (before obs+action concat). - # 2) no separate state-preprocessor: concat obs+actions right away. +def build_sac_model(policy: Policy, obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> ModelV2: + """Constructs the necessary ModelV2 for the Policy and returns it. + + Args: + policy (Policy): The TFPolicy that will use the models. + obs_space (gym.spaces.Space): The observation space. + action_space (gym.spaces.Space): The action space. + config (TrainerConfigDict): The SAC trainer's config dict. + + Returns: + ModelV2: The ModelV2 to be used by the Policy. Note: An additional + target model will be created in this function and assigned to + `policy.target_model`. + """ + # With separate state-preprocessor (before obs+action concat). if config["use_state_preprocessor"]: num_outputs = 256 # Flatten last Conv2D to this many nodes. + # No separate state-preprocessor: concat obs+actions right away. else: num_outputs = 0 # No state preprocessor: fcnet_hiddens should be empty. @@ -44,7 +69,7 @@ def build_sac_model(policy, obs_space, action_space, config): # Force-ignore any additionally provided hidden layer sizes. # Everything should be configured using SAC's "Q_model" and "policy_model" # settings. - policy.model = ModelCatalog.get_model_v2( + model = ModelCatalog.get_model_v2( obs_space=obs_space, action_space=action_space, num_outputs=num_outputs, @@ -61,6 +86,10 @@ def build_sac_model(policy, obs_space, action_space, config): initial_alpha=config["initial_alpha"], target_entropy=config["target_entropy"]) + # Create an exact copy of the model and store it in `policy.target_model`. + # This will be used for tau-synched Q-target models that run behind the + # actual Q-networks and are used for target q-value calculations in the + # loss terms. policy.target_model = ModelCatalog.get_model_v2( obs_space=obs_space, action_space=action_space, @@ -78,17 +107,51 @@ def build_sac_model(policy, obs_space, action_space, config): initial_alpha=config["initial_alpha"], target_entropy=config["target_entropy"]) - return policy.model + return model -def postprocess_trajectory(policy, - sample_batch, - other_agent_batches=None, - episode=None): +def postprocess_trajectory( + policy: Policy, + sample_batch: SampleBatch, + other_agent_batches: Optional[Dict[AgentID, SampleBatch]] = None, + episode: Optional[MultiAgentEpisode] = None) -> SampleBatch: + """Postprocesses a trajectory and returns the processed trajectory. + + The trajectory contains only data from one episode and from one agent. + - If `config.batch_mode=truncate_episodes` (default), sample_batch may + contain a truncated (at-the-end) episode, in case the + `config.rollout_fragment_length` was reached by the sampler. + - If `config.batch_mode=complete_episodes`, sample_batch will contain + exactly one episode (no matter how long). + New columns can be added to sample_batch and existing ones may be altered. + + Args: + policy (Policy): The Policy used to generate the trajectory + (`sample_batch`) + sample_batch (SampleBatch): The SampleBatch to postprocess. + other_agent_batches (Optional[Dict[AgentID, SampleBatch]]): Optional + dict of AgentIDs mapping to other agents' trajectory data (from the + same episode). NOTE: The other agents use the same policy. + episode (Optional[MultiAgentEpisode]): Optional multi-agent episode + object in which the agents operated. + + Returns: + SampleBatch: The postprocessed, modified SampleBatch (or a new one). + """ return postprocess_nstep_and_prio(policy, sample_batch) -def get_dist_class(config, action_space): +def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space + ) -> Type[TFActionDistribution]: + """Helper function to return a dist class based on config and action space. + + Args: + config (TrainerConfigDict): The Trainer's config dict. + action_space (gym.spaces.Space): The action space used. + + Returns: + Type[TFActionDistribution]: A TF distribution class. + """ if isinstance(action_space, Discrete): return Categorical else: @@ -99,43 +162,89 @@ def get_dist_class(config, action_space): return DiagGaussian -def get_distribution_inputs_and_class(policy, - model, - obs_batch, - *, - explore=True, - **kwargs): - # Get base-model output. +def get_distribution_inputs_and_class( + policy: Policy, + model: ModelV2, + obs_batch: TensorType, + *, + explore: bool = True, + **kwargs) \ + -> Tuple[TensorType, Type[TFActionDistribution], List[TensorType]]: + """The action distribution function to be used the algorithm. + + An action distribution function is used to customize the choice of action + distribution class and the resulting action distribution inputs (to + parameterize the distribution object). + After parameterizing the distribution, a `sample()` call + will be made on it to generate actions. + + Args: + policy (Policy): The Policy being queried for actions and calling this + function. + model (SACTFModel): The SAC specific Model to use to generate the + distribution inputs (see sac_tf|torch_model.py). Must support the + `get_policy_output` method. + obs_batch (TensorType): The observations to be used as inputs to the + model. + explore (bool): Whether to activate exploration or not. + + Returns: + Tuple[TensorType, Type[TFActionDistribution], List[TensorType]]: The + dist inputs, dist class, and a list of internal state outputs + (in the RNN case). + """ + # Get base-model output (w/o the SAC specific parts of the network). model_out, state_out = model({ "obs": obs_batch, "is_training": policy._get_is_training_placeholder(), }, [], None) - # Get action model output from base-model output. + # Use the base output to get the policy outputs from the SAC model's + # policy components. distribution_inputs = model.get_policy_output(model_out) - action_dist_class = get_dist_class(policy.config, policy.action_space) + # Get a distribution class to be used with the just calculated dist-inputs. + action_dist_class = _get_dist_class(policy.config, policy.action_space) + return distribution_inputs, action_dist_class, state_out -def sac_actor_critic_loss(policy, model, _, train_batch): +def sac_actor_critic_loss( + policy: Policy, model: ModelV2, dist_class: Type[TFActionDistribution], + train_batch: SampleBatch) -> Union[TensorType, List[TensorType]]: + """Constructs the loss for the Soft Actor Critic. + + Args: + policy (Policy): The Policy to calculate the loss for. + model (ModelV2): The Model to calculate the loss for. + dist_class (Type[ActionDistribution]: The action distr. class. + train_batch (SampleBatch): The training data. + + Returns: + Union[TensorType, List[TensorType]]: A single loss tensor or a list + of loss tensors. + """ # Should be True only for debugging purposes (e.g. test cases)! deterministic = policy.config["_deterministic_loss"] + # Get the base model output from the train batch. model_out_t, _ = model({ "obs": train_batch[SampleBatch.CUR_OBS], "is_training": policy._get_is_training_placeholder(), }, [], None) + # Get the base model output from the next observations in the train batch. model_out_tp1, _ = model({ "obs": train_batch[SampleBatch.NEXT_OBS], "is_training": policy._get_is_training_placeholder(), }, [], None) + # Get the target model's base outputs from the next observations in the + # train batch. target_model_out_tp1, _ = policy.target_model({ "obs": train_batch[SampleBatch.NEXT_OBS], "is_training": policy._get_is_training_placeholder(), }, [], None) - # Discrete case. + # Discrete actions case. if model.discrete: # Get all action probs directly from pi and form their logp. log_pis_t = tf.nn.log_softmax(model.get_policy_output(model_out_t), -1) @@ -168,7 +277,7 @@ def sac_actor_critic_loss(policy, model, _, train_batch): # Continuous actions case. else: # Sample simgle actions from distribution. - action_dist_class = get_dist_class(policy.config, policy.action_space) + action_dist_class = _get_dist_class(policy.config, policy.action_space) action_dist_t = action_dist_class( model.get_policy_output(model_out_t), policy.model) policy_t = action_dist_t.sample() if not deterministic else \ @@ -212,7 +321,7 @@ def sac_actor_critic_loss(policy, model, _, train_batch): q_tp1_best_masked = (1.0 - tf.cast(train_batch[SampleBatch.DONES], tf.float32)) * q_tp1_best - # compute RHS of bellman equation + # Compute RHS of bellman equation for the Q-loss (critic(s)). q_t_selected_target = tf.stop_gradient( train_batch[SampleBatch.REWARDS] + policy.config["gamma"]**policy.config["n_step"] * q_tp1_best_masked) @@ -225,6 +334,7 @@ def sac_actor_critic_loss(policy, model, _, train_batch): else: td_error = base_td_error + # Calculate one or two critic losses (2 in the twin_q case). critic_loss = [ 0.5 * tf.keras.losses.MSE( y_true=q_t_selected_target, y_pred=q_t_selected) @@ -258,7 +368,7 @@ def sac_actor_critic_loss(policy, model, _, train_batch): tf.stop_gradient(log_pis_t + model.target_entropy)) actor_loss = tf.reduce_mean(model.alpha * log_pis_t - q_t_det_policy) - # save for stats function + # Save for stats function. policy.policy_t = policy_t policy.q_t = q_t policy.td_error = td_error @@ -268,13 +378,35 @@ def sac_actor_critic_loss(policy, model, _, train_batch): policy.alpha_value = model.alpha policy.target_entropy = model.target_entropy - # in a custom apply op we handle the losses separately, but return them - # combined in one loss for now + # In a custom apply op we handle the losses separately, but return them + # combined in one loss here. return actor_loss + tf.math.add_n(critic_loss) + alpha_loss -def gradients_fn(policy, optimizer, loss): - # Eager: Use GradientTape. +def compute_and_clip_gradients(policy: Policy, optimizer: LocalOptimizer, + loss: TensorType) -> ModelGradients: + """Gradients computing function (from loss tensor, using local optimizer). + + Note: For SAC, optimizer and loss are ignored b/c we have 3 + losses and 3 local optimizers (all stored in policy). + `optimizer` will be used, though, in the tf-eager case b/c it is then a + fake optimizer (OptimizerWrapper) object with a `tape` property to + generate a GradientTape object for gradient recording. + + Args: + policy (Policy): The Policy object that generated the loss tensor and + that holds the given local optimizer. + optimizer (LocalOptimizer): The tf (local) optimizer object to + calculate the gradients with. + loss (TensorType): The loss tensor for which gradients should be + calculated. + + Returns: + ModelGradients: List of the possibly clipped gradients- and variable + tuples. + """ + # Eager: Use GradientTape (which is a property of the `optimizer` object + # (an OptimizerWrapper): see rllib/policy/eager_tf_policy.py). if policy.config["framework"] in ["tf2", "tfe"]: tape = optimizer.tape pol_weights = policy.model.policy_variables() @@ -343,7 +475,26 @@ def gradients_fn(policy, optimizer, loss): return grads_and_vars -def apply_gradients(policy, optimizer, grads_and_vars): +def apply_gradients( + policy: Policy, optimizer: LocalOptimizer, + grads_and_vars: ModelGradients) -> Union["tf.Operation", None]: + """Gradients applying function (from list of "grad_and_var" tuples). + + Note: For SAC, optimizer and grads_and_vars are ignored b/c we have 3 + losses and optimizers (stored in policy). + + Args: + policy (Policy): The Policy object whose Model(s) the given gradients + should be applied to. + optimizer (LocalOptimizer): The tf (local) optimizer object through + which to apply the gradients. + grads_and_vars (ModelGradients): The list of grad_and_var tuples to + apply via the given optimizer. + + Returns: + Union[tf.Operation, None]: The tf op to be used to run the apply + operation. None for eager mode. + """ actor_apply_ops = policy._actor_optimizer.apply_gradients( policy._actor_grads_and_vars) @@ -359,9 +510,11 @@ def apply_gradients(policy, optimizer, grads_and_vars): policy._critic_optimizer[0].apply_gradients(cgrads) ] + # Eager mode -> Just apply and return None. if policy.config["framework"] in ["tf2", "tfe"]: policy._alpha_optimizer.apply_gradients(policy._alpha_grads_and_vars) return + # Tf static graph -> Return op. else: alpha_apply_ops = policy._alpha_optimizer.apply_gradients( policy._alpha_grads_and_vars, @@ -369,10 +522,17 @@ def apply_gradients(policy, optimizer, grads_and_vars): return tf.group([actor_apply_ops, alpha_apply_ops] + critic_apply_ops) -def stats(policy, train_batch): +def stats(policy: Policy, train_batch: SampleBatch) -> Dict[str, TensorType]: + """Stats function for SAC. Returns a dict with important loss stats. + + Args: + policy (Policy): The Policy to generate stats for. + train_batch (SampleBatch): The SampleBatch (already) used for training. + + Returns: + Dict[str, TensorType]: The stats dict. + """ return { - # "policy_t": policy.policy_t, - # "td_error": policy.td_error, "mean_td_error": tf.reduce_mean(policy.td_error), "actor_loss": tf.reduce_mean(policy.actor_loss), "critic_loss": tf.reduce_mean(policy.critic_loss), @@ -386,9 +546,14 @@ def stats(policy, train_batch): class ActorCriticOptimizerMixin: + """Mixin class to generate the necessary optimizers for actor-critic algos. + + - Creates global step for counting the number of update operations. + - Creates separate optimizers for actor, critic, and alpha. + """ + def __init__(self, config): - # - Create global step for counting the number of update operations. - # - Use separate optimizers for actor & critic. + # Eager mode. if config["framework"] in ["tf2", "tfe"]: self.global_step = get_variable(0, tf_name="global_step") self._actor_optimizer = tf.keras.optimizers.Adam( @@ -403,6 +568,7 @@ class ActorCriticOptimizerMixin: "optimization"]["critic_learning_rate"])) self._alpha_optimizer = tf.keras.optimizers.Adam( learning_rate=config["optimization"]["entropy_learning_rate"]) + # Static graph mode. else: self.global_step = tf1.train.get_or_create_global_step() self._actor_optimizer = tf1.train.AdamOptimizer( @@ -419,31 +585,91 @@ class ActorCriticOptimizerMixin: learning_rate=config["optimization"]["entropy_learning_rate"]) -def setup_early_mixins(policy, obs_space, action_space, config): +def setup_early_mixins(policy: Policy, obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> None: + """Call mixin classes' constructors before Policy's initialization. + + Adds the necessary optimizers to the given Policy. + + Args: + policy (Policy): The Policy object. + obs_space (gym.spaces.Space): The Policy's observation space. + action_space (gym.spaces.Space): The Policy's action space. + config (TrainerConfigDict): The Policy's config. + """ ActorCriticOptimizerMixin.__init__(policy, config) -def setup_mid_mixins(policy, obs_space, action_space, config): +def setup_mid_mixins(policy: Policy, obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> None: + """Call mixin classes' constructors before Policy's loss initialization. + + Adds the `compute_td_error` method to the given policy. + Calling `compute_td_error` with batch data will re-calculate the loss + on that batch AND return the per-batch-item TD-error for prioritized + replay buffer record weight updating (in case a prioritized replay buffer + is used). + + Args: + policy (Policy): The Policy object. + obs_space (gym.spaces.Space): The Policy's observation space. + action_space (gym.spaces.Space): The Policy's action space. + config (TrainerConfigDict): The Policy's config. + """ ComputeTDErrorMixin.__init__(policy, sac_actor_critic_loss) -def setup_late_mixins(policy, obs_space, action_space, config): +def setup_late_mixins(policy: Policy, obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> None: + """Call mixin classes' constructors after Policy initialization. + + Adds the `update_target` method to the given policy. + Calling `update_target` updates all target Q-networks' weights from their + respective "main" Q-metworks, based on tau (smooth, partial updating). + + Args: + policy (Policy): The Policy object. + obs_space (gym.spaces.Space): The Policy's observation space. + action_space (gym.spaces.Space): The Policy's action space. + config (TrainerConfigDict): The Policy's config. + """ TargetNetworkMixin.__init__(policy, config) -def validate_spaces(pid, observation_space, action_space, config): +def validate_spaces(policy: Policy, observation_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> None: + """Validates the observation- and action spaces used for the Policy. + + Args: + policy (Policy): The policy, whose spaces are being validated. + observation_space (gym.spaces.Space): The observation space to + validate. + action_space (gym.spaces.Space): The action space to validate. + config (TrainerConfigDict): The Policy's config dict. + + Raises: + UnsupportedSpaceException: If one of the spaces is not supported. + """ + # Only support single Box or single Discreete spaces. if not isinstance(action_space, (Box, Discrete)): raise UnsupportedSpaceException( "Action space ({}) of {} is not supported for " - "SAC.".format(action_space, pid)) - if isinstance(action_space, Box) and len(action_space.shape) > 1: + "SAC.".format(action_space, policy)) + # If Box, make sure it's a 1D vector space. + elif isinstance(action_space, Box) and len(action_space.shape) > 1: raise UnsupportedSpaceException( "Action space ({}) of {} has multiple dimensions " - "{}. ".format(action_space, pid, action_space.shape) + + "{}. ".format(action_space, policy, action_space.shape) + "Consider reshaping this into a single dimension, " "using a Tuple action space, or the multi-agent API.") +# Build a child class of `DynamicTFPolicy`, given the custom functions defined +# above. SACTFPolicy = build_tf_policy( name="SACTFPolicy", get_default_config=lambda: ray.rllib.agents.sac.sac.DEFAULT_CONFIG, @@ -452,7 +678,7 @@ SACTFPolicy = build_tf_policy( action_distribution_fn=get_distribution_inputs_and_class, loss_fn=sac_actor_critic_loss, stats_fn=stats, - gradients_fn=gradients_fn, + gradients_fn=compute_and_clip_gradients, apply_gradients_fn=apply_gradients, extra_learn_fetches_fn=lambda policy: {"td_error": policy.td_error}, mixins=[ diff --git a/rllib/agents/sac/sac_torch_model.py b/rllib/agents/sac/sac_torch_model.py index 7a8f5772f..a756ceb22 100644 --- a/rllib/agents/sac/sac_torch_model.py +++ b/rllib/agents/sac/sac_torch_model.py @@ -1,52 +1,59 @@ +import gym from gym.spaces import Discrete import numpy as np +from typing import Optional, Tuple from ray.rllib.models.torch.misc import SlimFC from ray.rllib.models.torch.torch_modelv2 import TorchModelV2 from ray.rllib.utils.framework import get_activation_fn, try_import_torch +from ray.rllib.utils.typing import ModelConfigDict, TensorType torch, nn = try_import_torch() class SACTorchModel(TorchModelV2, nn.Module): - """Extension of standard TorchModelV2 for SAC. + """Extension of the standard TorchModelV2 for SAC. + + Instances of this Model get created via wrapping this class around another + default- or custom model (inside + rllib/agents/sac/sac_torch_policy.py::build_sac_model). Doing so simply + adds this class' methods (`get_q_values`, etc..) to the wrapped model, such + that the wrapped model can be used by the SAC algorithm. Data flow: - obs -> forward() -> model_out - model_out -> get_policy_output() -> pi(s) - model_out, actions -> get_q_values() -> Q(s, a) - model_out, actions -> get_twin_q_values() -> Q_twin(s, a) - - Note that this class by itself is not a valid model unless you - implement forward() in a subclass.""" + `obs` -> forward() -> `model_out` + `model_out` -> get_policy_output() -> pi(actions|obs) + `model_out`, `actions` -> get_q_values() -> Q(s, a) + `model_out`, `actions` -> get_twin_q_values() -> Q_twin(s, a) + """ def __init__(self, - obs_space, - action_space, - num_outputs, - model_config, - name, - actor_hidden_activation="relu", - actor_hiddens=(256, 256), - critic_hidden_activation="relu", - critic_hiddens=(256, 256), - twin_q=False, - initial_alpha=1.0, - target_entropy=None): - """Initialize variables of this model. - - Extra model kwargs: - actor_hidden_activation (str): activation for actor network - actor_hiddens (list): hidden layers sizes for actor network - critic_hidden_activation (str): activation for critic network - critic_hiddens (list): hidden layers sizes for critic network - twin_q (bool): build twin Q networks. + obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + num_outputs: Optional[int], + model_config: ModelConfigDict, + name: str, + actor_hidden_activation: str = "relu", + actor_hiddens: Tuple[int] = (256, 256), + critic_hidden_activation: str = "relu", + critic_hiddens: Tuple[int] = (256, 256), + twin_q: bool = False, + initial_alpha: float = 1.0, + target_entropy: Optional[float] = None): + """Initializes a SACTorchModel instance. +7 + Args: + actor_hidden_activation (str): Activation for the actor network. + actor_hiddens (list): Hidden layers sizes for the actor network. + critic_hidden_activation (str): Activation for the critic network. + critic_hiddens (list): Hidden layers sizes for the critic network. + twin_q (bool): Build twin Q networks (Q-net and target) for more + stable Q-learning. initial_alpha (float): The initial value for the to-be-optimized alpha parameter (default: 1.0). - target_entropy (Optional[float]): An optional fixed value for the - SAC alpha loss term. None or "auto" for automatic calculation - of this value according to [1] (cont. actions) or [2] - (discrete actions). + target_entropy (Optional[float]): A target entropy value for + the to-be-optimized alpha parameter. If None, will use the + defaults described in the papers for SAC (and discrete SAC). Note that the core layers for forward() are not defined here, this only defines the layers for the output heads. Those layers for @@ -142,58 +149,69 @@ class SACTorchModel(TorchModelV2, nn.Module): self.target_entropy = torch.tensor( data=[target_entropy], dtype=torch.float32, requires_grad=False) - def get_q_values(self, model_out, actions=None): - """Return the Q estimates for the most recent forward pass. + def get_q_values(self, + model_out: TensorType, + actions: Optional[TensorType] = None) -> TensorType: + """Returns Q-values, given the output of self.__call__(). - This implements Q(s, a). + This implements Q(s, a) -> [single Q-value] for the continuous case and + Q(s) -> [Q-values for all actions] for the discrete case. - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. - actions (Optional[Tensor]): Actions to return the Q-values for. - Shape: [BATCH_SIZE, action_dim]. If None (discrete action - case), return Q-values for all actions. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). + actions (Optional[TensorType]): Continuous action batch to return + Q-values for. Shape: [BATCH_SIZE, action_dim]. If None + (discrete action case), return Q-values for all actions. Returns: - tensor of shape [BATCH_SIZE]. + TensorType: Q-values tensor of shape [BATCH_SIZE, 1]. """ + # Continuous case -> concat actions to model_out. if actions is not None: return self.q_net(torch.cat([model_out, actions], -1)) + # Discrete case -> return q-vals for all actions. else: return self.q_net(model_out) - def get_twin_q_values(self, model_out, actions=None): + def get_twin_q_values(self, + model_out: TensorType, + actions: Optional[TensorType] = None) -> TensorType: """Same as get_q_values but using the twin Q net. This implements the twin Q(s, a). - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). actions (Optional[Tensor]): Actions to return the Q-values for. Shape: [BATCH_SIZE, action_dim]. If None (discrete action case), return Q-values for all actions. Returns: - tensor of shape [BATCH_SIZE]. + TensorType: Q-values tensor of shape [BATCH_SIZE, 1]. """ + # Continuous case -> concat actions to model_out. if actions is not None: return self.twin_q_net(torch.cat([model_out, actions], -1)) + # Discrete case -> return q-vals for all actions. else: return self.twin_q_net(model_out) - def get_policy_output(self, model_out): - """Return the action output for the most recent forward pass. + def get_policy_output(self, model_out: TensorType) -> TensorType: + """Returns policy outputs, given the output of self.__call__(). - This outputs the support for pi(s). For continuous action spaces, this - is the action directly. For discrete, is is the mean / std dev. + For continuous action spaces, these will be the mean/stddev + distribution inputs for the (SquashedGaussian) action distribution. + For discrete action spaces, these will be the logits for a categorical + distribution. - Arguments: - model_out (Tensor): obs embeddings from the model layers, of shape - [BATCH_SIZE, num_outputs]. + Args: + model_out (TensorType): Feature outputs from the model layers + (result of doing `self.__call__(obs)`). Returns: - tensor of shape [BATCH_SIZE, action_out_size] + TensorType: Distribution inputs for sampling actions. """ return self.action_model(model_out) diff --git a/rllib/agents/sac/sac_torch_policy.py b/rllib/agents/sac/sac_torch_policy.py index 7d9b748fd..3b9d0f6e9 100644 --- a/rllib/agents/sac/sac_torch_policy.py +++ b/rllib/agents/sac/sac_torch_policy.py @@ -1,5 +1,11 @@ +""" +PyTorch policy class used for SAC. +""" + +import gym from gym.spaces import Discrete import logging +from typing import Dict, List, Optional, Tuple, Type, Union import ray import ray.experimental.tf_utils @@ -7,11 +13,16 @@ from ray.rllib.agents.a3c.a3c_torch_policy import apply_grad_clipping from ray.rllib.agents.sac.sac_tf_policy import build_sac_model, \ postprocess_trajectory, validate_spaces from ray.rllib.agents.dqn.dqn_tf_policy import PRIO_WEIGHTS +from ray.rllib.models.modelv2 import ModelV2 +from ray.rllib.models.torch.torch_action_dist import TorchDistributionWrapper +from ray.rllib.policy.policy import Policy from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.torch_policy_template import build_torch_policy from ray.rllib.models.torch.torch_action_dist import ( TorchCategorical, TorchSquashedGaussian, TorchDiagGaussian, TorchBeta) from ray.rllib.utils.framework import try_import_torch +from ray.rllib.utils.typing import LocalOptimizer, TensorType, \ + TrainerConfigDict torch, nn = try_import_torch() F = nn.functional @@ -19,13 +30,41 @@ F = nn.functional logger = logging.getLogger(__name__) -def build_sac_model_and_action_dist(policy, obs_space, action_space, config): +def build_sac_model_and_action_dist( + policy: Policy, + obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> \ + Tuple[ModelV2, Type[TorchDistributionWrapper]]: + """Constructs the necessary ModelV2 and action dist class for the Policy. + + Args: + policy (Policy): The TFPolicy that will use the models. + obs_space (gym.spaces.Space): The observation space. + action_space (gym.spaces.Space): The action space. + config (TrainerConfigDict): The SAC trainer's config dict. + + Returns: + ModelV2: The ModelV2 to be used by the Policy. Note: An additional + target model will be created in this function and assigned to + `policy.target_model`. + """ model = build_sac_model(policy, obs_space, action_space, config) - action_dist_class = get_dist_class(config, action_space) + action_dist_class = _get_dist_class(config, action_space) return model, action_dist_class -def get_dist_class(config, action_space): +def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space + ) -> Type[TorchDistributionWrapper]: + """Helper function to return a dist class based on config and action space. + + Args: + config (TrainerConfigDict): The Trainer's config dict. + action_space (gym.spaces.Space): The action space used. + + Returns: + Type[TFActionDistribution]: A TF distribution class. + """ if isinstance(action_space, Discrete): return TorchCategorical else: @@ -36,28 +75,83 @@ def get_dist_class(config, action_space): return TorchDiagGaussian -def action_distribution_fn(policy, - model, - obs_batch, - *, - state_batches=None, - seq_lens=None, - prev_action_batch=None, - prev_reward_batch=None, - explore=None, - timestep=None, - is_training=None): +def action_distribution_fn( + policy: Policy, + model: ModelV2, + obs_batch: TensorType, + *, + state_batches: Optional[List[TensorType]] = None, + seq_lens: Optional[TensorType] = None, + prev_action_batch: Optional[TensorType] = None, + prev_reward_batch=None, + explore: Optional[bool] = None, + timestep: Optional[int] = None, + is_training: Optional[bool] = None) -> \ + Tuple[TensorType, Type[TorchDistributionWrapper], List[TensorType]]: + """The action distribution function to be used the algorithm. + + An action distribution function is used to customize the choice of action + distribution class and the resulting action distribution inputs (to + parameterize the distribution object). + After parameterizing the distribution, a `sample()` call + will be made on it to generate actions. + + Args: + policy (Policy): The Policy being queried for actions and calling this + function. + model (TorchModelV2): The SAC specific Model to use to generate the + distribution inputs (see sac_tf|torch_model.py). Must support the + `get_policy_output` method. + obs_batch (TensorType): The observations to be used as inputs to the + model. + state_batches (Optional[List[TensorType]]): The list of internal state + tensor batches. + seq_lens (Optional[TensorType]): The tensor of sequence lengths used + in RNNs. + prev_action_batch (Optional[TensorType]): Optional batch of prev + actions used by the model. + prev_reward_batch (Optional[TensorType]): Optional batch of prev + rewards used by the model. + explore (Optional[bool]): Whether to activate exploration or not. If + None, use value of `config.explore`. + timestep (Optional[int]): An optional timestep. + is_training (Optional[bool]): An optional is-training flag. + + Returns: + Tuple[TensorType, Type[TorchDistributionWrapper], List[TensorType]]: + The dist inputs, dist class, and a list of internal state outputs + (in the RNN case). + """ + # Get base-model output (w/o the SAC specific parts of the network). model_out, _ = model({ "obs": obs_batch, "is_training": is_training, }, [], None) + # Use the base output to get the policy outputs from the SAC model's + # policy components. distribution_inputs = model.get_policy_output(model_out) - action_dist_class = get_dist_class(policy.config, policy.action_space) + # Get a distribution class to be used with the just calculated dist-inputs. + action_dist_class = _get_dist_class(policy.config, policy.action_space) return distribution_inputs, action_dist_class, [] -def actor_critic_loss(policy, model, _, train_batch): +def actor_critic_loss( + policy: Policy, model: ModelV2, + dist_class: Type[TorchDistributionWrapper], + train_batch: SampleBatch) -> Union[TensorType, List[TensorType]]: + """Constructs the loss for the Soft Actor Critic. + + Args: + policy (Policy): The Policy to calculate the loss for. + model (ModelV2): The Model to calculate the loss for. + dist_class (Type[TorchDistributionWrapper]: The action distr. class. + train_batch (SampleBatch): The training data. + + Returns: + Union[TensorType, List[TensorType]]: A single loss tensor or a list + of loss tensors. + """ # Should be True only for debugging purposes (e.g. test cases)! deterministic = policy.config["_deterministic_loss"] @@ -110,7 +204,7 @@ def actor_critic_loss(policy, model, _, train_batch): # Continuous actions case. else: # Sample single actions from distribution. - action_dist_class = get_dist_class(policy.config, policy.action_space) + action_dist_class = _get_dist_class(policy.config, policy.action_space) action_dist_t = action_dist_class( model.get_policy_output(model_out_t), policy.model) policy_t = action_dist_t.sample() if not deterministic else \ @@ -218,7 +312,16 @@ def actor_critic_loss(policy, model, _, train_batch): [policy.alpha_loss]) -def stats(policy, train_batch): +def stats(policy: Policy, train_batch: SampleBatch) -> Dict[str, TensorType]: + """Stats function for SAC. Returns a dict with important loss stats. + + Args: + policy (Policy): The Policy to generate stats for. + train_batch (SampleBatch): The SampleBatch (already) used for training. + + Returns: + Dict[str, TensorType]: The stats dict. + """ return { "td_error": policy.td_error, "mean_td_error": torch.mean(policy.td_error), @@ -235,11 +338,19 @@ def stats(policy, train_batch): } -def optimizer_fn(policy, config): +def optimizer_fn(policy: Policy, config: TrainerConfigDict) -> \ + Tuple[LocalOptimizer]: """Creates all necessary optimizers for SAC learning. The 3 or 4 (twin_q=True) optimizers returned here correspond to the number of loss terms returned by the loss function. + + Args: + policy (Policy): The policy object to be trained. + config (TrainerConfigDict): The Trainer's config dict. + + Returns: + Tuple[LocalOptimizer]: The local optimizers to use for policy training. """ policy.actor_optim = torch.optim.Adam( params=policy.model.policy_variables(), @@ -276,6 +387,12 @@ def optimizer_fn(policy, config): class ComputeTDErrorMixin: + """Mixin class calculating TD-error (part of critic loss) per batch item. + + - Adds `policy.compute_td_error()` method for TD-error calculation from a + batch of observations/actions/rewards/etc.. + """ + def __init__(self): def compute_td_error(obs_t, act_t, rew_t, obs_tp1, done_mask, importance_weights): @@ -291,13 +408,22 @@ class ComputeTDErrorMixin: # (one TD-error value per item in batch to update PR weights). actor_critic_loss(self, self.model, None, input_dict) - # Self.td_error is set within actor_critic_loss call. + # `self.td_error` is set within actor_critic_loss call. Return + # its updated value here. return self.td_error + # Assign the method to policy (self) for later usage. self.compute_td_error = compute_td_error class TargetNetworkMixin: + """Mixin class adding a method for (soft) target net(s) synchronizations. + + - Adds the `update_target` method to the policy. + Calling `update_target` updates all target Q-networks' weights from their + respective "main" Q-metworks, based on tau (smooth, partial updating). + """ + def __init__(self): # Hard initial update from Q-net(s) to target Q-net(s). self.update_target(tau=1.0) @@ -318,7 +444,27 @@ class TargetNetworkMixin: self.target_model.load_state_dict(model_state_dict) -def setup_late_mixins(policy, obs_space, action_space, config): +def setup_late_mixins(policy: Policy, obs_space: gym.spaces.Space, + action_space: gym.spaces.Space, + config: TrainerConfigDict) -> None: + """Call mixin classes' constructors after Policy initialization. + + - Moves the target model(s) to the GPU, if necessary. + - Adds the `compute_td_error` method to the given policy. + Calling `compute_td_error` with batch data will re-calculate the loss + on that batch AND return the per-batch-item TD-error for prioritized + replay buffer record weight updating (in case a prioritized replay buffer + is used). + - Also adds the `update_target` method to the given policy. + Calling `update_target` updates all target Q-networks' weights from their + respective "main" Q-metworks, based on tau (smooth, partial updating). + + Args: + policy (Policy): The Policy object. + obs_space (gym.spaces.Space): The Policy's observation space. + action_space (gym.spaces.Space): The Policy's action space. + config (TrainerConfigDict): The Policy's config. + """ policy.target_model = policy.target_model.to(policy.device) policy.model.log_alpha = policy.model.log_alpha.to(policy.device) policy.model.target_entropy = policy.model.target_entropy.to(policy.device) @@ -326,6 +472,8 @@ def setup_late_mixins(policy, obs_space, action_space, config): TargetNetworkMixin.__init__(policy) +# Build a child class of `DynamicTFPolicy`, given the custom functions defined +# above. SACTorchPolicy = build_torch_policy( name="SACTorchPolicy", loss_fn=actor_critic_loss, diff --git a/rllib/agents/trainer.py b/rllib/agents/trainer.py index b3a4de3cc..1c48e2dfb 100644 --- a/rllib/agents/trainer.py +++ b/rllib/agents/trainer.py @@ -770,7 +770,7 @@ class Trainer(Trainable): Note that you can also access the policy object through self.get_policy(policy_id) and call compute_actions() on it directly. - Arguments: + Args: observation (TensorStructType): observation from the environment. state (List[TensorStructType]): RNN hidden state, if any. If state is not None, then all of compute_single_action(...) is returned @@ -831,7 +831,7 @@ class Trainer(Trainable): Note that you can also access the policy object through self.get_policy(policy_id) and call compute_actions() on it directly. - Arguments: + Args: observation (obj): observation from the environment. state (dict): RNN hidden state, if any. If state is not None, then all of compute_single_action(...) is returned @@ -920,7 +920,7 @@ class Trainer(Trainable): def get_policy(self, policy_id: PolicyID = DEFAULT_POLICY_ID) -> Policy: """Return policy for the specified id, or None. - Arguments: + Args: policy_id (str): id of policy to return. """ return self.workers.local_worker().get_policy(policy_id) @@ -929,7 +929,7 @@ class Trainer(Trainable): def get_weights(self, policies: List[PolicyID] = None) -> dict: """Return a dictionary of policy ids to weights. - Arguments: + Args: policies (list): Optional list of policies to return weights for, or None for all policies. """ @@ -939,7 +939,7 @@ class Trainer(Trainable): def set_weights(self, weights: Dict[PolicyID, dict]): """Set policy weights by policy id. - Arguments: + Args: weights (dict): Map of policy ids to weights to set. """ self.workers.local_worker().set_weights(weights) @@ -950,7 +950,7 @@ class Trainer(Trainable): policy_id: PolicyID = DEFAULT_POLICY_ID): """Export policy model with given policy_id to local directory. - Arguments: + Args: export_dir (string): Writable local directory. policy_id (string): Optional policy id to export. @@ -969,7 +969,7 @@ class Trainer(Trainable): policy_id: PolicyID = DEFAULT_POLICY_ID): """Export tensorflow policy model checkpoint to local directory. - Arguments: + Args: export_dir (string): Writable local directory. filename_prefix (string): file name prefix of checkpoint files. policy_id (string): Optional policy id to export. @@ -989,7 +989,7 @@ class Trainer(Trainable): policy_id: PolicyID = DEFAULT_POLICY_ID): """Imports a policy's model with given policy_id from a local h5 file. - Arguments: + Args: import_file (str): The h5 file to import from. policy_id (string): Optional policy id to import into. diff --git a/rllib/agents/trainer_template.py b/rllib/agents/trainer_template.py index 2241621d9..26ae72158 100644 --- a/rllib/agents/trainer_template.py +++ b/rllib/agents/trainer_template.py @@ -155,11 +155,22 @@ def build_trainer( @staticmethod @override(Trainer) def with_updates(**overrides) -> Type[Trainer]: - """Build a copy of this trainer with the specified overrides. + """Build a copy of this trainer class with the specified overrides. Keyword Args: overrides (dict): use this to override any of the arguments originally passed to build_trainer() for this policy. + + Returns: + Type[Trainer]: A the Trainer sub-class using `original_kwargs` + and `overrides`. + + Examples: + >>> MyClass = SomeOtherClass.with_updates({"name": "Mine"}) + >>> issubclass(MyClass, SomeOtherClass) + ... False + >>> issubclass(MyClass, Trainer) + ... True """ return build_trainer(**dict(original_kwargs, **overrides)) diff --git a/rllib/env/base_env.py b/rllib/env/base_env.py index 20444ec49..6c9544833 100644 --- a/rllib/env/base_env.py +++ b/rllib/env/base_env.py @@ -165,7 +165,7 @@ class BaseEnv: Actions should be sent for each ready agent that returned observations in the previous poll() call. - Arguments: + Args: action_dict (dict): Actions values keyed by env_id and agent_id. """ raise NotImplementedError @@ -364,7 +364,7 @@ class _MultiAgentEnvToBaseEnv(BaseEnv): existing_envs: List[MultiAgentEnv], num_envs: int): """Wrap existing multi-agent envs. - Arguments: + Args: make_env (func|None): Factory that produces a new multiagent env. Must be defined if the number of existing envs is less than num_envs. diff --git a/rllib/env/external_multi_agent_env.py b/rllib/env/external_multi_agent_env.py index 42cd11c46..b4a5d6f6d 100644 --- a/rllib/env/external_multi_agent_env.py +++ b/rllib/env/external_multi_agent_env.py @@ -87,7 +87,7 @@ class ExternalMultiAgentEnv(ExternalEnv): observation_dict is expected to contain the observation of all agents acting in this episode step. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation_dict (dict): Current environment observation. @@ -104,7 +104,7 @@ class ExternalMultiAgentEnv(ExternalEnv): action_dict: MultiAgentDict) -> None: """Record an observation and (off-policy) action taken. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation_dict (dict): Current environment observation. action_dict (dict): Action for the observation. @@ -126,7 +126,7 @@ class ExternalMultiAgentEnv(ExternalEnv): episode. Rewards accumulate until the next action. If no reward is logged before the next action, a reward of 0.0 is assumed. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). reward_dict (dict): Reward from the environment agents. info_dict (dict): Optional info dict. @@ -156,7 +156,7 @@ class ExternalMultiAgentEnv(ExternalEnv): observation_dict: MultiAgentDict) -> None: """Record the end of an episode. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation_dict (dict): Current environment observation. """ diff --git a/rllib/env/group_agents_wrapper.py b/rllib/env/group_agents_wrapper.py index 375033f14..eab11b8ee 100644 --- a/rllib/env/group_agents_wrapper.py +++ b/rllib/env/group_agents_wrapper.py @@ -18,7 +18,7 @@ class _GroupAgentsWrapper(MultiAgentEnv): See MultiAgentEnv.with_agent_groups() for usage info. - Arguments: + Args: env (MultiAgentEnv): env to wrap groups (dict): Grouping spec as documented in MultiAgentEnv obs_space (Space): Optional observation space for the grouped diff --git a/rllib/env/multi_agent_env.py b/rllib/env/multi_agent_env.py index 96db0637d..1000a6a0d 100644 --- a/rllib/env/multi_agent_env.py +++ b/rllib/env/multi_agent_env.py @@ -100,7 +100,7 @@ class MultiAgentEnv: This API is experimental. - Arguments: + Args: groups (dict): Mapping from group id to a list of the agent ids of group members. If an agent id is not present in any group value, it will be left ungrouped. diff --git a/rllib/env/policy_client.py b/rllib/env/policy_client.py index b8efa2eb9..7b6ce4306 100644 --- a/rllib/env/policy_client.py +++ b/rllib/env/policy_client.py @@ -102,7 +102,7 @@ class PolicyClient: ) -> Union[EnvActionType, MultiAgentDict]: """Record an observation and get the on-policy action. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation (obj): Current environment observation. @@ -133,7 +133,7 @@ class PolicyClient: action: Union[EnvActionType, MultiAgentDict]) -> None: """Record an observation and (off-policy) action taken. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation (obj): Current environment observation. action (obj): Action for the observation. @@ -163,7 +163,7 @@ class PolicyClient: episode. Rewards accumulate until the next action. If no reward is logged before the next action, a reward of 0.0 is assumed. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). reward (float): Reward from the environment. info (dict): Extra info dict. @@ -191,7 +191,7 @@ class PolicyClient: observation: Union[EnvObsType, MultiAgentDict]) -> None: """Record the end of an episode. - Arguments: + Args: episode_id (str): Episode id returned from start_episode(). observation (obj): Current environment observation. """ @@ -323,7 +323,7 @@ def _auto_wrap_external(real_env_creator): def _create_embedded_rollout_worker(kwargs, send_fn): """Create a local rollout worker and a thread that samples from it. - Arguments: + Args: kwargs (dict): args for the RolloutWorker constructor. send_fn (fn): function to send a JSON request to the server. """ diff --git a/rllib/env/vector_env.py b/rllib/env/vector_env.py index b97050933..10b14bfbb 100644 --- a/rllib/env/vector_env.py +++ b/rllib/env/vector_env.py @@ -68,7 +68,7 @@ class VectorEnv: ) -> Tuple[List[EnvObsType], List[float], List[bool], List[EnvInfoDict]]: """Performs a vectorized step on all sub environments using `actions`. - Arguments: + Args: actions (List[any]): List of actions (one for each sub-env). Returns: diff --git a/rllib/evaluation/metrics.py b/rllib/evaluation/metrics.py index 5ac1fa74e..6ed723b15 100644 --- a/rllib/evaluation/metrics.py +++ b/rllib/evaluation/metrics.py @@ -91,7 +91,7 @@ def summarize_episodes( ) -> ResultDict: """Summarizes a set of episode metrics tuples. - Arguments: + Args: episodes: smoothed set of episodes including historical ones new_episodes: just the new episodes in this iteration. This must be a subset of `episodes`. If None, assumes all episodes are new. diff --git a/rllib/evaluation/rollout_worker.py b/rllib/evaluation/rollout_worker.py index 899e2ef03..f982b3f90 100644 --- a/rllib/evaluation/rollout_worker.py +++ b/rllib/evaluation/rollout_worker.py @@ -782,7 +782,7 @@ class RolloutWorker(ParallelIteratorWorker): This is typically used in combination with distributed allreduce. - Arguments: + Args: expected_batch_size (int): Expected number of samples to learn on. num_sgd_iter (int): Number of SGD iterations. sgd_minibatch_size (int): SGD minibatch size. @@ -827,7 +827,7 @@ class RolloutWorker(ParallelIteratorWorker): self, policy_id: Optional[PolicyID] = DEFAULT_POLICY_ID) -> Policy: """Return policy for the specified id, or None. - Arguments: + Args: policy_id (str): id of policy to return. """ diff --git a/rllib/evaluation/sample_batch_builder.py b/rllib/evaluation/sample_batch_builder.py index c725a9248..9930abc81 100644 --- a/rllib/evaluation/sample_batch_builder.py +++ b/rllib/evaluation/sample_batch_builder.py @@ -141,7 +141,7 @@ class MultiAgentSampleBatchBuilder: **values: Any) -> None: """Add the given dictionary (row) of values to this batch. - Arguments: + Args: agent_id (obj): Unique id for the agent we are adding values for. policy_id (obj): Unique id for policy controlling the agent. values (dict): Row of values to add for this agent. diff --git a/rllib/evaluation/sampler.py b/rllib/evaluation/sampler.py index b3755d8e3..52702c180 100644 --- a/rllib/evaluation/sampler.py +++ b/rllib/evaluation/sampler.py @@ -1477,7 +1477,7 @@ def _get_or_raise(mapping: Dict[PolicyID, Union[Policy, Preprocessor, Filter]], Returns: Union[Policy, Preprocessor, Filter]: The found object. - Throws: + Raises: ValueError: If `policy_id` cannot be found in `mapping`. """ if policy_id not in mapping: diff --git a/rllib/evaluation/worker_set.py b/rllib/evaluation/worker_set.py index 8d278f2d4..a89b918c2 100644 --- a/rllib/evaluation/worker_set.py +++ b/rllib/evaluation/worker_set.py @@ -39,7 +39,7 @@ class WorkerSet: _setup: bool = True): """Create a new WorkerSet and initialize its workers. - Arguments: + Args: env_creator (Optional[Callable[[EnvContext], EnvType]]): Function that returns env given env config. policy (Optional[Type[Policy]]): A rllib.policy.Policy class. diff --git a/rllib/examples/custom_eval.py b/rllib/examples/custom_eval.py index d8e18ce8c..d0637c29a 100644 --- a/rllib/examples/custom_eval.py +++ b/rllib/examples/custom_eval.py @@ -82,7 +82,7 @@ parser.add_argument("--no-custom-eval", action="store_true") def custom_eval_function(trainer, eval_workers): """Example of a custom evaluation function. - Arguments: + Args: trainer (Trainer): trainer class to evaluate. eval_workers (WorkerSet): evaluation workers. diff --git a/rllib/execution/concurrency_ops.py b/rllib/execution/concurrency_ops.py index 9dfa07c13..3979cb65c 100644 --- a/rllib/execution/concurrency_ops.py +++ b/rllib/execution/concurrency_ops.py @@ -12,7 +12,7 @@ def Concurrently(ops: List[LocalIterator], round_robin_weights=None): """Operator that runs the given parent iterators concurrently. - Arguments: + Args: mode (str): One of {'round_robin', 'async'}. - In 'round_robin' mode, we alternate between pulling items from each parent iterator in order deterministically. @@ -105,7 +105,7 @@ def Dequeue(input_queue: queue.Queue, check=lambda: True): The dequeue is non-blocking, so Dequeue operations can executed with Enqueue via the Concurrently() operator. - Arguments: + Args: input_queue (Queue): queue to pull items from. check (fn): liveness check. When this function returns false, Dequeue() will raise an error to halt execution. diff --git a/rllib/execution/learner_thread.py b/rllib/execution/learner_thread.py index dd9b4f469..9e905148d 100644 --- a/rllib/execution/learner_thread.py +++ b/rllib/execution/learner_thread.py @@ -25,7 +25,7 @@ class LearnerThread(threading.Thread): learner_queue_size, learner_queue_timeout): """Initialize the learner thread. - Arguments: + Args: local_worker (RolloutWorker): process local rollout worker holding policies this thread will call learn_on_batch() on minibatch_buffer_size (int): max number of train batches to store diff --git a/rllib/execution/metric_ops.py b/rllib/execution/metric_ops.py index f14258e5d..374ff047d 100644 --- a/rllib/execution/metric_ops.py +++ b/rllib/execution/metric_ops.py @@ -15,7 +15,7 @@ def StandardMetricsReporting( selected_workers: List["ActorHandle"] = None) -> LocalIterator[dict]: """Operator to periodically collect and report metrics. - Arguments: + Args: train_op (LocalIterator): Operator for executing training steps. We ignore the output values. workers (WorkerSet): Rollout workers to collect metrics from. diff --git a/rllib/execution/minibatch_buffer.py b/rllib/execution/minibatch_buffer.py index 5271378ca..4cd41fcc7 100644 --- a/rllib/execution/minibatch_buffer.py +++ b/rllib/execution/minibatch_buffer.py @@ -7,7 +7,7 @@ class MinibatchBuffer: def __init__(self, inqueue, size, timeout, num_passes, init_num_passes=1): """Initialize a minibatch buffer. - Arguments: + Args: inqueue: Queue to populate the internal ring buffer from. size: Max number of data items to buffer. timeout: Queue timeout diff --git a/rllib/execution/multi_gpu_learner.py b/rllib/execution/multi_gpu_learner.py index 4e1b2a342..b1965863b 100644 --- a/rllib/execution/multi_gpu_learner.py +++ b/rllib/execution/multi_gpu_learner.py @@ -38,7 +38,7 @@ class TFMultiGPULearner(LearnerThread): _fake_gpus=False): """Initialize a multi-gpu learner thread. - Arguments: + Args: local_worker (RolloutWorker): process local rollout worker holding policies this thread will call learn_on_batch() on num_gpus (int): number of GPUs to use for data-parallel SGD diff --git a/rllib/execution/replay_ops.py b/rllib/execution/replay_ops.py index 7c14a4ef4..9be190a74 100644 --- a/rllib/execution/replay_ops.py +++ b/rllib/execution/replay_ops.py @@ -60,7 +60,7 @@ def Replay(*, This should be combined with the StoreToReplayActors operation using the Concurrently() operator. - Arguments: + Args: local_buffer (LocalReplayBuffer): Local buffer to use. Only one of this and replay_actors can be specified. actors (list): List of replay actors. Only one of this and diff --git a/rllib/execution/rollout_ops.py b/rllib/execution/rollout_ops.py index 7af144b81..ed07c7480 100644 --- a/rllib/execution/rollout_ops.py +++ b/rllib/execution/rollout_ops.py @@ -25,7 +25,7 @@ def ParallelRollouts(workers: WorkerSet, *, mode="bulk_sync", If there are no remote workers, experiences will be collected serially from the local worker instance instead. - Arguments: + Args: workers (WorkerSet): set of rollout workers to use. mode (str): One of {'async', 'bulk_sync', 'raw'}. - In 'async' mode, batches are returned as soon as they are @@ -94,7 +94,7 @@ def AsyncGradients( workers: WorkerSet) -> LocalIterator[Tuple[ModelGradients, int]]: """Operator to compute gradients in parallel from rollout workers. - Arguments: + Args: workers (WorkerSet): set of rollout workers to use. Returns: diff --git a/rllib/execution/train_ops.py b/rllib/execution/train_ops.py index f5ac24e0e..274208aa1 100644 --- a/rllib/execution/train_ops.py +++ b/rllib/execution/train_ops.py @@ -273,7 +273,7 @@ class ApplyGradients: update_all=True): """Creates an ApplyGradients instance. - Arguments: + Args: workers (WorkerSet): workers to apply gradients to. update_all (bool): If true, updates all workers. Otherwise, only update the worker that produced the sample batch we are diff --git a/rllib/models/modelv2.py b/rllib/models/modelv2.py index 15dbfb8c1..b87904117 100644 --- a/rllib/models/modelv2.py +++ b/rllib/models/modelv2.py @@ -142,7 +142,7 @@ class ModelV2: You can find an runnable example in examples/custom_loss.py. - Arguments: + Args: policy_loss (Union[List[Tensor],Tensor]): List of or single policy loss(es) from the policy. loss_inputs (dict): map of input placeholders for rollout data. @@ -181,7 +181,7 @@ class ModelV2: Custom models should override forward() instead of __call__. - Arguments: + Args: input_dict (dict): dictionary of input tensors, including "obs", "prev_action", "prev_reward", "is_training" state (list): list of state tensors with sizes matching those diff --git a/rllib/models/tf/recurrent_net.py b/rllib/models/tf/recurrent_net.py index d6c9bfab8..80a8f56d4 100644 --- a/rllib/models/tf/recurrent_net.py +++ b/rllib/models/tf/recurrent_net.py @@ -65,7 +65,7 @@ class RecurrentNetwork(TFModelV2): def forward_rnn(self, inputs, state, seq_lens): """Call the model with the given input tensors and state. - Arguments: + Args: inputs (dict): observation tensor with shape [B, T, obs_size]. state (list): list of state tensors, each with shape [B, T, size]. seq_lens (Tensor): 1d tensor holding input sequence lengths. diff --git a/rllib/offline/input_reader.py b/rllib/offline/input_reader.py index cf18288c4..e6db9e8b0 100644 --- a/rllib/offline/input_reader.py +++ b/rllib/offline/input_reader.py @@ -37,7 +37,7 @@ class InputReader: This method creates a queue runner thread that will call next() on this reader repeatedly to feed the TensorFlow queue. - Arguments: + Args: queue_size (int): Max elements to allow in the TF queue. Example: diff --git a/rllib/offline/json_reader.py b/rllib/offline/json_reader.py index e6315f561..22e9d593c 100644 --- a/rllib/offline/json_reader.py +++ b/rllib/offline/json_reader.py @@ -34,7 +34,7 @@ class JsonReader(InputReader): def __init__(self, inputs: List[str], ioctx: IOContext = None): """Initialize a JsonReader. - Arguments: + Args: inputs (str|list): either a glob expression for files, e.g., "/tmp/**/*.json", or a list of single file paths or URIs, e.g., ["s3://bucket/file.json", "s3://bucket/file2.json"]. diff --git a/rllib/offline/json_writer.py b/rllib/offline/json_writer.py index f34545ffa..619d9e154 100644 --- a/rllib/offline/json_writer.py +++ b/rllib/offline/json_writer.py @@ -36,7 +36,7 @@ class JsonWriter(OutputWriter): compress_columns: List[str] = frozenset(["obs", "new_obs"])): """Initialize a JsonWriter. - Arguments: + Args: path (str): a path/URI of the output directory to save files in. ioctx (IOContext): current IO context object. max_file_size (int): max size of single files before rolling over. diff --git a/rllib/offline/mixed_input.py b/rllib/offline/mixed_input.py index 039e9239c..9f23ba952 100644 --- a/rllib/offline/mixed_input.py +++ b/rllib/offline/mixed_input.py @@ -24,7 +24,7 @@ class MixedInput(InputReader): def __init__(self, dist: Dict[JsonReader, float], ioctx: IOContext): """Initialize a MixedInput. - Arguments: + Args: dist (dict): dict mapping JSONReader paths or "sampler" to probabilities. The probabilities must sum to 1.0. ioctx (IOContext): current IO context object. diff --git a/rllib/offline/off_policy_estimator.py b/rllib/offline/off_policy_estimator.py index ea163c1cd..ac2a04ed6 100644 --- a/rllib/offline/off_policy_estimator.py +++ b/rllib/offline/off_policy_estimator.py @@ -25,7 +25,7 @@ class OffPolicyEstimator: def __init__(self, policy: Policy, gamma: float): """Creates an off-policy estimator. - Arguments: + Args: policy (Policy): Policy to evaluate. gamma (float): Discount of the MDP. """ diff --git a/rllib/offline/output_writer.py b/rllib/offline/output_writer.py index 3e528fb87..8d168dfb4 100644 --- a/rllib/offline/output_writer.py +++ b/rllib/offline/output_writer.py @@ -11,7 +11,7 @@ class OutputWriter: def write(self, sample_batch: SampleBatchType): """Save a batch of experiences. - Arguments: + Args: sample_batch: SampleBatch or MultiAgentBatch to save. """ raise NotImplementedError diff --git a/rllib/offline/shuffled_input.py b/rllib/offline/shuffled_input.py index b829ee44a..24522c87a 100644 --- a/rllib/offline/shuffled_input.py +++ b/rllib/offline/shuffled_input.py @@ -20,7 +20,7 @@ class ShuffledInput(InputReader): def __init__(self, child: InputReader, n: int = 0): """Initialize a MixedInput. - Arguments: + Args: child (InputReader): child input reader to shuffle. n (int): if positive, shuffle input over this many batches. """ diff --git a/rllib/policy/dynamic_tf_policy.py b/rllib/policy/dynamic_tf_policy.py index f4dc78192..dcec4ba33 100644 --- a/rllib/policy/dynamic_tf_policy.py +++ b/rllib/policy/dynamic_tf_policy.py @@ -80,7 +80,7 @@ class DynamicTFPolicy(TFPolicy): obs_include_prev_action_reward: bool = True): """Initialize a dynamic TF policy. - Arguments: + Args: observation_space (gym.spaces.Space): Observation space of the policy. action_space (gym.spaces.Space): Action space of the policy. diff --git a/rllib/policy/torch_policy_template.py b/rllib/policy/torch_policy_template.py index ec91a58d1..3e4426106 100644 --- a/rllib/policy/torch_policy_template.py +++ b/rllib/policy/torch_policy_template.py @@ -1,5 +1,5 @@ import gym -from typing import Callable, Dict, List, Optional, Tuple, Type, Union +from typing import Any, Callable, Dict, List, Optional, Tuple, Type, Union from ray.rllib.models.catalog import ModelCatalog from ray.rllib.models.modelv2 import ModelV2 @@ -29,8 +29,9 @@ def build_torch_policy( stats_fn: Optional[Callable[[Policy, SampleBatch], Dict[ str, TensorType]]] = None, postprocess_fn: Optional[Callable[[ - Policy, SampleBatch, List[SampleBatch], "MultiAgentEpisode" - ], None]] = None, + Policy, SampleBatch, Optional[Dict[Any, SampleBatch]], Optional[ + "MultiAgentEpisode"] + ], SampleBatch]] = None, extra_action_out_fn: Optional[Callable[[ Policy, Dict[str, TensorType], List[TensorType], ModelV2, TorchDistributionWrapper @@ -59,7 +60,7 @@ def build_torch_policy( ], ModelV2]] = None, make_model_and_action_dist: Optional[Callable[[ Policy, gym.spaces.Space, gym.spaces.Space, TrainerConfigDict - ], Tuple[ModelV2, TorchDistributionWrapper]]] = None, + ], Tuple[ModelV2, Type[TorchDistributionWrapper]]]] = None, apply_gradients_fn: Optional[Callable[ [Policy, "torch.optim.Optimizer"], None]] = None, mixins: Optional[List[type]] = None, @@ -77,9 +78,9 @@ def build_torch_policy( overrides. If None, uses only(!) the user-provided PartialTrainerConfigDict as dict for this Policy. postprocess_fn (Optional[Callable[[Policy, SampleBatch, - List[SampleBatch], MultiAgentEpisode], None]]): Optional callable - for post-processing experience batches (called after the - super's `postprocess_trajectory` method). + Optional[Dict[Any, SampleBatch]], Optional["MultiAgentEpisode"]], + SampleBatch]]): Optional callable for post-processing experience + batches (called after the super's `postprocess_trajectory` method). stats_fn (Optional[Callable[[Policy, SampleBatch], Dict[str, TensorType]]]): Optional callable that returns a dict of values given the policy and training batch. If None, @@ -143,9 +144,10 @@ def build_torch_policy( a default Model will be created. make_model_and_action_dist (Optional[Callable[[Policy, gym.spaces.Space, gym.spaces.Space, TrainerConfigDict], - Tuple[ModelV2, TorchDistributionWrapper]]]): Optional callable that - takes the same arguments as Policy.__init__ and returns a tuple - of model instance and torch action distribution class. + Tuple[ModelV2, Type[TorchDistributionWrapper]]]]): Optional + callable that takes the same arguments as Policy.__init__ and + returns a tuple of model instance and torch action distribution + class. Note: Only one of `make_model` or `make_model_and_action_dist` should be provided. If both are None, a default Model will be created. diff --git a/rllib/utils/annotations.py b/rllib/utils/annotations.py index 929c2c52f..a52df829d 100644 --- a/rllib/utils/annotations.py +++ b/rllib/utils/annotations.py @@ -1,7 +1,7 @@ def override(cls): """Annotation for documenting method overrides. - Arguments: + Args: cls (type): The superclass that provides the overriden method. If this cls does not actually have the method, an error is raised. """ diff --git a/rllib/utils/sgd.py b/rllib/utils/sgd.py index 3a2b711fc..b267d04f0 100644 --- a/rllib/utils/sgd.py +++ b/rllib/utils/sgd.py @@ -18,7 +18,7 @@ def averaged(kv, axis=None): For non-scalar values, we simply pick the first value. - Arguments: + Args: kv (dict): dictionary with values that are lists of floats. Returns: @@ -36,7 +36,7 @@ def averaged(kv, axis=None): def standardized(array): """Normalize the values in an array. - Arguments: + Args: array (np.ndarray): Array of values to normalize. Returns: @@ -48,7 +48,7 @@ def standardized(array): def minibatches(samples, sgd_minibatch_size): """Return a generator yielding minibatches from a sample batch. - Arguments: + Args: samples (SampleBatch): batch of samples to split up. sgd_minibatch_size (int): size of minibatches to return. @@ -98,7 +98,7 @@ def do_minibatch_sgd(samples, policies, local_worker, num_sgd_iter, sgd_minibatch_size, standardize_fields): """Execute minibatch SGD. - Arguments: + Args: samples (SampleBatch): batch of samples to optimize. policies (dict): dictionary of policies to optimize. local_worker (RolloutWorker): master rollout worker instance. diff --git a/rllib/utils/test_utils.py b/rllib/utils/test_utils.py index 007bed40a..37395206f 100644 --- a/rllib/utils/test_utils.py +++ b/rllib/utils/test_utils.py @@ -253,7 +253,7 @@ def check_learning_achieved(tune_results, min_reward): tune_results: The tune.run returned results object. min_reward (float): The min reward that must be reached. - Throws: + Raises: ValueError: If `min_reward` not reached. """ if tune_results.trials[0].last_result["episode_reward_mean"] < min_reward: @@ -273,7 +273,7 @@ def check_compute_single_action(trainer, include_prev_action_reward (bool): Whether to include the prev-action and -reward in the `compute_action` call. - Throws: + Raises: ValueError: If anything unexpected happens. """ try: diff --git a/rllib/utils/tf_ops.py b/rllib/utils/tf_ops.py index 5e429ea95..df695eae1 100644 --- a/rllib/utils/tf_ops.py +++ b/rllib/utils/tf_ops.py @@ -55,7 +55,7 @@ def make_tf_callable(session_or_none, dynamic_shape=False): will build a function that executes a session run with placeholders internally. - Arguments: + Args: session_or_none (tf.Session): tf.Session if in graph mode, else None. dynamic_shape (bool): True if the placeholders should have a dynamic batch dimension. Otherwise they will be fixed shape.