def train(self, num_episodes, save_model=SAVE_MODEL, save_video=SAVE_VIDEO): total_rewards = [] if save_video: self.video = VideoRecorder(self.env, f'{os.path.dirname(__file__)}/training.mp4', enabled=True) for episode in range(num_episodes): state, _ = self.env.reset() done = False total_reward = 0 log_probs = [] rewards = [] while not done: if RENDER_MODE == 'human': self.env.render() if save_video: self.video.capture_frame() state = torch.FloatTensor(state).unsqueeze(0).to(device) action_probs = self.policy_network(state) action = torch.multinomial(action_probs, 1).item() log_prob = torch.log(action_probs.squeeze(0)[action]) log_probs.append(log_prob) next_state, reward, done, _, _ = self.env.step(action) rewards.append(reward) state = next_state total_reward += reward total_rewards.append(total_reward) self.update_policy(log_probs, rewards) print(f"Episode {episode}, Total Reward: {total_reward}") if episode % 5 == 0 and episode > 0: print(f"Episode {episode}, Average Reward: {sum(total_rewards) / len(total_rewards)}") if save_model: torch.save(self.policy_network.state_dict(), self.model_path) print("Saved model to disk") del log_probs, rewards, state, action_probs, action gc.collect() if save_model: torch.save(self.policy_network.state_dict(), self.model_path) print("Saved model to disk") if save_video: self.video.close() self.env.close() return sum(total_rewards) / len(total_rewards)