Added more tests
This commit is contained in:
@@ -2,3 +2,6 @@ class BaseController:
|
||||
|
||||
def get_action(self, des_pos, des_vel, c_pos, c_vel):
|
||||
raise NotImplementedError
|
||||
|
||||
def __call__(self, des_pos, des_vel, c_pos, c_vel):
|
||||
return self.get_action(des_pos, des_vel, c_pos, c_vel)
|
||||
|
||||
@@ -18,7 +18,8 @@ class MetaWorldController(BaseController):
|
||||
cur_pos = c_pos[:-1]
|
||||
xyz_pos = des_pos[:-1]
|
||||
|
||||
assert xyz_pos.shape == cur_pos.shape, \
|
||||
f"Mismatch in dimension between desired position {xyz_pos.shape} and current position {cur_pos.shape}"
|
||||
if xyz_pos.shape != cur_pos.shape:
|
||||
raise ValueError(f"Mismatch in dimension between desired position"
|
||||
f" {xyz_pos.shape} and current position {cur_pos.shape}")
|
||||
trq = np.hstack([(xyz_pos - cur_pos), gripper_pos])
|
||||
return trq
|
||||
|
||||
@@ -8,7 +8,6 @@ class PDController(BaseController):
|
||||
A PD-Controller. Using position and velocity information from a provided environment,
|
||||
the tracking_controller calculates a response based on the desired position and velocity
|
||||
|
||||
:param env: A position environment
|
||||
:param p_gains: Factors for the proportional gains
|
||||
:param d_gains: Factors for the differential gains
|
||||
"""
|
||||
@@ -20,9 +19,11 @@ class PDController(BaseController):
|
||||
self.d_gains = d_gains
|
||||
|
||||
def get_action(self, des_pos, des_vel, c_pos, c_vel):
|
||||
assert des_pos.shape == c_pos.shape, \
|
||||
f"Mismatch in dimension between desired position {des_pos.shape} and current position {c_pos.shape}"
|
||||
assert des_vel.shape == c_vel.shape, \
|
||||
f"Mismatch in dimension between desired velocity {des_vel.shape} and current velocity {c_vel.shape}"
|
||||
if des_pos.shape != c_pos.shape:
|
||||
raise ValueError(f"Mismatch in dimension between desired position "
|
||||
f"{des_pos.shape} and current position {c_pos.shape}")
|
||||
if des_vel.shape != c_vel.shape:
|
||||
raise ValueError(f"Mismatch in dimension between desired velocity"
|
||||
f" {des_vel.shape} and current velocity {c_vel.shape}")
|
||||
trq = self.p_gains * (des_pos - c_pos) + self.d_gains * (des_vel - c_vel)
|
||||
return trq
|
||||
|
||||
Reference in New Issue
Block a user