fix bug in states; the states were empty

This commit is contained in:
Brian Delhaisse
2019-06-16 17:37:10 +02:00
parent aa5009f6d1
commit 0cd4fc12ff
8 changed files with 85 additions and 45 deletions
+2 -2
View File
@@ -4,7 +4,7 @@
import numpy as np
from itertools import count
from pyrobolearn.utils.orientation import get_rpy_from_quaternion
from pyrobolearn.utils.transformation import get_rpy_from_quaternion
from pyrobolearn.simulators import BulletSim
from pyrobolearn.worlds import BasicWorld
from pyrobolearn.robots import Cubli
@@ -19,7 +19,7 @@ world = BasicWorld(sim)
scale = 1. # Warning: this does not scale the mass...
position = [0., 0., np.sqrt(2) / 2. * scale + 0.001]
orientation = [0.383, 0, 0, 0.924]
robot = Cubli(sim, position, orientation, scaling=scale)
robot = Cubli(sim, position, orientation, scale=scale)
# print information about the robot
robot.print_info()
+2 -1
View File
@@ -19,7 +19,8 @@ class WAM(ManipulatorRobot):
r"""Wam robot
References:
[1] https://github.com/jhu-lcsr/barrett_model
[1] https://advanced.barrett.com/wam-arm-1
[2] https://github.com/jhu-lcsr/barrett_model
"""
def __init__(self,
+1 -1
View File
@@ -131,7 +131,7 @@ class CounterState(State):
def _read(self):
"""Read the next counter state."""
self.data = self._data + 1
self.data = self.last_data + 1
class PreviousActionState(State):
+63 -28
View File
@@ -65,15 +65,14 @@ class State(object):
"""
def __init__(self, states=(), data=None, space=None, name=None, window_size=1, axis=None, ticks=1):
def __init__(self, states=(), data=None, space=None, window_size=1, axis=None, ticks=1, name=None):
"""
Initialize the state. The state contains some kind of data, or is a state combined of other states.
Args:
states (list/tuple of State): list of states to be combined together (if given, we can not specified data)
data (np.ndarray): data associated to this state
data (np.array): data associated to this state
space (gym.space): space associated with the given data
name (str, None): name of the state. If None, by default, it will have the name of the class.
window_size (int): window size of the state. This is the total number of states we should remember. That
is, if the user wants to remember the current state :math:`s_t` and the previous state :math:`s_{t-1}`,
the window size is 2. By default, the :attr:`window_size` is one which means we only remember the
@@ -87,6 +86,7 @@ class State(object):
(w,n), and for axis=-1 or 1, it will have a shape of (n,w). The :attr:`axis` attribute is only when the
state is not a combination of states, but is given some :attr:`data`.
ticks (int): number of ticks to sleep before getting the next state data.
name (str, None): name of the state. If None, by default, it will have the name of the class.
Warning:
Both arguments can not be provided to the state.
@@ -95,16 +95,20 @@ class State(object):
if states is None:
states = tuple()
# check that the given `states` is a list of states
if not isinstance(states, (list, tuple, set, OrderedSet)):
# # TODO: should check that states is a list of state, however O(N)
# if data is None: # this is in the case someone calls `State(data)`
# data = states
# else:
raise TypeError("Expecting a list, tuple, or (ordered) set of states.")
# TODO: should check that states is a list of state, however O(N)
if data is None and isinstance(states, np.ndarray): # this is in the case someone calls `State(data)`
data = states
states = tuple()
else:
raise TypeError("Expecting a list, tuple, or (ordered) set of states.")
# check that the list of states and the data are not provided together
if len(states) > 0 and data is not None:
raise ValueError("Please specify only one of the argument `states` xor `data`, but not both.")
# Check if data is given
# Check if the data is given, and convert it to a numpy array if necessary
if data is not None:
if not isinstance(data, np.ndarray):
if isinstance(data, (list, tuple)):
@@ -129,15 +133,14 @@ class State(object):
if self._data is None:
self.add(states)
# set window
self._window = None
self._torch_window = None
self.window_size = window_size
# set data windows
self._window, self._torch_window = None, None
self.window_size = window_size # this initializes the windows (FIFO queues)
self.axis = axis
# set ticks and counter
self.cnt = 0
self.ticks = int(ticks)
self._cnt = 0
self.ticks = ticks
# reset state
self.reset()
@@ -190,7 +193,7 @@ class State(object):
# stack the data in the window
return [np.stack(self.window.queue, axis=self.axis)] # stack
# if multiple states
# if multiple states, return the combined data associated to each state
return [state.data[0] for state in self._states]
@data.setter
@@ -261,6 +264,13 @@ class State(object):
# return the data
return fused_state.data
@property
def last_data(self):
"""Return the last provided data."""
if self.has_data():
return self.window[-1]
return [state.last_data for state in self._states]
@property
def torch_data(self):
"""
@@ -348,6 +358,13 @@ class State(object):
# return the data
return fused_state.torch_data
@property
def last_torch_data(self):
"""Return the last provided torch data."""
if self.has_data():
return self.torch_window[-1]
return [state.last_torch_data for state in self._states]
@property
def vec_data(self):
"""
@@ -527,6 +544,19 @@ class State(object):
"{}".format(type(axis)))
self._axis = axis
@property
def ticks(self):
"""Return the number of ticks to sleep before getting the next state data."""
return self._ticks
@ticks.setter
def ticks(self, ticks):
"""Set the number of ticks to sleep before getting the next state data."""
ticks = int(ticks)
if ticks < 1:
ticks = 1
self._ticks = ticks
###########
# Methods #
###########
@@ -554,6 +584,7 @@ class State(object):
def has_data(self):
"""Check if the state has data."""
return self._data is not None
# return len(self._states) == 0
def has_space(self):
"""Check if the state has a space."""
@@ -595,16 +626,17 @@ class State(object):
Read the state values from the simulator for each state, set it and return their values.
"""
# if time to read
if self.cnt % self.ticks == 0:
if self._cnt % self.ticks == 0:
if self.has_data(): # read the current state
self._read()
else: # read each state
# if multiple states, read each state
if self.has_states(): # read each state
for state in self.states:
state._read()
else: # else, read the current state
self._read()
# increment counter
self.cnt += 1
self._cnt += 1
# return the data
return self.data
@@ -613,19 +645,21 @@ class State(object):
"""
Reset the state. This has to be overwritten in the child class.
"""
self.cnt = 0
self._cnt = 0
self._read()
def reset(self):
"""
Some states need to be reset. It returns the initial state.
"""
self.cnt = 0
if self.has_data(): # reset the current state
self._reset()
else: # reset each state
self._cnt = 0
# if multiple states, reset each state
if self.has_states():
for state in self.states:
state._reset()
else: # else, reset this state
self._reset()
# return the first state data
return self.data # self.read()
@@ -817,12 +851,13 @@ class State(object):
# if there is one state
if self.has_data():
if self.__class__ == class_type or self.__class__.__name__.lower() == class_type:
if self.__class__ == class_type or self.__class__.__name__.lower() == class_type or self.name == class_type:
return self
# the state has multiple states, thus we go through each state
for state in self.states:
if state.__class__ == class_type or state.__class__.__name__.lower() == class_type:
if state.__class__ == class_type or state.__class__.__name__.lower() == class_type or \
state.name == class_type:
return state
#############
+10 -10
View File
@@ -136,7 +136,7 @@ class CumulativeTimeState(TimeState):
def _read(self):
"""Read the next cumulative time state."""
next_time = time.time()
self.data = self._data + (next_time - self.current_time)
self.data = self.last_data + (next_time - self.current_time)
self.current_time = next_time
@@ -187,10 +187,10 @@ class PhaseState(TimeState):
def _read(self):
"""Read the next linear phase state."""
if self.sign > 0 and self._data[0] < self.end_value:
self.data = np.minimum(self._data + self.dphase, self.end_value)
elif self.sign < 0 and self._data[0] > self.end_value:
self.data = np.maximum(self._data + self.dphase, self.end_value)
if self.sign > 0 and self.last_data[0] < self.end_value:
self.data = np.minimum(self.last_data + self.dphase, self.end_value)
elif self.sign < 0 and self.last_data[0] > self.end_value:
self.data = np.maximum(self.last_data + self.dphase, self.end_value)
class ExponentialPhaseState(TimeState):
@@ -251,9 +251,9 @@ class ExponentialPhaseState(TimeState):
self.data = np.array([self.s0]) * np.exp(self.a * self.t)
if self.sf is not None:
if self.a < 0:
self.data = np.maximum(self._data, self.sf)
self.data = np.maximum(self.last_data, self.sf)
elif self.a > 0:
self.data = np.minimum(self._data, self.sf)
self.data = np.minimum(self.last_data, self.sf)
# alias
@@ -296,10 +296,10 @@ class RhythmicPhase(PhaseState):
def _read(self):
"""Read the next rhythmic phase state."""
self.data = self._data + self.dphase
if self.sign > 0 and self._data[0] >= self.end_value:
self.data = self.last_data + self.dphase
if self.sign > 0 and self.last_data[0] >= self.end_value:
self.data = np.array([self.start_value])
if self.sign < 0 and self._data[0] <= self.end_value:
if self.sign < 0 and self.last_data[0] <= self.end_value:
self.data = np.array([self.start_value])
+3 -2
View File
@@ -276,8 +276,9 @@ class ILTask(Task):
Perform one step in the training of the policy(ies) using the recorded data.
"""
if self.training_enabled:
# print(self.recorders[0].data)
# print(np.array(self.recorders[0].data).shape)
# print("data: ", self.recorders[0].data)
# print("shape: ", np.array(self.recorders[0].data).shape)
# take data from recorder
data = np.array([data[0] for data in self.recorders[0].data]).T
@@ -337,7 +337,7 @@ class BridgeMouseKeyboardWorld(Bridge):
if object_id != -1 and self.world.is_robot_id(object_id): # valid object
# Set robot and link_id
self.robot, self.link_id = self.world.get_robot(object_id), link_id
self.robot, self.link_id = self.world.get_body(object_id), link_id
width, height = camera[:2]
x_screen = np.array([width/2, height/10, 0.95, 1])
@@ -311,6 +311,9 @@ if __name__ == '__main__':
for i, item in enumerate(q):
print("Item {}: {}".format(i, item))
# print last item put in the queue
print("Last item in queue: {}".format(q[-1]))
# create stack
stack = LIFOQueue(maxsize=2)
print("\nInitial stack (maxsize={}): {}".format(stack.maxsize, stack))