class REINFORCE: def __init__(self, env, policy_network, optimizer, model_path='model/model.pth', gamma=0.99): self.env = env self.policy_network = policy_network self.optimizer = optimizer self.model_path = model_path self.gamma = gamma if os.path.exists(os.path.dirname(self.model_path)): if os.path.isfile(self.model_path): self.policy_network.load_state_dict(torch.load(self.model_path)) print("Loaded model from disk") else: os.makedirs(os.path.dirname(self.model_path)) 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) def update_policy(self, log_probs, rewards): discounted_rewards = [] for t in range(len(rewards)): Gt = sum(self.gamma ** i * rewards[t + i] for i in range(len(rewards) - t)) discounted_rewards.append(Gt) discounted_rewards = torch.FloatTensor(discounted_rewards).to(device) discounted_rewards = (discounted_rewards - discounted_rewards.mean()) / (discounted_rewards.std() + 1e-9) policy_loss = [] for log_prob, Gt in zip(log_probs, discounted_rewards): policy_loss.append(-log_prob * Gt) self.optimizer.zero_grad() policy_loss = torch.stack(policy_loss).sum() policy_loss.backward() self.optimizer.step()