update policies to account for refactored values + add predict()

This commit is contained in:
Brian Delhaisse
2019-04-12 02:52:13 +02:00
parent d0073aaaa5
commit a4bfe94b28
2 changed files with 58 additions and 10 deletions
+28 -10
View File
@@ -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
+30
View File
@@ -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.