From 40d95ecde8ba046f943d648daf20669cd0d2840e Mon Sep 17 00:00:00 2001 From: Brian Delhaisse Date: Wed, 17 Apr 2019 18:05:59 +0200 Subject: [PATCH] correct few errors in policies --- pyrobolearn/policies/basic_policy.py | 3 +++ pyrobolearn/policies/policy.py | 13 +++++++++---- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/pyrobolearn/policies/basic_policy.py b/pyrobolearn/policies/basic_policy.py index 6c04a90..7628d86 100644 --- a/pyrobolearn/policies/basic_policy.py +++ b/pyrobolearn/policies/basic_policy.py @@ -42,6 +42,9 @@ class RandomPolicy(Policy): np.random.seed(seed) self._seed = seed + def reset(self): + pass + def predict(self, state=None, to_numpy=True): spaces = self.actions.space return [space.sample() for space in spaces] diff --git a/pyrobolearn/policies/policy.py b/pyrobolearn/policies/policy.py index 714b323..9192a32 100644 --- a/pyrobolearn/policies/policy.py +++ b/pyrobolearn/policies/policy.py @@ -504,6 +504,10 @@ class Policy(object): action_data[idx] = self.__convert_to_numpy(discrete_data, to_numpy=to_numpy) else: action_data[idx] = self.__convert_to_numpy(data, to_numpy=to_numpy) + + elif isinstance(data, (float, int)): + action.data = int(data) + action_data[idx] = int(data) # elif isinstance(data, (float, int)): # discrete_data = np.argmax(data) # action.data = discrete_data @@ -511,7 +515,7 @@ class Policy(object): # action_data[idx] = discrete_data else: raise TypeError( - "Expecting the `data` action to be a numpy array or torch.Tensor, instead got: " + "Expecting the `data` action to be an int, numpy array, torch.Tensor, instead got: " "{}".format(type(data))) else: # continuous action if isinstance(data, np.ndarray): @@ -519,8 +523,8 @@ class Policy(object): elif isinstance(data, torch.Tensor): action.torch_data = data action_data[idx] = self.__convert_to_numpy(data, to_numpy=to_numpy) - # elif isinstance(data, (float, int)): - # pass + elif isinstance(data, (float, int)): + action.data = data else: raise TypeError( "Expecting `data` to be a numpy array or torch.Tensor, instead got: " @@ -682,7 +686,8 @@ class Policy(object): Returns: Action: action """ - return self.act(*args, **kwargs) + return self.act(state=state, deterministic=deterministic, to_numpy=to_numpy, return_logits=return_logits, + apply_action=apply_action) def __repr__(self): """Return representation of python object."""