Spaces:
Running on Zero
Running on Zero
| import gymnasium as gym | |
| import numpy as np | |
| from envs.tasks.dmcontrol import cartpole, cheetah, walker, hopper, reacher, ball_in_cup, pendulum, fish, giraffe, spinner, jumper, finger | |
| from dm_control import suite | |
| suite._DOMAINS['giraffe'] = giraffe | |
| suite._DOMAINS['spinner'] = spinner | |
| suite._DOMAINS['jumper'] = jumper | |
| suite.ALL_TASKS = suite.ALL_TASKS + suite._get_tasks('custom') | |
| suite.TASKS_BY_DOMAIN = suite._get_tasks_by_domain(suite.ALL_TASKS) | |
| from dm_control.suite.wrappers import action_scale | |
| from envs.wrappers.timeout import Timeout | |
| from envs.wrappers.pixels import Pixels | |
| def get_obs_shape(env): | |
| obs_shp = [] | |
| for v in env.observation_spec().values(): | |
| try: | |
| shp = np.prod(v.shape) | |
| except: | |
| shp = 1 | |
| obs_shp.append(shp) | |
| return (int(np.sum(obs_shp)),) | |
| class DMControlWrapper: | |
| def __init__(self, env, domain): | |
| self.env = env | |
| self.camera_id = 2 if domain == 'quadruped' else 0 | |
| obs_shape = get_obs_shape(env) | |
| action_shape = env.action_spec().shape | |
| self.observation_space = gym.spaces.Box( | |
| low=np.full(obs_shape, -np.inf, dtype=np.float32), | |
| high=np.full(obs_shape, np.inf, dtype=np.float32), | |
| dtype=np.float32) | |
| self.action_space = gym.spaces.Box( | |
| low=np.full(action_shape, env.action_spec().minimum), | |
| high=np.full(action_shape, env.action_spec().maximum), | |
| dtype=env.action_spec().dtype) | |
| self.action_spec_dtype = env.action_spec().dtype | |
| self._cumulative_reward = 0 | |
| def unwrapped(self): | |
| return self.env | |
| def metadata(self): | |
| return None | |
| def info(self): | |
| return { | |
| 'terminated': False, | |
| 'truncated': False, | |
| 'success': float('nan'), | |
| 'score': self._cumulative_reward/1000, | |
| } | |
| def _obs_to_array(self, obs): | |
| return np.concatenate([v.flatten() for v in obs.values()], dtype=np.float32) | |
| def reset(self): | |
| self._cumulative_reward = 0 | |
| return self._obs_to_array(self.env.reset().observation), self.info | |
| def step(self, action): | |
| reward = 0 | |
| action = action.astype(self.action_spec_dtype) | |
| for _ in range(2): | |
| step = self.env.step(action) | |
| reward += step.reward | |
| self._cumulative_reward += reward | |
| return self._obs_to_array(step.observation), reward, False, False, self.info | |
| def render(self, width=224, height=224, camera_id=None): | |
| return self.env.physics.render(height, width, camera_id or self.camera_id) | |
| def close(self): | |
| self.env.close() | |
| def make_env(cfg): | |
| """ | |
| Make DMControl environment. | |
| Adapted from https://github.com/facebookresearch/drqv2 | |
| """ | |
| domain, task = cfg.task.replace('-', '_').split('_', 1) | |
| domain = dict(cup='ball_in_cup', pointmass='point_mass').get(domain, domain) | |
| if (domain, task) not in suite.ALL_TASKS: | |
| raise ValueError('Unknown task:', task) | |
| assert cfg.obs in {'state', 'rgb'}, 'This task only supports state and rgb observations.' | |
| env = suite.load(domain, | |
| task, | |
| task_kwargs={'random': cfg.seed}, | |
| visualize_reward=False) | |
| env = action_scale.Wrapper(env, minimum=-1., maximum=1.) | |
| env = DMControlWrapper(env, domain) | |
| if cfg.obs == 'rgb': | |
| env = Pixels(env, cfg) | |
| env = Timeout(env, max_episode_steps=500) | |
| return env | |