[RLlib] Add type annotations for agents/dqn (#10626)

This commit is contained in:
desktable
2020-09-09 18:55:26 +02:00
committed by GitHub
parent 153813936b
commit 799318d7d7
7 changed files with 183 additions and 94 deletions
+17 -12
View File
@@ -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
+35 -20
View File
@@ -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"],
+44 -24
View File
@@ -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}
+31 -11
View File
@@ -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"],
+35 -18
View File
@@ -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)
+20 -8
View File
@@ -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)
+1 -1
View File
@@ -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