mirror of
https://github.com/wassname/ray.git
synced 2026-08-11 11:24:51 +08:00
[RLlib] Add type annotations for agents/dqn (#10626)
This commit is contained in:
+17
-12
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user