Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eccbf32794 | ||
|
|
26fe6082fc | ||
|
|
e3a2932c42 | ||
|
|
12a3558950 |
@@ -0,0 +1,32 @@
|
||||
name: Run pytest
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
# - master
|
||||
- release
|
||||
pull_request:
|
||||
branches:
|
||||
# - master
|
||||
- release
|
||||
|
||||
jobs:
|
||||
pytest:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.x'
|
||||
|
||||
- name: Install package dependencies
|
||||
run: pip install .[testing]
|
||||
|
||||
- name: Install Metaworld
|
||||
run: pip install metaworld@git+https://github.com/Farama-Foundation/Metaworld.git@d155d0051630bb365ea6a824e02c66c068947439#egg=metaworld
|
||||
|
||||
- name: Run pytest
|
||||
run: pytest
|
||||
@@ -115,7 +115,6 @@ class AntJumpEnv(AntEnvCustomXML):
|
||||
contact_force_range=contact_force_range,
|
||||
reset_noise_scale=reset_noise_scale,
|
||||
exclude_current_positions_from_observation=exclude_current_positions_from_observation, **kwargs)
|
||||
self.render_active = False
|
||||
|
||||
def step(self, action):
|
||||
self.current_step += 1
|
||||
@@ -154,15 +153,8 @@ class AntJumpEnv(AntEnvCustomXML):
|
||||
}
|
||||
truncated = False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return obs, reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _get_obs(self):
|
||||
return np.append(super()._get_obs(), self.goal)
|
||||
|
||||
|
||||
@@ -44,7 +44,6 @@ class BeerPongEnv(MujocoEnv, utils.EzPickle):
|
||||
}
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
utils.EzPickle.__init__(self)
|
||||
self._steps = 0
|
||||
# Small Context -> Easier. Todo: Should we do different versions?
|
||||
# self.xml_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "assets", "beerpong_wo_cup.xml")
|
||||
@@ -90,7 +89,7 @@ class BeerPongEnv(MujocoEnv, utils.EzPickle):
|
||||
observation_space=self.observation_space,
|
||||
**kwargs
|
||||
)
|
||||
self.render_active = False
|
||||
utils.EzPickle.__init__(self)
|
||||
|
||||
@property
|
||||
def start_pos(self):
|
||||
@@ -170,15 +169,8 @@ class BeerPongEnv(MujocoEnv, utils.EzPickle):
|
||||
|
||||
truncated = False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return ob, reward, terminated, truncated, infos
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _get_obs(self):
|
||||
theta = self.data.qpos.flat[:7].copy()
|
||||
theta_dot = self.data.qvel.flat[:7].copy()
|
||||
|
||||
@@ -4,7 +4,6 @@ import numpy as np
|
||||
from gymnasium import utils, spaces
|
||||
from gymnasium.envs.mujoco import MujocoEnv
|
||||
from fancy_gym.envs.mujoco.box_pushing.box_pushing_utils import rot_to_quat, get_quaternion_error, rotation_distance
|
||||
from fancy_gym.envs.mujoco.box_pushing.box_pushing_utils import rot_to_quat, get_quaternion_error, rotation_distance
|
||||
from fancy_gym.envs.mujoco.box_pushing.box_pushing_utils import q_max, q_min, q_dot_max, q_torque_max
|
||||
from fancy_gym.envs.mujoco.box_pushing.box_pushing_utils import desired_rod_quat
|
||||
|
||||
@@ -61,7 +60,6 @@ class BoxPushingEnvBase(MujocoEnv, utils.EzPickle):
|
||||
frame_skip=self.frame_skip,
|
||||
observation_space=self.observation_space, **kwargs)
|
||||
self.action_space = spaces.Box(low=-1, high=1, shape=(7,))
|
||||
self.render_active = False
|
||||
|
||||
def step(self, action):
|
||||
action = 10 * np.clip(action, self.action_space.low, self.action_space.high)
|
||||
@@ -110,15 +108,8 @@ class BoxPushingEnvBase(MujocoEnv, utils.EzPickle):
|
||||
terminated = episode_end and infos['is_success']
|
||||
truncated = episode_end and not infos['is_success']
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return obs, reward, terminated, truncated, infos
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def reset_model(self):
|
||||
# rest box to initial position
|
||||
self.set_state(self.init_qpos_box_pushing, self.init_qvel_box_pushing)
|
||||
|
||||
@@ -60,11 +60,7 @@ class HalfCheetahEnvCustomXML(HalfCheetahEnv):
|
||||
default_camera_config=DEFAULT_CAMERA_CONFIG,
|
||||
**kwargs,
|
||||
)
|
||||
self.render_active = False
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
class HalfCheetahJumpEnv(HalfCheetahEnvCustomXML):
|
||||
"""
|
||||
@@ -124,9 +120,6 @@ class HalfCheetahJumpEnv(HalfCheetahEnvCustomXML):
|
||||
'max_height': self.max_height
|
||||
}
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def _get_obs(self):
|
||||
|
||||
@@ -88,12 +88,6 @@ class HopperEnvCustomXML(HopperEnv):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.render_active = False
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
|
||||
class HopperJumpEnv(HopperEnvCustomXML):
|
||||
"""
|
||||
@@ -207,10 +201,6 @@ class HopperJumpEnv(HopperEnvCustomXML):
|
||||
healthy=self.is_healthy,
|
||||
contact_dist=self.contact_dist or 0
|
||||
)
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def _get_obs(self):
|
||||
|
||||
@@ -140,9 +140,6 @@ class HopperJumpOnBoxEnv(HopperEnvCustomXML):
|
||||
|
||||
truncated = self.current_step >= self.max_episode_steps and not terminated
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def _get_obs(self):
|
||||
|
||||
@@ -61,8 +61,6 @@ class HopperThrowEnv(HopperEnvCustomXML):
|
||||
exclude_current_positions_from_observation=exclude_current_positions_from_observation,
|
||||
**kwargs)
|
||||
|
||||
self.render_active = False
|
||||
|
||||
def step(self, action):
|
||||
self.current_step += 1
|
||||
self.do_simulation(action, self.frame_skip)
|
||||
@@ -96,15 +94,8 @@ class HopperThrowEnv(HopperEnvCustomXML):
|
||||
}
|
||||
truncated = False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _get_obs(self):
|
||||
return np.append(super()._get_obs(), self.goal)
|
||||
|
||||
|
||||
@@ -68,7 +68,6 @@ class HopperThrowInBasketEnv(HopperEnvCustomXML):
|
||||
reset_noise_scale=reset_noise_scale,
|
||||
exclude_current_positions_from_observation=exclude_current_positions_from_observation,
|
||||
**kwargs)
|
||||
self.render_active = False
|
||||
|
||||
def step(self, action):
|
||||
|
||||
@@ -119,15 +118,8 @@ class HopperThrowInBasketEnv(HopperEnvCustomXML):
|
||||
}
|
||||
truncated = False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _get_obs(self):
|
||||
return np.append(super()._get_obs(), self.basket_x)
|
||||
|
||||
|
||||
@@ -47,8 +47,6 @@ class ReacherEnv(MujocoEnv, utils.EzPickle):
|
||||
**kwargs
|
||||
)
|
||||
|
||||
self.render_active = False
|
||||
|
||||
def step(self, action):
|
||||
self._steps += 1
|
||||
|
||||
@@ -79,15 +77,8 @@ class ReacherEnv(MujocoEnv, utils.EzPickle):
|
||||
goal=self.goal if hasattr(self, "goal") else None
|
||||
)
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return ob, reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def distance_reward(self):
|
||||
vec = self.get_body_com("fingertip") - self.get_body_com("target")
|
||||
return -self._reward_weight * np.linalg.norm(vec)
|
||||
|
||||
@@ -71,8 +71,6 @@ class TableTennisEnv(MujocoEnv, utils.EzPickle):
|
||||
observation_space=self.observation_space,
|
||||
**kwargs)
|
||||
|
||||
self.render_active = False
|
||||
|
||||
if ctxt_dim == 2:
|
||||
self.context_bounds = CONTEXT_BOUNDS_2DIMS
|
||||
elif ctxt_dim == 4:
|
||||
@@ -160,15 +158,8 @@ class TableTennisEnv(MujocoEnv, utils.EzPickle):
|
||||
|
||||
terminated, truncated = self._terminated, False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return self._get_obs(), reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _contact_checker(self, id_1, id_2):
|
||||
for coni in range(0, self.data.ncon):
|
||||
con = self.data.contact[coni]
|
||||
|
||||
@@ -79,8 +79,6 @@ class Walker2dEnvCustomXML(Walker2dEnv):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.render_active = False
|
||||
|
||||
|
||||
class Walker2dJumpEnv(Walker2dEnvCustomXML):
|
||||
"""
|
||||
@@ -147,15 +145,8 @@ class Walker2dJumpEnv(Walker2dEnvCustomXML):
|
||||
}
|
||||
truncated = False
|
||||
|
||||
if self.render_active and self.render_mode=='human':
|
||||
self.render()
|
||||
|
||||
return observation, reward, terminated, truncated, info
|
||||
|
||||
def render(self):
|
||||
self.render_active = True
|
||||
return super().render()
|
||||
|
||||
def _get_obs(self):
|
||||
return np.append(super()._get_obs(), self.goal)
|
||||
|
||||
|
||||
@@ -3,14 +3,14 @@ import fancy_gym
|
||||
|
||||
|
||||
def example_run_replanning_env(env_name="fancy_ProDMP/BoxPushingDenseReplan-v0", seed=1, iterations=1, render=False):
|
||||
env = gym.make(env_name, render_mode='human' if render else None)
|
||||
env = gym.make(env_name)
|
||||
env.reset(seed=seed)
|
||||
for i in range(iterations):
|
||||
while True:
|
||||
ac = env.action_space.sample()
|
||||
obs, reward, terminated, truncated, info = env.step(ac)
|
||||
if render:
|
||||
env.render()
|
||||
env.render(mode="human")
|
||||
if terminated or truncated:
|
||||
env.reset()
|
||||
break
|
||||
@@ -38,13 +38,13 @@ def example_custom_replanning_envs(seed=0, iteration=100, render=True):
|
||||
'replanning_schedule': lambda pos, vel, obs, action, t: t % 25 == 0,
|
||||
'condition_on_desired': True}
|
||||
|
||||
base_env = gym.make(base_env_id, render_mode='human' if render else None)
|
||||
base_env = gym.make(base_env_id)
|
||||
env = fancy_gym.make_bb(env=base_env, wrappers=wrappers, black_box_kwargs=black_box_kwargs,
|
||||
traj_gen_kwargs=trajectory_generator_kwargs, controller_kwargs=controller_kwargs,
|
||||
phase_kwargs=phase_generator_kwargs, basis_kwargs=basis_generator_kwargs,
|
||||
seed=seed)
|
||||
if render:
|
||||
env.render()
|
||||
env.render(mode="human")
|
||||
|
||||
obs = env.reset()
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ def example_dmc(env_id="dm_control/fish-swim", seed=1, iterations=1000, render=T
|
||||
Returns:
|
||||
|
||||
"""
|
||||
env = gym.make(env_id, render_mode='human' if render else None)
|
||||
env = gym.make(env_id)
|
||||
rewards = 0
|
||||
obs = env.reset(seed=seed)
|
||||
print("observation shape:", env.observation_space.shape)
|
||||
@@ -26,7 +26,7 @@ def example_dmc(env_id="dm_control/fish-swim", seed=1, iterations=1000, render=T
|
||||
for i in range(iterations):
|
||||
ac = env.action_space.sample()
|
||||
if render:
|
||||
env.render()
|
||||
env.render(mode="human")
|
||||
obs, reward, terminated, truncated, info = env.step(ac)
|
||||
rewards += reward
|
||||
|
||||
@@ -84,7 +84,7 @@ def example_custom_dmc_and_mp(seed=1, iterations=1, render=True):
|
||||
# basis_generator_kwargs = {'basis_generator_type': 'rbf',
|
||||
# 'num_basis': 5
|
||||
# }
|
||||
base_env = gym.make(base_env_id, render_mode='human' if render else None)
|
||||
base_env = gym.make(base_env_id)
|
||||
env = fancy_gym.make_bb(env=base_env, wrappers=wrappers, black_box_kwargs={},
|
||||
traj_gen_kwargs=trajectory_generator_kwargs, controller_kwargs=controller_kwargs,
|
||||
phase_kwargs=phase_generator_kwargs, basis_kwargs=basis_generator_kwargs,
|
||||
@@ -96,7 +96,7 @@ def example_custom_dmc_and_mp(seed=1, iterations=1, render=True):
|
||||
# It is also possible to change them mode multiple times when
|
||||
# e.g. only every nth trajectory should be displayed.
|
||||
if render:
|
||||
env.render()
|
||||
env.render(mode="human")
|
||||
|
||||
rewards = 0
|
||||
obs = env.reset()
|
||||
@@ -115,7 +115,7 @@ def example_custom_dmc_and_mp(seed=1, iterations=1, render=True):
|
||||
env.close()
|
||||
del env
|
||||
|
||||
def main(render = False):
|
||||
def main(render = True):
|
||||
# # Standard DMC Suite tasks
|
||||
example_dmc("dm_control/fish-swim", seed=10, iterations=1000, render=render)
|
||||
#
|
||||
|
||||
@@ -21,7 +21,7 @@ def example_general(env_id="Pendulum-v1", seed=1, iterations=1000, render=True):
|
||||
|
||||
"""
|
||||
|
||||
env = gym.make(env_id, render_mode='human' if render else None)
|
||||
env = gym.make(env_id)
|
||||
rewards = 0
|
||||
obs = env.reset(seed=seed)
|
||||
print("Observation shape: ", env.observation_space.shape)
|
||||
@@ -85,7 +85,7 @@ def example_async(env_id="fancy/HoleReacher-v0", n_cpu=4, seed=int('533D', 16),
|
||||
# do not return values above threshold
|
||||
return *map(lambda v: np.stack(v)[:n_samples], buffer.values()),
|
||||
|
||||
def main(render = False):
|
||||
def main(render = True):
|
||||
# Basic gym task
|
||||
example_general("Pendulum-v1", seed=10, iterations=200, render=render)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import gymnasium as gym
|
||||
import fancy_gym
|
||||
|
||||
|
||||
def example_meta(env_id="metaworld/button-press-v2", seed=1, iterations=1000, render=True):
|
||||
def example_meta(env_id="fish-swim", seed=1, iterations=1000, render=True):
|
||||
"""
|
||||
Example for running a MetaWorld based env in the step based setting.
|
||||
The env_id has to be specified as `task_name-v2`. V1 versions are not supported and we always
|
||||
@@ -18,7 +18,7 @@ def example_meta(env_id="metaworld/button-press-v2", seed=1, iterations=1000, re
|
||||
Returns:
|
||||
|
||||
"""
|
||||
env = gym.make(env_id, render_mode='human' if render else None)
|
||||
env = gym.make(env_id)
|
||||
rewards = 0
|
||||
obs = env.reset(seed=seed)
|
||||
print("observation shape:", env.observation_space.shape)
|
||||
@@ -27,7 +27,9 @@ def example_meta(env_id="metaworld/button-press-v2", seed=1, iterations=1000, re
|
||||
for i in range(iterations):
|
||||
ac = env.action_space.sample()
|
||||
if render:
|
||||
env.render()
|
||||
# THIS NEEDS TO BE SET TO FALSE FOR NOW, BECAUSE THE INTERFACE FOR RENDERING IS DIFFERENT TO BASIC GYM
|
||||
# TODO: Remove this, when Metaworld fixes its interface.
|
||||
env.render(False)
|
||||
obs, reward, terminated, truncated, info = env.step(ac)
|
||||
rewards += reward
|
||||
if terminated or truncated:
|
||||
@@ -79,7 +81,7 @@ def example_custom_meta_and_mp(seed=1, iterations=1, render=True):
|
||||
basis_generator_kwargs = {'basis_generator_type': 'rbf',
|
||||
'num_basis': 5
|
||||
}
|
||||
base_env = gym.make(base_env_id, render_mode='human' if render else None)
|
||||
base_env = gym.make(base_env_id)
|
||||
env = fancy_gym.make_bb(env=base_env, wrappers=wrappers, black_box_kwargs={},
|
||||
traj_gen_kwargs=trajectory_generator_kwargs, controller_kwargs=controller_kwargs,
|
||||
phase_kwargs=phase_generator_kwargs, basis_kwargs=basis_generator_kwargs,
|
||||
@@ -91,7 +93,7 @@ def example_custom_meta_and_mp(seed=1, iterations=1, render=True):
|
||||
# It is also possible to change them mode multiple times when
|
||||
# e.g. only every nth trajectory should be displayed.
|
||||
if render:
|
||||
env.render()
|
||||
env.render(mode="human")
|
||||
|
||||
rewards = 0
|
||||
obs = env.reset(seed=seed)
|
||||
|
||||
@@ -13,13 +13,15 @@ def example_mp(env_name, seed=1, render=True):
|
||||
Returns:
|
||||
|
||||
"""
|
||||
env = gym.make(env_name, render_mode='human' if render else None)
|
||||
env = gym.make(env_name)
|
||||
|
||||
returns = 0
|
||||
obs = env.reset(seed=seed)
|
||||
# number of samples/full trajectories (multiple environment steps)
|
||||
for i in range(10):
|
||||
if render and i % 2 == 0:
|
||||
env.render(mode="human")
|
||||
else:
|
||||
env.render()
|
||||
ac = env.action_space.sample()
|
||||
obs, reward, terminated, truncated, info = env.step(ac)
|
||||
|
||||
@@ -21,15 +21,21 @@ GYM_MP_IDS = fancy_gym.ALL_DMC_MOVEMENT_PRIMITIVE_ENVIRONMENTS['all']
|
||||
SEED = 1
|
||||
|
||||
|
||||
known_fail_functionality = ['LunarLander-v2', 'Blackjack-v1', 'CliffWalking-v0']
|
||||
@pytest.mark.parametrize('env_id', GYM_IDS)
|
||||
def test_step_gym_functionality(env_id: str):
|
||||
"""Tests that step environments run without errors using random actions."""
|
||||
if env_id in known_fail_functionality:
|
||||
pytest.xfail(f"{env_id} is expected to fail the functionality test")
|
||||
run_env(env_id)
|
||||
|
||||
|
||||
known_fail_deteminism = ['LunarLanderContinuous-v2', 'CliffWalking-v0']
|
||||
@pytest.mark.parametrize('env_id', GYM_IDS)
|
||||
def test_step_gym_determinism(env_id: str):
|
||||
"""Tests that for step environments identical seeds produce identical trajectories."""
|
||||
if env_id in known_fail_deteminism:
|
||||
pytest.xfail(f"{env_id} is expected to fail the determinism test")
|
||||
run_env_determinism(env_id, SEED)
|
||||
|
||||
|
||||
|
||||
@@ -15,15 +15,21 @@ DMC_MP_IDS = fancy_gym.ALL_DMC_MOVEMENT_PRIMITIVE_ENVIRONMENTS['all']
|
||||
SEED = 1
|
||||
|
||||
|
||||
known_fail_functionality = []
|
||||
@pytest.mark.parametrize('env_id', DMC_IDS)
|
||||
def test_step_dm_control_functionality(env_id: str):
|
||||
"""Tests that suite step environments run without errors using random actions."""
|
||||
if env_id in known_fail_functionality:
|
||||
pytest.xfail(f"{env_id} is expected to fail the functionality test")
|
||||
run_env(env_id, 5000, wrappers=[gym.wrappers.FlattenObservation])
|
||||
|
||||
|
||||
known_fail_deteminism = ['dm_control/CmuHumanoidMazeForage-v0', 'dm_control/CmuHumanoidHeterogeneousForage-v0', 'dm_control/RodentMazeForage-v0', 'dm_control/RodentTwoTouch-v0']
|
||||
@pytest.mark.parametrize('env_id', DMC_IDS)
|
||||
def test_step_dm_control_determinism(env_id: str):
|
||||
"""Tests that for step environments identical seeds produce identical trajectories."""
|
||||
if env_id in known_fail_deteminism:
|
||||
pytest.xfail(f"{env_id} is expected to fail the determinism test")
|
||||
run_env_determinism(env_id, SEED, 5000, wrappers=[gym.wrappers.FlattenObservation])
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user