| 17 | |
| 18 | # Base class for RL tasks |
| 19 | class BaseTask(): |
| 20 | def __init__(self, config, device): |
| 21 | self.config = config |
| 22 | # optimization flags for pytorch JIT |
| 23 | torch._C._jit_set_profiling_mode(False) |
| 24 | torch._C._jit_set_profiling_executor(False) |
| 25 | |
| 26 | # self.simulator = instantiate(config=self.config.simulator, device=device) |
| 27 | SimulatorClass = get_class(self.config.simulator._target_) |
| 28 | self.simulator: BaseSimulator = SimulatorClass(config=self.config, device=device) |
| 29 | |
| 30 | self.headless = config.headless |
| 31 | self.simulator.set_headless(self.headless) |
| 32 | self.simulator.setup() |
| 33 | self.device = self.simulator.sim_device |
| 34 | self.sim_dt = self.simulator.sim_dt |
| 35 | self.up_axis_idx = 2 |
| 36 | |
| 37 | self.dt = self.config.simulator.config.sim.control_decimation * self.sim_dt |
| 38 | self.max_episode_length_s = self.config.max_episode_length_s |
| 39 | self.max_episode_length = np.ceil(self.max_episode_length_s / self.dt) |
| 40 | |
| 41 | self.num_envs = self.config.num_envs |
| 42 | self.dim_obs = self.config.robot.policy_obs_dim |
| 43 | self.dim_critic_obs = self.config.robot.critic_obs_dim |
| 44 | self.dim_actions = self.config.robot.actions_dim |
| 45 | |
| 46 | terrain_mesh_type = self.config.terrain.mesh_type |
| 47 | self.simulator.setup_terrain(terrain_mesh_type) |
| 48 | self.setup_visualize_entities() |
| 49 | |
| 50 | # create envs, sim and viewer |
| 51 | self._load_assets() |
| 52 | self._get_env_origins() |
| 53 | self._create_envs() |
| 54 | self.dof_pos_limits, self.dof_vel_limits, self.torque_limits = self.simulator.get_dof_limits_properties() |
| 55 | self._setup_robot_body_indices() |
| 56 | # self._create_sim() |
| 57 | self.simulator.prepare_sim() |
| 58 | # if running with a viewer, set up keyboard shortcuts and camera |
| 59 | self.viewer = None |
| 60 | if self.headless == False: |
| 61 | self.debug_viz = False |
| 62 | self.simulator.setup_viewer() |
| 63 | self.viewer = self.simulator.viewer |
| 64 | self._init_buffers() |
| 65 | |
| 66 | if self.headless == False: |
| 67 | self.viewer = self.simulator.viewer |
| 68 | |
| 69 | |
| 70 | def _init_buffers(self): |
| 71 | self.obs_buf_dict = {} |
| 72 | self.rew_buf = torch.zeros(self.num_envs, device=self.device, dtype=torch.float) |
| 73 | self.reset_buf = torch.ones(self.num_envs, device=self.device, dtype=torch.long) |
| 74 | self.episode_length_buf = torch.zeros(self.num_envs, device=self.device, dtype=torch.long) |
| 75 | self.time_out_buf = torch.zeros(self.num_envs, device=self.device, dtype=torch.bool) |
| 76 | self.extras = {} |
nothing calls this directly
no outgoing calls
no test coverage detected