(self, start_epoch=0)
| 160 | pass |
| 161 | |
| 162 | def start_training(self, start_epoch=0): |
| 163 | # self._current_path_builder = PathBuilder() |
| 164 | self.ready_env_ids = np.arange(self.env_num) |
| 165 | observations = self._start_new_rollout( |
| 166 | self.ready_env_ids |
| 167 | ) # Do it for support vec env |
| 168 | |
| 169 | self._current_path_builder = [ |
| 170 | PathBuilder() for _ in range(len(self.ready_env_ids)) |
| 171 | ] |
| 172 | |
| 173 | for epoch in gt.timed_for( |
| 174 | range(start_epoch, self.num_epochs), |
| 175 | save_itrs=True, |
| 176 | ): |
| 177 | self._start_epoch(epoch) |
| 178 | total_rews = np.array([0.0 for _ in range(len(self.ready_env_ids))]) |
| 179 | for steps_this_epoch in range(self.num_env_steps_per_epoch // self.env_num): |
| 180 | actions = self._get_action_and_info(observations) |
| 181 | |
| 182 | if type(actions) is tuple: |
| 183 | actions = actions[0] |
| 184 | |
| 185 | if self.render: |
| 186 | self.training_env.render() |
| 187 | |
| 188 | next_obs, raw_rewards, terminals, env_infos = self.training_env.step( |
| 189 | actions, self.ready_env_ids |
| 190 | ) |
| 191 | if self.no_terminal: |
| 192 | terminals = [False for _ in range(len(self.ready_env_ids))] |
| 193 | # self._n_env_steps_total += 1 |
| 194 | self._n_env_steps_total += len(self.ready_env_ids) |
| 195 | |
| 196 | rewards = raw_rewards |
| 197 | total_rews += raw_rewards |
| 198 | |
| 199 | self._handle_vec_step( |
| 200 | observations, |
| 201 | actions, |
| 202 | rewards, |
| 203 | next_obs, |
| 204 | np.array([False for _ in range(len(self.ready_env_ids))]) |
| 205 | if self.no_terminal |
| 206 | else terminals, |
| 207 | absorbings=[ |
| 208 | np.array([0.0, 0.0]) for _ in range(len(self.ready_env_ids)) |
| 209 | ], |
| 210 | env_infos=env_infos, |
| 211 | ) |
| 212 | if np.any(terminals): |
| 213 | env_ind_local = np.where(terminals)[0] |
| 214 | total_rews[env_ind_local] = 0.0 |
| 215 | if self.wrap_absorbing: |
| 216 | # raise NotImplementedError() |
| 217 | """ |
| 218 | If we wrap absorbing states, two additional |
| 219 | transitions must be added: (s_T, s_abs) and |
no test coverage detected