From 26176ec5706cc7b3c153676554e6f51f071bf545 Mon Sep 17 00:00:00 2001 From: bcahlit Date: Mon, 2 Nov 2020 19:22:33 +0800 Subject: [PATCH] [RLlib] Fix epsilon_greedy on nested_action_spaces only in pytorch (#11453) * [RLlib] Fix epsilon_greedy on nested_action_spaces only in pytorch * epsilon_greedy on Continuous action * formatt * Fix error * fix format * fix bug * increase speed * Update rllib/utils/exploration/epsilon_greedy.py * Update rllib/utils/exploration/epsilon_greedy.py * Update rllib/utils/exploration/epsilon_greedy.py Co-authored-by: Sven Mika --- rllib/utils/exploration/epsilon_greedy.py | 85 +++++++++++++++-------- 1 file changed, 57 insertions(+), 28 deletions(-) diff --git a/rllib/utils/exploration/epsilon_greedy.py b/rllib/utils/exploration/epsilon_greedy.py index c45f712d2..5ce114249 100644 --- a/rllib/utils/exploration/epsilon_greedy.py +++ b/rllib/utils/exploration/epsilon_greedy.py @@ -1,6 +1,10 @@ import numpy as np +import tree +import random from typing import Union, Optional +from ray.rllib.models.torch.torch_action_dist \ + import TorchMultiActionDistribution from ray.rllib.models.action_dist import ActionDistribution from ray.rllib.utils.annotations import override from ray.rllib.utils.exploration.exploration import Exploration, TensorType @@ -71,25 +75,29 @@ class EpsilonGreedy(Exploration): timestep: Union[int, TensorType], explore: bool = True): - q_values = action_distribution.inputs if self.framework in ["tf2", "tf", "tfe"]: - return self._get_tf_exploration_action_op(q_values, explore, - timestep) + return self._get_tf_exploration_action_op(action_distribution, + explore, timestep) else: - return self._get_torch_exploration_action(q_values, explore, - timestep) + return self._get_torch_exploration_action(action_distribution, + explore, timestep) - def _get_tf_exploration_action_op(self, q_values: TensorType, + def _get_tf_exploration_action_op(self, + action_distribution: ActionDistribution, explore: Union[bool, TensorType], timestep: Union[int, TensorType]): """TF method to produce the tf op for an epsilon exploration action. Args: - q_values (Tensor): The Q-values coming from some q-model. + action_distribution (ActionDistribution): The instantiated + ActionDistribution object to work with when creating + exploration actions. Returns: tf.Tensor: The tf exploration-action op. """ + # TODO: Support MultiActionDistr for tf. + q_values = action_distribution.inputs epsilon = self.epsilon_schedule(timestep if timestep is not None else self.last_timestep) @@ -125,41 +133,62 @@ class EpsilonGreedy(Exploration): with tf1.control_dependencies([assign_op]): return action, tf.zeros_like(action, dtype=tf.float32) - def _get_torch_exploration_action(self, q_values: TensorType, - explore: bool, - timestep: Union[int, TensorType]): + def _get_torch_exploration_action( + self, action_distribution: ActionDistribution, explore: bool, + timestep: Union[int, TensorType]): """Torch method to produce an epsilon exploration action. Args: - q_values (Tensor): The Q-values coming from some Q-model. + action_distribution (ActionDistribution): The instantiated + ActionDistribution object to work with when creating + exploration actions. Returns: torch.Tensor: The exploration-action. """ + q_values = action_distribution.inputs self.last_timestep = timestep - _, exploit_action = torch.max(q_values, 1) - action_logp = torch.zeros_like(exploit_action) + exploit_action = action_distribution.deterministic_sample() + batch_size = q_values.size()[0] + action_logp = torch.zeros(batch_size, dtype=torch.float) # Explore. if explore: # Get the current epsilon. epsilon = self.epsilon_schedule(self.last_timestep) - batch_size = q_values.size()[0] - # Mask out actions, whose Q-values are -inf, so that we don't - # even consider them for exploration. - random_valid_action_logits = torch.where( - q_values <= FLOAT_MIN, - torch.ones_like(q_values) * 0.0, torch.ones_like(q_values)) - # A random action. - random_actions = torch.squeeze( - torch.multinomial(random_valid_action_logits, 1), axis=1) - # Pick either random or greedy. - action = torch.where( - torch.empty( - (batch_size, )).uniform_().to(self.device) < epsilon, - random_actions, exploit_action) + if isinstance(action_distribution, TorchMultiActionDistribution): + exploit_action = tree.flatten(exploit_action) + for i in range(batch_size): + if random.random() < epsilon: + # TODO: (bcahlit) Mask out actions + random_action = tree.flatten( + self.action_space.sample()) + for j in range(len(exploit_action)): + exploit_action[j][i] = torch.tensor( + random_action[j]) + exploit_action = tree.unflatten_as( + action_distribution.action_space_struct, + exploit_action) - return action, action_logp + return exploit_action, action_logp + + else: + # Mask out actions, whose Q-values are -inf, so that we don't + # even consider them for exploration. + random_valid_action_logits = torch.where( + q_values <= FLOAT_MIN, + torch.ones_like(q_values) * 0.0, torch.ones_like(q_values)) + # A random action. + random_actions = torch.squeeze( + torch.multinomial(random_valid_action_logits, 1), axis=1) + + # Pick either random or greedy. + action = torch.where( + torch.empty( + (batch_size, )).uniform_().to(self.device) < epsilon, + random_actions, exploit_action) + + return action, action_logp # Return the deterministic "sample" (argmax) over the logits. else: return exploit_action, action_logp