diff --git a/rllib/agents/dqn/dqn.py b/rllib/agents/dqn/dqn.py index d07e7e9a7..6a4ba288b 100644 --- a/rllib/agents/dqn/dqn.py +++ b/rllib/agents/dqn/dqn.py @@ -1,16 +1,20 @@ import logging +from typing import Type -from ray.rllib.agents.trainer import with_common_config -from ray.rllib.agents.trainer_template import build_trainer from ray.rllib.agents.dqn.dqn_tf_policy import DQNTFPolicy from ray.rllib.agents.dqn.simple_q_tf_policy import SimpleQTFPolicy -from ray.rllib.policy.policy import LEARNER_STATS_KEY -from ray.rllib.execution.replay_buffer import LocalReplayBuffer -from ray.rllib.execution.rollout_ops import ParallelRollouts +from ray.rllib.agents.trainer import with_common_config +from ray.rllib.agents.trainer_template import build_trainer +from ray.rllib.evaluation.worker_set import WorkerSet from ray.rllib.execution.concurrency_ops import Concurrently -from ray.rllib.execution.replay_ops import StoreToReplayBuffer, Replay -from ray.rllib.execution.train_ops import TrainOneStep, UpdateTargetNetwork from ray.rllib.execution.metric_ops import StandardMetricsReporting +from ray.rllib.execution.replay_buffer import LocalReplayBuffer +from ray.rllib.execution.replay_ops import Replay, StoreToReplayBuffer +from ray.rllib.execution.rollout_ops import ParallelRollouts +from ray.rllib.execution.train_ops import TrainOneStep, UpdateTargetNetwork +from ray.rllib.policy.policy import LEARNER_STATS_KEY, Policy +from ray.rllib.utils.typing import TrainerConfigDict +from ray.util.iter import LocalIterator logger = logging.getLogger(__name__) @@ -122,7 +126,7 @@ DEFAULT_CONFIG = with_common_config({ # yapf: enable -def validate_config(config): +def validate_config(config: TrainerConfigDict) -> None: """Checks and updates the config based on settings. Rewrites rollout_fragment_length to take into account n_step truncation. @@ -152,7 +156,8 @@ def validate_config(config): "replay_sequence_length > 1.") -def execution_plan(workers, config): +def execution_plan(workers: WorkerSet, + config: TrainerConfigDict) -> LocalIterator[dict]: if config.get("prioritized_replay"): prio_args = { "prioritized_replay_alpha": config["prioritized_replay_alpha"], @@ -217,7 +222,7 @@ def execution_plan(workers, config): return StandardMetricsReporting(train_op, workers, config) -def calculate_rr_weights(config): +def calculate_rr_weights(config: TrainerConfigDict): if not config["training_intensity"]: return [1, 1] # e.g., 32 / 4 -> native ratio of 8.0 @@ -229,7 +234,7 @@ def calculate_rr_weights(config): return weights -def get_policy_class(config): +def get_policy_class(config: TrainerConfigDict) -> Type[Policy]: if config["framework"] == "torch": from ray.rllib.agents.dqn.dqn_torch_policy import DQNTorchPolicy return DQNTorchPolicy @@ -237,7 +242,7 @@ def get_policy_class(config): return DQNTFPolicy -def get_simple_policy_class(config): +def get_simple_policy_class(config: TrainerConfigDict) -> Type[Policy]: if config["framework"] == "torch": from ray.rllib.agents.dqn.simple_q_torch_policy import \ SimpleQTorchPolicy diff --git a/rllib/agents/dqn/dqn_tf_policy.py b/rllib/agents/dqn/dqn_tf_policy.py index ddce5b332..177129f20 100644 --- a/rllib/agents/dqn/dqn_tf_policy.py +++ b/rllib/agents/dqn/dqn_tf_policy.py @@ -1,22 +1,26 @@ -from gym.spaces import Discrete -import numpy as np +from typing import Dict +import gym +import numpy as np import ray from ray.rllib.agents.dqn.distributional_q_tf_model import \ DistributionalQTFModel from ray.rllib.agents.dqn.simple_q_tf_policy import TargetNetworkMixin from ray.rllib.models import ModelCatalog +from ray.rllib.models.modelv2 import ModelV2 from ray.rllib.models.tf.tf_action_dist import Categorical +from ray.rllib.policy.policy import Policy from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy import LearningRateSchedule from ray.rllib.policy.tf_policy_template import build_tf_policy from ray.rllib.utils.error import UnsupportedSpaceException from ray.rllib.utils.exploration import ParameterNoise -from ray.rllib.utils.numpy import convert_to_numpy from ray.rllib.utils.framework import try_import_tf -from ray.rllib.utils.tf_ops import huber_loss, reduce_mean_ignore_inf, \ - minimize_and_clip -from ray.rllib.utils.tf_ops import make_tf_callable +from ray.rllib.utils.numpy import convert_to_numpy +from ray.rllib.utils.tf_ops import (huber_loss, make_tf_callable, + minimize_and_clip, reduce_mean_ignore_inf) +from ray.rllib.utils.typing import (ModelGradients, TensorType, + TrainerConfigDict) tf1, tf, tfv = try_import_tf() @@ -126,9 +130,11 @@ class ComputeTDErrorMixin: self.compute_td_error = compute_td_error -def build_q_model(policy, obs_space, action_space, config): +def build_q_model(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> ModelV2: - if not isinstance(action_space, Discrete): + if not isinstance(action_space, gym.spaces.Discrete): raise UnsupportedSpaceException( "Action space {} is not supported for DQN.".format(action_space)) @@ -184,9 +190,9 @@ def build_q_model(policy, obs_space, action_space, config): return policy.q_model -def get_distribution_inputs_and_class(policy, - model, - obs_batch, +def get_distribution_inputs_and_class(policy: Policy, + model: ModelV2, + obs_batch: TensorType, *, explore=True, **kwargs): @@ -198,7 +204,8 @@ def get_distribution_inputs_and_class(policy, return policy.q_values, Categorical, [] # state-out -def build_q_losses(policy, model, _, train_batch): +def build_q_losses(policy: Policy, model, _, + train_batch: SampleBatch) -> TensorType: config = policy.config # q network evaluation q_t, q_logits_t, q_dist_t = compute_q_values( @@ -253,7 +260,8 @@ def build_q_losses(policy, model, _, train_batch): return policy.q_loss.loss -def adam_optimizer(policy, config): +def adam_optimizer(policy: Policy, config: TrainerConfigDict + ) -> "tf.keras.optimizers.Optimizer": if policy.config["framework"] in ["tf2", "tfe"]: return tf.keras.optimizers.Adam( learning_rate=policy.cur_lr, epsilon=config["adam_epsilon"]) @@ -262,7 +270,8 @@ def adam_optimizer(policy, config): learning_rate=policy.cur_lr, epsilon=config["adam_epsilon"]) -def clip_gradients(policy, optimizer, loss): +def clip_gradients(policy: Policy, optimizer: "tf.keras.optimizers.Optimizer", + loss: TensorType) -> ModelGradients: if policy.config["grad_clip"] is not None: grads_and_vars = minimize_and_clip( optimizer, @@ -276,25 +285,28 @@ def clip_gradients(policy, optimizer, loss): return grads_and_vars -def build_q_stats(policy, batch): +def build_q_stats(policy: Policy, batch) -> Dict[str, TensorType]: return dict({ "cur_lr": tf.cast(policy.cur_lr, tf.float64), }, **policy.q_loss.stats) -def setup_early_mixins(policy, obs_space, action_space, config): +def setup_early_mixins(policy: Policy, obs_space, action_space, + config: TrainerConfigDict) -> None: LearningRateSchedule.__init__(policy, config["lr"], config["lr_schedule"]) -def setup_mid_mixins(policy, obs_space, action_space, config): +def setup_mid_mixins(policy: Policy, obs_space, action_space, config) -> None: ComputeTDErrorMixin.__init__(policy) -def setup_late_mixins(policy, obs_space, action_space, config): +def setup_late_mixins(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> None: TargetNetworkMixin.__init__(policy, obs_space, action_space, config) -def compute_q_values(policy, model, obs, explore): +def compute_q_values(policy: Policy, model: ModelV2, obs: TensorType, explore): config = policy.config model_out, state = model({ @@ -361,7 +373,10 @@ def _adjust_nstep(n_step, gamma, obs, actions, rewards, new_obs, dones): rewards[i] += gamma**j * rewards[i + j] -def postprocess_nstep_and_prio(policy, batch, other_agent=None, episode=None): +def postprocess_nstep_and_prio(policy: Policy, + batch: SampleBatch, + other_agent=None, + episode=None) -> SampleBatch: # N-step Q adjustments. if policy.config["n_step"] > 1: _adjust_nstep(policy.config["n_step"], policy.config["gamma"], diff --git a/rllib/agents/dqn/dqn_torch_policy.py b/rllib/agents/dqn/dqn_torch_policy.py index cb6cf77ad..e400f6b24 100644 --- a/rllib/agents/dqn/dqn_torch_policy.py +++ b/rllib/agents/dqn/dqn_torch_policy.py @@ -1,21 +1,27 @@ -from gym.spaces import Discrete +from typing import Dict, List, Tuple +import gym import ray -from ray.rllib.agents.dqn.dqn_tf_policy import postprocess_nstep_and_prio, \ - PRIO_WEIGHTS, Q_SCOPE, Q_TARGET_SCOPE from ray.rllib.agents.a3c.a3c_torch_policy import apply_grad_clipping +from ray.rllib.agents.dqn.dqn_tf_policy import ( + PRIO_WEIGHTS, Q_SCOPE, Q_TARGET_SCOPE, postprocess_nstep_and_prio) from ray.rllib.agents.dqn.dqn_torch_model import DQNTorchModel from ray.rllib.agents.dqn.simple_q_torch_policy import TargetNetworkMixin -from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.models.catalog import ModelCatalog -from ray.rllib.models.torch.torch_action_dist import TorchCategorical +from ray.rllib.models.modelv2 import ModelV2 +from ray.rllib.models.torch.torch_action_dist import (TorchCategorical, + TorchDistributionWrapper) +from ray.rllib.policy.policy import Policy +from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.torch_policy import LearningRateSchedule from ray.rllib.policy.torch_policy_template import build_torch_policy from ray.rllib.utils.error import UnsupportedSpaceException from ray.rllib.utils.exploration.parameter_noise import ParameterNoise from ray.rllib.utils.framework import try_import_torch -from ray.rllib.utils.torch_ops import huber_loss, reduce_mean_ignore_inf, \ - softmax_cross_entropy_with_logits, FLOAT_MIN +from ray.rllib.utils.torch_ops import (FLOAT_MIN, huber_loss, + reduce_mean_ignore_inf, + softmax_cross_entropy_with_logits) +from ray.rllib.utils.typing import TensorType, TrainerConfigDict torch, nn = try_import_torch() F = None @@ -115,9 +121,11 @@ class ComputeTDErrorMixin: self.compute_td_error = compute_td_error -def build_q_model_and_distribution(policy, obs_space, action_space, config): +def build_q_model_and_distribution( + policy: Policy, obs_space: gym.Space, action_space: gym.Space, + config: TrainerConfigDict) -> Tuple[ModelV2, TorchDistributionWrapper]: - if not isinstance(action_space, Discrete): + if not isinstance(action_space, gym.spaces.Discrete): raise UnsupportedSpaceException( "Action space {} is not supported for DQN.".format(action_space)) @@ -179,13 +187,14 @@ def build_q_model_and_distribution(policy, obs_space, action_space, config): return policy.q_model, TorchCategorical -def get_distribution_inputs_and_class(policy, - model, - obs_batch, - *, - explore=True, - is_training=False, - **kwargs): +def get_distribution_inputs_and_class( + policy: Policy, + model: ModelV2, + obs_batch: TensorType, + *, + explore: bool = True, + is_training: bool = False, + **kwargs) -> Tuple[TensorType, type, List[TensorType]]: q_vals = compute_q_values(policy, model, obs_batch, explore, is_training) q_vals = q_vals[0] if isinstance(q_vals, tuple) else q_vals @@ -193,7 +202,8 @@ def get_distribution_inputs_and_class(policy, return policy.q_values, TorchCategorical, [] # state-out -def build_q_losses(policy, model, _, train_batch): +def build_q_losses(policy: Policy, model, _, + train_batch: SampleBatch) -> TensorType: config = policy.config # Q-network evaluation. q_t, q_logits_t, q_probs_t = compute_q_values( @@ -259,22 +269,25 @@ def build_q_losses(policy, model, _, train_batch): return policy.q_loss.loss -def adam_optimizer(policy, config): +def adam_optimizer(policy: Policy, + config: TrainerConfigDict) -> "torch.optim.Optimizer": return torch.optim.Adam( policy.q_func_vars, lr=policy.cur_lr, eps=config["adam_epsilon"]) -def build_q_stats(policy, batch): +def build_q_stats(policy: Policy, batch) -> Dict[str, TensorType]: return dict({ "cur_lr": policy.cur_lr, }, **policy.q_loss.stats) -def setup_early_mixins(policy, obs_space, action_space, config): +def setup_early_mixins(policy: Policy, obs_space, action_space, + config: TrainerConfigDict) -> None: LearningRateSchedule.__init__(policy, config["lr"], config["lr_schedule"]) -def after_init(policy, obs_space, action_space, config): +def after_init(policy: Policy, obs_space: gym.Space, action_space: gym.Space, + config: TrainerConfigDict) -> None: ComputeTDErrorMixin.__init__(policy) TargetNetworkMixin.__init__(policy, obs_space, action_space, config) # Move target net to device (this is done autoatically for the @@ -282,7 +295,11 @@ def after_init(policy, obs_space, action_space, config): policy.target_q_model = policy.target_q_model.to(policy.device) -def compute_q_values(policy, model, obs, explore, is_training=False): +def compute_q_values(policy: Policy, + model: ModelV2, + obs: TensorType, + explore, + is_training: bool = False): config = policy.config model_out, state = model({ @@ -323,12 +340,15 @@ def compute_q_values(policy, model, obs, explore, is_training=False): return value, logits, probs_or_logits -def grad_process_and_td_error_fn(policy, optimizer, loss): +def grad_process_and_td_error_fn(policy: Policy, + optimizer: "torch.optim.Optimizer", + loss: TensorType) -> Dict[str, TensorType]: # Clip grads if configured. return apply_grad_clipping(policy, optimizer, loss) -def extra_action_out_fn(policy, input_dict, state_batches, model, action_dist): +def extra_action_out_fn(policy: Policy, input_dict, state_batches, model, + action_dist) -> Dict[str, TensorType]: return {"q_values": policy.q_values} diff --git a/rllib/agents/dqn/simple_q.py b/rllib/agents/dqn/simple_q.py index d24bb786a..443daf7f8 100644 --- a/rllib/agents/dqn/simple_q.py +++ b/rllib/agents/dqn/simple_q.py @@ -1,14 +1,28 @@ +""" +Simple Q (simple_q) +=================== + +This file defines the distributed Trainer class for the simple Q learning. +See `simple_q_[tf|torch]_policy.py` for the definition of the policy loss. +""" + import logging +from typing import Optional, Type -from ray.rllib.agents.trainer import with_common_config -from ray.rllib.agents.dqn.simple_q_tf_policy import SimpleQTFPolicy from ray.rllib.agents.dqn.dqn import DQNTrainer +from ray.rllib.agents.dqn.simple_q_tf_policy import SimpleQTFPolicy +from ray.rllib.agents.dqn.simple_q_torch_policy import SimpleQTorchPolicy +from ray.rllib.agents.trainer import with_common_config +from ray.rllib.evaluation.worker_set import WorkerSet from ray.rllib.execution.concurrency_ops import Concurrently -from ray.rllib.execution.replay_ops import StoreToReplayBuffer, Replay -from ray.rllib.execution.rollout_ops import ParallelRollouts -from ray.rllib.execution.train_ops import TrainOneStep, UpdateTargetNetwork from ray.rllib.execution.metric_ops import StandardMetricsReporting from ray.rllib.execution.replay_buffer import LocalReplayBuffer +from ray.rllib.execution.replay_ops import Replay, StoreToReplayBuffer +from ray.rllib.execution.rollout_ops import ParallelRollouts +from ray.rllib.execution.train_ops import TrainOneStep, UpdateTargetNetwork +from ray.rllib.policy.policy import Policy +from ray.rllib.utils.typing import TrainerConfigDict +from ray.util.iter import LocalIterator logger = logging.getLogger(__name__) @@ -78,16 +92,22 @@ DEFAULT_CONFIG = with_common_config({ # yapf: enable -def get_policy_class(config): +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 PGTrainer. + If None, use `default_policy` provided in build_trainer(). + """ if config["framework"] == "torch": - from ray.rllib.agents.dqn.simple_q_torch_policy import \ - SimpleQTorchPolicy return SimpleQTorchPolicy - else: - return SimpleQTFPolicy -def execution_plan(workers, config): +def execution_plan(workers: WorkerSet, + config: TrainerConfigDict) -> LocalIterator[dict]: local_replay_buffer = LocalReplayBuffer( num_shards=1, learning_starts=config["learning_starts"], diff --git a/rllib/agents/dqn/simple_q_tf_policy.py b/rllib/agents/dqn/simple_q_tf_policy.py index c6a70615b..526980c1a 100644 --- a/rllib/agents/dqn/simple_q_tf_policy.py +++ b/rllib/agents/dqn/simple_q_tf_policy.py @@ -1,19 +1,24 @@ """Basic example of a DQN policy without any optimizations.""" -from gym.spaces import Discrete import logging +from typing import List, Tuple, Type +import gym import ray -from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.models import ModelCatalog +from ray.rllib.models.modelv2 import ModelV2 +from ray.rllib.models.tf.tf_action_dist import (Categorical, + TFActionDistribution) from ray.rllib.models.torch.torch_action_dist import TorchCategorical -from ray.rllib.models.tf.tf_action_dist import Categorical -from ray.rllib.utils.annotations import override -from ray.rllib.utils.error import UnsupportedSpaceException +from ray.rllib.policy import Policy +from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy import TFPolicy from ray.rllib.policy.tf_policy_template import build_tf_policy +from ray.rllib.utils.annotations import override +from ray.rllib.utils.error import UnsupportedSpaceException from ray.rllib.utils.framework import try_import_tf from ray.rllib.utils.tf_ops import huber_loss, make_tf_callable +from ray.rllib.utils.typing import TensorType, TrainerConfigDict tf1, tf, tfv = try_import_tf() logger = logging.getLogger(__name__) @@ -23,7 +28,8 @@ Q_TARGET_SCOPE = "target_q_func" class TargetNetworkMixin: - def __init__(self, obs_space, action_space, config): + def __init__(self, obs_space: gym.Space, action_space: gym.Space, + config: TrainerConfigDict): @make_tf_callable(self.get_session()) def do_update(): # update_target_fn will be called periodically to copy Q network to @@ -44,9 +50,11 @@ class TargetNetworkMixin: return self.q_func_vars + self.target_q_func_vars -def build_q_models(policy, obs_space, action_space, config): +def build_q_models(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> ModelV2: - if not isinstance(action_space, Discrete): + if not isinstance(action_space, gym.spaces.Discrete): raise UnsupportedSpaceException( "Action space {} is not supported for DQN.".format(action_space)) @@ -72,13 +80,14 @@ def build_q_models(policy, obs_space, action_space, config): return policy.q_model -def get_distribution_inputs_and_class(policy, - q_model, - obs_batch, - *, - explore=True, - is_training=True, - **kwargs): +def get_distribution_inputs_and_class( + policy: Policy, + q_model: ModelV2, + obs_batch: TensorType, + *, + explore=True, + is_training=True, + **kwargs) -> Tuple[TensorType, type, List[TensorType]]: q_vals = compute_q_values(policy, q_model, obs_batch, explore, is_training) q_vals = q_vals[0] if isinstance(q_vals, tuple) else q_vals @@ -88,7 +97,9 @@ def get_distribution_inputs_and_class(policy, Categorical), [] # state-outs -def build_q_losses(policy, model, dist_class, train_batch): +def build_q_losses(policy: Policy, model: ModelV2, + dist_class: Type[TFActionDistribution], + train_batch: SampleBatch) -> TensorType: # q network evaluation q_t = compute_q_values( policy, @@ -131,7 +142,11 @@ def build_q_losses(policy, model, dist_class, train_batch): return loss -def compute_q_values(policy, model, obs, explore, is_training=None): +def compute_q_values(policy: Policy, + model: ModelV2, + obs: TensorType, + explore, + is_training=None) -> TensorType: model_out, _ = model({ SampleBatch.CUR_OBS: obs, "is_training": is_training @@ -141,7 +156,9 @@ def compute_q_values(policy, model, obs, explore, is_training=None): return model_out -def setup_late_mixins(policy, obs_space, action_space, config): +def setup_late_mixins(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> None: TargetNetworkMixin.__init__(policy, obs_space, action_space, config) diff --git a/rllib/agents/dqn/simple_q_torch_policy.py b/rllib/agents/dqn/simple_q_torch_policy.py index 941bacb0e..fbdcc05ae 100644 --- a/rllib/agents/dqn/simple_q_torch_policy.py +++ b/rllib/agents/dqn/simple_q_torch_policy.py @@ -1,15 +1,20 @@ """Basic example of a DQN policy without any optimizations.""" import logging +from typing import Dict +import gym import ray -from ray.rllib.agents.dqn.simple_q_tf_policy import build_q_models, \ - get_distribution_inputs_and_class, compute_q_values -from ray.rllib.policy.sample_batch import SampleBatch +from ray.rllib.agents.dqn.simple_q_tf_policy import ( + build_q_models, compute_q_values, get_distribution_inputs_and_class) +from ray.rllib.models.modelv2 import ModelV2 from ray.rllib.models.torch.torch_action_dist import TorchCategorical +from ray.rllib.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.utils.framework import try_import_torch from ray.rllib.utils.torch_ops import huber_loss +from ray.rllib.utils.typing import TensorType, TrainerConfigDict torch, nn = try_import_torch() F = None @@ -19,7 +24,8 @@ logger = logging.getLogger(__name__) class TargetNetworkMixin: - def __init__(self, obs_space, action_space, config): + def __init__(self, obs_space: gym.Space, action_space: gym.Space, + config: TrainerConfigDict): def do_update(): # Update_target_fn will be called periodically to copy Q network to # target Q network. @@ -30,12 +36,15 @@ class TargetNetworkMixin: self.update_target = do_update -def build_q_model_and_distribution(policy, obs_space, action_space, config): +def build_q_model_and_distribution(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> ModelV2: return build_q_models(policy, obs_space, action_space, config), \ TorchCategorical -def build_q_losses(policy, model, dist_class, train_batch): +def build_q_losses(policy: Policy, model, dist_class, + train_batch: SampleBatch) -> TensorType: # q network evaluation q_t = compute_q_values( policy, @@ -78,12 +87,15 @@ def build_q_losses(policy, model, dist_class, train_batch): return loss -def extra_action_out_fn(policy, input_dict, state_batches, model, action_dist): +def extra_action_out_fn(policy: Policy, input_dict, state_batches, model, + action_dist) -> Dict[str, TensorType]: """Adds q-values to action out dict.""" return {"q_values": policy.q_values} -def setup_late_mixins(policy, obs_space, action_space, config): +def setup_late_mixins(policy: Policy, obs_space: gym.Space, + action_space: gym.Space, + config: TrainerConfigDict) -> None: TargetNetworkMixin.__init__(policy, obs_space, action_space, config) diff --git a/rllib/policy/torch_policy_template.py b/rllib/policy/torch_policy_template.py index 1e1fd4806..ec91a58d1 100644 --- a/rllib/policy/torch_policy_template.py +++ b/rllib/policy/torch_policy_template.py @@ -85,7 +85,7 @@ def build_torch_policy( values given the policy and training batch. If None, will use `TorchPolicy.extra_grad_info()` instead. The stats dict is used for logging (e.g. in TensorBoard). - extra_action_out_fn (Optional[Callable[[Policy, Dict[str, TensorType, + extra_action_out_fn (Optional[Callable[[Policy, Dict[str, TensorType], List[TensorType], ModelV2, TorchDistributionWrapper]], Dict[str, TensorType]]]): Optional callable that returns a dict of extra values to include in experiences. If None, no extra computations