diff --git a/pyrobolearn/policies/basic_policy.py b/pyrobolearn/policies/basic_policy.py index 59bc737..6c04a90 100644 --- a/pyrobolearn/policies/basic_policy.py +++ b/pyrobolearn/policies/basic_policy.py @@ -8,8 +8,8 @@ import numpy as np import torch from pyrobolearn.policies.policy import Policy -from pyrobolearn.values.value import ParametrizedStateOutputActionValue from pyrobolearn.approximators import LinearApproximator +from pyrobolearn.values.value import ParametrizedQValueOutput __author__ = "Brian Delhaisse" @@ -108,7 +108,7 @@ class LinearPolicy(Policy): class PolicyFromQValue(Policy): - r"""Policy from state-action value function approximator + r"""Policy from Q-value function approximator This computes the optimal discrete action :math:`a` using the underlying value function approximator :math:`Q(s,a)` which given the state as input computes the Q-value for each discrete action. The policy select @@ -124,7 +124,7 @@ class PolicyFromQValue(Policy): Initialize the Policy from the value function approximator. Args: - value (ParametrizedStateOutputActionValue): trainable value function approximator. + value (ParametrizedQValueOutput): trainable value function approximator. rate (int, float): rate (float) at which the policy operates if we are operating in real-time. If we are stepping deterministically in the simulator, it represents the number of ticks (int) to sleep before executing the model. @@ -134,7 +134,7 @@ class PolicyFromQValue(Policy): **kwargs (dict): dictionary of arguments """ self.value = value - super(PolicyFromQValue, self).__init__(value.state, value.action, model=value, rate=rate, + super(PolicyFromQValue, self).__init__(states=value.state, actions=value.action, model=value, rate=rate, preprocessors=preprocessors, postprocessors=postprocessors, *args, **kwargs) @@ -152,20 +152,21 @@ class PolicyFromQValue(Policy): """Set the value function approximator.""" # TODO: need to check that input and output dimensions of the new value function approximator match the ones # from the previous model. - if not isinstance(value, ParametrizedStateOutputActionValue): + if not isinstance(value, ParametrizedQValueOutput): raise TypeError("Expecting the given `value` function approximator to be an instance of " - "`ParametrizedStateOutputActionValue`, instead got: {}".format(type(value))) + "`ParametrizedQValueOutput`, instead got: {}".format(type(value))) self._value = value ########### # Methods # ########### - def inner_predict(self, state, to_numpy=False, return_logits=True, set_output_data=False): + def inner_predict(self, state, deterministic=True, to_numpy=False, return_logits=True, set_output_data=False): """Inner prediction step. Args: state ((list of) torch.Tensor, (list of) np.array): state data. + deterministic (bool): True by default. It can only be set to False, if the policy is stochastic. to_numpy (bool): If True, it will convert the data (torch.Tensors) to numpy arrays. return_logits (bool): If True, in the case of discrete outputs, it will return the logits. set_output_data (bool): If True, it will set the predicted output data to the outputs given to the @@ -174,10 +175,27 @@ class PolicyFromQValue(Policy): Returns: (list of) torch.Tensor, (list of) np.array: predicted action data. """ - action = self.model.evaluate(state, to_numpy=to_numpy) + values = self.value.evaluate(state, to_numpy=to_numpy) + if return_logits: + return values if to_numpy: - return np.argmax(action) - return torch.argmax(action, dim=0, keepdim=True) + return np.argmax(values) + return torch.argmax(values, dim=0, keepdim=True) + + # def evaluate(self, state, to_numpy=False): + # """ + # Evaluate the state. + # + # Args: + # state (None, State, (list of) np.array, (list of) torch.Tensor): state input data. If None, it will get + # the data from the inputs that were given at the initialization. + # to_numpy (bool): If True, it will convert the data (torch.Tensors) to numpy arrays. + # + # Returns: + # torch.Tensor, np.array: + # """ + # values = self.value.evaluate(state, to_numpy=to_numpy) + # return values # def act(self, state=None, deterministic=True, to_numpy=True, return_logits=False, apply_action=True): # pass diff --git a/pyrobolearn/policies/policy.py b/pyrobolearn/policies/policy.py index 2d8e205..1aac9da 100644 --- a/pyrobolearn/policies/policy.py +++ b/pyrobolearn/policies/policy.py @@ -521,6 +521,36 @@ class Policy(object): return action_data + def predict(self, state=None, deterministic=True, to_numpy=False, return_logits=True): + """Predict the action given the state. + + This does not set the action data in the action instances, nor apply the actions in the simulator. Instead, + it gets the state data, preprocess it, predict using the actions using the inner model, then post-process + the actions, and return the resulting action data. + + Args: + state (State): current state + deterministic (bool): True by default. It can only be set to False, if the policy is stochastic. + to_numpy (bool): If True, it will convert the data (torch.Tensors) to numpy arrays. + return_logits (bool): If True, in the case of discrete outputs, it will return the logits. + + Returns: + (list of) torch.Tensor: action data + """ + # get the state data + state_data = self.get_state_data(state=state) + + # pre-process the state data + state_data = self.preprocess(state_data) + + # predict the output using the inner model + action_data = self.inner_predict(state_data, to_numpy=False, return_logits=True, set_output_data=False) + + # post-process the action data + action_data = self.postprocess(action_data) + + return action_data + def act(self, state=None, deterministic=True, to_numpy=True, return_logits=False, apply_action=True): """Perform the action given the state.